Skip to main content

rill_ml/metrics/
rolling.rs

1//! Rolling metrics with fixed-size windows.
2//!
3//! These store per-sample contributions and correctly maintain running sums
4//! when the oldest contribution is evicted. Space complexity: `O(window_size)`.
5
6use std::collections::VecDeque;
7
8use crate::error::{RillError, checked_finite_add, ensure_finite, ensure_finite_target};
9use crate::traits::Metric;
10
11/// Rolling Mean Absolute Error.
12#[derive(Debug, Clone)]
13#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
14pub struct RollingMae {
15    errors: VecDeque<f64>,
16    sum: f64,
17    capacity: usize,
18}
19
20impl RollingMae {
21    /// Create a new rolling MAE with the given window capacity.
22    pub fn new(capacity: usize) -> Result<Self, RillError> {
23        if capacity == 0 {
24            return Err(RillError::InvalidWindowSize);
25        }
26        Ok(Self {
27            errors: VecDeque::with_capacity(capacity),
28            sum: 0.0,
29            capacity,
30        })
31    }
32}
33
34impl Metric for RollingMae {
35    type Truth = f64;
36    type Prediction = f64;
37
38    fn update(&mut self, truth: f64, prediction: f64) -> Result<(), RillError> {
39        ensure_finite_target(truth)?;
40        ensure_finite("prediction", prediction)?;
41        let err = (truth - prediction).abs();
42        ensure_finite("rolling absolute error", err)?;
43        let base_sum = if self.errors.len() == self.capacity {
44            checked_finite_add(
45                self.sum,
46                -self.errors.front().copied().unwrap_or(0.0),
47                "rolling MAE sum",
48            )?
49        } else {
50            self.sum
51        };
52        let next_sum = checked_finite_add(base_sum, err, "rolling MAE sum")?;
53        if self.errors.len() == self.capacity {
54            self.errors.pop_front();
55        }
56        self.errors.push_back(err);
57        self.sum = next_sum;
58        Ok(())
59    }
60
61    fn value(&self) -> Option<f64> {
62        if self.errors.is_empty() {
63            None
64        } else {
65            Some(self.sum / self.errors.len() as f64)
66        }
67    }
68
69    fn samples_seen(&self) -> u64 {
70        self.errors.len() as u64
71    }
72
73    fn reset(&mut self) {
74        self.errors.clear();
75        self.sum = 0.0;
76    }
77}
78
79/// Rolling Mean Squared Error.
80#[derive(Debug, Clone)]
81#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
82pub struct RollingMse {
83    errors: VecDeque<f64>,
84    sum: f64,
85    capacity: usize,
86}
87
88impl RollingMse {
89    /// Create a new rolling MSE.
90    pub fn new(capacity: usize) -> Result<Self, RillError> {
91        if capacity == 0 {
92            return Err(RillError::InvalidWindowSize);
93        }
94        Ok(Self {
95            errors: VecDeque::with_capacity(capacity),
96            sum: 0.0,
97            capacity,
98        })
99    }
100}
101
102impl Metric for RollingMse {
103    type Truth = f64;
104    type Prediction = f64;
105
106    fn update(&mut self, truth: f64, prediction: f64) -> Result<(), RillError> {
107        ensure_finite_target(truth)?;
108        ensure_finite("prediction", prediction)?;
109        let difference = truth - prediction;
110        ensure_finite("rolling squared error input", difference)?;
111        let err = difference.powi(2);
112        ensure_finite("rolling squared error", err)?;
113        let base_sum = if self.errors.len() == self.capacity {
114            checked_finite_add(
115                self.sum,
116                -self.errors.front().copied().unwrap_or(0.0),
117                "rolling MSE sum",
118            )?
119        } else {
120            self.sum
121        };
122        let next_sum = checked_finite_add(base_sum, err, "rolling MSE sum")?;
123        if self.errors.len() == self.capacity {
124            self.errors.pop_front();
125        }
126        self.errors.push_back(err);
127        self.sum = next_sum;
128        Ok(())
129    }
130
131    fn value(&self) -> Option<f64> {
132        if self.errors.is_empty() {
133            None
134        } else {
135            Some(self.sum / self.errors.len() as f64)
136        }
137    }
138
139    fn samples_seen(&self) -> u64 {
140        self.errors.len() as u64
141    }
142
143    fn reset(&mut self) {
144        self.errors.clear();
145        self.sum = 0.0;
146    }
147}
148
149/// Rolling Accuracy for binary classification.
150#[derive(Debug, Clone)]
151#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
152pub struct RollingAccuracy {
153    correct: VecDeque<bool>,
154    sum: u64,
155    capacity: usize,
156}
157
158impl RollingAccuracy {
159    /// Create a new rolling accuracy.
160    pub fn new(capacity: usize) -> Result<Self, RillError> {
161        if capacity == 0 {
162            return Err(RillError::InvalidWindowSize);
163        }
164        Ok(Self {
165            correct: VecDeque::with_capacity(capacity),
166            sum: 0,
167            capacity,
168        })
169    }
170}
171
172impl Metric for RollingAccuracy {
173    type Truth = bool;
174    type Prediction = bool;
175
176    fn update(&mut self, truth: bool, prediction: bool) -> Result<(), RillError> {
177        let is_correct = truth == prediction;
178        if self.correct.len() == self.capacity
179            && let Some(old) = self.correct.pop_front()
180            && old
181        {
182            self.sum -= 1;
183        }
184        self.correct.push_back(is_correct);
185        if is_correct {
186            self.sum += 1;
187        }
188        Ok(())
189    }
190
191    fn value(&self) -> Option<f64> {
192        if self.correct.is_empty() {
193            None
194        } else {
195            Some(self.sum as f64 / self.correct.len() as f64)
196        }
197    }
198
199    fn samples_seen(&self) -> u64 {
200        self.correct.len() as u64
201    }
202
203    fn reset(&mut self) {
204        self.correct.clear();
205        self.sum = 0;
206    }
207}
208
209#[cfg(test)]
210mod tests {
211    use super::*;
212
213    #[test]
214    fn rolling_mae_evicts_correctly() {
215        let mut m = RollingMae::new(2).unwrap();
216        m.update(0.0, 2.0).unwrap(); // err=2
217        m.update(0.0, 4.0).unwrap(); // err=4, window=[2,4], sum=6, mean=3
218        assert!((m.value().unwrap() - 3.0).abs() < 1e-12);
219        m.update(0.0, 6.0).unwrap(); // err=6, window=[4,6], sum=10, mean=5
220        assert!((m.value().unwrap() - 5.0).abs() < 1e-12);
221    }
222
223    #[test]
224    fn rolling_mse_evicts_correctly() {
225        let mut m = RollingMse::new(2).unwrap();
226        m.update(0.0, 2.0).unwrap(); // sq_err=4
227        m.update(0.0, 4.0).unwrap(); // sq_err=16, mean=(4+16)/2=10
228        assert!((m.value().unwrap() - 10.0).abs() < 1e-12);
229        m.update(0.0, 6.0).unwrap(); // sq_err=36, mean=(16+36)/2=26
230        assert!((m.value().unwrap() - 26.0).abs() < 1e-12);
231    }
232
233    #[test]
234    fn rolling_accuracy_evicts_correctly() {
235        let mut m = RollingAccuracy::new(2).unwrap();
236        m.update(true, true).unwrap(); // correct
237        m.update(false, false).unwrap(); // correct
238        assert!((m.value().unwrap() - 1.0).abs() < 1e-12);
239        m.update(true, false).unwrap(); // incorrect, window=[correct, incorrect]
240        assert!((m.value().unwrap() - 0.5).abs() < 1e-12);
241    }
242
243    #[test]
244    fn rolling_zero_capacity_rejected() {
245        assert!(RollingMae::new(0).is_err());
246        assert!(RollingMse::new(0).is_err());
247        assert!(RollingAccuracy::new(0).is_err());
248    }
249
250    #[test]
251    fn rolling_metrics_reject_overflow_without_mutating_state() {
252        let mut mae = RollingMae::new(2).unwrap();
253        let mut mse = RollingMse::new(2).unwrap();
254        mae.update(0.0, 1.0).unwrap();
255        mse.update(0.0, 1.0).unwrap();
256
257        assert!(mae.update(f64::MAX, -f64::MAX).is_err());
258        assert!(mse.update(f64::MAX, 0.0).is_err());
259
260        assert_eq!(mae.samples_seen(), 1);
261        assert_eq!(mse.samples_seen(), 1);
262        assert_eq!(mae.value(), Some(1.0));
263        assert_eq!(mse.value(), Some(1.0));
264    }
265
266    #[test]
267    fn rolling_empty_returns_none() {
268        assert!(RollingMae::new(5).unwrap().value().is_none());
269    }
270}