vstat
Loading...
Searching...
No Matches
univariate.hpp
1// SPDX-License-Identifier: MIT
2// SPDX-FileCopyrightText: Copyright 2020-2024 Heal Research
3
4#ifndef VSTAT_UNIVARIATE_HPP
5#define VSTAT_UNIVARIATE_HPP
6
7#include <limits>
8#include <ostream>
9#include <type_traits>
10
11#include "combine.hpp"
12
13namespace VSTAT_NAMESPACE
14{
15
25enum class stats { sum, mean, variance };
26
35template<typename T, stats Stats = stats::variance>
37{
38 static auto load_state(T sw, T sx, T sxx) noexcept -> univariate_accumulator
39 {
41 acc.sum_w = sw;
42 acc.sum_w_old = sw;
43 acc.sum_x = sx;
44 acc.sum_xx = sxx;
45 return acc;
46 }
47
48 static auto load_state(std::tuple<T, T, T> state) noexcept -> univariate_accumulator
49 {
50 auto [sw, sx, sxx] = state;
51 return load_state(sw, sx, sxx);
52 }
53
54 void operator()(T x) noexcept
55 {
56 if constexpr (Stats == stats::variance) {
57 T dx = (sum_w * x) - sum_x;
58 sum_x += x;
59 sum_w += 1;
60 // guards 0/0: a prior weighted zero-weight call can leave sum_w_old at 0
61 T denom = sum_w * sum_w_old;
62 sum_xx += eve::if_else(denom != T{0}, (dx * dx) / denom, T{0});
63 sum_w_old = sum_w;
64 } else if constexpr (Stats == stats::mean) {
65 sum_x += x;
66 sum_w += 1;
67 } else {
68 sum_x += x;
69 }
70 }
71
72 void operator()(T x, T w) noexcept
73 {
74 if constexpr (Stats == stats::variance) {
75 x *= w;
76 T dx = (sum_w * x) - (sum_x * w);
77 sum_x += x;
78 sum_w += w;
79 T denom = w * sum_w * sum_w_old;
80 sum_xx += eve::if_else(denom != T{0}, (dx * dx) / denom, T{0});
81 sum_w_old = sum_w;
82 } else if constexpr (Stats == stats::mean) {
83 sum_x += x * w;
84 sum_w += w;
85 } else {
86 sum_x += x * w;
87 }
88 }
89
90 template<typename U>
91 requires eve::simd_value<T> && eve::simd_compatible_ptr<U, T>
92 void operator()(U const* x) noexcept
93 {
94 (*this)(T {x});
95 }
96
97 template<typename U>
98 requires eve::simd_value<T> && eve::simd_compatible_ptr<U, T>
99 void operator()(U const* x, U const* w) noexcept
100 {
101 (*this)(T {x}, T {w});
102 }
103
104 // Returns { sum_w, sum_x, sum_xx }.
105 // Fields not tracked by Stats are 0: sum_w when Stats==sum, sum_xx when Stats!=variance.
106 [[nodiscard]] auto stats() const noexcept -> std::tuple<double, double, double>
107 {
108 if constexpr (std::is_floating_point_v<T>) {
109 return {sum_w, sum_x, sum_xx};
110 } else if constexpr (Stats == stats::variance) {
111 return {eve::reduce(sum_w), eve::reduce(sum_x), combine(sum_w, sum_x, sum_xx)};
112 } else if constexpr (Stats == stats::mean) {
113 return {eve::reduce(sum_w), eve::reduce(sum_x), 0.0};
114 } else {
115 // Stats == stats::sum: sum_w was never written, skip the reduce
116 return {0.0, eve::reduce(sum_x), 0.0};
117 }
118 }
119
120 private:
121 T sum_w {0};
122 T sum_w_old {1};
123 T sum_x {0};
124 T sum_xx {0};
125};
126
131{
132 double count;
133 double sum;
134 double ssr;
135 double mean;
136 double variance;
137 double sample_variance;
138
139 template<typename T, stats Stats>
141 {
142 auto [sw, sx, sxx] = accumulator.stats();
143 if constexpr (Stats == stats::sum) {
144 count = std::numeric_limits<double>::quiet_NaN();
145 } else {
146 count = sw;
147 }
148 sum = sx;
149 ssr = sxx;
150 if constexpr (Stats != stats::sum) {
151 mean = sx / sw;
152 } else {
153 mean = std::numeric_limits<double>::quiet_NaN();
154 }
155 if constexpr (Stats == stats::variance) {
156 variance = sxx / sw;
157 sample_variance = sxx / (sw - 1);
158 } else {
159 variance = std::numeric_limits<double>::quiet_NaN();
160 sample_variance = std::numeric_limits<double>::quiet_NaN();
161 }
162 }
163};
164
165inline auto operator<<(std::ostream& os, univariate_statistics const& stats) -> std::ostream&
166{
167 os << "count: \t" << stats.count << "\nsum: \t" << stats.sum << "\nssr: \t"
168 << stats.ssr << "\nmean: \t" << stats.mean << "\nvariance: \t" << stats.variance
169 << "\nsample variance:\t" << stats.sample_variance << "\n";
170 return os;
171}
172
173} // namespace VSTAT_NAMESPACE
174
175#endif
Univariate accumulator object.
Definition univariate.hpp:37
Univariate statistics.
Definition univariate.hpp:131