Skip to main content

rill_ml/metrics/
regression.rs

1//! Regression metrics: MAE, MSE, RMSE, R².
2
3use crate::error::{
4    RillError, checked_finite_add, checked_increment, ensure_finite, ensure_finite_target,
5};
6use crate::traits::Metric;
7
8/// Mean Absolute Error.
9#[derive(Debug, Clone, Default)]
10#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
11pub struct Mae {
12    sum_abs_error: f64,
13    count: u64,
14}
15
16impl Mae {
17    /// Create a new MAE accumulator.
18    pub const fn new() -> Self {
19        Self {
20            sum_abs_error: 0.0,
21            count: 0,
22        }
23    }
24}
25
26impl Metric for Mae {
27    type Truth = f64;
28    type Prediction = f64;
29
30    fn update(&mut self, truth: f64, prediction: f64) -> Result<(), RillError> {
31        ensure_finite_target(truth)?;
32        ensure_finite("prediction", prediction)?;
33        let error = truth - prediction;
34        ensure_finite("absolute error", error)?;
35        let next_sum = checked_finite_add(self.sum_abs_error, error.abs(), "MAE sum")?;
36        let next_count = checked_increment(self.count, "MAE sample")?;
37        self.sum_abs_error = next_sum;
38        self.count = next_count;
39        Ok(())
40    }
41
42    fn value(&self) -> Option<f64> {
43        if self.count == 0 {
44            None
45        } else {
46            Some(self.sum_abs_error / self.count as f64)
47        }
48    }
49
50    fn samples_seen(&self) -> u64 {
51        self.count
52    }
53
54    fn reset(&mut self) {
55        self.sum_abs_error = 0.0;
56        self.count = 0;
57    }
58}
59
60/// Mean Squared Error.
61#[derive(Debug, Clone, Default)]
62#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
63pub struct Mse {
64    sum_sq_error: f64,
65    count: u64,
66}
67
68impl Mse {
69    /// Create a new MSE accumulator.
70    pub const fn new() -> Self {
71        Self {
72            sum_sq_error: 0.0,
73            count: 0,
74        }
75    }
76}
77
78impl Metric for Mse {
79    type Truth = f64;
80    type Prediction = f64;
81
82    fn update(&mut self, truth: f64, prediction: f64) -> Result<(), RillError> {
83        ensure_finite_target(truth)?;
84        ensure_finite("prediction", prediction)?;
85        let err = truth - prediction;
86        ensure_finite("squared error input", err)?;
87        let squared_error = err * err;
88        ensure_finite("squared error", squared_error)?;
89        let next_sum = checked_finite_add(self.sum_sq_error, squared_error, "MSE sum")?;
90        let next_count = checked_increment(self.count, "MSE sample")?;
91        self.sum_sq_error = next_sum;
92        self.count = next_count;
93        Ok(())
94    }
95
96    fn value(&self) -> Option<f64> {
97        if self.count == 0 {
98            None
99        } else {
100            Some(self.sum_sq_error / self.count as f64)
101        }
102    }
103
104    fn samples_seen(&self) -> u64 {
105        self.count
106    }
107
108    fn reset(&mut self) {
109        self.sum_sq_error = 0.0;
110        self.count = 0;
111    }
112}
113
114/// Root Mean Squared Error.
115#[derive(Debug, Clone, Default)]
116#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
117pub struct Rmse {
118    mse: Mse,
119}
120
121impl Rmse {
122    /// Create a new RMSE accumulator.
123    pub const fn new() -> Self {
124        Self { mse: Mse::new() }
125    }
126}
127
128impl Metric for Rmse {
129    type Truth = f64;
130    type Prediction = f64;
131
132    fn update(&mut self, truth: f64, prediction: f64) -> Result<(), RillError> {
133        self.mse.update(truth, prediction)
134    }
135
136    fn value(&self) -> Option<f64> {
137        self.mse.value().map(|v| v.sqrt())
138    }
139
140    fn samples_seen(&self) -> u64 {
141        self.mse.samples_seen()
142    }
143
144    fn reset(&mut self) {
145        self.mse.reset();
146    }
147}
148
149/// R² (coefficient of determination).
150///
151/// Uses Welford's online algorithm for the variance of the truth, avoiding
152/// the catastrophic cancellation that `sum(y²) - n · mean(y)²` suffers on
153/// large-offset, small-variance data. Returns `None` when fewer than 2
154/// samples have been seen or when the truth variance is zero (constant
155/// truth).
156#[derive(Debug, Clone, Default)]
157#[cfg_attr(feature = "serde", derive(serde::Serialize))]
158pub struct R2 {
159    ss_res: f64,
160    mean_truth: f64,
161    /// Running `M2 = sum((y_i - mean)^2)` from Welford's algorithm.
162    m2_truth: f64,
163    count: u64,
164}
165
166#[cfg(feature = "serde")]
167impl<'de> serde::Deserialize<'de> for R2 {
168    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
169    where
170        D: serde::Deserializer<'de>,
171    {
172        #[derive(serde::Deserialize)]
173        struct R2State {
174            ss_res: f64,
175            mean_truth: f64,
176            m2_truth: f64,
177            count: u64,
178        }
179
180        let state = R2State::deserialize(deserializer)?;
181        if !state.ss_res.is_finite() || state.ss_res < 0.0 {
182            return Err(serde::de::Error::custom(
183                "r2 ss_res must be finite and non-negative",
184            ));
185        }
186        if !state.mean_truth.is_finite() {
187            return Err(serde::de::Error::custom("r2 mean_truth must be finite"));
188        }
189        if !state.m2_truth.is_finite() || state.m2_truth < 0.0 {
190            return Err(serde::de::Error::custom(
191                "r2 m2_truth must be finite and non-negative",
192            ));
193        }
194        // count == 0: no samples seen → all accumulators must be exactly 0.
195        // The normal paths (new() / reset()) guarantee this; a non-zero
196        // value indicates a corrupted or malicious payload.
197        if state.count == 0 {
198            if state.ss_res != 0.0 {
199                return Err(serde::de::Error::custom(format!(
200                    "r2 ss_res must be 0 when count == 0, got {}",
201                    state.ss_res
202                )));
203            }
204            if state.mean_truth != 0.0 {
205                return Err(serde::de::Error::custom(format!(
206                    "r2 mean_truth must be 0 when count == 0, got {}",
207                    state.mean_truth
208                )));
209            }
210            if state.m2_truth != 0.0 {
211                return Err(serde::de::Error::custom(format!(
212                    "r2 m2_truth must be 0 when count == 0, got {}",
213                    state.m2_truth
214                )));
215            }
216        }
217        // count == 1: after a single Welford update, M2 is exactly 0
218        // (delta2 = truth - mean = 0, so m2_delta = 0). The normal update
219        // path guarantees exact 0, so no floating-point tolerance is
220        // introduced here. ss_res may be non-zero because the single
221        // prediction can have an error.
222        if state.count == 1 && state.m2_truth != 0.0 {
223            return Err(serde::de::Error::custom(format!(
224                "r2 m2_truth must be 0 when count == 1, got {}",
225                state.m2_truth
226            )));
227        }
228        Ok(R2 {
229            ss_res: state.ss_res,
230            mean_truth: state.mean_truth,
231            m2_truth: state.m2_truth,
232            count: state.count,
233        })
234    }
235}
236
237impl R2 {
238    /// Create a new R² accumulator.
239    pub const fn new() -> Self {
240        Self {
241            ss_res: 0.0,
242            mean_truth: 0.0,
243            m2_truth: 0.0,
244            count: 0,
245        }
246    }
247}
248
249impl Metric for R2 {
250    type Truth = f64;
251    type Prediction = f64;
252
253    fn update(&mut self, truth: f64, prediction: f64) -> Result<(), RillError> {
254        ensure_finite_target(truth)?;
255        ensure_finite("prediction", prediction)?;
256        let err = truth - prediction;
257        ensure_finite("R2 error", err)?;
258        let squared_error = err * err;
259        ensure_finite("R2 squared error", squared_error)?;
260
261        // Welford update for the truth variance.
262        let next_count = checked_increment(self.count, "R2 sample")?;
263        let delta = truth - self.mean_truth;
264        ensure_finite("R2 welford delta", delta)?;
265        let next_mean = self.mean_truth + delta / next_count as f64;
266        ensure_finite("R2 welford mean", next_mean)?;
267        let delta2 = truth - next_mean;
268        ensure_finite("R2 welford delta2", delta2)?;
269        let m2_delta = delta * delta2;
270        ensure_finite("R2 welford m2_delta", m2_delta)?;
271        let next_m2 = checked_finite_add(self.m2_truth, m2_delta, "R2 m2_truth")?;
272
273        let next_ss_res = checked_finite_add(self.ss_res, squared_error, "R2 residual sum")?;
274
275        // Commit atomically.
276        self.count = next_count;
277        self.mean_truth = next_mean;
278        self.m2_truth = next_m2;
279        self.ss_res = next_ss_res;
280        Ok(())
281    }
282
283    fn value(&self) -> Option<f64> {
284        if self.count < 2 {
285            return None;
286        }
287        // M2 is the sum of squared deviations; treat floating-point noise
288        // that drives it slightly negative as zero (constant truth).
289        if self.m2_truth <= 0.0 {
290            return None;
291        }
292        Some(1.0 - self.ss_res / self.m2_truth)
293    }
294
295    fn samples_seen(&self) -> u64 {
296        self.count
297    }
298
299    fn reset(&mut self) {
300        self.ss_res = 0.0;
301        self.mean_truth = 0.0;
302        self.m2_truth = 0.0;
303        self.count = 0;
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310
311    #[test]
312    fn mae_basic() {
313        let mut m = Mae::new();
314        m.update(3.0, 5.0).unwrap(); // err=2
315        m.update(5.0, 4.0).unwrap(); // err=1
316        assert!((m.value().unwrap() - 1.5).abs() < 1e-12);
317    }
318
319    #[test]
320    fn mse_basic() {
321        let mut m = Mse::new();
322        m.update(3.0, 5.0).unwrap(); // err=2, sq=4
323        m.update(5.0, 4.0).unwrap(); // err=1, sq=1
324        assert!((m.value().unwrap() - 2.5).abs() < 1e-12);
325    }
326
327    #[test]
328    fn metrics_reject_overflow_without_mutating_state() {
329        let mut mae = Mae::new();
330        let mut mse = Mse::new();
331        let mut r2 = R2::new();
332
333        assert!(mae.update(f64::MAX, -f64::MAX).is_err());
334        assert!(mse.update(f64::MAX, 0.0).is_err());
335        assert!(r2.update(f64::MAX, 0.0).is_err());
336
337        assert_eq!(mae.samples_seen(), 0);
338        assert_eq!(mse.samples_seen(), 0);
339        assert_eq!(r2.samples_seen(), 0);
340    }
341
342    #[test]
343    fn rmse_basic() {
344        let mut m = Rmse::new();
345        m.update(3.0, 5.0).unwrap();
346        m.update(5.0, 4.0).unwrap();
347        assert!((m.value().unwrap() - 2.5_f64.sqrt()).abs() < 1e-12);
348    }
349
350    #[test]
351    fn r2_perfect_prediction_is_one() {
352        let mut m = R2::new();
353        m.update(1.0, 1.0).unwrap();
354        m.update(2.0, 2.0).unwrap();
355        m.update(3.0, 3.0).unwrap();
356        assert!((m.value().unwrap() - 1.0).abs() < 1e-9);
357    }
358
359    #[test]
360    fn r2_mean_prediction_is_zero() {
361        let mut m = R2::new();
362        // predict the mean every time
363        m.update(1.0, 2.0).unwrap();
364        m.update(3.0, 2.0).unwrap();
365        // mean=2, ss_res = 1+1=2, ss_tot = 1+1=2 -> R2=0
366        assert!((m.value().unwrap()).abs() < 1e-9);
367    }
368
369    #[test]
370    fn r2_insufficient_data_returns_none() {
371        let mut m = R2::new();
372        m.update(1.0, 1.0).unwrap();
373        assert!(m.value().is_none());
374    }
375
376    #[test]
377    fn r2_constant_truth_returns_none() {
378        let mut m = R2::new();
379        m.update(5.0, 3.0).unwrap();
380        m.update(5.0, 4.0).unwrap();
381        assert!(m.value().is_none());
382    }
383
384    #[test]
385    fn r2_welford_large_offset_small_variance() {
386        // Large offset, tiny variance: the old `sum(y²) - n · mean(y)²`
387        // formula lost all precision. Welford must remain accurate.
388        let truths = [
389            1_000_000_000_001.0,
390            1_000_000_000_002.0,
391            1_000_000_000_003.0,
392        ];
393        let mut m = R2::new();
394        for y in truths {
395            // Perfect prediction: R² should be 1.0.
396            m.update(y, y).unwrap();
397        }
398        assert!((m.value().unwrap() - 1.0).abs() < 1e-9);
399
400        // Predict the (known) mean every time: R² should be ~0.0.
401        let mean = truths.iter().sum::<f64>() / truths.len() as f64;
402        let mut m = R2::new();
403        for y in truths {
404            m.update(y, mean).unwrap();
405        }
406        assert!(m.value().unwrap().abs() < 1e-6);
407    }
408
409    #[test]
410    #[cfg(feature = "serde")]
411    fn r2_partial_update_is_atomic() {
412        // Restore a counter near overflow; the Welford update must fail
413        // without mutating any state.
414        let json = format!(
415            "{{\"ss_res\":1.0,\"mean_truth\":1.0,\"m2_truth\":1.0,\"count\":{}}}",
416            u64::MAX
417        );
418        let mut m: R2 = serde_json::from_str(&json).unwrap();
419        let result = m.update(1.0, 1.0);
420        assert!(result.is_err(), "expected counter overflow");
421        assert_eq!(m.count, u64::MAX);
422        assert_eq!(m.ss_res, 1.0);
423        assert_eq!(m.mean_truth, 1.0);
424        assert_eq!(m.m2_truth, 1.0);
425    }
426
427    #[test]
428    #[cfg(feature = "serde")]
429    fn r2_serde_rejects_negative_m2() {
430        let json = "{\"ss_res\":0.0,\"mean_truth\":0.0,\"m2_truth\":-1.0,\"count\":2}";
431        assert!(serde_json::from_str::<R2>(json).is_err());
432    }
433
434    #[test]
435    #[cfg(feature = "serde")]
436    fn r2_serde_rejects_negative_ss_res() {
437        // Regression: a malicious or corrupted state with negative residual
438        // sum of squares must be rejected. ``ss_res`` is a sum of squares and
439        // cannot be negative outside of floating-point corruption.
440        let json = "{\"ss_res\":-0.5,\"mean_truth\":0.0,\"m2_truth\":1.0,\"count\":2}";
441        assert!(serde_json::from_str::<R2>(json).is_err());
442    }
443
444    #[test]
445    #[cfg(feature = "serde")]
446    fn r2_serde_roundtrip_preserves_state() {
447        let mut m = R2::new();
448        m.update(1.0, 1.0).unwrap();
449        m.update(2.0, 1.5).unwrap();
450        m.update(3.0, 2.5).unwrap();
451        let before = m.value();
452        let json = serde_json::to_string(&m).unwrap();
453        let restored: R2 = serde_json::from_str(&json).unwrap();
454        assert_eq!(restored.count, 3);
455        assert_eq!(restored.value(), before);
456    }
457
458    #[test]
459    fn non_finite_rejected() {
460        let mut m = Mae::new();
461        assert!(m.update(f64::NAN, 1.0).is_err());
462        assert!(m.update(1.0, f64::INFINITY).is_err());
463    }
464
465    #[test]
466    fn empty_metric_returns_none() {
467        assert!(Mae::new().value().is_none());
468        assert!(Mse::new().value().is_none());
469        assert!(Rmse::new().value().is_none());
470        assert!(R2::new().value().is_none());
471    }
472
473    // -----------------------------------------------------------------
474    // §6.3: R² count state consistency
475    // -----------------------------------------------------------------
476
477    #[test]
478    #[cfg(feature = "serde")]
479    fn r2_serde_rejects_count_zero_with_nonzero_ss_res() {
480        let json = "{\"ss_res\":1.0,\"mean_truth\":0.0,\"m2_truth\":0.0,\"count\":0}";
481        assert!(serde_json::from_str::<R2>(json).is_err());
482    }
483
484    #[test]
485    #[cfg(feature = "serde")]
486    fn r2_serde_rejects_count_zero_with_nonzero_mean() {
487        let json = "{\"ss_res\":0.0,\"mean_truth\":5.0,\"m2_truth\":0.0,\"count\":0}";
488        assert!(serde_json::from_str::<R2>(json).is_err());
489    }
490
491    #[test]
492    #[cfg(feature = "serde")]
493    fn r2_serde_rejects_count_zero_with_nonzero_m2() {
494        let json = "{\"ss_res\":0.0,\"mean_truth\":0.0,\"m2_truth\":3.0,\"count\":0}";
495        assert!(serde_json::from_str::<R2>(json).is_err());
496    }
497
498    #[test]
499    #[cfg(feature = "serde")]
500    fn r2_serde_rejects_count_one_with_nonzero_m2() {
501        // After a single Welford update, M2 is exactly 0. A non-zero M2 at
502        // count == 1 indicates a corrupted or malicious payload.
503        let json = "{\"ss_res\":0.5,\"mean_truth\":3.0,\"m2_truth\":0.25,\"count\":1}";
504        assert!(serde_json::from_str::<R2>(json).is_err());
505    }
506
507    #[test]
508    #[cfg(feature = "serde")]
509    fn r2_serde_accepts_count_one_with_nonzero_ss_res() {
510        // A single sample can have a non-zero prediction error, so ss_res
511        // at count == 1 is legitimately non-zero. Only m2_truth must be 0.
512        let json = "{\"ss_res\":2.5,\"mean_truth\":3.0,\"m2_truth\":0.0,\"count\":1}";
513        let m: R2 = serde_json::from_str(json).unwrap();
514        assert_eq!(m.count, 1);
515        // value() returns None for count < 2.
516        assert!(m.value().is_none());
517    }
518}