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_COLUMN_ROW_TO_ROW_HPP
2#define TATAMI_MULT_SPARSE_MATRIX_DENSE_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#include "../../utils.hpp"
13
19namespace tatami_mult {
20
21/* See https://github.com/tatami-inc/test-multiplication/tree/master/dense_column/sparse_matrix
22 * for an explanation of the choice of algorithm.
23 */
24
41
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
71 const bool do_parallel = options.num_threads > 1;
72 std::optional<std::vector<std::optional<std::vector<Output_> > > > tmp_results;
73 if (do_parallel) {
74 tmp_results.emplace(sanisizer::cast<I<decltype(tmp_results->size())> >(options.num_threads - 1));
75 }
76
77 std::fill_n(output, sanisizer::product_unsafe<std::size_t>(left_NR, right_NC), 0);
78
79 const int num_used = tatami::parallelize([&](int t, LeftIndex_ start, LeftIndex_ length) -> void {
80 auto left_ext = tatami::consecutive_extractor<false>(left, false, start, length);
81 auto right_ext = tatami::consecutive_extractor<true>(right, true, start, length);
82
83 std::optional<std::vector<Output_> > tmp_output;
84 Output_* outptr;
85 if (!do_parallel || t == 0) {
86 outptr = output;
87 } else {
88 tmp_output.emplace(sanisizer::product<I<decltype(tmp_output->size())> >(left_NR, right_NC));
89 outptr = tmp_output->data();
90 }
91
92 if (options.block_size == 1) {
96
97 for (LeftIndex_ cd = 0; cd < length; ++cd) {
98 const auto lptr = left_ext->fetch(dbuffer.data());
99 const auto rrange = right_ext->fetch(vbuffer.data(), ibuffer.data());
100
101 // Make sure this skip is done after all fetch() calls, otherwise the extractors will not be in sync with the common dimension.
102 if (rrange.number == 0) {
103 continue;
104 }
105
106 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
107 const Output_ mult = lptr[lr];
108 for (RightIndex_ x = 0; x < rrange.number; ++x) {
109 outptr[sanisizer::nd_offset<std::size_t>(rrange.index[x], right_NC, lr)] += mult * static_cast<Output_>(rrange.value[x]);
110 }
111 }
112 }
113
114 } else {
115 std::vector<std::vector<LeftValue_> > left_dbuffers;
116 std::vector<const LeftValue_*> left_ptrs;
117 std::vector<std::vector<RightValue_> > right_vbuffers;
118 std::vector<std::vector<RightIndex_> > right_ibuffers;
119 std::vector<tatami::SparseRange<RightValue_, RightIndex_> > right_ranges;
120 {
121 const LeftIndex_ max_block_cols = sanisizer::min(length, options.block_size);
122 left_dbuffers.reserve(max_block_cols);
123 right_vbuffers.reserve(max_block_cols);
124 right_ibuffers.reserve(max_block_cols);
125 for (LeftIndex_ cd = 0; cd < max_block_cols; ++cd) {
126 left_dbuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<LeftValue_> >(left_NR));
127 right_vbuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<RightValue_> >(right_NC));
128 right_ibuffers.emplace_back(tatami::cast_Index_to_container_size<std::vector<RightIndex_> >(right_NC));
129 }
130 sanisizer::resize(left_ptrs, max_block_cols);
131 sanisizer::resize(right_ranges, max_block_cols);
132 }
133
134 LeftIndex_ cd = 0;
135 while (cd < length) {
136 // Only processing LHS columns if the corresponding RHS row has some structural non-zeros.
137 // If not, we just skip it altogether; no need to zero or do anything else, as we're skipping the corresponding RHS row too.
138 LeftIndex_ cd_num = 0;
139 do {
140 auto lptr = left_ext->fetch(left_dbuffers[cd_num].data());
141 auto rrange = right_ext->fetch(right_vbuffers[cd_num].data(), right_ibuffers[cd_num].data());
142
143 // Again, this skip must be done after the LHS row is fetched, otherwise the extractors will be out of sync.
144 if (rrange.number == 0) {
145 ++cd;
146 continue;
147 }
148
149 left_ptrs[cd_num] = lptr;
150 right_ranges[cd_num] = std::move(rrange);
151 ++cd_num;
152 ++cd;
153
154 if (sanisizer::is_equal(cd_num, options.block_size)) {
155 break;
156 }
157 } while (cd < length);
158
159 for (LeftIndex_ lr = 0; lr < left_NR; ++lr) {
160 for (LeftIndex_ cd_counter = 0; cd_counter < cd_num; ++cd_counter) {
161 const auto mult = left_ptrs[cd_counter][lr];
162 const auto rrange = right_ranges[cd_counter];
163 for (RightIndex_ x = 0; x < rrange.number; ++x) {
164 outptr[sanisizer::nd_offset<std::size_t>(rrange.index[x], right_NC, lr)] += mult * static_cast<Output_>(rrange.value[x]);
165 }
166 }
167 }
168 }
169 }
170
171 if (do_parallel && t > 0) {
172 (*tmp_results)[t - 1] = std::move(tmp_output);
173 }
174 }, common_dim, options.num_threads);
175
176 if (do_parallel) {
177 for (int u = 1; u < num_used; ++u) {
178 const auto& tmp = *((*tmp_results)[u - 1]);
179 const auto N = tmp.size();
180 for (I<decltype(N)> x = 0; x < N; ++x) {
181 output[x] += tmp[x];
182 }
183 }
184 }
185}
186
187}
188
189#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_sparse_row_matrix_to_row_output(const tatami::Matrix< LeftValue_, LeftIndex_ > &left, const tatami::Matrix< RightValue_, RightIndex_ > &right, Output_ *const output, const MultiplyDenseColumnWithSparseRowMatrixToRowOutputOptions &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_column_with_sparse_row_matrix_to_row_output().
Definition row_to_row.hpp:28