1#ifndef TATAMI_MULT_MULTIPLE_VECTORS_DENSE_COLUMN_HPP
2#define TATAMI_MULT_MULTIPLE_VECTORS_DENSE_COLUMN_HPP
10#include "sanisizer/sanisizer.hpp"
11#include "jiwoo/jiwoo.hpp"
13#include "../utils.hpp"
54template<
typename LeftValue_,
typename LeftIndex_,
typename RightVectors_,
typename GetRightVector_,
typename GetOutputVector_>
55void multiply_dense_column_with_multiple_vectors_internal(
57 const LeftIndex_ start,
58 const LeftIndex_ length,
59 const LeftIndex_ left_NR,
60 const RightVectors_ right_vectors,
61 GetRightVector_ get_right_vector,
62 GetOutputVector_ get_output_vector,
66 typedef I<
decltype(get_output_vector(0)[0])> Output;
70 for (LeftIndex_ cd = 0; cd < length; ++cd) {
71 const auto ptr = ext->fetch(buffer.data());
72 for (RightVectors_ rv = 0; rv < right_vectors; ++rv) {
73 const auto optr = get_output_vector(rv);
74 const Output mult = get_right_vector(rv)[start + cd];
75 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
76 optr[lr] += mult *
static_cast<Output
>(ptr[lr]);
82 std::vector<std::vector<LeftValue_> > left_buffers;
83 std::vector<const LeftValue_*> left_ptrs;
86 left_buffers.reserve(max_block_cols);
87 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
90 sanisizer::resize(left_ptrs, max_block_cols);
96 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
97 left_ptrs[cd_counter] = ext->fetch(left_buffers[cd_counter].data());
100 RightVectors_ rv = 0;
101 while (rv < right_vectors) {
102 const RightVectors_ rv_end = rv + sanisizer::min(options.
primary_block_size, right_vectors - rv);
104 while (lr < left_NR) {
107 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
108 const auto matcol = left_ptrs[cd_counter];
109 for (
auto rv_copy = rv; rv_copy < rv_end; ++rv_copy) {
110 const Output mult = get_right_vector(rv_copy)[start + cd + cd_counter];
111 const auto outvec = get_output_vector(rv_copy);
112 for (
auto lr_copy = lr; lr_copy < lr_end; ++lr_copy) {
113 outvec[lr_copy] += mult *
static_cast<Output
>(matcol[lr_copy]);
147template<
typename LeftValue_,
typename LeftIndex_,
typename RightVectors_,
typename GetRightVector_,
typename GetOutput_>
150 const RightVectors_ right_vectors,
151 GetRightVector_ get_right_vector,
152 GetOutput_ get_output_vector,
155 const auto left_NR = left.
nrow();
156 const auto common_dim = left.
ncol();
157 for (RightVectors_ rv = 0; rv < right_vectors; ++rv) {
158 std::fill_n(get_output_vector(rv), left_NR, 0);
162 typedef I<
decltype(get_output_vector(0)[0])> Output;
163 std::optional<std::vector<std::optional<std::vector<std::vector<Output> > > > > tmp_results;
165 tmp_results.emplace(sanisizer::cast<I<
decltype(tmp_results->size())> >(options.
num_threads - 1));
168 const auto num_used =
tatami::parallelize([&](
int t, LeftIndex_ start, LeftIndex_ length) ->
void {
169 if (!do_parallel || t == 0) {
170 multiply_dense_column_with_multiple_vectors_internal(
182 std::vector<std::vector<Output> > tmp_output;
183 tmp_output.reserve(right_vectors);
184 for (RightVectors_ rv = 0; rv < right_vectors; ++rv) {
187 multiply_dense_column_with_multiple_vectors_internal(
194 [&](
const RightVectors_ rv) -> Output* {
195 return tmp_output[rv].data();
199 (*tmp_results)[t - 1] = std::move(tmp_output);
204 for (
int u = 1; u < num_used; ++u) {
205 const auto& tmp = *((*tmp_results)[u - 1]);
206 for (RightVectors_ rv = 0; rv < right_vectors; ++rv) {
207 const auto& tmpvec = tmp[rv];
208 const auto outptr = get_output_vector(rv);
209 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
210 outptr[lr] += tmpvec[lr];
234template<
typename LeftValue_,
typename LeftIndex_,
typename RightValue_,
typename Output_>
237 const std::vector<RightValue_*>& right,
238 const std::vector<Output_*>& output,
241 const auto left_NR = left.
nrow();
242 const auto common_dim = left.
ncol();
243 const auto right_vectors = right.size();
244 typedef I<
decltype(right_vectors)> RightVectors;
245 for (RightVectors rv = 0; rv < right_vectors; ++rv) {
246 std::fill_n(output[rv], left_NR, 0);
250 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > tmp_results;
252 tmp_results.emplace(sanisizer::cast<I<
decltype(tmp_results->size())> >(options.
num_threads - 1));
255 const auto num_used =
tatami::parallelize([&](
int t, LeftIndex_ start, LeftIndex_ length) ->
void {
256 std::optional<jiwoo::EquilengthArrays<Output_> > tmp_output;
258 Output_*
const * output_ptrs;
259 if (!do_parallel || t == 0) {
260 output_ptrs = output.data();
263 sanisizer::cast<I<
decltype(tmp_output->size())> >(right_vectors),
264 static_cast<std::size_t
>(left_NR),
267 output_ptrs = tmp_output->get();
270 multiply_dense_column_with_multiple_vectors_internal(
276 [&](
const RightVectors rv) ->
const RightValue_* {
279 [&](
const RightVectors rv) -> Output_* {
280 return output_ptrs[rv];
285 if (do_parallel && t > 0) {
286 (*tmp_results)[t - 1] = std::move(tmp_output);
291 for (
int u = 1; u < num_used; ++u) {
292 const auto& tmp = *((*tmp_results)[u - 1]);
293 for (RightVectors rv = 0; rv < right_vectors; ++rv) {
294 const auto tmpvec = tmp[rv];
295 const auto outptr = output[rv];
296 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
297 outptr[lr] += tmpvec[lr];
virtual Index_ ncol() const=0
virtual Index_ nrow() const=0
Multiplication of tatami matrices.
Definition column_to_column.hpp:19
void multiply_dense_column_with_multiple_vectors(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const RightVectors_ right_vectors, GetRightVector_ get_right_vector, GetOutput_ get_output_vector, const MultiplyDenseColumnWithMultipleVectorsOptions &options)
Definition dense_column.hpp:148
int parallelize(Function_ fun, const Index_ tasks, const int workers)
I< decltype(std::declval< Container_ >().size())> cast_Index_to_container_size(const Index_ x)
Container_ create_container_of_Index_size(const Index_ x, Args_ &&... args)
auto consecutive_extractor(const Matrix< Value_, Index_ > &matrix, const bool row, const Index_ iter_start, const Index_ iter_length, Args_ &&... args)
Options for multiply_dense_column_with_multiple_vectors().
Definition dense_column.hpp:29
int primary_block_size
Definition dense_column.hpp:41
int secondary_block_size
Definition dense_column.hpp:48
int num_threads
Definition dense_column.hpp:34