tatami_mult
Multiply tatami matrices
Loading...
Searching...
No Matches
row_to_column.hpp
Go to the documentation of this file.
1#ifndef TATAMI_MULT_DENSE_MATRIX_DENSE_COLUMN_ROW_TO_COLUMN_HPP
2#define TATAMI_MULT_DENSE_MATRIX_DENSE_COLUMN_ROW_TO_COLUMN_HPP
3
4#include <cstddef>
5#include <vector>
6#include <optional>
7
8#include "tatami/tatami.hpp"
9#include "sanisizer/sanisizer.hpp"
10
11#include "../../utils.hpp"
12
18namespace tatami_mult {
19
20/* See https://github.com/tatami-inc/test-multiplication/tree/master/dense_column/dense_matrix
21 * for an explanation of the choice of algorithm.
22 */
23
48
66template<typename LeftValue_, typename LeftIndex_, typename RightColumns_, typename GetRightRow_, typename Output_>
69 const RightColumns_ right_columns,
70 GetRightRow_ get_right_row,
71 Output_* const output,
73) {
74 const auto left_NR = left.nrow();
75 const auto common_dim = left.ncol();
76
77 // Product must fit in a size_t in order for output to have been allocated correctly in the first place.
78 // Technically, right_columns could be larger than a size_t if left_NR == 0, but the product after wraparound would still be zero, so it's fine.
79 std::fill_n(output, sanisizer::product_unsafe<std::size_t>(left_NR, right_columns), 0);
80
81 const bool do_parallel = options.num_threads > 1;
82 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
83 if (do_parallel) {
84 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
85 }
86
87 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
88 auto left_ext = tatami::consecutive_extractor<false>(left, false, start, length);
89
90 std::optional<std::vector<Output_> > tmp_output;
91 Output_* outptr;
92 if (!do_parallel || t == 0) {
93 outptr = output;
94 } else {
95 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_columns));
96 outptr = tmp_output->data();
97 }
98
99 if (options.primary_block_size == 1) {
101 for (LeftIndex_ cd = 0; cd < length; ++cd) {
102 const auto left_ptr = left_ext->fetch(left_buffer.data());
103 const auto right_ptr = get_right_row(start + cd);
104 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
105 const Output_ mult = right_ptr[rc];
106 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
107 outptr[sanisizer::nd_offset<std::size_t>(lr, left_NR, rc)] += mult * static_cast<Output_>(left_ptr[lr]);
108 }
109 }
110 }
111
112 } else {
113 std::vector<std::vector<LeftValue_> > left_buffers;
114 std::vector<const LeftValue_*> left_ptrs;
115 {
116 const LeftIndex_ max_block_cols = sanisizer::min(length, options.primary_block_size);
117 left_buffers.reserve(max_block_cols);
118 tatami::resize_container_to_Index_size(left_ptrs, max_block_cols);
119 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
120 left_buffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(left_NR));
121 }
122 }
123
124 LeftIndex_ cd = 0;
125 while (cd < length) {
126 const auto cd_num = sanisizer::min(options.primary_block_size, length - cd);
127 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
128 left_ptrs[cd_counter] = left_ext->fetch(left_buffers[cd_counter].data());
129 }
130
131 RightColumns_ rc = 0;
132 while (rc < right_columns) {
133 const RightColumns_ rc_end = rc + sanisizer::min(options.primary_block_size, right_columns - rc);
134 LeftIndex_ lr = 0;
135 while (lr < left_NR) {
136 const LeftIndex_ lr_end = lr + sanisizer::min(options.secondary_block_size, left_NR - lr);
137
138 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
139 const auto leftcol = left_ptrs[cd_counter];
140 const auto rightrow = get_right_row(start + cd + cd_counter);
141 for (auto rc_copy = rc; rc_copy < rc_end; ++rc_copy) {
142 const Output_ mult = rightrow[rc_copy];
143 for (auto lr_copy = lr; lr_copy < lr_end; ++lr_copy) {
144 outptr[sanisizer::nd_offset<std::size_t>(lr_copy, left_NR, rc_copy)] += mult * static_cast<Output_>(leftcol[lr_copy]);
145 }
146 }
147 }
148
149 lr = lr_end;
150 }
151 rc = rc_end;
152 }
153 cd += cd_num;
154 }
155 }
156
157 if (do_parallel && t > 0) {
158 (*tmp_results)[t - 1] = std::move(tmp_output);
159 }
160 }, common_dim, options.num_threads);
161
162 if (do_parallel) {
163 for (int u = 1; u < num_used; ++u) {
164 const auto& tmp = *((*tmp_results)[u - 1]);
165 const auto N = tmp.size();
166 for (I<decltype(N)> x = 0; x < N; ++x) {
167 output[x] += tmp[x];
168 }
169 }
170 }
171}
172
192template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename RightIndex_, typename Output_>
196 Output_* const output,
198) {
199 const auto left_NR = left.nrow();
200 const auto common_dim = left.ncol();
201 const auto right_NC = right.ncol();
202 std::fill_n(output, sanisizer::product_unsafe<std::size_t>(left_NR, right_NC), 0);
203
204 const bool do_parallel = options.num_threads > 1;
205 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
206 if (do_parallel) {
207 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
208 }
209
210 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
211 auto left_ext = tatami::consecutive_extractor<false>(left, false, start, length);
212 auto right_ext = tatami::consecutive_extractor<false>(right, true, start, length);
213
214 std::optional<std::vector<Output_> > tmp_output;
215 Output_* outptr;
216 if (!do_parallel || t == 0) {
217 outptr = output;
218 } else {
219 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_NC));
220 outptr = tmp_output->data();
221 }
222
223 if (options.primary_block_size == 1) {
226
227 for (LeftIndex_ cd = 0; cd < length; ++cd) {
228 const auto left_ptr = left_ext->fetch(left_buffer.data());
229 const auto right_ptr = right_ext->fetch(right_buffer.data());
230 for (RightIndex_ rc = 0; rc < right_NC; ++rc) {
231 const Output_ mult = right_ptr[rc];
232 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
233 outptr[sanisizer::nd_offset<std::size_t>(lr, left_NR, rc)] += mult * static_cast<Output_>(left_ptr[lr]);
234 }
235 }
236 }
237
238 } else {
239 std::vector<std::vector<LeftValue_> > left_buffers;
240 std::vector<std::vector<RightValue_> > right_buffers;
241 std::vector<const LeftValue_*> left_ptrs;
242 std::vector<const RightValue_*> right_ptrs;
243 {
244 const LeftIndex_ max_block_cols = sanisizer::min(length, options.primary_block_size);
245 left_buffers.reserve(max_block_cols);
246 tatami::resize_container_to_Index_size(left_ptrs, max_block_cols);
247 right_buffers.reserve(max_block_cols);
248 tatami::resize_container_to_Index_size(right_ptrs, max_block_cols);
249 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
250 left_buffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(left_NR));
251 right_buffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<RightValue_> >(right_NC));
252 }
253 }
254
255 LeftIndex_ cd = 0;
256 while (cd < length) {
257 const auto cd_num = sanisizer::min(options.primary_block_size, length - cd);
258 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
259 left_ptrs[cd_counter] = left_ext->fetch(left_buffers[cd_counter].data());
260 right_ptrs[cd_counter] = right_ext->fetch(right_buffers[cd_counter].data());
261 }
262
263 RightIndex_ rc = 0;
264 while (rc < right_NC) {
265 const RightIndex_ rc_end = rc + sanisizer::min(options.primary_block_size, right_NC - rc);
266 LeftIndex_ lr = 0;
267 while (lr < left_NR) {
268 const LeftIndex_ lr_end = lr + sanisizer::min(options.secondary_block_size, left_NR - lr);
269
270 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
271 const auto leftcol = left_ptrs[cd_counter];
272 const auto rightrow = right_ptrs[cd_counter];
273 for (auto rc_copy = rc; rc_copy < rc_end; ++rc_copy) {
274 const Output_ mult = rightrow[rc_copy];
275 for (auto lr_copy = lr; lr_copy < lr_end; ++lr_copy) {
276 outptr[sanisizer::nd_offset<std::size_t>(lr_copy, left_NR, rc_copy)] += mult * static_cast<Output_>(leftcol[lr_copy]);
277 }
278 }
279 }
280
281 lr = lr_end;
282 }
283 rc = rc_end;
284 }
285 cd += cd_num;
286 }
287 }
288
289 if (do_parallel && t > 0) {
290 (*tmp_results)[t - 1] = std::move(tmp_output);
291 }
292 }, common_dim, options.num_threads);
293
294 if (do_parallel) {
295 for (int u = 1; u < num_used; ++u) {
296 const auto& tmp = *((*tmp_results)[u - 1]);
297 const auto N = tmp.size();
298 for (I<decltype(N)> x = 0; x < N; ++x) {
299 output[x] += tmp[x];
300 }
301 }
302 }
303}
304
305}
306
307#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_dense_row_matrix_to_column_output(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const RightColumns_ right_columns, GetRightRow_ get_right_row, Output_ *const output, const MultiplyDenseColumnWithDenseRowMatrixToColumnOutputOptions &options)
Definition row_to_column.hpp:67
void resize_container_to_Index_size(Container_ &container, const Index_ x, Args_ &&... args)
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_dense_row_matrix_to_column_output().
Definition row_to_column.hpp:27