tatami_mult
Multiply tatami matrices
Loading...
Searching...
No Matches
column_to_row.hpp
Go to the documentation of this file.
1#ifndef TATAMI_MULT_DENSE_MATRIX_DENSE_COLUMN_COLUMN_TO_ROW_HPP
2#define TATAMI_MULT_DENSE_MATRIX_DENSE_COLUMN_COLUMN_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#include "../../utils.hpp"
13
19namespace tatami_mult {
20
21/* See https://github.com/tatami-inc/test-multiplication/tree/master/dense_column/dense_matrix
22 * for an explanation of the choice of algorithm.
23 */
24
49
67template<typename LeftValue_, typename LeftIndex_, typename RightColumns_, class GetRightColumn_, typename Output_>
70 const RightColumns_ right_columns,
71 GetRightColumn_ get_right_column,
72 Output_* const output,
74) {
75 const auto left_NR = left.nrow();
76 const auto common_dim = left.ncol();
77
78 // Product must fit in a size_t in order for output to have been allocated correctly in the first place.
79 // 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.
80 std::fill_n(output, sanisizer::product_unsafe<std::size_t>(left_NR, right_columns), 0);
81
82 const bool do_parallel = options.num_threads > 1;
83 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
84 if (do_parallel) {
85 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
86 }
87
88 const auto num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
89 auto ext = tatami::consecutive_extractor<false>(left, false, start, length);
90
91 std::optional<std::vector<Output_> > tmp_output;
92 Output_* outptr;
93 if (!do_parallel || t == 0) {
94 outptr = output;
95 } else {
96 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_columns));
97 outptr = tmp_output->data();
98 }
99
100 if (options.primary_block_size == 1) {
102 for (LeftIndex_ cd = 0; cd < length; ++cd) {
103 const auto ptr = ext->fetch(buffer.data());
104 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
105 const auto mult = ptr[lr];
106 for (RightColumns_ rc = 0; rc < right_columns; ++rc) {
107 outptr[sanisizer::nd_offset<std::size_t>(rc, right_columns, lr)] += mult * static_cast<Output_>(get_right_column(rc)[start + cd]);
108 }
109 }
110 }
111
112 } else {
113 std::vector<std::vector<LeftValue_> > left_buffers;
114 std::vector<const LeftValue_*> left_ptrs;
115 std::vector<Output_> tmp_output;
116 {
117 const LeftIndex_ max_block_cols = sanisizer::min(length, options.primary_block_size);
118 left_buffers.reserve(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 sanisizer::resize(left_ptrs, max_block_cols);
123
124 // We create a temporary buffer so that the inner loop for each block operates on contiguous output.
125 // This improves the efficiency of the hot loop, though we will have to transpose it to the output rows eventually.
126 const RightColumns_ max_block_right_cols = sanisizer::min(right_columns, options.primary_block_size);
127 const LeftIndex_ max_block_rows = sanisizer::min(left_NR, options.secondary_block_size);
128 tmp_output.resize(sanisizer::product<I<decltype(tmp_output.size())> >(max_block_rows, max_block_right_cols));
129 }
130
131 LeftIndex_ cd = 0;
132 while (cd < length) {
133 const auto cd_num = sanisizer::min(options.primary_block_size, length - cd);
134 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
135 left_ptrs[cd_counter] = ext->fetch(left_buffers[cd_counter].data());
136 }
137
138 RightColumns_ rc = 0;
139 while (rc < right_columns) {
140 const RightColumns_ rc_num = sanisizer::min(options.primary_block_size, right_columns - rc);
141 LeftIndex_ lr = 0;
142 while (lr < left_NR) {
143 const LeftIndex_ lr_num = sanisizer::min(options.secondary_block_size, left_NR - lr);
144
145 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
146 const auto matcol = left_ptrs[cd_counter];
147 for (RightColumns_ rc_counter = 0; rc_counter < rc_num; ++rc_counter) {
148 const Output_ mult = get_right_column(rc + rc_counter)[start + cd + cd_counter];
149 for (LeftIndex_ lr_counter = 0; lr_counter < lr_num; ++lr_counter) {
150 tmp_output[sanisizer::nd_offset<std::size_t>(lr_counter, lr_num, rc_counter)] += mult * static_cast<Output_>(matcol[lr + lr_counter]);
151 }
152 }
153 }
154
155 // Transposition using square blocks of the smaller (primary) block size.
156 RightColumns_ lrt = 0;
157 while (lrt < lr_num) {
158 const LeftIndex_ lrt_end = lrt + sanisizer::min(options.primary_block_size, lr_num - lrt);
159 for (RightColumns_ rc_counter = 0; rc_counter < rc_num; ++rc_counter) {
160 for (LeftIndex_ lrt_copy = lrt; lrt_copy < lrt_end; ++lrt_copy) {
161 const auto val = tmp_output[sanisizer::nd_offset<std::size_t>(lrt_copy, lr_num, rc_counter)];
162 outptr[sanisizer::nd_offset<std::size_t>(rc + rc_counter, right_columns, lr + lrt_copy)] += val;
163 }
164 }
165 lrt = lrt_end;
166 }
167 std::fill_n(tmp_output.begin(), sanisizer::product_unsafe<std::size_t>(rc_num, lr_num), 0);
168
169 lr += lr_num;
170 }
171 rc += rc_num;
172 }
173 cd += cd_num;
174 }
175 }
176
177 if (do_parallel && t > 0) {
178 (*tmp_results)[t - 1] = std::move(tmp_output);
179 }
180 }, common_dim, options.num_threads);
181
182 if (do_parallel) {
183 for (int u = 1; u < num_used; ++u) {
184 const auto& tmp = *((*tmp_results)[u - 1]);
185 const auto N = tmp.size();
186 for (I<decltype(N)> x = 0; x < N; ++x) {
187 output[x] += tmp[x];
188 }
189 }
190 }
191}
192
213template<typename LeftValue_, typename LeftIndex_, typename RightValue_, typename RightIndex_, typename Output_>
217 Output_* const output,
219) {
220 const auto right_NC = right.ncol();
223 const auto common_dim = left.ncol();
224 populate_dense_buffers(false, right_NC, common_dim, right, right_buffers, right_ptrs, options.num_threads);
225
227 left,
228 right_NC,
229 [&](const RightIndex_ rc) -> const RightValue_* {
230 return right_ptrs[rc];
231 },
232 output,
233 options
234 );
235}
236
237}
238
239#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_column_matrix_to_row_output(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const RightColumns_ right_columns, GetRightColumn_ get_right_column, Output_ *const output, const MultiplyDenseColumnWithDenseColumnMatrixToRowOutputOptions &options)
Definition column_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_column_with_dense_column_matrix_to_row_output().
Definition column_to_row.hpp:28