tatami_stats
Matrix statistics for tatami
Loading...
Searching...
No Matches
group_rss.hpp
Go to the documentation of this file.
1#ifndef TATAMI_STATS_GROUP_RSS_HPP
2#define TATAMI_STATS_GROUP_RSS_HPP
3
4#include "utils.hpp"
5
6#include <vector>
7#include <algorithm>
8#include <cstddef>
9#include <optional>
10#include <cassert>
11#include <limits>
12#include <cmath>
13
14#include "tatami/tatami.hpp"
15#include "sanisizer/sanisizer.hpp"
17#include "jiwoo/jiwoo.hpp"
18
25namespace tatami_stats {
26
31template<typename Output_ = double>
45
51template<typename Output_>
58 std::vector<Output_*> mean;
59
65 std::vector<Output_*> rss;
66};
67
71template<typename Group_, typename Count_, typename Output_, typename Index_>
72void group_rss_finish_means(
73 const Group_ num_groups,
74 const Count_* const group_size,
75 std::vector<Output_>& means,
76 const Index_ i,
77 const std::vector<Output_*>& output_means,
78 const Output_ placeholder
79) {
80 for (Group_ g = 0; g < num_groups; ++g) {
81 if (group_size[g]) {
82 means[g] /= group_size[g];
83 } else {
84 means[g] = placeholder;
85 }
86 output_means[g][i] = means[g];
87 }
88}
89
90template<typename Value_, typename Index_, typename Group_, typename Count_, typename Output_>
91void group_rss_direct(
92 const bool row,
94 const Group_* const group,
95 const Group_ num_groups,
96 const Count_* const group_size,
97 const GroupRssBuffers<Output_>& output,
98 const GroupRssOptions<Output_>& opt
99) {
100 const auto dim = (row ? mat.nrow() : mat.ncol());
101 const auto otherdim = (row ? mat.ncol() : mat.nrow());
102
103 if (mat.sparse()) {
104 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
105 auto ext = tatami::consecutive_extractor<true>(mat, row, s, l);
108 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
109 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
110 auto cur_non_zeros = sanisizer::create<std::vector<Index_> >(num_groups);
111
112 for (Index_ x = 0; x < l; ++x) {
113 auto range = ext->fetch(vbuffer.data(), ibuffer.data());
114
115 // Computing the mean first.
116 for (Index_ i = 0; i < range.number; ++i) {
117 const auto g = group[range.index[i]];
118 cur_means[g] += range.value[i];
119 ++cur_non_zeros[g];
120 }
121 group_rss_finish_means(num_groups, group_size, cur_means, static_cast<Index_>(s + x), output.mean, opt.mean_placeholder);
122
123 // Now computing the RSS.
124 for (Index_ i = 0; i < range.number; ++i) {
125 const auto g = group[range.index[i]];
126 const auto delta = range.value[i] - cur_means[g];
127 cur_rss[g] += delta * delta;
128 }
129 for (Group_ g = 0; g < num_groups; ++g) {
130 if (group_size[g] > 0) { // preserve RSS = 0 if the group is empty, otherwise the placeholder mean might be a NaN that causes problems.
131 const Output_ my_rss = cur_rss[g] + cur_means[g] * cur_means[g] * (group_size[g] - cur_non_zeros[g]);
132 output.rss[g][s + x] = my_rss;
133 } else {
134 output.rss[g][s + x] = 0;
135 }
136 }
137
138 std::fill(cur_means.begin(), cur_means.end(), 0);
139 std::fill(cur_rss.begin(), cur_rss.end(), 0);
140 std::fill(cur_non_zeros.begin(), cur_non_zeros.end(), 0);
141 }
142 }, dim, opt.num_threads);
143
144 } else {
145 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
146 auto ext = tatami::consecutive_extractor<false>(mat, row, s, l);
148 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
149 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
150
151 for (Index_ x = 0; x < l; ++x) {
152 auto ptr = ext->fetch(buffer.data());
153
154 // Computing the mean first.
155 for (Index_ j = 0; j < otherdim; ++j) {
156 cur_means[group[j]] += ptr[j];
157 }
158 group_rss_finish_means(num_groups, group_size, cur_means, static_cast<Index_>(s + x), output.mean, opt.mean_placeholder);
159
160 // Now computing the RSS.
161 for (Index_ j = 0; j < otherdim; ++j) {
162 const auto g = group[j];
163 const auto delta = ptr[j] - cur_means[g];
164 cur_rss[g] += delta * delta;
165 }
166 for (Group_ g = 0; g < num_groups; ++g) {
167 output.rss[g][s + x] = cur_rss[g];
168 }
169
170 std::fill(cur_means.begin(), cur_means.end(), 0);
171 std::fill(cur_rss.begin(), cur_rss.end(), 0);
172 }
173 }, dim, opt.num_threads);
174 }
175}
176
177template<typename Value_, typename Index_, typename Group_, typename Count_, typename Output_>
178void group_rss_running_nonempty(
179 const bool row,
180 const Index_ dim,
181 const Index_ otherdim,
183 const Group_* const group,
184 const Group_ num_groups,
185 const Count_* const group_size,
186 const GroupRssBuffers<Output_>& output,
187 const GroupRssOptions<Output_>& opt
188) {
189 const bool do_parallel = opt.num_threads > 1;
190 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > all_partial_mean, all_partial_rss;
191 std::optional<std::vector<std::optional<std::vector<Count_> > > > all_partial_count;
192 if (do_parallel) {
193 // -1, as we'll repurpose the RSS output buffer to store the partial RSS of the first thread.
194 all_partial_rss.emplace(sanisizer::cast<I<decltype(all_partial_rss->size())> >(opt.num_threads - 1));
195 all_partial_mean.emplace(sanisizer::cast<I<decltype(all_partial_mean->size())> >(opt.num_threads));
196 all_partial_count.emplace(sanisizer::cast<I<decltype(all_partial_count->size())> >(opt.num_threads));
197 }
198
199 // All groups are assumed to be non-empty at this point,
200 // which allows us to skip some allocations.
201 for (Group_ g = 0; g < num_groups; ++g) {
202 assert(group_size[g] > 0);
203 }
204
205 // We overwrite any existing value in the array in the do_parallel=true situation.
206 // So, the initial value doesn't need to be zero.
207 if (!do_parallel) {
208 for (Group_ g = 0; g < num_groups; ++g) {
209 std::fill_n(output.mean[g], dim, 0);
210 }
211 }
212 for (Group_ g = 0; g < num_groups; ++g) {
213 std::fill_n(output.rss[g], dim, 0);
214 }
215
216 const bool is_sparse = mat.is_sparse();
217 const int nused = tatami::parallelize([&](int thread, Index_ s, Index_ l) -> void {
218 std::optional<jiwoo::EquilengthArrays<Output_> > cur_mean, cur_rss;
219
220 Output_* const * mean_ptrs;
221 Output_* const * rss_ptrs;
222 if (!do_parallel) {
223 // Storing mean and RSS directly in the output vector to cut down two allocations if we're not working in parallel.
224 mean_ptrs = output.mean.data();
225 rss_ptrs = output.rss.data();
226
227 } else {
228 // Storing the partial RSS directly in the output vectors to save ourselves an allocation if we're in the first thread.
229 if (thread == 0) {
230 rss_ptrs = output.rss.data();
231 } else {
232 cur_rss.emplace(
233 sanisizer::cast<I<decltype(cur_rss->size())> >(num_groups),
234 static_cast<std::size_t>(dim), // cast to size_t is safe, based on the tatami contract.
235 0
236 );
237 rss_ptrs = cur_rss->get();
238 }
239
240 // We can't do the same for the mean, though, as we need to keep the partial mean and the global mean separate for the reduction.
241 cur_mean.emplace(
242 sanisizer::cast<I<decltype(cur_mean->size())> >(num_groups),
243 static_cast<std::size_t>(dim), // cast to size_t is safe, based on the tatami contract.
244 0
245 );
246 mean_ptrs = cur_mean->get();
247 }
248
249 auto cur_count = sanisizer::create<std::vector<Count_> >(num_groups);
250
251 if (is_sparse) {
252 auto ext = tatami::consecutive_extractor<true>(mat, !row, s, l);
255 auto nonzeros = sanisizer::create<std::vector<std::vector<Index_> > >(num_groups);
256 for (Group_ g = 0; g < num_groups; ++g) {
258 }
259
260 for (Index_ x = 0; x < l; ++x) {
261 auto out = ext->fetch(vbuffer.data(), ibuffer.data());
262 const auto grp = group[s + x];
263 ++cur_count[grp]; // increment is safe as 'cur_count[grp] + 1 <= l' fits in an Index_.
264 const auto mptr = mean_ptrs[grp];
265 const auto rptr = rss_ptrs[grp];
266 auto& nnz = nonzeros[grp];
267 for (Index_ i = 0; i < out.number; ++i) {
268 const auto d = out.index[i];
269 quickstats::update_rss(mptr[d], rptr[d], out.value[i], ++nnz[d]); // increment is safe as 'nnz + 1 <= l' fits in an Index_.
270 }
271 }
272
273 for (Group_ g = 0; g < num_groups; ++g) {
274 const auto curtotal = cur_count[g];
275 if (curtotal) {
276 const auto mptr = mean_ptrs[g];
277 const auto rptr = rss_ptrs[g];
278 const auto& nnz = nonzeros[g];
279 for (Index_ d = 0; d < dim; ++d) {
280 // unsafe call is possible as we check for curtotal > 0.
281 quickstats::update_rss_with_zeros_unsafe(mptr[d], rptr[d], static_cast<Count_>(curtotal - nnz[d]), curtotal);
282 }
283 }
284 }
285
286 } else {
287 auto ext = tatami::consecutive_extractor<false>(mat, !row, s, l);
289
290 for (Index_ x = 0; x < l; ++x) {
291 auto out = ext->fetch(buffer.data());
292 const auto grp = group[s + x];
293 ++cur_count[grp]; // increment is safe as 'cur_count[grp] + 1 <= l' fits in an Index_.
294 const auto mptr = mean_ptrs[grp];
295 const auto rptr = rss_ptrs[grp];
296 for (Index_ d = 0; d < dim; ++d) {
297 quickstats::update_rss(mptr[d], rptr[d], out[d], cur_count[grp]);
298 }
299 }
300 }
301
302 if (do_parallel) {
303 (*all_partial_count)[thread] = std::move(cur_count);
304 (*all_partial_mean)[thread] = std::move(cur_mean);
305 if (thread > 0) {
306 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
307 }
308 }
309 }, otherdim, opt.num_threads);
310 assert(nused > 0);
311
312 if (do_parallel) {
313 const auto& ap_mean = *all_partial_mean;
314 const auto& ap_rss = *all_partial_rss;
315
316 for (Group_ g = 0; g < num_groups; ++g) {
317 const auto cur_output = output.mean[g];
318 const auto cur_global_count = group_size[g];
319 assert(cur_global_count > 0);
320 bool initialized = false;
321
322 for (int u = 0; u < nused; ++u) {
323 const auto cur_count = (*((*all_partial_count)[u]))[g];
324 if (cur_count == 0) {
325 continue;
326 }
327
328 const auto cur_mean = (*(ap_mean[u]))[g];
329 const Output_ mult = static_cast<Output_>(cur_count) / static_cast<Output_>(cur_global_count);
330 if (!initialized) { // Don't use u == 0, as the first non-empty 'g' might not occur in the first thread.
331 for (Index_ d = 0; d < dim; ++d) {
332 cur_output[d] = cur_mean[d] * mult;
333 }
334 initialized = true;
335 } else {
336 for (Index_ d = 0; d < dim; ++d) {
337 cur_output[d] += cur_mean[d] * mult;
338 }
339 }
340 }
341
342 assert(initialized);
343 }
344
345 // Combining the RSS.
346 for (Group_ g = 0; g < num_groups; ++g) {
347 const auto cur_global = output.mean[g];
348 const auto cur_output = output.rss[g];
349 bool initialized = false;
350
351 for (int u = 0; u < nused; ++u) {
352 const auto cur_count = (*((*all_partial_count)[u]))[g];
353 if (cur_count == 0) { // This check allows us to use the unsafe RSS centering below.
354 continue;
355 }
356
357 const auto cur_mean = (*(ap_mean[u]))[g];
358 if (u == 0) { // Special case to avoid trying to access u - 1.
359 for (Index_ d = 0; d < dim; ++d) {
360 cur_output[d] = quickstats::recenter_rss_unsafe(cur_count, cur_output[d], cur_mean[d], cur_global[d]);
361 }
362 initialized = true;
363 } else {
364 const auto cur_rss = (*(ap_rss[u - 1]))[g];
365 if (!initialized) { // Don't use u == 0, as the first non-empty 'g' might not occur in the first thread.
366 for (Index_ d = 0; d < dim; ++d) {
367 cur_output[d] = quickstats::recenter_rss_unsafe(cur_count, cur_rss[d], cur_mean[d], cur_global[d]);
368 }
369 initialized = true;
370 } else {
371 for (Index_ d = 0; d < dim; ++d) {
372 cur_output[d] += quickstats::recenter_rss_unsafe(cur_count, cur_rss[d], cur_mean[d], cur_global[d]);
373 }
374 }
375 }
376 }
377
378 assert(initialized);
379 }
380 }
381}
382
383template<typename Value_, typename Index_, typename Group_, typename Count_, typename Output_>
384void group_rss_running(
385 const bool row,
387 const Group_* const group,
388 const Group_ num_groups,
389 const Count_* const group_size,
390 const GroupRssBuffers<Output_>& output,
391 const GroupRssOptions<Output_>& opt
392) {
393 const auto dim = (row ? mat.nrow() : mat.ncol());
394 const auto otherdim = (row ? mat.ncol() : mat.nrow());
395 if (otherdim == 0) {
396 for (Group_ g = 0; g < num_groups; ++g) {
397 std::fill_n(output.mean[g], dim, opt.mean_placeholder);
398 std::fill_n(output.rss[g], dim, 0);
399 }
400 return;
401 }
402
403 std::size_t num_empty = 0;
404 for (Group_ g = 0; g < num_groups; ++g) {
405 num_empty += (group_size[g] == 0);
406 }
407
408 // We strip out all empty groups so that we can skip some allocations in each thread.
409 std::optional<std::vector<Count_> > new_group_size_store;
410 std::optional<GroupRssBuffers<Output_> > new_output_store;
411 std::optional<std::vector<Group_> > new_group_store;
412 Group_ num_non_empty;
413 const Count_* new_group_size;
414 const Group_* new_group;
415 const GroupRssBuffers<Output_>* new_output;
416
417 if (num_empty > 0) {
418 num_non_empty = num_groups - num_empty; // must be positive, otherwise otherdim == 0 if all groups are empty.
419 new_group_size_store.emplace();
420 new_group_size_store->reserve(num_non_empty);
421 new_output_store.emplace();
422 new_output_store->mean.reserve(num_non_empty);
423 new_output_store->rss.reserve(num_non_empty);
424
425 auto mapping = sanisizer::create<std::vector<std::size_t> >(num_groups);
426 for (Group_ g = 0; g < num_groups; ++g) {
427 if (group_size[g]) {
428 mapping[g] = new_group_size_store->size();
429 new_group_size_store->push_back(group_size[g]);
430 new_output_store->mean.push_back(output.mean[g]);
431 new_output_store->rss.push_back(output.rss[g]);
432 } else {
433 std::fill_n(output.mean[g], dim, opt.mean_placeholder);
434 std::fill_n(output.rss[g], dim, 0);
435 }
436 }
437
438 new_group_store.emplace(tatami::cast_Index_to_container_size<std::vector<Group_> >(otherdim));
439 for (Index_ i = 0; i < otherdim; ++i) {
440 (*new_group_store)[i] = mapping[group[i]];
441 }
442
443 new_group_size = new_group_size_store->data();
444 new_output = &(*new_output_store);
445 new_group = new_group_store->data();
446 } else {
447 num_non_empty = num_groups;
448 new_group_size = group_size;
449 new_output = &output;
450 new_group = group;
451 }
452
453 group_rss_running_nonempty(
454 row,
455 dim,
456 otherdim,
457 mat,
458 new_group,
459 num_non_empty,
460 new_group_size,
461 *new_output,
462 opt
463 );
464}
489template<typename Value_, typename Index_, typename Group_, typename Count_, typename Output_>
491 bool row,
493 const Group_* const group,
494 const Group_ num_groups,
495 const Count_* const group_size,
496 const GroupRssBuffers<Output_>& output,
497 const GroupRssOptions<Output_>& opt
498) {
499 assert(sanisizer::is_equal(num_groups, output.mean.size()));
500 assert(sanisizer::is_equal(num_groups, output.rss.size()));
501 if (mat.prefer_rows() == row) {
502 group_rss_direct(row, mat, group, num_groups, group_size, output, opt);
503 } else {
504 group_rss_running(row, mat, group, num_groups, group_size, output, opt);
505 }
506}
507
527template<typename Value_, typename Index_, typename Group_, typename Output_>
529 bool row,
531 const Group_* const group,
532 const Group_ num_groups,
533 const GroupRssBuffers<Output_>& output,
534 const GroupRssOptions<Output_>& opt
535) {
536 auto group_size = sanisizer::create<std::vector<Index_> >(num_groups);
537 const auto otherdim = (row ? mat.ncol() : mat.nrow());
538 for (Index_ o = 0; o < otherdim; ++o) {
539 group_size[group[o]] += 1;
540 }
541 group_rss(row, mat, group, num_groups, group_size.data(), output, opt);
542}
543
549template<typename Output_>
556 std::vector<std::vector<Output_> > mean;
557
563 std::vector<std::vector<Output_> > rss;
564};
565
584template<typename Output_, typename Value_, typename Index_, typename Group_>
586 bool row,
588 const Group_* const group,
589 const Group_ num_groups,
590 const GroupRssOptions<Output_>& opt
591) {
593 sanisizer::resize(output.mean, num_groups);
594 sanisizer::resize(output.rss, num_groups);
595
597 sanisizer::resize(buffers.mean, num_groups);
598 sanisizer::resize(buffers.rss, num_groups);
599
600 const auto dim = (row ? mat.nrow() : mat.ncol());
601 for (Group_ g = 0; g < num_groups; ++g) {
603#ifdef TATAMI_STATS_TEST_DIRTY
604 , -1
605#endif
606 );
607 buffers.mean[g] = output.mean[g].data();
609#ifdef TATAMI_STATS_TEST_DIRTY
610 , -1
611#endif
612 );
613 buffers.rss[g] = output.rss[g].data();
614 }
615
616 group_rss(row, mat, group, num_groups, buffers, opt);
617 return output;
618}
619
620}
621
622#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
Float_ recenter_rss_unsafe(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean)
void update_rss_with_zeros_unsafe(Output_ &mean, Output_ &rss, const Count_ num_zeros, const Count_ num_total)
constexpr Value_ nan_if_available_else_zero()
void update_rss(Output_ &mean, Output_ &rss, const Input_ value, const Count_ num_total)
Functions to compute statistics from a tatami::Matrix.
Definition count.hpp:20
void range(bool row, const tatami::Matrix< Value_, Index_ > &mat, RangeBuffers< Output_ > &output, const RangeOptions< Output_ > &opt)
Definition range.hpp:339
void group_rss(bool row, const tatami::Matrix< Value_, Index_ > &mat, const Group_ *const group, const Group_ num_groups, const Count_ *const group_size, const GroupRssBuffers< Output_ > &output, const GroupRssOptions< Output_ > &opt)
Definition group_rss.hpp:490
void resize_container_to_Index_size(Container_ &container, const Index_ x, Args_ &&... args)
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)
Result buffers for group_rss().
Definition group_rss.hpp:52
std::vector< Output_ * > rss
Definition group_rss.hpp:65
std::vector< Output_ * > mean
Definition group_rss.hpp:58
Options for group_rss().
Definition group_rss.hpp:32
int num_threads
Definition group_rss.hpp:37
Output_ mean_placeholder
Definition group_rss.hpp:43
Results of group_rss().
Definition group_rss.hpp:550
std::vector< std::vector< Output_ > > mean
Definition group_rss.hpp:556
std::vector< std::vector< Output_ > > rss
Definition group_rss.hpp:563