tatami_stats
Matrix statistics for tatami
Loading...
Searching...
No Matches
group_rss.hpp
Go to the documentation of this file.
1#ifndef TATAMI_STATS_SKIP_NAN_GROUP_RSS_HPP
2#define TATAMI_STATS_SKIP_NAN_GROUP_RSS_HPP
3
4#include <vector>
5#include <algorithm>
6#include <cstddef>
7#include <optional>
8#include <cassert>
9#include <limits>
10#include <cmath>
11
12#include "tatami/tatami.hpp"
13#include "sanisizer/sanisizer.hpp"
15#include "jiwoo/jiwoo.hpp"
16
17#include "../group_rss.hpp"
18
25namespace tatami_stats {
26
27namespace skip_nan {
28
33template<typename Output_ = double>
47
55template<typename Output_, typename Count_>
62 std::vector<Output_*> mean;
63
69 std::vector<Output_*> rss;
70
76 std::vector<Count_*> count;
77};
78
82template<typename Value_, typename Index_, typename Group_, typename Output_, typename Count_>
83void group_rss_direct(
84 const bool row,
86 const Group_* const group,
87 const Group_ num_groups,
90) {
91 const auto dim = (row ? mat.nrow() : mat.ncol());
92 const auto otherdim = (row ? mat.ncol() : mat.nrow());
93
94 if (mat.sparse()) {
95 auto full_group_sizes = sanisizer::create<std::vector<Index_> >(num_groups);
96 for (Index_ i = 0; i < otherdim; ++i) {
97 full_group_sizes[group[i]] += 1;
98 }
99
100 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
101 auto ext = tatami::consecutive_extractor<true>(mat, row, s, l);
104 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
105 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
106 auto cur_non_zeros = sanisizer::create<std::vector<Index_> >(num_groups);
107 auto cur_sizes = sanisizer::create<std::vector<Index_> >(num_groups);
108
109 for (Index_ x = 0; x < l; ++x) {
110 auto range = ext->fetch(vbuffer.data(), ibuffer.data());
111
112 // Computing the mean first.
113 for (Index_ i = 0; i < range.number; ++i) {
114 const auto val = range.value[i];
115 const auto b = group[range.index[i]];
116 if (!std::isnan(val)) {
117 ++cur_non_zeros[b];
118 cur_means[b] += val;
119 } else {
120 ++cur_sizes[b];
121 }
122 }
123 for (Group_ g = 0; g < num_groups; ++g) {
124 const auto actual_size = full_group_sizes[g] - cur_sizes[g];
125 cur_sizes[g] = actual_size;
126 output.count[g][s + x] = actual_size;
127 }
128 group_rss_finish_means(num_groups, cur_sizes.data(), cur_means, static_cast<Index_>(s + x), output.mean, opt.mean_placeholder);
129
130 // Now computing the RSS.
131 for (Index_ i = 0; i < range.number; ++i) {
132 const auto val = range.value[i];
133 if (!std::isnan(val)) {
134 const auto g = group[range.index[i]];
135 const auto delta = val - cur_means[g];
136 cur_rss[g] += delta * delta;
137 }
138 }
139 for (Group_ g = 0; g < num_groups; ++g) {
140 if (cur_sizes[g] > 0) { // preserve RSS = 0 for empty groups, otherwise placeholder mean might cause problems.
141 const Output_ my_rss = cur_rss[g] + cur_means[g] * cur_means[g] * (cur_sizes[g] - cur_non_zeros[g]);
142 output.rss[g][s + x] = my_rss;
143 } else {
144 output.rss[g][s + x] = 0;
145 }
146 }
147
148 std::fill(cur_means.begin(), cur_means.end(), 0);
149 std::fill(cur_rss.begin(), cur_rss.end(), 0);
150 std::fill(cur_non_zeros.begin(), cur_non_zeros.end(), 0);
151 std::fill(cur_sizes.begin(), cur_sizes.end(), 0);
152 }
153 }, dim, opt.num_threads);
154
155 } else {
156 tatami::parallelize([&](int, Index_ s, Index_ l) -> void {
157 auto ext = tatami::consecutive_extractor<false>(mat, row, s, l);
159 auto cur_means = sanisizer::create<std::vector<Output_> >(num_groups);
160 auto cur_rss = sanisizer::create<std::vector<Output_> >(num_groups);
161 auto cur_sizes = sanisizer::create<std::vector<Index_> >(num_groups);
162
163 for (Index_ x = 0; x < l; ++x) {
164 auto ptr = ext->fetch(buffer.data());
165
166 // Computing the mean first.
167 for (Index_ j = 0; j < otherdim; ++j) {
168 const auto val = ptr[j];
169 if (!std::isnan(val)) {
170 const auto g = group[j];
171 cur_means[g] += val;
172 ++cur_sizes[g];
173 }
174 }
175 for (Group_ g = 0; g < num_groups; ++g) {
176 output.count[g][s + x] = cur_sizes[g];
177 }
178 group_rss_finish_means(num_groups, cur_sizes.data(), cur_means, static_cast<Index_>(s + x), output.mean, opt.mean_placeholder);
179
180 // Now computing the RSS.
181 for (Index_ j = 0; j < otherdim; ++j) {
182 const auto val = ptr[j];
183 if (!std::isnan(val)) {
184 const auto g = group[j];
185 const auto delta = val - cur_means[g];
186 cur_rss[g] += delta * delta;
187 }
188 }
189 for (Group_ g = 0; g < num_groups; ++g) {
190 output.rss[g][s + x] = cur_rss[g];
191 }
192
193 std::fill(cur_means.begin(), cur_means.end(), 0);
194 std::fill(cur_rss.begin(), cur_rss.end(), 0);
195 std::fill(cur_sizes.begin(), cur_sizes.end(), 0);
196 }
197 }, dim, opt.num_threads);
198 }
199}
200
201template<typename Value_, typename Index_, typename Group_, typename Output_, typename Count_>
202void group_rss_running(
203 const bool row,
205 const Group_* const group,
206 const Group_ num_groups,
207 const GroupRssBuffers<Output_, Count_>& output,
208 const GroupRssOptions<Output_>& opt
209) {
210 const auto dim = (row ? mat.nrow() : mat.ncol());
211 for (Group_ g = 0; g < num_groups; ++g) {
212 std::fill_n(output.rss[g], dim, 0);
213 std::fill_n(output.count[g], dim, 0);
214 }
215
216 const auto otherdim = (row ? mat.ncol() : mat.nrow());
217 if (otherdim == 0) {
218 for (Group_ g = 0; g < num_groups; ++g) {
219 std::fill_n(output.mean[g], dim, opt.mean_placeholder);
220 }
221 return;
222 } else {
223 for (Group_ g = 0; g < num_groups; ++g) {
224 std::fill_n(output.mean[g], dim, 0);
225 }
226 }
227
228 const bool do_parallel = opt.num_threads > 1;
229 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Output_> > > > all_partial_mean, all_partial_rss;
230 std::optional<std::vector<std::optional<jiwoo::EquilengthArrays<Count_> > > > all_partial_count;
231 if (do_parallel) {
232 // -1, as we'll repurpose the RSS output buffer to store the partial RSS of the first thread.
233 all_partial_rss.emplace(sanisizer::cast<I<decltype(all_partial_rss->size())> >(opt.num_threads - 1));
234 all_partial_mean.emplace(sanisizer::cast<I<decltype(all_partial_mean->size())> >(opt.num_threads));
235 all_partial_count.emplace(sanisizer::cast<I<decltype(all_partial_count->size())> >(opt.num_threads));
236 }
237
238 const bool is_sparse = mat.is_sparse();
239 const int nused = tatami::parallelize([&](int thread, Index_ s, Index_ l) -> void {
240 std::optional<jiwoo::EquilengthArrays<Output_> > cur_mean, cur_rss;
241 std::optional<jiwoo::EquilengthArrays<Count_> > cur_count;
242
243 Output_* const * mean_ptrs;
244 Output_* const * rss_ptrs;
245 Count_* const * count_ptrs;
246 if (!do_parallel) {
247 // Storing mean and RSS directly in the output vector to cut down two allocations if we're not working in parallel.
248 mean_ptrs = output.mean.data();
249 rss_ptrs = output.rss.data();
250 count_ptrs = output.count.data();
251
252 } else {
253 // Storing the partial RSS directly in the output vectors to save ourselves an allocation if we're in the first thread.
254 if (thread == 0) {
255 rss_ptrs = output.rss.data();
256 } else {
257 cur_rss.emplace(
258 sanisizer::cast<I<decltype(cur_rss->size())> >(num_groups),
259 static_cast<std::size_t>(dim), // cast to size_t is safe due to the tatami contract.
260 0
261 );
262 rss_ptrs = cur_rss->get();
263 }
264
265 // 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.
266 cur_mean.emplace(
267 sanisizer::cast<I<decltype(cur_mean->size())> >(num_groups),
268 static_cast<std::size_t>(dim),
269 0
270 );
271 mean_ptrs = cur_mean->get();
272
273 // Similarly, we need to keep the global count separate from the partial count for reduction.
274 cur_count.emplace(
275 sanisizer::cast<I<decltype(cur_count->size())> >(num_groups),
276 static_cast<std::size_t>(dim),
277 0
278 );
279 count_ptrs = cur_count->get();
280 }
281
282 if (is_sparse) {
283 auto ext = tatami::consecutive_extractor<true>(mat, !row, s, l);
286 auto nonzeros = sanisizer::create<std::vector<std::vector<Count_> > >(num_groups);
287 for (Group_ g = 0; g < num_groups; ++g) {
289 }
290
291 auto cur_group_size = sanisizer::create<std::vector<Count_> >(num_groups);
292 for (Index_ x = 0; x < l; ++x) {
293 auto out = ext->fetch(vbuffer.data(), ibuffer.data());
294 const auto grp = group[s + x];
295 ++cur_group_size[grp];
296
297 const auto mptr = mean_ptrs[grp];
298 const auto rptr = rss_ptrs[grp];
299 const auto cptr = count_ptrs[grp];
300 auto& nnz = nonzeros[grp];
301
302 for (Index_ i = 0; i < out.number; ++i) {
303 const auto d = out.index[i];
304 const auto val = out.value[i];
305 if (!std::isnan(val)) {
306 quickstats::update_rss(mptr[d], rptr[d], val, ++nnz[d]); // increment is safe as 'nnz + 1 <= l' fits in an Index_.
307 } else {
308 ++cptr[d];
309 }
310 }
311 }
312
313 for (Group_ g = 0; g < num_groups; ++g) {
314 const auto mptr = mean_ptrs[g];
315 const auto rptr = rss_ptrs[g];
316 const auto cptr = count_ptrs[g];
317 const auto& nnz = nonzeros[g];
318 const auto curtotal = cur_group_size[g];
319
320 for (Index_ d = 0; d < dim; ++d) {
321 auto& unskipped_total = cptr[d];
322 unskipped_total = curtotal - unskipped_total; // could be zero, so the update with zeros needs to be safe.
323 quickstats::update_rss_with_zeros(mptr[d], rptr[d], static_cast<Count_>(unskipped_total - nnz[d]), unskipped_total);
324 }
325 }
326
327 } else {
328 auto ext = tatami::consecutive_extractor<false>(mat, !row, s, l);
330
331 for (Index_ x = 0; x < l; ++x) {
332 auto out = ext->fetch(buffer.data());
333 const auto grp = group[s + x];
334 const auto mptr = mean_ptrs[grp];
335 const auto rptr = rss_ptrs[grp];
336 const auto cptr = count_ptrs[grp];
337
338 for (Index_ d = 0; d < dim; ++d) {
339 const auto val = out[d];
340 if (!std::isnan(val)) {
341 quickstats::update_rss(mptr[d], rptr[d], val, ++cptr[d]); // increment is safe as 'cptr[d] + 1 <= l' fits in an Index_.
342 }
343 }
344 }
345 }
346
347 if (do_parallel) {
348 (*all_partial_count)[thread] = std::move(cur_count);
349 (*all_partial_mean)[thread] = std::move(cur_mean);
350 if (thread > 0) {
351 (*all_partial_rss)[thread - 1] = std::move(cur_rss);
352 }
353 }
354 }, otherdim, opt.num_threads);
355 assert(nused > 0);
356
357 if (do_parallel) {
358 const auto& ap_mean = *all_partial_mean;
359 const auto& ap_rss = *all_partial_rss;
360 const auto& ap_count = *all_partial_count;
361
362 for (Group_ g = 0; g < num_groups; ++g) {
363 const auto cur_global_count = output.count[g];
364 for (int u = 0; u < nused; ++u) {
365 const auto& cur_count = (*(ap_count[u]))[g];
366 for (Index_ d = 0; d < dim; ++d) {
367 cur_global_count[d] += cur_count[d];
368 }
369 }
370 }
371
372 // Computing the global mean.
373 for (Group_ g = 0; g < num_groups; ++g) {
374 const auto cur_global_count = output.count[g];
375 const auto cur_global_mean = output.mean[g];
376
377 for (int u = 0; u < nused; ++u) {
378 const auto& cur_mean = (*(ap_mean[u]))[g];
379 const auto& cur_count = (*(ap_count[u]))[g];
380 for (Index_ d = 0; d < dim; ++d) {
381 if (cur_count[d] > 0) {
382 const auto mult = static_cast<Output_>(cur_count[d]) / static_cast<Output_>(cur_global_count[d]);
383 cur_global_mean[d] += cur_mean[d] * mult;
384 }
385 }
386 }
387 }
388
389 // Combining the RSS. We need to use the safe variant of recenter_rss(), just to protect against the
390 // case where a group has no observations within a particular thread.
391 for (Group_ g = 0; g < num_groups; ++g) {
392 const auto cur_global_mean = output.mean[g];
393 const auto cur_output = output.rss[g];
394 for (int u = 0; u < nused; ++u) {
395 const auto& cur_mean = (*(ap_mean[u]))[g];
396 const auto& cur_count = (*(ap_count[u]))[g];
397 if (u == 0) {
398 for (Index_ d = 0; d < dim; ++d) {
399 cur_output[d] = quickstats::recenter_rss(cur_count[d], cur_output[d], cur_mean[d], cur_global_mean[d]);
400 }
401 } else {
402 const auto& cur_rss = (*(ap_rss[u - 1]))[g];
403 for (Index_ d = 0; d < dim; ++d) {
404 cur_output[d] += quickstats::recenter_rss(cur_count[d], cur_rss[d], cur_mean[d], cur_global_mean[d]);
405 }
406 }
407 }
408 }
409 }
410
411 for (Group_ g = 0; g < num_groups; ++g) {
412 const auto mptr = output.mean[g];
413 const auto cptr = output.count[g];
414 for (Index_ d = 0; d < dim; ++d) {
415 if (cptr[d] == 0) {
416 mptr[d] = opt.mean_placeholder;
417 }
418 }
419 }
420}
445template<typename Value_, typename Index_, typename Group_, typename Output_, typename Count_>
447 bool row,
449 const Group_* const group,
450 const Group_ num_groups,
452 const GroupRssOptions<Output_>& opt
453) {
454 assert(sanisizer::is_equal(num_groups, output.mean.size()));
455 assert(sanisizer::is_equal(num_groups, output.rss.size()));
456 assert(sanisizer::is_equal(num_groups, output.count.size()));
457 if (mat.prefer_rows() == row) {
458 group_rss_direct(row, mat, group, num_groups, output, opt);
459 } else {
460 group_rss_running(row, mat, group, num_groups, output, opt);
461 }
462}
463
471template<typename Output_, typename Count_>
478 std::vector<std::vector<Output_> > mean;
479
485 std::vector<std::vector<Output_> > rss;
486
492 std::vector<std::vector<Count_> > count;
493};
494
515template<typename Output_, typename Count_, typename Value_, typename Index_, typename Group_>
517 bool row,
519 const Group_* const group,
520 const Group_ num_groups,
521 const GroupRssOptions<Output_>& opt
522) {
524 sanisizer::resize(output.mean, num_groups);
525 sanisizer::resize(output.rss, num_groups);
526 sanisizer::resize(output.count, num_groups);
527
529 sanisizer::resize(buffers.mean, num_groups);
530 sanisizer::resize(buffers.rss, num_groups);
531 sanisizer::resize(buffers.count, num_groups);
532
533 const auto dim = (row ? mat.nrow() : mat.ncol());
534 for (Group_ g = 0; g < num_groups; ++g) {
536#ifdef TATAMI_STATS_TEST_DIRTY
537 , -1
538#endif
539 );
540 buffers.mean[g] = output.mean[g].data();
541
543#ifdef TATAMI_STATS_TEST_DIRTY
544 , -1
545#endif
546 );
547 buffers.rss[g] = output.rss[g].data();
548
550#ifdef TATAMI_STATS_TEST_DIRTY
551 , -1
552#endif
553 );
554 buffers.count[g] = output.count[g].data();
555 }
556
557 group_rss(row, mat, group, num_groups, buffers, opt);
558 return output;
559}
560
561}
562
563}
564
565#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
Compute group-wise residual sum of squares from a tatami::Matrix.
Float_ recenter_rss(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean)
constexpr Value_ nan_if_available_else_zero()
void update_rss_with_zeros(Output_ &mean, Output_ &rss, const Count_ num_zeros, const Count_ num_total)
void update_rss(Output_ &mean, Output_ &rss, const Input_ value, const Count_ num_total)
void range(bool row, const tatami::Matrix< Value_, Index_ > &mat, RangeBuffers< Output_, Count_ > &output, const RangeOptions< Output_ > &opt)
Definition range.hpp:419
void group_rss(bool row, const tatami::Matrix< Value_, Index_ > &mat, const Group_ *const group, const Group_ num_groups, const GroupRssBuffers< Output_, Count_ > &output, const GroupRssOptions< Output_ > &opt)
Definition group_rss.hpp:446
Functions to compute statistics from a tatami::Matrix.
Definition count.hpp:20
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)
Result buffers for skip_nan::group_rss().
Definition group_rss.hpp:56
std::vector< Output_ * > mean
Definition group_rss.hpp:62
std::vector< Output_ * > rss
Definition group_rss.hpp:69
std::vector< Count_ * > count
Definition group_rss.hpp:76
Options for skip_nan::group_rss().
Definition group_rss.hpp:34
int num_threads
Definition group_rss.hpp:39
Output_ mean_placeholder
Definition group_rss.hpp:45
Results of skip_nan::group_rss().
Definition group_rss.hpp:472
std::vector< std::vector< Output_ > > rss
Definition group_rss.hpp:485
std::vector< std::vector< Count_ > > count
Definition group_rss.hpp:492
std::vector< std::vector< Output_ > > mean
Definition group_rss.hpp:478