tatami_mult
Multiply tatami matrices
Loading...
Searching...
No Matches
dense_column.hpp
Go to the documentation of this file.
1#ifndef TATAMI_MULT_MULTIPLE_VECTORS_DENSE_COLUMN_HPP
2#define TATAMI_MULT_MULTIPLE_VECTORS_DENSE_COLUMN_HPP
3
4#include <cstddef>
5#include <vector>
6#include <optional>
7#include <algorithm>
8
9#include "tatami/tatami.hpp"
10#include "sanisizer/sanisizer.hpp"
11#include "jiwoo/jiwoo.hpp"
12
13#include "../utils.hpp"
14
20namespace tatami_mult {
21
22/* See https://github.com/tatami-inc/test-multiplication/tree/master/dense_column/multiple_vectors
23 * for an explanation of the choice of algorithm.
24 */
25
50
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,
64) {
65 auto ext = tatami::consecutive_extractor<false>(left, false, start, length);
66 typedef I<decltype(get_output_vector(0)[0])> Output;
67
68 if (options.primary_block_size == 1) {
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]);
77 }
78 }
79 }
80
81 } else {
82 std::vector<std::vector<LeftValue_> > left_buffers;
83 std::vector<const LeftValue_*> left_ptrs;
84 {
85 const LeftIndex_ max_block_cols = sanisizer::min(length, options.primary_block_size);
86 left_buffers.reserve(max_block_cols);
87 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
88 left_buffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(left_NR));
89 }
90 sanisizer::resize(left_ptrs, max_block_cols);
91 }
92
93 LeftIndex_ cd = 0;
94 while (cd < length) {
95 const LeftIndex_ cd_num = sanisizer::min(options.primary_block_size, length - cd);
96 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
97 left_ptrs[cd_counter] = ext->fetch(left_buffers[cd_counter].data());
98 }
99
100 RightVectors_ rv = 0;
101 while (rv < right_vectors) {
102 const RightVectors_ rv_end = rv + sanisizer::min(options.primary_block_size, right_vectors - rv);
103 LeftIndex_ lr = 0;
104 while (lr < left_NR) {
105 const LeftIndex_ lr_end = lr + sanisizer::min(options.secondary_block_size, left_NR - lr);
106
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]);
114 }
115 }
116 }
117
118 lr = lr_end;
119 }
120 rv = rv_end;
121 }
122 cd += cd_num;
123 }
124 }
125}
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,
154) {
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);
159 }
160
161 const bool do_parallel = options.num_threads > 1;
162 typedef I<decltype(get_output_vector(0)[0])> Output;
163 std::optional<std::vector<std::optional<std::vector<std::vector<Output> > > > > tmp_results;
164 if (do_parallel) {
165 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
166 }
167
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(
171 left,
172 start,
173 length,
174 left_NR,
175 right_vectors,
176 get_right_vector,
177 get_output_vector,
178 options
179 );
180
181 } else {
182 std::vector<std::vector<Output> > tmp_output;
183 tmp_output.reserve(right_vectors);
184 for (RightVectors_ rv = 0; rv < right_vectors; ++rv) {
185 tmp_output.emplace_back(tatami::cast_Index_to_container_size<std::vector<Output> >(left_NR));
186 }
187 multiply_dense_column_with_multiple_vectors_internal(
188 left,
189 start,
190 length,
191 left_NR,
192 right_vectors,
193 get_right_vector,
194 [&](const RightVectors_ rv) -> Output* {
195 return tmp_output[rv].data();
196 },
197 options
198 );
199 (*tmp_results)[t - 1] = std::move(tmp_output);
200 }
201 }, common_dim, options.num_threads);
202
203 if (do_parallel) {
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];
211 }
212 }
213 }
214 }
215}
216
234template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename Output_>
237 const std::vector<RightValue_*>& right,
238 const std::vector<Output_*>& output,
240) {
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);
247 }
248
249 const bool do_parallel = options.num_threads > 1;
250 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > tmp_results;
251 if (do_parallel) {
252 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
253 }
254
255 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
256 std::optional<jiwoo::EquilengthArrays<Output_> > tmp_output;
257
258 Output_* const * output_ptrs;
259 if (!do_parallel || t == 0) {
260 output_ptrs = output.data();
261 } else {
262 tmp_output.emplace(
263 sanisizer::cast<I<decltype(tmp_output->size())> >(right_vectors),
264 static_cast<std::size_t>(left_NR), // cast to size_t is safe due to the tatami contract.
265 0
266 );
267 output_ptrs = tmp_output->get();
268 }
269
270 multiply_dense_column_with_multiple_vectors_internal(
271 left,
272 start,
273 length,
274 left_NR,
275 right_vectors,
276 [&](const RightVectors rv) -> const RightValue_* {
277 return right[rv];
278 },
279 [&](const RightVectors rv) -> Output_* {
280 return output_ptrs[rv];
281 },
282 options
283 );
284
285 if (do_parallel && t > 0) {
286 (*tmp_results)[t - 1] = std::move(tmp_output);
287 }
288 }, common_dim, options.num_threads);
289
290 if (do_parallel) {
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];
298 }
299 }
300 }
301 }
302}
303
304}
305
306#endif
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