tatami_mult
Multiply tatami matrices
Loading...
Searching...
No Matches
row_to_row.hpp
Go to the documentation of this file.
1#ifndef TATAMI_MULT_DENSE_MATRIX_SPARSE_COLUMN_ROW_TO_ROW_HPP
2#define TATAMI_MULT_DENSE_MATRIX_SPARSE_COLUMN_ROW_TO_ROW_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
33
51template<typename LeftValue_, typename LeftIndex_, typename RightColumns_, typename GetRightRow_, typename Output_>
54 const RightColumns_ right_columns,
55 GetRightRow_ get_right_row,
56 Output_* const output,
58) {
59 const auto left_NR = left.nrow();
60 const auto common_dim = left.ncol();
61
62 // Product must fit in a size_t in order for output to have been allocated correctly in the first place.
63 // 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.
64 std::fill_n(output, sanisizer::product<std::size_t>(left_NR, right_columns), 0);
65
66 const bool do_parallel = options.num_threads > 1;
67 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
68 if (do_parallel) {
69 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
70 }
71
72 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
73 auto left_ext = tatami::consecutive_extractor<true>(left, false, start, length);
76
77 std::optional<std::vector<Output_> > tmp_output;
78 Output_* outptr;
79 if (!do_parallel || t == 0) {
80 outptr = output;
81 } else {
82 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_columns));
83 outptr = tmp_output->data();
84 }
85
86 for (LeftIndex_ cd = 0; cd < length; ++cd) {
87 const auto lrange = left_ext->fetch(vbuffer.data(), ibuffer.data());
88 const auto rptr = get_right_row(start + cd);
89 for (LeftIndex_ x = 0; x < lrange.number; ++x) {
90 const Output_ mult = lrange.value[x];
91 const auto curout = outptr + sanisizer::product_unsafe<std::size_t>(lrange.index[x], right_columns);
92 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
93 curout[rc] += mult * static_cast<Output_>(rptr[rc]);
94 }
95 }
96 }
97
98 if (do_parallel && t > 0) {
99 (*tmp_results)[t - 1] = std::move(tmp_output);
100 }
101 }, common_dim, options.num_threads);
102
103 if (do_parallel) {
104 for (int u = 1; u < num_used; ++u) {
105 const auto& tmp = *((*tmp_results)[u - 1]);
106 const auto N = tmp.size();
107 for (I<decltype(N)> x = 0; x < N; ++x) {
108 output[x] += tmp[x];
109 }
110 }
111 }
112}
113
133template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename RightIndex_, typename Output_>
137 Output_* const output,
139) {
140 const auto left_NR = left.nrow();
141 const auto common_dim = left.ncol();
142 const auto right_NC = right.ncol();
143
144 const bool do_parallel = options.num_threads > 1;
145 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
146 if (do_parallel) {
147 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
148 }
149
150 std::fill_n(output, sanisizer::product<std::size_t>(left_NR, right_NC), 0);
151
152 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
153 auto left_ext = tatami::consecutive_extractor<true>(left, false, start, length);
154 auto right_ext = tatami::consecutive_extractor<false>(right, true, start, length);
155
159
160 std::optional<std::vector<Output_> > tmp_output;
161 Output_* outptr;
162 if (!do_parallel || t == 0) {
163 outptr = output;
164 } else {
165 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_NC));
166 outptr = tmp_output->data();
167 }
168
169 for (LeftIndex_ cd = 0; cd < length; ++cd) {
170 const auto lrange = left_ext->fetch(vbuffer.data(), ibuffer.data());
171 const auto rptr = right_ext->fetch(rbuffer.data());
172 for (LeftIndex_ x = 0; x < lrange.number; ++x) {
173 const Output_ mult = lrange.value[x];
174 const auto curout = outptr + sanisizer::product_unsafe<std::size_t>(lrange.index[x], right_NC);
175 for (RightIndex_ rc = 0; rc < right_NC; ++rc) {
176 curout[rc] += mult * static_cast<Output_>(rptr[rc]);
177 }
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
197}
198
199#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_row_output(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const RightColumns_ right_columns, GetRightRow_ get_right_row, Output_ *const output, const MultiplySparseColumnWithDenseRowMatrixToRowOutputOptions &options)
Definition row_to_row.hpp:52
int parallelize(Function_ fun, const Index_ tasks, const int workers)
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_row_output().
Definition row_to_row.hpp:27