use std::f64;
#[derive(Debug, Clone)]
pub struct StatisticsAccumulator {
sum_w: f64,
sum_wf: f64,
sum_wf2: f64,
integral: f64,
error: f64,
chi_square: f64,
iterations: usize,
}
impl StatisticsAccumulator {
pub fn new() -> Self {
Self {
sum_w: 0.0,
sum_wf: 0.0,
sum_wf2: 0.0,
integral: 0.0,
error: f64::INFINITY,
chi_square: 0.0,
iterations: 0,
}
}
pub fn add_sample(&mut self, weight: f64, f: f64) {
self.sum_w += weight;
self.sum_wf += weight * f;
self.sum_wf2 += weight * f * f;
}
pub fn samples(&self) -> usize {
self.sum_w as usize
}
pub fn finalize_iteration(&mut self) {
if self.sum_w <= 0.0 || self.sum_wf2 < 0.0 {
self.reset_iteration();
return;
}
let mean = self.sum_wf / self.sum_w;
let var = (self.sum_wf2 / self.sum_w) - mean * mean;
let sig2 = if var > 0.0 { var } else { 0.0 };
let iter_err = sig2.sqrt();
self.combine_iteration(mean, iter_err);
self.reset_iteration();
}
fn combine_iteration(&mut self, mean: f64, err: f64) {
let err = if err > 1e-150 { err } else { 1e-150 };
let w = 1.0 / (err * err);
if self.iterations == 0 {
self.integral = mean;
self.error = err;
self.chi_square = 0.0;
} else {
let prev_w = 1.0 / (self.error * self.error);
let new_w = prev_w + w;
let new_integral = (prev_w * self.integral + w * mean) / new_w;
let delta_prev = self.integral - new_integral;
let delta_cur = mean - new_integral;
self.chi_square += prev_w * delta_prev * delta_prev + w * delta_cur * delta_cur;
self.integral = new_integral;
self.error = new_w.sqrt().recip();
}
self.iterations += 1;
}
fn reset_iteration(&mut self) {
self.sum_w = 0.0;
self.sum_wf = 0.0;
self.sum_wf2 = 0.0;
}
pub fn integral(&self) -> f64 {
self.integral
}
pub fn error(&self) -> f64 {
self.error
}
pub fn chi_square(&self) -> f64 {
self.chi_square
}
pub fn iterations(&self) -> usize {
self.iterations
}
}
impl Default for StatisticsAccumulator {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn constant_integrand_converges_to_constant() {
let mut acc = StatisticsAccumulator::new();
for _ in 0..1000 {
acc.add_sample(1.0, 5.0);
}
acc.finalize_iteration();
assert!((acc.integral() - 5.0).abs() < 1e-12);
assert!(acc.error().abs() < 1e-9);
assert!(acc.chi_square().abs() < 1e-9);
}
#[test]
fn linear_integrand_matches_analytic() {
let mut acc = StatisticsAccumulator::new();
let n = 50_000u32;
for i in 0..n {
let x = (i as f64 + 0.5) / n as f64;
acc.add_sample(1.0, x);
}
acc.finalize_iteration();
assert!(
(acc.integral() - 0.5).abs() < 1e-3,
"got {}",
acc.integral()
);
}
#[test]
fn combine_two_iterations_uses_inverse_variance_weighting() {
let mut acc = StatisticsAccumulator::new();
for _ in 0..1000 {
acc.add_sample(1.0, 1.0);
}
acc.finalize_iteration();
for _ in 0..1000 {
acc.add_sample(1.0, 3.0);
}
acc.finalize_iteration();
assert!((acc.integral() - 2.0).abs() < 1e-9);
assert_eq!(acc.iterations(), 2);
}
}