rill_ml/metrics/
rolling.rs1use std::collections::VecDeque;
7
8use crate::error::{RillError, checked_finite_add, ensure_finite, ensure_finite_target};
9use crate::traits::Metric;
10
11#[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 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#[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 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#[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 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(); m.update(0.0, 4.0).unwrap(); assert!((m.value().unwrap() - 3.0).abs() < 1e-12);
219 m.update(0.0, 6.0).unwrap(); 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(); m.update(0.0, 4.0).unwrap(); assert!((m.value().unwrap() - 10.0).abs() < 1e-12);
229 m.update(0.0, 6.0).unwrap(); 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(); m.update(false, false).unwrap(); assert!((m.value().unwrap() - 1.0).abs() < 1e-12);
239 m.update(true, false).unwrap(); 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}