Skip to main content

rill_ml/stats/
mean.rs

1//! Online mean using a numerically stable incremental update.
2//!
3//! Time complexity per update: `O(1)`. Space complexity: `O(1)`.
4
5use crate::error::{RillError, checked_increment, ensure_finite};
6#[cfg(feature = "serde")]
7use crate::persistence::ValidateState;
8use crate::traits::OnlineStatistic;
9
10/// Incremental mean computed with the delta method to minimize floating-point
11/// accumulation error.
12///
13/// # Examples
14///
15/// ```
16/// use rill_ml::stats::Mean;
17/// use rill_ml::OnlineStatistic;
18///
19/// let mut m = Mean::new();
20/// m.update(1.0).unwrap();
21/// m.update(2.0).unwrap();
22/// m.update(3.0).unwrap();
23/// assert_eq!(m.value(), 2.0);
24/// assert_eq!(m.samples_seen(), 3);
25/// ```
26#[derive(Debug, Clone, Default)]
27#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
28pub struct Mean {
29    count: u64,
30    mean: f64,
31}
32
33impl Mean {
34    /// Create a new empty mean accumulator.
35    pub const fn new() -> Self {
36        Self {
37            count: 0,
38            mean: 0.0,
39        }
40    }
41
42    /// Current mean, or `0.0` if no observations have been seen.
43    pub const fn value(&self) -> f64 {
44        self.mean
45    }
46
47    /// Number of observations seen so far.
48    pub const fn count(&self) -> u64 {
49        self.count
50    }
51}
52
53#[cfg(feature = "serde")]
54impl ValidateState for Mean {
55    fn validate_state(&self) -> Result<(), RillError> {
56        ensure_finite("mean", self.mean)?;
57        Ok(())
58    }
59}
60
61impl OnlineStatistic for Mean {
62    fn update(&mut self, value: f64) -> Result<(), RillError> {
63        ensure_finite("value", value)?;
64        let next_count = checked_increment(self.count, "mean sample")?;
65        let delta = value - self.mean;
66        ensure_finite("mean delta", delta)?;
67        let next_mean = self.mean + delta / next_count as f64;
68        ensure_finite("mean", next_mean)?;
69
70        self.count = next_count;
71        self.mean = next_mean;
72        Ok(())
73    }
74
75    fn samples_seen(&self) -> u64 {
76        self.count
77    }
78
79    fn reset(&mut self) {
80        self.count = 0;
81        self.mean = 0.0;
82    }
83}
84
85#[cfg(test)]
86mod tests {
87    use super::*;
88    use rand::SeedableRng;
89
90    #[test]
91    fn mean_of_simple_sequence() {
92        let mut m = Mean::new();
93        for x in [1.0, 2.0, 3.0, 4.0, 5.0] {
94            m.update(x).unwrap();
95        }
96        assert_eq!(m.value(), 3.0);
97        assert_eq!(m.count(), 5);
98    }
99
100    #[test]
101    fn mean_empty_is_zero() {
102        let m = Mean::new();
103        assert_eq!(m.value(), 0.0);
104        assert_eq!(m.count(), 0);
105    }
106
107    #[test]
108    fn mean_rejects_non_finite() {
109        let mut m = Mean::new();
110        assert!(m.update(f64::NAN).is_err());
111        assert!(m.update(f64::INFINITY).is_err());
112        assert_eq!(m.count(), 0);
113    }
114
115    #[test]
116    fn mean_rejects_overflow_without_mutating_state() {
117        let mut m = Mean::new();
118        m.update(f64::MAX).unwrap();
119        let before = m.clone();
120        assert!(m.update(-f64::MAX).is_err());
121        assert_eq!(m.count(), before.count());
122        assert_eq!(m.value(), before.value());
123    }
124
125    #[test]
126    fn mean_reset() {
127        let mut m = Mean::new();
128        m.update(10.0).unwrap();
129        m.update(20.0).unwrap();
130        m.reset();
131        assert_eq!(m.count(), 0);
132        assert_eq!(m.value(), 0.0);
133    }
134
135    #[test]
136    fn mean_matches_batch_formula() {
137        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
138        let mut m = Mean::new();
139        let mut data = Vec::new();
140        for _ in 0..1000 {
141            let x = rand::Rng::gen_range(&mut rng, -100.0..100.0);
142            m.update(x).unwrap();
143            data.push(x);
144        }
145        let batch: f64 = data.iter().sum::<f64>() / data.len() as f64;
146        assert!((m.value() - batch).abs() < 1e-9);
147    }
148}