vstat
Loading...
Searching...
No Matches
bivariate.hpp
1// SPDX-License-Identifier: MIT
2// SPDX-FileCopyrightText: Copyright 2020-2024 Heal Research
3
4#ifndef VSTAT_BIVARIATE_HPP
5#define VSTAT_BIVARIATE_HPP
6
7#include "combine.hpp"
8#include "eve/module/core/regular/diff_of_prod.hpp"
9
10namespace VSTAT_NAMESPACE
11{
17template<typename T>
19{
20 static auto load_state(T sx, T sy, T sw, T sxx, T syy, T sxy) noexcept -> bivariate_accumulator // NOLINT
21 {
23 acc.sum_w = sw;
24 acc.sum_w_old = sw;
25 acc.sum_x = sx;
26 acc.sum_y = sy;
27 acc.sum_xx = sxx;
28 acc.sum_yy = syy;
29 acc.sum_xy = sxy;
30 return acc;
31 }
32
33 static auto load_state(std::tuple<T, T, T, T, T, T> state) noexcept -> bivariate_accumulator
34 {
35 auto [sx, sy, sw, sxx, syy, sxy] = state;
36 return load_state(sx, sy, sw, sxx, syy, sxy);
37 }
38
39 void operator()(T x, T y) noexcept
40 {
41 // Route through the weighted update with unit weight: the weighted
42 // overload already has the same zero-denominator guard the univariate
43 // univariate accumulator got in d843f76, while the original
44 // unweighted bivariate overload here computed `1/(sum_w * sum_w_old)`
45 // unconditionally and produced 0/0 -> NaN whenever a prior masked
46 // zero-weight call left sum_w_old at 0 (a real path now that
47 // bivariate::accumulate_finite exists). Delegating instead of
48 // duplicating the guard avoids the eve::if_else scalar-arg trap where
49 // `eve::if_else(cond, 1./d, T{0})` with mixed double/float args
50 // pathologically returns `1`, and keeps one implementation of the
51 // Welford update, not two.
52 (*this)(x, y, T{1});
53 }
54
55 void operator()(T x, T y, T w) noexcept // NOLINT
56 {
57 T dx = (x * sum_w) - sum_x;
58 T dy = (y * sum_w) - sum_y;
59
60 sum_x += x * w;
61 sum_y += y * w;
62 sum_w += w;
63
64 T denom = sum_w * sum_w_old;
65 T f = eve::if_else(denom != T{0}, w / denom, T{0});
66 sum_xx += f * dx * dx;
67 sum_yy += f * dy * dy;
68 sum_xy += f * dx * dy;
69
70 sum_w_old = sum_w;
71 }
72
73 template<typename U>
74 requires eve::simd_value<T> && eve::simd_compatible_ptr<U, T>
75 inline void operator()(U const* x, U const* y) noexcept
76 {
77 (*this)(T {x}, T {y});
78 }
79
80 template<typename U>
81 requires eve::simd_value<T> && eve::simd_compatible_ptr<U, T>
82 inline void operator()(U const* x, U const* y, U const* w) noexcept
83 {
84 (*this)(T {x}, T {y}, T {w});
85 }
86
87 // performs a reduction on the vector types and returns the sums and the squared residuals sums
88 [[nodiscard]] auto stats() const noexcept -> std::tuple<double, double, double, double, double, double>
89 {
90 if constexpr (std::is_floating_point_v<T>) {
91 return {sum_w, sum_x, sum_y, sum_xx, sum_yy, sum_xy};
92 } else {
93 auto [sxx, syy, sxy] = combine(sum_w, sum_x, sum_y, sum_xx, sum_yy, sum_xy);
94 return {eve::reduce(sum_w), eve::reduce(sum_x), eve::reduce(sum_y), sxx, syy, sxy};
95 }
96 }
97
98 private:
99 T sum_w {0};
100 T sum_w_old {1};
101 T sum_x {0};
102 T sum_y {0};
103 T sum_xx {0};
104 T sum_yy {0};
105 T sum_xy {0};
106};
107
112{
113 double count;
114 double sum_x;
115 double sum_y;
116 double ssr_x;
117 double ssr_y;
118 double sum_xy;
119 double mean_x;
120 double mean_y;
121 double variance_x;
122 double variance_y;
123 double sample_variance_x;
124 double sample_variance_y;
125 double correlation;
126 double covariance;
127 double sample_covariance;
128
129 template<typename T>
130 explicit bivariate_statistics(T accumulator)
131 {
132 auto [sw, sx, sy, sxx, syy, sxy] = accumulator.stats();
133 count = sw;
134 sum_x = sx;
135 sum_y = sy;
136 ssr_x = sxx;
137 ssr_y = syy;
138 sum_xy = sxy;
139 mean_x = sx / sw;
140 mean_y = sy / sw;
141 variance_x = sxx / sw;
142 variance_y = syy / sw;
143 sample_variance_x = sxx / (sw - 1);
144 sample_variance_y = syy / (sw - 1);
145
146 if (!(sxx > 0 && syy > 0)) {
147 correlation = static_cast<double>(sxx == syy);
148 } else {
149 correlation = sxy / std::sqrt(sxx * syy);
150 }
151
152 covariance = sxy / sw;
153 sample_covariance = sxy / (sw - 1);
154 }
155};
156
157inline auto operator<<(std::ostream& os, bivariate_statistics const& stats) -> std::ostream&
158{
159 os << "count: \t" << stats.count << "\nsum_x: \t" << stats.sum_x << "\nssr_x: \t"
160 << stats.ssr_x << "\nmean_x: \t" << stats.mean_x << "\nvariance_x: \t" << stats.variance_x
161 << "\nsample variance_x:\t" << stats.sample_variance_x << "\nsum_y: \t" << stats.sum_y
162 << "\nssr_y: \t" << stats.ssr_y << "\nmean_y: \t" << stats.mean_y
163 << "\nvariance_y: \t" << stats.variance_y << "\nsample variance_y:\t" << stats.sample_variance_y
164 << "\ncorrelation: \t" << stats.correlation << "\ncovariance: \t" << stats.covariance
165 << "\nsample covariance:\t" << stats.sample_covariance << "\n";
166 return os;
167}
168} // namespace VSTAT_NAMESPACE
169
170#endif
Bivariate accumulator object.
Definition bivariate.hpp:19
Bivariate statistics.
Definition bivariate.hpp:112