Skip to main content

rill_ml/stats/
ew_mean.rs

1//! Exponentially weighted mean.
2//!
3//! Time complexity per update: `O(1)`. Space complexity: `O(1)`.
4//!
5//! The update rule is `mean = alpha * x + (1 - alpha) * mean`. The first
6//! observation seeds the mean directly.
7
8use crate::error::{RillError, checked_increment, ensure_finite};
9#[cfg(feature = "serde")]
10use crate::persistence::ValidateState;
11use crate::traits::OnlineStatistic;
12
13/// Exponentially weighted moving average.
14///
15/// `alpha` must satisfy `0 < alpha <= 1`. Smaller values give more weight to
16/// the past; `alpha = 1` reduces to a `LastValue`-like tracker.
17///
18/// # Examples
19///
20/// ```
21/// use rill_ml::stats::ExponentiallyWeightedMean;
22/// use rill_ml::OnlineStatistic;
23///
24/// let mut ew = ExponentiallyWeightedMean::new(0.5).unwrap();
25/// ew.update(10.0).unwrap();
26/// ew.update(20.0).unwrap();
27/// // 10.0, then 0.5*20 + 0.5*10 = 15.0
28/// assert!((ew.value() - 15.0).abs() < 1e-12);
29/// ```
30#[derive(Debug, Clone)]
31#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
32pub struct ExponentiallyWeightedMean {
33    alpha: f64,
34    mean: f64,
35    count: u64,
36}
37
38impl ExponentiallyWeightedMean {
39    /// Create a new exponentially weighted mean accumulator.
40    ///
41    /// Returns an error if `alpha` is not in `(0, 1]`.
42    pub fn new(alpha: f64) -> Result<Self, RillError> {
43        ensure_finite("alpha", alpha)?;
44        if alpha <= 0.0 || alpha > 1.0 {
45            return Err(RillError::InvalidParameter {
46                name: "alpha",
47                value: alpha,
48            });
49        }
50        Ok(Self {
51            alpha,
52            mean: 0.0,
53            count: 0,
54        })
55    }
56
57    /// The configured alpha.
58    pub const fn alpha(&self) -> f64 {
59        self.alpha
60    }
61
62    /// Current weighted mean, or `0.0` if no observations have been seen.
63    pub const fn value(&self) -> f64 {
64        self.mean
65    }
66
67    /// Number of observations seen so far.
68    pub const fn count(&self) -> u64 {
69        self.count
70    }
71}
72
73#[cfg(feature = "serde")]
74impl ValidateState for ExponentiallyWeightedMean {
75    fn validate_state(&self) -> Result<(), RillError> {
76        ensure_finite("alpha", self.alpha)?;
77        ensure_finite("ew mean", self.mean)?;
78        if self.alpha <= 0.0 || self.alpha > 1.0 {
79            return Err(RillError::InvalidParameter {
80                name: "alpha",
81                value: self.alpha,
82            });
83        }
84        Ok(())
85    }
86}
87
88impl OnlineStatistic for ExponentiallyWeightedMean {
89    fn update(&mut self, value: f64) -> Result<(), RillError> {
90        ensure_finite("value", value)?;
91        let next_count = checked_increment(self.count, "EW mean sample")?;
92        let next_mean = if self.count == 0 {
93            value
94        } else {
95            self.alpha * value + (1.0 - self.alpha) * self.mean
96        };
97        ensure_finite("EW mean", next_mean)?;
98        self.mean = next_mean;
99        self.count = next_count;
100        Ok(())
101    }
102
103    fn samples_seen(&self) -> u64 {
104        self.count
105    }
106
107    fn reset(&mut self) {
108        self.mean = 0.0;
109        self.count = 0;
110    }
111}
112
113#[cfg(test)]
114mod tests {
115    use super::*;
116
117    #[test]
118    fn first_sample_seeds_mean() {
119        let mut ew = ExponentiallyWeightedMean::new(0.3).unwrap();
120        ew.update(10.0).unwrap();
121        assert!((ew.value() - 10.0).abs() < 1e-12);
122    }
123
124    #[test]
125    fn weighted_update_matches_formula() {
126        let mut ew = ExponentiallyWeightedMean::new(0.5).unwrap();
127        ew.update(10.0).unwrap();
128        ew.update(20.0).unwrap();
129        assert!((ew.value() - 15.0).abs() < 1e-12);
130    }
131
132    #[test]
133    fn alpha_one_tracks_last_value() {
134        let mut ew = ExponentiallyWeightedMean::new(1.0).unwrap();
135        ew.update(3.0).unwrap();
136        ew.update(7.0).unwrap();
137        assert!((ew.value() - 7.0).abs() < 1e-12);
138    }
139
140    #[test]
141    fn invalid_alpha_rejected() {
142        assert!(ExponentiallyWeightedMean::new(0.0).is_err());
143        assert!(ExponentiallyWeightedMean::new(-0.1).is_err());
144        assert!(ExponentiallyWeightedMean::new(1.5).is_err());
145        assert!(ExponentiallyWeightedMean::new(f64::NAN).is_err());
146    }
147
148    #[test]
149    fn reset_clears_state() {
150        let mut ew = ExponentiallyWeightedMean::new(0.5).unwrap();
151        ew.update(10.0).unwrap();
152        ew.reset();
153        assert_eq!(ew.count(), 0);
154        assert_eq!(ew.value(), 0.0);
155    }
156}