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_SPARSE_COLUMN_ROW_TO_COLUMN_HPP
2#define TATAMI_MULT_DENSE_MATRIX_SPARSE_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/sparse_column/dense_matrix
21 * for an explanation of the choice of algorithm.
22 */
23
40
58template<typename LeftValue_, typename LeftIndex_, typename RightColumns_, typename GetRightRow_, typename Output_>
61 const RightColumns_ right_columns,
62 GetRightRow_ get_right_row,
63 Output_* const output,
65) {
66 const auto left_NR = left.nrow();
67 const auto common_dim = left.ncol();
68
69 // Product must fit in a size_t in order for output to have been allocated correctly in the first place.
70 // 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.
71 std::fill_n(output, sanisizer::product<std::size_t>(left_NR, right_columns), 0);
72
73 const bool do_parallel = options.num_threads > 1;
74 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
75 if (do_parallel) {
76 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
77 }
78
79 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
80 auto left_ext = tatami::consecutive_extractor<true>(left, false, start, length);
81
82 std::optional<std::vector<Output_> > tmp_output;
83 Output_* outptr;
84 if (!do_parallel || t == 0) {
85 outptr = output;
86 } else {
87 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_columns));
88 outptr = tmp_output->data();
89 }
90
91 if (options.block_size == 1) {
94 for (LeftIndex_ cd = 0; cd < length; ++cd) {
95 const auto lrange = left_ext->fetch(vbuffer.data(), ibuffer.data());
96 if (lrange.number == 0) {
97 continue;
98 }
99 const auto rptr = get_right_row(start + cd);
100 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
101 const Output_ mult = rptr[rc];
102 for (LeftIndex_ x = 0; x < lrange.number; ++x) {
103 outptr[sanisizer::nd_offset<std::size_t>(lrange.index[x], left_NR, rc)] += mult * static_cast<Output_>(lrange.value[x]);
104 }
105 }
106 }
107
108 } else {
109 // Our blocking strategy is to collect multiple LHS columns so that, for each RHS vector,
110 // we can keep the corresponding output vector in cache for re-use with each LHS column.
111 std::vector<std::vector<LeftValue_> > left_vbuffers;
112 std::vector<std::vector<LeftIndex_> > left_ibuffers;
113 std::vector<tatami::SparseRange<LeftValue_, LeftIndex_> > left_ranges;
114 std::vector<LeftIndex_> left_non_empty;
115 {
116 const LeftIndex_ max_block_cols = sanisizer::min(length, options.block_size);
117 left_vbuffers.reserve(max_block_cols);
118 left_ibuffers.reserve(max_block_cols);
119 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
120 left_vbuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(left_NR));
121 left_ibuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftIndex_> >(left_NR));
122 }
123 tatami::resize_container_to_Index_size(left_ranges, max_block_cols);
124 left_non_empty.reserve(max_block_cols);
125 }
126
127 LeftIndex_ cd = 0;
128 while (cd < length) {
129 // Only processing LHS columns (and the corresponding RHS rows) if the LHS column has some structural non-zeros.
130 // If not, we just skip it altogether; no need to zero or do anything else, as we're skipping the corresponding RHS row too.
131 LeftIndex_ cd_num = 0, cd_copy = cd;
132 bool left_all_non_empty = true;
133 left_non_empty.clear();
134 do {
135 auto lrange = left_ext->fetch(left_vbuffers[cd_num].data(), left_ibuffers[cd_num].data());
136 if (lrange.number == 0) {
137 ++cd_copy;
138 left_all_non_empty = false;
139 continue;
140 }
141
142 left_ranges[cd_num] = std::move(lrange);
143 left_non_empty.push_back(cd_copy);
144 ++cd_num;
145 ++cd_copy;
146
147 if (sanisizer::is_equal(cd_num, options.block_size)) {
148 break;
149 }
150 } while (cd_copy < length);
151
152 if (left_all_non_empty) {
153 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
154 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
155 const auto& currange = left_ranges[cd_counter];
156 const Output_ mult = get_right_row(start + cd + cd_counter)[rc];
157 for (LeftIndex_ x = 0; x < currange.number; ++x) {
158 outptr[sanisizer::nd_offset<std::size_t>(currange.index[x], left_NR, rc)] += mult * static_cast<Output_>(currange.value[x]);
159 }
160 }
161 }
162 } else {
163 for (auto& cdne : left_non_empty) {
164 cdne += start;
165 }
166 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
167 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
168 const auto& currange = left_ranges[cd_counter];
169 const Output_ mult = get_right_row(left_non_empty[cd_counter])[rc];
170 for (LeftIndex_ x = 0; x < currange.number; ++x) {
171 outptr[sanisizer::nd_offset<std::size_t>(currange.index[x], left_NR, rc)] += mult * static_cast<Output_>(currange.value[x]);
172 }
173 }
174 }
175 }
176
177 cd = cd_copy;
178 }
179 }
180
181 if (do_parallel && t > 0) {
182 (*tmp_results)[t - 1] = std::move(tmp_output);
183 }
184 }, common_dim, options.num_threads);
185
186 if (do_parallel) {
187 for (int u = 1; u < num_used; ++u) {
188 const auto& tmp = *((*tmp_results)[u - 1]);
189 const auto N = tmp.size();
190 for (I<decltype(N)> x = 0; x < N; ++x) {
191 output[x] += tmp[x];
192 }
193 }
194 }
195}
196
216template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename RightIndex_, typename Output_>
220 Output_* const output,
222) {
223 const auto left_NR = left.nrow();
224 const auto common_dim = left.ncol();
225 const auto right_NC = right.ncol();
226 std::fill_n(output, sanisizer::product<std::size_t>(left_NR, right_NC), 0);
227
228 const bool do_parallel = options.num_threads > 1;
229 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
230 if (do_parallel) {
231 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
232 }
233
234 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
235 auto left_ext = tatami::consecutive_extractor<true>(left, false, start, length);
236 auto right_ext = tatami::consecutive_extractor<false>(right, true, start, length);
237
238 std::optional<std::vector<Output_> > tmp_output;
239 Output_* outptr;
240 if (!do_parallel || t == 0) {
241 outptr = output;
242 } else {
243 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_NC));
244 outptr = tmp_output->data();
245 }
246
247 if (options.block_size == 1) {
251
252 for (LeftIndex_ cd = 0; cd < length; ++cd) {
253 const auto lrange = left_ext->fetch(vbuffer.data(), ibuffer.data());
254 const auto rptr = right_ext->fetch(dbuffer.data());
255
256 // This skip must be done after right_ext->fetch(), otherwise the two extractors won't be in sync along the common dimension.
257 // No need to zero anything as the output buffer should already be zeroed at this point.
258 if (lrange.number == 0) {
259 continue;
260 }
261
262 for (RightIndex_ rc = 0; rc < right_NC; ++rc) {
263 const Output_ mult = rptr[rc];
264 for (LeftIndex_ x = 0; x < lrange.number; ++x) {
265 outptr[sanisizer::nd_offset<std::size_t>(lrange.index[x], left_NR, rc)] += mult * static_cast<Output_>(lrange.value[x]);
266 }
267 }
268 }
269
270 } else {
271 // Our blocking strategy is to collect multiple LHS columns so that, for each RHS vector,
272 // we can keep the corresponding output vector in cache for re-use with each LHS column.
273 std::vector<std::vector<LeftValue_> > left_vbuffers;
274 std::vector<std::vector<LeftIndex_> > left_ibuffers;
275 std::vector<tatami::SparseRange<LeftValue_, LeftIndex_> > left_ranges;
276 std::vector<std::vector<RightValue_> > right_dbuffers;
277 std::vector<const RightValue_*> right_ptrs;
278 {
279 const LeftIndex_ max_block_cols = sanisizer::min(length, options.block_size);
280 left_vbuffers.reserve(max_block_cols);
281 left_ibuffers.reserve(max_block_cols);
282 right_dbuffers.reserve(max_block_cols);
283 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
284 left_vbuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(left_NR));
285 left_ibuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftIndex_> >(left_NR));
286 right_dbuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<RightValue_> >(right_NC));
287 }
288 tatami::resize_container_to_Index_size(left_ranges, max_block_cols);
289 tatami::resize_container_to_Index_size(right_ptrs, max_block_cols);
290 }
291
292 LeftIndex_ cd = 0;
293 while (cd < length) {
294 // Only processing LHS columns (and the corresponding RHS rows) if the LHS column has some structural non-zeros.
295 // If not, we just skip it altogether; no need to zero or do anything else, as we're skipping the corresponding RHS row too.
296 LeftIndex_ cd_num = 0;
297 do {
298 auto lrange = left_ext->fetch(left_vbuffers[cd_num].data(), left_ibuffers[cd_num].data());
299 auto rptr = right_ext->fetch(right_dbuffers[cd_num].data());
300
301 // Again, this skip must be done after the RHS row is fetched, otherwise the extractors will be out of sync.
302 if (lrange.number == 0) {
303 ++cd;
304 continue;
305 }
306
307 left_ranges[cd_num] = std::move(lrange);
308 right_ptrs[cd_num] = rptr;
309 ++cd_num;
310 ++cd;
311
312 if (sanisizer::is_equal(cd_num, options.block_size)) {
313 break;
314 }
315 } while (cd < length);
316
317 for (RightIndex_ rc = 0; rc < right_NC; ++rc) {
318 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
319 const auto& currange = left_ranges[cd_counter];
320 const Output_ mult = right_ptrs[cd_counter][rc];
321 for (LeftIndex_ x = 0; x < currange.number; ++x) {
322 outptr[sanisizer::nd_offset<std::size_t>(currange.index[x], left_NR, rc)] += mult * static_cast<Output_>(currange.value[x]);
323 }
324 }
325 }
326 }
327 }
328
329 if (do_parallel && t > 0) {
330 (*tmp_results)[t - 1] = std::move(tmp_output);
331 }
332 }, common_dim, options.num_threads);
333
334 if (do_parallel) {
335 for (int u = 1; u < num_used; ++u) {
336 const auto& tmp = *((*tmp_results)[u - 1]);
337 const auto N = tmp.size();
338 for (I<decltype(N)> x = 0; x < N; ++x) {
339 output[x] += tmp[x];
340 }
341 }
342 }
343}
344
345}
346
347#endif
virtual Index_ ncol() const=0
virtual Index_ nrow() const=0
Multiplication of tatami matrices.
Definition column_to_column.hpp:19
void multiply_sparse_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 MultiplySparseColumnWithDenseRowMatrixToColumnOutputOptions &options)
Definition row_to_column.hpp:59
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_sparse_column_with_dense_row_matrix_to_column_output().
Definition row_to_column.hpp:27