tatami_mult
Multiply tatami matrices
Loading...
Searching...
No Matches
row_to_row.hpp
Go to the documentation of this file.
1#ifndef TATAMI_MULT_SPARSE_MATRIX_DENSE_ROW_ROW_TO_ROW_HPP
2#define TATAMI_MULT_SPARSE_MATRIX_DENSE_ROW_ROW_TO_ROW_HPP
3
4#include <cstddef>
5#include <vector>
6
7#include "tatami/tatami.hpp"
8#include "sanisizer/sanisizer.hpp"
9
10#include "../utils.hpp"
11#include "../../utils.hpp"
12
18namespace tatami_mult {
19
20/* See https://github.com/tatami-inc/test-multiplication/tree/master/dense_row/sparse_matrix
21 * for an explanation of the choice of algorithm.
22 */
23
40
60template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename RightIndex_, typename Output_>
64 Output_* const output,
66) {
67 const auto left_NR = left.nrow();
68 const auto common_dim = left.ncol();
69 const auto right_NC = right.ncol();
70
74 populate_sparse_buffers(true, common_dim, right_NC, right, right_vbuffers, right_ibuffers, right_ranges, options.num_threads);
75
76 // If there are any empty RHS rows, we only iterate over the non-empty ones in the loop for each LHS row.
77 auto right_non_empty = filter_non_empty_sparse(
78 right_ranges,
79 [&](RightIndex_) -> void {}
80 );
81
82 const bool do_parallel = options.num_threads > 1;
83 if (!do_parallel) {
84 std::fill_n(output, sanisizer::product_unsafe<std::size_t>(left_NR, right_NC), 0);
85 }
86
87 if (options.block_size == 1) {
88 tatami::parallelize([&](int, LeftIndex_ start, LeftIndex_ length) -> void {
89 auto ext = tatami::consecutive_extractor<false>(left, true, start, length);
91
92 std::optional<std::vector<Output_> > tmp_row;
93 if (do_parallel) {
94 tmp_row.emplace(tatami::cast_Index_to_container_size<std::vector<Output_> >(right_NC));
95 }
96
97 for (LeftIndex_ lr = 0; lr < length; ++lr) {
98 const auto lptr = ext->fetch(dbuffer.data());
99 const auto optr = output + sanisizer::product_unsafe<std::size_t>(start + lr, right_NC);
100 const auto tmp_optr = (do_parallel ? tmp_row->data() : optr);
101
102 auto loop_body = [&](LeftIndex_ cd) -> void {
103 const auto rrange = right_ranges[cd];
104 const Output_ mult = lptr[cd];
105 for (RightIndex_ x = 0; x < rrange.number; ++x) {
106 tmp_optr[rrange.index[x]] += mult * static_cast<Output_>(rrange.value[x]);
107 }
108 };
109
110 if (right_non_empty.has_value()) {
111 for (const auto cd : *right_non_empty) {
112 loop_body(cd);
113 }
114 } else {
115 for (LeftIndex_ cd = 0; cd < common_dim; ++cd) {
116 loop_body(cd);
117 }
118 }
119
120 if (do_parallel) {
121 std::copy_n(tmp_optr, right_NC, optr);
122
123 // Technically, we only have to reset the positions at which there is at least one non-zero across all RHS rows.
124 // However, the union of all non-zero positions across all RHS rows is probably quite dense.
125 // It'll likely be faster to just zero the entire buffer rather than trying to zero specific positions;
126 // for example, one 64-byte cache line contains 8 doubles, so you'd need a density below ~10% to even avoid loading every cache line.
127 // And that's not even considering further optimizations in the memset call.
128 std::fill_n(tmp_optr, right_NC, 0);
129 }
130 }
131 }, left_NR, options.num_threads);
132
133 } else {
134 tatami::parallelize([&](int, LeftIndex_ start, LeftIndex_ length) -> void {
135 auto ext = tatami::consecutive_extractor<false>(left, true, start, length);
136
137 const LeftIndex_ max_block_rows = sanisizer::min(length, options.block_size);
138 std::vector<std::vector<LeftValue_> > lbuffers;
139 lbuffers.reserve(max_block_rows);
140 for (LeftIndex_ b = 0; b < max_block_rows; ++b) {
141 lbuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(common_dim));
142 }
144
145 std::optional<std::vector<Output_> > tmp_rows;
146 if (do_parallel) {
147 tmp_rows.emplace(sanisizer::product<I<decltype(tmp_rows->size())> >(max_block_rows, right_NC));
148 }
149
150 LeftIndex_ lr = 0;
151 while (lr < length) {
152 const LeftIndex_ lr_num = sanisizer::min(options.block_size, length - lr);
153 for (LeftIndex_ lr_counter = 0; lr_counter < lr_num; ++lr_counter) {
154 lptrs[lr_counter] = ext->fetch(lbuffers[lr_counter].data());
155 }
156 const auto optr = output + sanisizer::product_unsafe<std::size_t>(start + lr, right_NC);
157 const auto tmp_optr = (do_parallel ? tmp_rows->data() : optr);
158
159 auto loop_body = [&](LeftIndex_ cd) -> void {
160 const auto rrange = right_ranges[cd];
161 for (LeftIndex_ lr_counter = 0; lr_counter < lr_num; ++lr_counter) {
162 const Output_ mult = lptrs[lr_counter][cd];
163 for (RightIndex_ x = 0; x < rrange.number; ++x) {
164 tmp_optr[sanisizer::nd_offset<std::size_t>(rrange.index[x], right_NC, lr_counter)] += mult * static_cast<Output_>(rrange.value[x]);
165 }
166 }
167 };
168
169 if (right_non_empty.has_value()) {
170 for (const auto cd : *right_non_empty) {
171 loop_body(cd);
172 }
173 } else {
174 for (LeftIndex_ cd = 0; cd < common_dim; ++cd) {
175 loop_body(cd);
176 }
177 }
178
179 if (do_parallel) {
180 const auto output_size = sanisizer::product_unsafe<std::size_t>(lr_num, right_NC);
181 std::copy_n(tmp_optr, output_size, optr);
182 std::fill_n(tmp_optr, output_size, 0);
183 }
184
185 lr += lr_num;
186 }
187 }, left_NR, options.num_threads);
188 }
189}
190
191}
192
193#endif
virtual Index_ ncol() const=0
virtual Index_ nrow() const=0
Multiplication of tatami matrices.
Definition column_to_column.hpp:19
void multiply_dense_row_with_sparse_row_matrix_to_row_output(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const tatami::Matrix< RightValue_, RightIndex_ > &right, Output_ *const output, const MultiplyDenseRowWithSparseRowMatrixToRowOutputOptions &options)
Definition row_to_row.hpp:61
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_row_with_sparse_row_matrix_to_row_output().
Definition row_to_row.hpp:27