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_DENSE_ROW_ROW_TO_ROW_HPP
2#define TATAMI_MULT_DENSE_MATRIX_DENSE_ROW_ROW_TO_ROW_HPP
3
4#include <vector>
5#include <cstddef>
6#include <optional>
7
8#include "tatami/tatami.hpp"
9#include "sanisizer/sanisizer.hpp"
10
11#include "../utils.hpp"
12#include "../../utils.hpp"
13
19namespace tatami_mult {
20
21/* See https://github.com/tatami-inc/test-multiplication/tree/master/dense_row/dense_matrix
22 * for an explanation of the choice of algorithm.
23 */
24
49
67template<typename LeftValue_, typename LeftIndex_, typename RightColumns_, class GetRightRow_, typename Output_>
70 const RightColumns_ right_columns,
71 GetRightRow_ get_right_row,
72 Output_* const output,
74) {
75 const auto left_NR = left.nrow();
76 const auto common_dim = left.ncol();
77
78 const bool do_parallel = options.num_threads > 1;
79 if (!do_parallel) {
80 // Product must fit in a size_t in order for output to have been allocated correctly in the first place.
81 // 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.
82 std::fill_n(output, sanisizer::product_unsafe<std::size_t>(right_columns, left_NR), 0);
83 }
84
85 if (options.primary_block_size == 1) {
86 tatami::parallelize([&](int, LeftIndex_ start, LeftIndex_ length) -> void {
87 auto ext = tatami::consecutive_extractor<false>(left, true, start, length);
89
90 // Use a temporary buffer for each output row to mitigate false sharing during updates across all 'c'.
91 // There is still some false sharing when we transfer the results to the output row,
92 // but this is fine as it is outside of the innermost loop.
93 std::optional<std::vector<Output_> > tmp_output;
94 if (do_parallel) {
95 tmp_output.emplace(tatami::cast_Index_to_container_size<std::vector<Output_> >(right_columns));
96 }
97
98 for (LeftIndex_ lr = 0; lr < length; ++lr) {
99 const auto left_ptr = ext->fetch(buffer.data());
100 const auto optr = output + sanisizer::product_unsafe<std::size_t>(start + lr, right_columns);
101
102 Output_* tmp_optr;
103 if (!do_parallel) {
104 tmp_optr = optr;
105 } else {
106 tmp_optr = tmp_output->data();
107 }
108
109 for (LeftIndex_ cd = 0; cd < common_dim; ++cd) {
110 const Output_ mult = left_ptr[cd];
111 const auto rightrow = get_right_row(cd);
112 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
113 tmp_optr[rc] += static_cast<Output_>(rightrow[rc]) * mult;
114 }
115 }
116
117 if (do_parallel) {
118 std::copy_n(tmp_optr, right_columns, optr);
119 std::fill_n(tmp_optr, right_columns, 0);
120 }
121 }
122 }, left_NR, options.num_threads);
123
124 } else {
125 tatami::parallelize([&](int, LeftIndex_ start, LeftIndex_ length) -> void {
126 auto left_ext = tatami::consecutive_extractor<false>(left, true, start, length);
127 std::vector<std::vector<LeftValue_> > left_buffers;
128 std::vector<const LeftValue_*> left_ptrs;
129
130 std::optional<std::vector<Output_> > tmp_output;
131 {
132 const LeftIndex_ max_block_rows = sanisizer::min(length, options.primary_block_size);
133 left_buffers.reserve(max_block_rows);
134 for (LeftIndex_ b = 0; b < max_block_rows ; ++b) {
135 left_buffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(common_dim));
136 }
137 sanisizer::resize(left_ptrs, max_block_rows);
138
139 // Creating a block to hold the output during the updates over all 'c', to avoid false sharing.
140 // We should be able to hold the block size in a size_t safely here, as this block is no larger than the array referenced by 'output'.
141 // Of course, we might get an error if the vector's size_type is smaller than size_t but that seems a bit pathological.
142 if (do_parallel) {
143 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(max_block_rows, right_columns));
144 }
145 }
146
147 LeftIndex_ lr = 0;
148 while (lr < length) {
149 const LeftIndex_ lr_num = sanisizer::min(options.primary_block_size, length - lr);
150 for (LeftIndex_ lr_counter = 0; lr_counter < lr_num; ++lr_counter) {
151 left_ptrs[lr_counter] = left_ext->fetch(left_buffers[lr_counter].data());
152 }
153
154 Output_* const optr = output + sanisizer::product_unsafe<std::size_t>(start + lr, right_columns);
155 Output_* tmp_optr;
156 if (!do_parallel) {
157 tmp_optr = optr;
158 } else {
159 tmp_optr = tmp_output->data();
160 }
161
162 LeftIndex_ cd = 0;
163 while (cd < common_dim) {
164 const LeftIndex_ cd_end = cd + sanisizer::min(options.primary_block_size, common_dim - cd);
165 RightColumns_ rc = 0;
166 while (rc < right_columns) {
167 const RightColumns_ rc_end = rc + sanisizer::min(options.secondary_block_size, right_columns - rc);
168
169 for (LeftIndex_ lr_counter = 0; lr_counter < lr_num; ++lr_counter) {
170 const auto matrow = left_ptrs[lr_counter];
171 const auto prod = tmp_optr + sanisizer::product_unsafe<std::size_t>(lr_counter, right_columns);
172 for (auto ccopy = cd; ccopy < cd_end; ++ccopy) {
173 const auto mult = matrow[ccopy];
174 const auto& rightrow = get_right_row(ccopy);
175 for (auto rc_copy = rc; rc_copy < rc_end; ++rc_copy) {
176 prod[rc_copy] += mult * rightrow[rc_copy];
177 }
178 }
179 }
180
181 rc = rc_end;
182 }
183 cd = cd_end;
184 }
185
186 if (do_parallel) {
187 const auto out_space = sanisizer::product_unsafe<std::size_t>(lr_num, right_columns);
188 std::copy_n(tmp_optr, out_space, optr);
189 std::fill_n(tmp_optr, out_space, 0);
190 }
191 lr += lr_num;
192 }
193 }, left_NR, options.num_threads);
194 }
195}
196
217template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename RightIndex_, typename Output_>
221 Output_* const output,
223) {
224 const auto common_dim = left.ncol();
227 const auto right_NC = right.ncol();
228 populate_dense_buffers(true, common_dim, right_NC, right, right_buffers, right_ptrs, options.num_threads);
229
231 left,
232 right_NC,
233 [&](const LeftIndex_ cd) -> const RightValue_* {
234 return right_ptrs[cd];
235 },
236 output,
237 options
238 );
239}
240
241}
242
243#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_dense_row_matrix_to_row_output(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const RightColumns_ right_columns, GetRightRow_ get_right_row, Output_ *const output, const MultiplyDenseRowWithDenseRowMatrixToRowOutputOptions &options)
Definition row_to_row.hpp:68
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_dense_row_matrix_to_row_output().
Definition row_to_row.hpp:28