1#ifndef QUICKSTATS_RSS_HPP
2#define QUICKSTATS_RSS_HPP
9#include "sanisizer/sanisizer.hpp"
26template<
typename Output_ =
double>
46template<
typename Output_ =
double>
62template<
typename Output_ =
double>
97template<std::
size_t accumulators_ = 4,
typename Input_,
typename Output_>
99 static_assert(std::is_floating_point<Output_>::value);
102 if (num_total == 0) {
107 Output_& mean = output.
mean;
113 Output_& ssd = output.
rss;
116 [&](
const std::size_t i) -> Output_ {
117 const auto delta =
static_cast<Output_
>(ptr[i]) - mean;
118 return delta * delta;
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;
151template<std::
size_t accumulators_ = 4,
typename Input_,
typename Output_>
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) {
168 RssOptions<Output_> options;
169 options.max_sum_length = limit_;
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) {
182 RssOptions<Output_> options;
183 options.max_sum_length = limit_;
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);
235template<
typename Output_ =
double,
typename Count_>
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;
262template<
typename Output_ =
double,
typename Count_>
264 assert(num_total >= 0);
265 assert(num_total >= num_zeros);
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);
274 rss += mean * mean * ratio * num_zeros;
281template<
typename Input_,
typename Output_ =
double>
282class RssRunningDense {
284 RssRunningDense(
const std::size_t num_obj, Output_*
const mean, Output_*
const rss) :
289 assert(check_zeroed(num_obj, mean));
290 assert(check_zeroed(num_obj, rss));
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);
302 finish(nan_if_available_else_zero<Output_>());
305 void finish(
const Output_ mean_placeholder) {
306 if (my_count == 0 && mean_placeholder != 0) {
307 std::fill_n(my_mean, my_num_obj, mean_placeholder);
311 std::size_t num_obs()
const {
316 std::size_t my_num_obj;
319 std::size_t my_count = 0;
321 static_assert(std::is_floating_point<Output_>::value);
324template<
typename Count_,
typename Input_,
typename Output_ =
double>
325class RssRunningDenseSkip {
327 RssRunningDenseSkip(
const std::size_t num_obj, Output_* mean, Output_* rss, Count_* num_unskipped) :
331 my_num_unskipped(num_unskipped)
333 assert(check_zeroed(num_obj, mean));
334 assert(check_zeroed(num_obj, rss));
335 assert(check_zeroed(num_obj, num_unskipped));
339 template<
class Skip_>
340 void add(
const Input_* ptr, Skip_ skip) {
342 my_count = sanisizer::sum<Count_>(my_count, 1);
344 for (std::size_t i = 0; i < my_num_obj; ++i) {
345 const auto val = ptr[i];
347 update_rss(my_mean[i], my_rss[i], val, ++(my_num_unskipped[i]));
353 finish(nan_if_available_else_zero<Output_>());
356 void finish(
const Output_ mean_placeholder) {
357 if (mean_placeholder != 0) {
359 std::fill_n(my_mean, my_num_obj, mean_placeholder);
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;
370 Count_ num_obs()
const {
375 std::size_t my_num_obj;
379 Count_* my_num_unskipped;
381 static_assert(std::is_integral<Count_>::value);
382 static_assert(std::is_floating_point<Output_>::value);
385template<
typename Count_,
typename Input_,
typename Output_ =
double>
386class RssRunningSparse {
388 RssRunningSparse(
const std::size_t num_obj, Output_*
const mean, Output_*
const rss, Count_*
const num_non_zero) :
392 my_num_non_zero(num_non_zero)
394 assert(check_zeroed(num_obj, mean));
395 assert(check_zeroed(num_obj, rss));
396 assert(check_zeroed(num_obj, num_non_zero));
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);
404 my_count = sanisizer::sum<Count_>(my_count, 1);
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]));
413 finish(nan_if_available_else_zero<Output_>());
416 void finish(
const Output_ mean_placeholder) {
418 if (mean_placeholder != 0) {
419 std::fill_n(my_mean, my_num_obj, mean_placeholder);
422 for (std::size_t i = 0; i < my_num_obj; ++i) {
428 Count_ num_obs()
const {
433 std::size_t my_num_obj;
436 Count_* my_num_non_zero;
439 static_assert(std::is_integral<Count_>::value);
440 static_assert(std::is_floating_point<Output_>::value);
443template<
typename Count_,
typename Input_,
typename Output_ =
double>
444class RssRunningSparseSkip {
446 RssRunningSparseSkip(
const std::size_t num_obj, Output_*
const mean, Output_*
const rss, Count_*
const num_non_zero, Count_*
const num_unskipped) :
450 my_num_non_zero(num_non_zero),
451 my_num_unskipped(num_unskipped)
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));
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);
464 my_count = sanisizer::sum<Count_>(my_count, 1);
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];
470 ++my_num_unskipped[ri];
472 update_rss(my_mean[ri], my_rss[ri], val, ++(my_num_non_zero[ri]));
478 finish(nan_if_available_else_zero<Output_>());
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];
487 if (mean_placeholder != 0) {
488 std::fill_n(my_mean, my_num_obj, mean_placeholder);
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;
495 const auto num_unskipped = my_num_unskipped[i];
502 Count_ num_obs()
const {
507 std::size_t my_num_obj;
510 Count_* my_num_non_zero;
511 Count_* my_num_unskipped;
514 static_assert(std::is_integral<Count_>::value);
515 static_assert(std::is_floating_point<Output_>::value);
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;
565template<
typename Count_,
typename Float_>
566Float_
recenter_rss(
const Count_ num_total,
const Float_ old_rss,
const Float_ old_mean,
const Float_ new_mean) {
569 const Float_ delta = (num_total ? old_mean : new_mean) - new_mean;
570 return old_rss + num_total * delta * delta;
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();
581 return rss / (num_total - 1);
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());
592 for (std::size_t i = 0; i < num_obj; ++i) {
593 rss[i] /= num_total - 1;
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]);
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
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