tatami_stats
Matrix statistics for tatami
Loading...
Searching...
No Matches
group_sum.hpp
Go to the documentation of this file.
1#ifndef TATAMI_STATS_GROUPED_SUMS_HPP
2#define TATAMI_STATS_GROUPED_SUMS_HPP
3
4#include "utils.hpp"
5#include "sum.hpp"
6
7#include <vector>
8#include <algorithm>
9#include <cstddef>
10
11#include "tatami/tatami.hpp"
12#include "sanisizer/sanisizer.hpp"
13#include "jiwoo/jiwoo.hpp"
14
21namespace tatami_stats {
22
31 bool skip_nan = false;
32
37 int num_threads = 1;
38};
39
43template<typename Value_, typename Index_, typename Group_, typename Output_>
44void group_sum_direct(
45 bool row,
47 const Group_* group,
48 const Group_ num_groups,
49 const std::vector<Output_*>& output,
50 const GroupSumOptions& opt
51) {
52 const Index_ dim = (row ? mat.nrow() : mat.ncol());
53 const Index_ otherdim = (row ? mat.ncol() : mat.nrow());
54
55 if (mat.sparse()) {
56 tatami::parallelize([&](int, Index_ start, Index_ len) -> void {
57 auto ext = tatami::consecutive_extractor<true>(mat, row, start, len);
60 auto tmp = sanisizer::create<std::vector<Output_> >(num_groups);
61
62 for (Index_ x = 0; x < len; ++x) {
63 auto range = ext->fetch(xbuffer.data(), ibuffer.data());
64 std::fill(tmp.begin(), tmp.end(), static_cast<Output_>(0));
65
66 nanable_ifelse<Value_>(
67 opt.skip_nan,
68 [&]() -> void {
69 for (Index_ j = 0; j < range.number; ++j) {
70 const auto val = range.value[j];
71 if (!std::isnan(val)) {
72 tmp[group[range.index[j]]] += val;
73 }
74 }
75 },
76 [&]() -> void {
77 for (Index_ j = 0; j < range.number; ++j) {
78 tmp[group[range.index[j]]] += range.value[j];
79 }
80 }
81 );
82
83 for (I<decltype(num_groups)> g = 0; g < num_groups; ++g) {
84 output[g][start + x] = tmp[g];
85 }
86 }
87 }, dim, opt.num_threads);
88
89 } else {
90 tatami::parallelize([&](int, Index_ start, Index_ len) -> void {
91 auto ext = tatami::consecutive_extractor<false>(mat, row, start, len);
93 auto tmp = sanisizer::create<std::vector<Output_> >(num_groups);
94
95 for (Index_ x = 0; x < len; ++x) {
96 auto ptr = ext->fetch(xbuffer.data());
97 std::fill(tmp.begin(), tmp.end(), static_cast<Output_>(0));
98
99 nanable_ifelse<Value_>(
100 opt.skip_nan,
101 [&]() -> void {
102 for (Index_ j = 0; j < otherdim; ++j) {
103 const auto val = ptr[j];
104 if (!std::isnan(val)) {
105 tmp[group[j]] += val;
106 }
107 }
108 },
109 [&]() -> void {
110 for (Index_ j = 0; j < otherdim; ++j) {
111 tmp[group[j]] += ptr[j];
112 }
113 }
114 );
115
116 for (I<decltype(num_groups)> g = 0; g < num_groups; ++g) {
117 output[g][start + x] = tmp[g];
118 }
119 }
120 }, dim, opt.num_threads);
121 }
122}
123
124template<typename Value_, typename Index_, typename Group_, typename Output_>
125void group_sum_running(
126 bool row,
128 const Group_* group,
129 const Group_ num_groups,
130 const std::vector<Output_*>& output,
131 const GroupSumOptions& opt
132) {
133 const Index_ dim = (row ? mat.nrow() : mat.ncol());
134 const Index_ otherdim = (row ? mat.ncol() : mat.nrow());
135 const bool is_sparse = mat.is_sparse();
136
137 const auto do_parallel = opt.num_threads > 1;
138 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > all_partial_sums;
139 if (do_parallel) {
140 all_partial_sums.emplace(sanisizer::cast<I<decltype(all_partial_sums->size())> >(opt.num_threads - 1));
141 }
142
143 for (Group_ g = 0; g < num_groups; ++g) {
144 std::fill_n(output[g], dim, 0);
145 }
146
147 const auto nused = tatami::parallelize([&](int thread, Index_ start, Index_ len) -> void {
148 // If we can, directly dump the sum to the output pointers, otherwise put it into a temporary.
149 std::optional<jiwoo::EquilengthArrays<Output_> > cur_sums;
150
151 Output_* const * sum_ptrs;
152 if (!do_parallel) {
153 sum_ptrs = output.data();
154 } else {
155 if (thread == 0) {
156 sum_ptrs = output.data();
157 } else {
158 cur_sums.emplace(
159 sanisizer::Cast(num_groups),
160 static_cast<std::size_t>(dim), // cast is safe due to the tatami contract.
161 0
162 );
163 sum_ptrs = cur_sums->get();
164 }
165 }
166
167 if (is_sparse) {
168 // Order within each observed vector doesn't affect numerical precision of the outcome,
169 // as addition order for each objective vector is already well-defined for a running calculation.
170 tatami::Options topt;
171 topt.sparse_ordered_index = false;
172 auto ext = tatami::consecutive_extractor<true>(mat, !row, start, len, topt);
175
176 for (Index_ x = 0; x < len; ++x) {
177 auto range = ext->fetch(xbuffer.data(), ibuffer.data());
178 const auto sum_ptr = sum_ptrs[group[start + x]];
179
180 nanable_ifelse<Value_>(
181 opt.skip_nan,
182 [&]() -> void {
183 for (Index_ i = 0; i < range.number; ++i) {
184 const auto val = range.value[i];
185 sum_ptr[range.index[i]] += (std::isnan(val) ? 0 : val); // slightly nicer for vectorization.
186 }
187 },
188 [&]() -> void {
189 for (Index_ i = 0; i < range.number; ++i) {
190 sum_ptr[range.index[i]] += range.value[i];
191 }
192 }
193 );
194 }
195
196 } else {
197 auto ext = tatami::consecutive_extractor<false>(mat, !row, start, len);
199
200 for (Index_ x = 0; x < len; ++x) {
201 auto ptr = ext->fetch(buffer.data());
202 const auto sum_ptr = sum_ptrs[group[start + x]];
203
204 nanable_ifelse<Value_>(
205 opt.skip_nan,
206 [&]() -> void {
207 for (Index_ d = 0; d < dim; ++d) {
208 const auto val = ptr[d];
209 sum_ptr[d] += (std::isnan(val) ? 0 : val); // slightly nicer for vectorization.
210 }
211 },
212 [&]() -> void {
213 for (Index_ d = 0; d < dim; ++d) {
214 sum_ptr[d] += ptr[d];
215 }
216 }
217 );
218 }
219 }
220
221 if (do_parallel) {
222 if (thread > 0) {
223 (*all_partial_sums)[thread - 1] = std::move(cur_sums);
224 }
225 }
226 }, otherdim, opt.num_threads);
227
228 if (do_parallel) {
229 for (Group_ g = 0; g < num_groups; ++g) {
230 const auto cur_out = output[g];
231 for (int u = 1; u < nused; ++u) {
232 const auto cur_sum = (*((*all_partial_sums)[u - 1]))[g];
233 for (Index_ d = 0; d < dim; ++d) {
234 cur_out[d] += cur_sum[d];
235 }
236 }
237 }
238 }
239}
266template<typename Value_, typename Index_, typename Group_, typename Output_>
268 bool row,
270 const Group_* group,
271 const Group_ num_groups,
272 const std::vector<Output_*>& output,
273 const GroupSumOptions& opt
274) {
275 if (mat.prefer_rows() == row) {
276 group_sum_direct(row, mat, group, num_groups, output, opt);
277 } else {
278 group_sum_running(row, mat, group, num_groups, output, opt);
279 }
280}
281
305template<typename Output_ = double, typename Value_, typename Index_, typename Group_>
306std::vector<std::vector<Output_> > group_sum(
307 bool row,
309 const Group_* group,
310 const Group_ num_groups,
311 const GroupSumOptions& opt
312) {
313 auto output = sanisizer::create<std::vector<std::vector<Output_> > >(num_groups);
314 auto ptrs = sanisizer::create<std::vector<Output_*> >(num_groups);
315 const Index_ dim = (row ? mat.nrow() : mat.ncol());
316 for (Group_ g = 0; g < num_groups; ++g) {
318#ifdef TATAMI_STATS_TEST_DIRTY
319 , -1
320#endif
321 );
322 ptrs[g] = output[g].data();
323 }
324 group_sum(row, mat, group, num_groups, ptrs, opt);
325 return output;
326}
327
328}
329
330#endif
virtual Index_ ncol() const=0
virtual Index_ nrow() const=0
virtual bool prefer_rows() const=0
virtual bool is_sparse() const=0
virtual std::unique_ptr< MyopicSparseExtractor< Value_, Index_ > > sparse(bool row, const Options &opt) const=0
Functions to compute statistics from a tatami::Matrix.
Definition count.hpp:20
void group_sum(bool row, const tatami::Matrix< Value_, Index_ > &mat, const Group_ *group, const Group_ num_groups, const std::vector< Output_ * > &output, const GroupSumOptions &opt)
Definition group_sum.hpp:267
void range(bool row, const tatami::Matrix< Value_, Index_ > &mat, RangeBuffers< Output_ > &output, const RangeOptions< Output_ > &opt)
Definition range.hpp:339
void resize_container_to_Index_size(Container_ &container, const Index_ x, Args_ &&... args)
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)
bool sparse_ordered_index
Options for group_sum().
Definition group_sum.hpp:26
bool skip_nan
Definition group_sum.hpp:31
int num_threads
Definition group_sum.hpp:37
Compute row and column sums from a tatami::Matrix.