quickstats
Quickly compute simple statistics
Loading...
Searching...
No Matches
rss.hpp
Go to the documentation of this file.
1#ifndef QUICKSTATS_RSS_HPP
2#define QUICKSTATS_RSS_HPP
3
4#include <cassert>
5#include <limits>
6#include <cstddef>
7#include <algorithm>
8
9#include "sanisizer/sanisizer.hpp"
10
11#include "utils.hpp"
12#include "pairwise_sum.hpp"
13
19namespace quickstats {
20
26template<typename Output_ = double>
27struct RssResult {
32 Output_ mean = 0;
33
38 Output_ rss = 0;
39};
40
46template<typename Output_ = double>
56
62template<typename Output_ = double>
74
97template<std::size_t accumulators_ = 4, typename Input_, typename Output_>
98RssResult<Output_> rss(const std::size_t num_total, const std::size_t num_non_zero, const Input_* const ptr, RssWorkspace<Output_>& work, const RssOptions<Output_>& options) {
99 static_assert(std::is_floating_point<Output_>::value);
100
101 RssResult<Output_> output;
102 if (num_total == 0) {
103 output.mean = options.mean_placeholder;
104 return output;
105 }
106
107 Output_& mean = output.mean;
108 PairwiseSumOptions psopt;
109 psopt.max_sum_length = options.max_sum_length;
110 mean = pairwise_sum<accumulators_>(num_non_zero, ptr, work.pswork, psopt);
111 mean /= num_total;
112
113 Output_& ssd = output.rss;
115 num_non_zero,
116 [&](const std::size_t i) -> Output_ {
117 const auto delta = static_cast<Output_>(ptr[i]) - mean;
118 return delta * delta;
119 },
120 work.pswork,
121 psopt
122 );
123
124 assert(num_non_zero <= num_total);
125 if (num_non_zero < num_total) {
126 ssd += static_cast<Output_>(num_total - num_non_zero) * mean * mean;
127 }
128
129 return output;
130}
131
151template<std::size_t accumulators_ = 4, typename Input_, typename Output_>
152RssResult<Output_> rss(const std::size_t num_total, const Input_* const ptr, RssWorkspace<Output_>& work, const RssOptions<Output_>& options) {
153 return rss<accumulators_>(num_total, num_total, ptr, work, options);
154}
155
159// For back-compatibility.
160template<std::size_t limit_ = 128, std::size_t accumulators_ = 4, typename Input_, typename Output_>
161RssResult<Output_> rss(const std::size_t num_total, const std::size_t num_non_zero, const Input_* const ptr, RssWorkspace<Output_>& work) {
162 return rss<accumulators_>(
163 num_total,
164 num_non_zero,
165 ptr,
166 work,
167 [&]{
168 RssOptions<Output_> options;
169 options.max_sum_length = limit_;
170 return options;
171 }()
172 );
173}
174
175template<std::size_t limit_ = 128, std::size_t accumulators_ = 4, typename Input_, typename Output_>
176RssResult<Output_> rss(const std::size_t num_total, const Input_* const ptr, RssWorkspace<Output_>& work) {
177 return rss<accumulators_>(
178 num_total,
179 ptr,
180 work,
181 [&]{
182 RssOptions<Output_> options;
183 options.max_sum_length = limit_;
184 return options;
185 }()
186 );
187}
209template<typename Output_ = double, typename Input_, typename Count_>
210void update_rss(Output_& mean, Output_& rss, const Input_ value, const Count_ num_total) {
211 assert(num_total > 0);
212 Output_ delta = static_cast<Output_>(value) - mean;
213 mean += delta / num_total;
214 rss += delta * (static_cast<Output_>(value) - mean);
215}
216
235template<typename Output_ = double, typename Count_>
236void update_rss_with_zeros_unsafe(Output_& mean, Output_& rss, const Count_ num_zeros, const Count_ num_total) {
237 assert(num_total > 0);
238 assert(num_total >= num_zeros);
239 const auto ratio = static_cast<Output_>(num_total - num_zeros) / static_cast<Output_>(num_total);
240 rss += mean * mean * ratio * num_zeros;
241 mean *= ratio;
242}
243
262template<typename Output_ = double, typename Count_>
263void update_rss_with_zeros(Output_& mean, Output_& rss, const Count_ num_zeros, const Count_ num_total) {
264 assert(num_total >= 0);
265 assert(num_total >= num_zeros);
266
267 // We add '1' to both the numerator and denominator if it's empty, which ensures we get a ratio of 1 and a no-op to the mean and rss.
268 // This is guaranteed to not overflow Count_ as the sum will just be 1 if empty == 1.
269 // We use this approach to avoid introducing a 'if (num_total == 0)' conditional that interferes with autovectorization,
270 // at the cost of doing unnecessary work if there are many 'num_total == 0' cases.
271 const Count_ empty = (num_total == 0);
272 const auto ratio = static_cast<Output_>(num_total - num_zeros + empty) / static_cast<Output_>(num_total + empty);
273
274 rss += mean * mean * ratio * num_zeros;
275 mean *= ratio;
276}
277
281template<typename Input_, typename Output_ = double>
282class RssRunningDense {
283public:
284 RssRunningDense(const std::size_t num_obj, Output_* const mean, Output_* const rss) :
285 my_num_obj(num_obj),
286 my_mean(mean),
287 my_rss(rss)
288 {
289 assert(check_zeroed(num_obj, mean));
290 assert(check_zeroed(num_obj, rss));
291 }
292
293public:
294 void add(const Input_* const ptr) {
295 my_count = sanisizer::sum<std::size_t>(my_count, 1);
296 for (std::size_t i = 0; i < my_num_obj; ++i) {
297 update_rss(my_mean[i], my_rss[i], ptr[i], my_count);
298 }
299 }
300
301 void finish() {
302 finish(nan_if_available_else_zero<Output_>());
303 }
304
305 void finish(const Output_ mean_placeholder) {
306 if (my_count == 0 && mean_placeholder != 0) { // my_mean should already be zeroed, so no need to fill if our placeholder is also zero.
307 std::fill_n(my_mean, my_num_obj, mean_placeholder);
308 }
309 }
310
311 std::size_t num_obs() const {
312 return my_count;
313 }
314
315private:
316 std::size_t my_num_obj;
317 Output_* my_mean;
318 Output_* my_rss;
319 std::size_t my_count = 0;
320
321 static_assert(std::is_floating_point<Output_>::value);
322};
323
324template<typename Count_, typename Input_, typename Output_ = double>
325class RssRunningDenseSkip {
326public:
327 RssRunningDenseSkip(const std::size_t num_obj, Output_* mean, Output_* rss, Count_* num_unskipped) :
328 my_num_obj(num_obj),
329 my_mean(mean),
330 my_rss(rss),
331 my_num_unskipped(num_unskipped)
332 {
333 assert(check_zeroed(num_obj, mean));
334 assert(check_zeroed(num_obj, rss));
335 assert(check_zeroed(num_obj, num_unskipped));
336 }
337
338public:
339 template<class Skip_>
340 void add(const Input_* ptr, Skip_ skip) {
341 // my_count is the upper bound of all my_num_unskipped, so we check it once here to avoid having to check it in the loop.
342 my_count = sanisizer::sum<Count_>(my_count, 1);
343
344 for (std::size_t i = 0; i < my_num_obj; ++i) {
345 const auto val = ptr[i];
346 if (!skip(i, val)) {
347 update_rss(my_mean[i], my_rss[i], val, ++(my_num_unskipped[i]));
348 }
349 }
350 }
351
352 void finish() {
353 finish(nan_if_available_else_zero<Output_>());
354 }
355
356 void finish(const Output_ mean_placeholder) {
357 if (mean_placeholder != 0) { // my_mean should already be zeroed, so no need to fill if our placeholder is also zero.
358 if (my_count == 0) {
359 std::fill_n(my_mean, my_num_obj, mean_placeholder);
360 } else {
361 for (std::size_t i = 0; i < my_num_obj; ++i) {
362 if (my_num_unskipped[i] == 0) {
363 my_mean[i] = mean_placeholder;
364 }
365 }
366 }
367 }
368 }
369
370 Count_ num_obs() const {
371 return my_count;
372 }
373
374private:
375 std::size_t my_num_obj;
376 Output_* my_mean;
377 Output_* my_rss;
378 Count_ my_count = 0;
379 Count_* my_num_unskipped;
380
381 static_assert(std::is_integral<Count_>::value);
382 static_assert(std::is_floating_point<Output_>::value);
383};
384
385template<typename Count_, typename Input_, typename Output_ = double>
386class RssRunningSparse {
387public:
388 RssRunningSparse(const std::size_t num_obj, Output_* const mean, Output_* const rss, Count_* const num_non_zero) :
389 my_num_obj(num_obj),
390 my_mean(mean),
391 my_rss(rss),
392 my_num_non_zero(num_non_zero)
393 {
394 assert(check_zeroed(num_obj, mean));
395 assert(check_zeroed(num_obj, rss));
396 assert(check_zeroed(num_obj, num_non_zero));
397 }
398
399 template<typename Index_>
400 void add(const std::size_t num_non_zero_obs, const Input_* const value, const Index_* const index) {
401 static_assert(std::is_integral<Index_>::value);
402
403 // my_count is the upper bound of all my_num_non_zero, so no need to check individual increments.
404 my_count = sanisizer::sum<Count_>(my_count, 1);
405
406 for (std::size_t i = 0; i < num_non_zero_obs; ++i) {
407 const auto ri = index[i];
408 update_rss(my_mean[ri], my_rss[ri], value[i], ++(my_num_non_zero[ri]));
409 }
410 }
411
412 void finish() {
413 finish(nan_if_available_else_zero<Output_>());
414 }
415
416 void finish(const Output_ mean_placeholder) {
417 if (my_count == 0) {
418 if (mean_placeholder != 0) { // my_mean should already be zeroed, so no need to fill again if the placeholder is also zero.
419 std::fill_n(my_mean, my_num_obj, mean_placeholder);
420 }
421 } else {
422 for (std::size_t i = 0; i < my_num_obj; ++i) {
423 update_rss_with_zeros_unsafe(my_mean[i], my_rss[i], static_cast<Count_>(my_count - my_num_non_zero[i]), my_count);
424 }
425 }
426 }
427
428 Count_ num_obs() const {
429 return my_count;
430 }
431
432private:
433 std::size_t my_num_obj;
434 Output_* my_mean;
435 Output_* my_rss;
436 Count_* my_num_non_zero;
437 Count_ my_count = 0;
438
439 static_assert(std::is_integral<Count_>::value);
440 static_assert(std::is_floating_point<Output_>::value);
441};
442
443template<typename Count_, typename Input_, typename Output_ = double>
444class RssRunningSparseSkip {
445public:
446 RssRunningSparseSkip(const std::size_t num_obj, Output_* const mean, Output_* const rss, Count_* const num_non_zero, Count_* const num_unskipped) :
447 my_num_obj(num_obj),
448 my_mean(mean),
449 my_rss(rss),
450 my_num_non_zero(num_non_zero),
451 my_num_unskipped(num_unskipped)
452 {
453 assert(check_zeroed(num_obj, mean));
454 assert(check_zeroed(num_obj, rss));
455 assert(check_zeroed(num_obj, num_non_zero));
456 assert(check_zeroed(num_obj, num_unskipped));
457 }
458
459 template<typename Index_, class Skip_>
460 void add(const std::size_t num_non_zero_obs, const Input_* value, const Index_* index, Skip_ skip) {
461 static_assert(std::is_integral<Index_>::value);
462
463 // my_count is the upper bound of all my_num_non_zero, so no need to check individual increments.
464 my_count = sanisizer::sum<Count_>(my_count, 1);
465
466 for (std::size_t i = 0; i < num_non_zero_obs; ++i) {
467 const auto val = value[i];
468 const auto ri = index[i];
469 if (skip(ri, val)) {
470 ++my_num_unskipped[ri]; // storing the number that was skipped so we don't have to add the zeros later.
471 } else {
472 update_rss(my_mean[ri], my_rss[ri], val, ++(my_num_non_zero[ri]));
473 }
474 }
475 }
476
477 void finish() {
478 finish(nan_if_available_else_zero<Output_>());
479 }
480
481 void finish(const Output_ mean_placeholder) {
482 for (std::size_t i = 0; i < my_num_obj; ++i) {
483 my_num_unskipped[i] = my_count - my_num_unskipped[i];
484 }
485
486 if (my_count == 0) {
487 if (mean_placeholder != 0) { // my_mean should already be zeroed, so no need to do a fill if the placeholder is zero.
488 std::fill_n(my_mean, my_num_obj, mean_placeholder);
489 }
490 } else {
491 for (std::size_t i = 0; i < my_num_obj; ++i) {
492 if (my_num_unskipped[i] == 0) {
493 my_mean[i] = mean_placeholder;
494 } else {
495 const auto num_unskipped = my_num_unskipped[i];
496 update_rss_with_zeros_unsafe(my_mean[i], my_rss[i], static_cast<Count_>(num_unskipped - my_num_non_zero[i]), num_unskipped);
497 }
498 }
499 }
500 }
501
502 Count_ num_obs() const {
503 return my_count;
504 }
505
506private:
507 std::size_t my_num_obj;
508 Output_* my_mean;
509 Output_* my_rss;
510 Count_* my_num_non_zero;
511 Count_* my_num_unskipped;
512 Count_ my_count = 0;
513
514 static_assert(std::is_integral<Count_>::value);
515 static_assert(std::is_floating_point<Output_>::value);
516};
542template<typename Count_, typename Float_>
543Float_ recenter_rss_unsafe(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean) {
544 assert(num_total > 0 || old_mean == 0);
545 const Float_ delta = old_mean - new_mean;
546 return old_rss + num_total * delta * delta;
547}
548
565template<typename Count_, typename Float_>
566Float_ recenter_rss(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean) {
567 // If num_total == 0, we avoid the old_mean of NaN and just return old_rss by setting delta == 0.
568 // We minimize the scope of the conditional to make it easier to auto-vectorize, based on Godbolt experiments with GCC and clang.
569 const Float_ delta = (num_total ? old_mean : new_mean) - new_mean;
570 return old_rss + num_total * delta * delta;
571}
572
576template<typename Float_>
577Float_ rss_to_variance(const std::size_t num_total, const Float_ rss) {
578 if (num_total <= 1) {
579 return std::numeric_limits<Float_>::quiet_NaN();
580 } else {
581 return rss / (num_total - 1);
582 }
583}
584
585template<typename Float_>
586void rss_to_variance(const std::size_t num_obj, const std::size_t num_total, Float_* const rss) {
587 if (num_total <= 1) {
588 std::fill_n(rss, num_obj, std::numeric_limits<Float_>::quiet_NaN());
589 } else {
590 // For consistency with the other overloads, we won't do the '* (1/denom)' trick.
591 // It shouldn't have much effect on throughput anyway as the bottleneck should be reading from memory.
592 for (std::size_t i = 0; i < num_obj; ++i) {
593 rss[i] /= num_total - 1;
594 }
595 }
596}
597
598template<typename Count_, typename Float_>
599void rss_to_variance(const std::size_t num_obj, const Count_* const num_total, Float_* const rss) {
600 for (std::size_t i = 0; i < num_obj; ++i) {
601 rss[i] = rss_to_variance(num_total[i], rss[i]);
602 }
603}
608}
609
610#endif
Quickly compute simple statistics.
Definition mad.hpp:15
Output_ pairwise_sum_abstract(const std::size_t num_total, Input_ input, PairwiseSumWorkspace< Output_ > &work, const PairwiseSumOptions &options)
Definition pairwise_sum.hpp:102
RssResult< Output_ > rss(const std::size_t num_total, const std::size_t num_non_zero, const Input_ *const ptr, RssWorkspace< Output_ > &work, const RssOptions< Output_ > &options)
Definition rss.hpp:98
Float_ recenter_rss_unsafe(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean)
Definition rss.hpp:543
Float_ recenter_rss(const Count_ num_total, const Float_ old_rss, const Float_ old_mean, const Float_ new_mean)
Definition rss.hpp:566
void update_rss_with_zeros_unsafe(Output_ &mean, Output_ &rss, const Count_ num_zeros, const Count_ num_total)
Definition rss.hpp:236
constexpr Value_ nan_if_available_else_zero()
Definition utils.hpp:63
Output_ pairwise_sum(const std::size_t num_total, const Input_ *const ptr, PairwiseSumWorkspace< Output_ > &work, const PairwiseSumOptions &options)
Definition pairwise_sum.hpp:215
void update_rss_with_zeros(Output_ &mean, Output_ &rss, const Count_ num_zeros, const Count_ num_total)
Definition rss.hpp:263
void update_rss(Output_ &mean, Output_ &rss, const Input_ value, const Count_ num_total)
Definition rss.hpp:210
Pairwise summation.
Options for pairwise_sum() and pairwise_sum_abstract().
Definition pairwise_sum.hpp:42
std::size_t max_sum_length
Definition pairwise_sum.hpp:48
Re-usable workspace for pairwise_sum() and pairwise_sum_abstract().
Definition pairwise_sum.hpp:24
Options for rss().
Definition rss.hpp:63
Output_ mean_placeholder
Definition rss.hpp:72
std::size_t max_sum_length
Definition rss.hpp:67
Result of rss().
Definition rss.hpp:27
Output_ rss
Definition rss.hpp:38
Output_ mean
Definition rss.hpp:32
Re-usable workspace for rss().
Definition rss.hpp:47
Miscellaneous utilities.