Skip to main content

rill_ml/metrics/
classification.rs

1//! Classification metrics: Accuracy, Precision, Recall, F1, LogLoss.
2
3use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
4use crate::loss::log_loss::BinaryLogLoss;
5use crate::traits::Metric;
6
7/// Accuracy for binary classification.
8#[derive(Debug, Clone, Default)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10pub struct Accuracy {
11    correct: u64,
12    count: u64,
13}
14
15impl Metric for Accuracy {
16    type Truth = bool;
17    type Prediction = bool;
18
19    fn update(&mut self, truth: bool, prediction: bool) -> Result<(), RillError> {
20        let next_count = checked_increment(self.count, "accuracy sample")?;
21        let next_correct = if truth == prediction {
22            checked_increment(self.correct, "accuracy correct")?
23        } else {
24            self.correct
25        };
26        self.count = next_count;
27        self.correct = next_correct;
28        Ok(())
29    }
30
31    fn value(&self) -> Option<f64> {
32        if self.count == 0 {
33            None
34        } else {
35            Some(self.correct as f64 / self.count as f64)
36        }
37    }
38
39    fn samples_seen(&self) -> u64 {
40        self.count
41    }
42
43    fn reset(&mut self) {
44        self.correct = 0;
45        self.count = 0;
46    }
47}
48
49/// Precision for the positive class.
50///
51/// `samples_seen()` reports the total number of successfully incorporated
52/// observations (including true negatives), not `TP + FP`. The confusion
53/// counts are kept separately so the metric remains computable.
54#[derive(Debug, Clone, Default)]
55#[cfg_attr(feature = "serde", derive(serde::Serialize))]
56pub struct Precision {
57    true_positive: u64,
58    false_positive: u64,
59    /// Total observations successfully incorporated via `update`.
60    /// Restored from serde as-is; older states without this field are
61    /// rejected because the true-negative count cannot be reconstructed
62    /// from `TP`/`FP` alone.
63    samples_seen: u64,
64}
65
66#[cfg(feature = "serde")]
67impl<'de> serde::Deserialize<'de> for Precision {
68    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
69    where
70        D: serde::Deserializer<'de>,
71    {
72        #[derive(serde::Deserialize)]
73        struct PrecisionState {
74            true_positive: u64,
75            false_positive: u64,
76            samples_seen: u64,
77        }
78
79        let state = PrecisionState::deserialize(deserializer)?;
80        // Internal consistency: samples_seen must be at least the
81        // confusion counts, since every TP/FP contributes one observation.
82        if state.samples_seen < state.true_positive.saturating_add(state.false_positive) {
83            return Err(serde::de::Error::custom("precision samples_seen < tp + fp"));
84        }
85        Ok(Precision {
86            true_positive: state.true_positive,
87            false_positive: state.false_positive,
88            samples_seen: state.samples_seen,
89        })
90    }
91}
92
93impl Metric for Precision {
94    type Truth = bool;
95    type Prediction = bool;
96
97    fn update(&mut self, truth: bool, prediction: bool) -> Result<(), RillError> {
98        let next_samples = checked_increment(self.samples_seen, "precision samples_seen")?;
99        let next_tp = if truth && prediction {
100            checked_increment(self.true_positive, "precision true positive")?
101        } else {
102            self.true_positive
103        };
104        let next_fp = if !truth && prediction {
105            checked_increment(self.false_positive, "precision false positive")?
106        } else {
107            self.false_positive
108        };
109        self.samples_seen = next_samples;
110        self.true_positive = next_tp;
111        self.false_positive = next_fp;
112        Ok(())
113    }
114
115    fn value(&self) -> Option<f64> {
116        let denominator = self.true_positive as f64 + self.false_positive as f64;
117        if denominator == 0.0 {
118            None
119        } else {
120            Some(self.true_positive as f64 / denominator)
121        }
122    }
123
124    fn samples_seen(&self) -> u64 {
125        self.samples_seen
126    }
127
128    fn reset(&mut self) {
129        self.true_positive = 0;
130        self.false_positive = 0;
131        self.samples_seen = 0;
132    }
133}
134
135/// Recall for the positive class.
136///
137/// `samples_seen()` reports the total number of successfully incorporated
138/// observations (including true negatives), not `TP + FN`.
139#[derive(Debug, Clone, Default)]
140#[cfg_attr(feature = "serde", derive(serde::Serialize))]
141pub struct Recall {
142    true_positive: u64,
143    false_negative: u64,
144    samples_seen: u64,
145}
146
147#[cfg(feature = "serde")]
148impl<'de> serde::Deserialize<'de> for Recall {
149    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
150    where
151        D: serde::Deserializer<'de>,
152    {
153        #[derive(serde::Deserialize)]
154        struct RecallState {
155            true_positive: u64,
156            false_negative: u64,
157            samples_seen: u64,
158        }
159
160        let state = RecallState::deserialize(deserializer)?;
161        if state.samples_seen < state.true_positive.saturating_add(state.false_negative) {
162            return Err(serde::de::Error::custom("recall samples_seen < tp + fn"));
163        }
164        Ok(Recall {
165            true_positive: state.true_positive,
166            false_negative: state.false_negative,
167            samples_seen: state.samples_seen,
168        })
169    }
170}
171
172impl Metric for Recall {
173    type Truth = bool;
174    type Prediction = bool;
175
176    fn update(&mut self, truth: bool, prediction: bool) -> Result<(), RillError> {
177        let next_samples = checked_increment(self.samples_seen, "recall samples_seen")?;
178        let next_tp = if truth && prediction {
179            checked_increment(self.true_positive, "recall true positive")?
180        } else {
181            self.true_positive
182        };
183        let next_fn = if truth && !prediction {
184            checked_increment(self.false_negative, "recall false negative")?
185        } else {
186            self.false_negative
187        };
188        self.samples_seen = next_samples;
189        self.true_positive = next_tp;
190        self.false_negative = next_fn;
191        Ok(())
192    }
193
194    fn value(&self) -> Option<f64> {
195        let denominator = self.true_positive as f64 + self.false_negative as f64;
196        if denominator == 0.0 {
197            None
198        } else {
199            Some(self.true_positive as f64 / denominator)
200        }
201    }
202
203    fn samples_seen(&self) -> u64 {
204        self.samples_seen
205    }
206
207    fn reset(&mut self) {
208        self.true_positive = 0;
209        self.false_negative = 0;
210        self.samples_seen = 0;
211    }
212}
213
214/// F1 score, the harmonic mean of precision and recall.
215///
216/// `samples_seen()` reports the total number of successfully incorporated
217/// observations (including true negatives), not `TP + FP + FN`.
218#[derive(Debug, Clone, Default)]
219#[cfg_attr(feature = "serde", derive(serde::Serialize))]
220pub struct F1Score {
221    true_positive: u64,
222    false_positive: u64,
223    false_negative: u64,
224    samples_seen: u64,
225}
226
227#[cfg(feature = "serde")]
228impl<'de> serde::Deserialize<'de> for F1Score {
229    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
230    where
231        D: serde::Deserializer<'de>,
232    {
233        #[derive(serde::Deserialize)]
234        struct F1State {
235            true_positive: u64,
236            false_positive: u64,
237            false_negative: u64,
238            samples_seen: u64,
239        }
240
241        let state = F1State::deserialize(deserializer)?;
242        let confusion = state
243            .true_positive
244            .saturating_add(state.false_positive)
245            .saturating_add(state.false_negative);
246        if state.samples_seen < confusion {
247            return Err(serde::de::Error::custom("f1 samples_seen < tp + fp + fn"));
248        }
249        Ok(F1Score {
250            true_positive: state.true_positive,
251            false_positive: state.false_positive,
252            false_negative: state.false_negative,
253            samples_seen: state.samples_seen,
254        })
255    }
256}
257
258impl Metric for F1Score {
259    type Truth = bool;
260    type Prediction = bool;
261
262    fn update(&mut self, truth: bool, prediction: bool) -> Result<(), RillError> {
263        let next_samples = checked_increment(self.samples_seen, "F1 samples_seen")?;
264        let next_tp = if truth && prediction {
265            checked_increment(self.true_positive, "F1 true positive")?
266        } else {
267            self.true_positive
268        };
269        let next_fp = if !truth && prediction {
270            checked_increment(self.false_positive, "F1 false positive")?
271        } else {
272            self.false_positive
273        };
274        let next_fn = if truth && !prediction {
275            checked_increment(self.false_negative, "F1 false negative")?
276        } else {
277            self.false_negative
278        };
279        self.samples_seen = next_samples;
280        self.true_positive = next_tp;
281        self.false_positive = next_fp;
282        self.false_negative = next_fn;
283        Ok(())
284    }
285
286    fn value(&self) -> Option<f64> {
287        let denominator = 2.0 * self.true_positive as f64
288            + self.false_positive as f64
289            + self.false_negative as f64;
290        if denominator == 0.0 {
291            None
292        } else {
293            Some(2.0 * self.true_positive as f64 / denominator)
294        }
295    }
296
297    fn samples_seen(&self) -> u64 {
298        self.samples_seen
299    }
300
301    fn reset(&mut self) {
302        self.true_positive = 0;
303        self.false_positive = 0;
304        self.false_negative = 0;
305        self.samples_seen = 0;
306    }
307}
308
309/// Binary log loss (cross-entropy).
310#[derive(Debug, Clone)]
311#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
312pub struct LogLoss {
313    loss: BinaryLogLoss,
314    sum_loss: f64,
315    count: u64,
316}
317
318impl Default for LogLoss {
319    fn default() -> Self {
320        Self {
321            loss: BinaryLogLoss::new(),
322            sum_loss: 0.0,
323            count: 0,
324        }
325    }
326}
327
328impl Metric for LogLoss {
329    type Truth = bool;
330    type Prediction = f64;
331
332    fn update(&mut self, truth: bool, prediction: f64) -> Result<(), RillError> {
333        ensure_finite("probability", prediction)?;
334        if !(0.0..=1.0).contains(&prediction) {
335            return Err(RillError::InvalidProbability(prediction));
336        }
337        let loss = self.loss.loss(prediction, truth);
338        ensure_finite("log loss", loss)?;
339        let next_sum = checked_finite_add(self.sum_loss, loss, "log loss sum")?;
340        let next_count = checked_increment(self.count, "log loss sample")?;
341        self.sum_loss = next_sum;
342        self.count = next_count;
343        Ok(())
344    }
345
346    fn value(&self) -> Option<f64> {
347        if self.count == 0 {
348            None
349        } else {
350            Some(self.sum_loss / self.count as f64)
351        }
352    }
353
354    fn samples_seen(&self) -> u64 {
355        self.count
356    }
357
358    fn reset(&mut self) {
359        self.sum_loss = 0.0;
360        self.count = 0;
361    }
362}
363
364#[cfg(test)]
365mod tests {
366    use super::*;
367
368    #[test]
369    fn accuracy_basic() {
370        let mut m = Accuracy::default();
371        m.update(true, true).unwrap();
372        m.update(false, false).unwrap();
373        m.update(true, false).unwrap();
374        assert!((m.value().unwrap() - 2.0 / 3.0).abs() < 1e-12);
375    }
376
377    #[test]
378    fn precision_basic() {
379        let mut m = Precision::default();
380        m.update(true, true).unwrap(); // tp
381        m.update(false, true).unwrap(); // fp
382        m.update(true, false).unwrap(); // fn
383        assert!((m.value().unwrap() - 0.5).abs() < 1e-12);
384    }
385
386    #[test]
387    fn recall_basic() {
388        let mut m = Recall::default();
389        m.update(true, true).unwrap(); // tp
390        m.update(false, true).unwrap(); // fp
391        m.update(true, false).unwrap(); // fn
392        assert!((m.value().unwrap() - 0.5).abs() < 1e-12);
393    }
394
395    #[test]
396    fn f1_basic() {
397        let mut m = F1Score::default();
398        m.update(true, true).unwrap(); // tp=1
399        m.update(false, true).unwrap(); // fp=1
400        m.update(true, false).unwrap(); // fn=1
401        // F1 = 2*1 / (2*1 + 1 + 1) = 0.5
402        assert!((m.value().unwrap() - 0.5).abs() < 1e-12);
403    }
404
405    #[test]
406    fn f1_perfect_is_one() {
407        let mut m = F1Score::default();
408        m.update(true, true).unwrap();
409        m.update(false, false).unwrap();
410        assert!((m.value().unwrap() - 1.0).abs() < 1e-12);
411    }
412
413    #[test]
414    fn log_loss_basic() {
415        let mut m = LogLoss::default();
416        m.update(true, 0.9).unwrap();
417        m.update(false, 0.1).unwrap();
418        let expected = (-0.9_f64.ln() + -0.9_f64.ln()) / 2.0;
419        assert!((m.value().unwrap() - expected).abs() < 1e-9);
420    }
421
422    #[test]
423    fn log_loss_rejects_invalid_probability() {
424        let mut m = LogLoss::default();
425        assert!(m.update(true, 1.5).is_err());
426        assert!(m.update(true, -0.1).is_err());
427        assert!(m.update(true, f64::NAN).is_err());
428    }
429
430    #[test]
431    fn empty_metrics_return_none() {
432        assert!(Accuracy::default().value().is_none());
433        assert!(Precision::default().value().is_none());
434        assert!(Recall::default().value().is_none());
435        assert!(F1Score::default().value().is_none());
436        assert!(LogLoss::default().value().is_none());
437    }
438
439    #[test]
440    fn precision_no_predictions_returns_none() {
441        let mut m = Precision::default();
442        m.update(true, false).unwrap();
443        m.update(false, false).unwrap();
444        assert!(m.value().is_none());
445    }
446
447    // -----------------------------------------------------------------
448    // Metric::samples_seen() contract: every successful update must
449    // increment the count by exactly one, including true negatives.
450    // -----------------------------------------------------------------
451
452    #[test]
453    fn samples_seen_counts_all_observations() {
454        let mut p = Precision::default();
455        let mut r = Recall::default();
456        let mut f = F1Score::default();
457        let mut a = Accuracy::default();
458
459        // All four confusion-matrix cells.
460        let cases = [(true, true), (true, false), (false, true), (false, false)];
461        for (truth, pred) in cases {
462            p.update(truth, pred).unwrap();
463            r.update(truth, pred).unwrap();
464            f.update(truth, pred).unwrap();
465            a.update(truth, pred).unwrap();
466        }
467
468        assert_eq!(p.samples_seen(), 4);
469        assert_eq!(r.samples_seen(), 4);
470        assert_eq!(f.samples_seen(), 4);
471        assert_eq!(a.samples_seen(), 4);
472    }
473
474    #[test]
475    #[cfg(feature = "serde")]
476    fn samples_seen_overflow_is_atomic() {
477        // Restore a near-overflow Precision via serde, then attempt one
478        // more update. The counter must overflow without mutating state.
479        let json = format!(
480            "{{\"true_positive\":1,\"false_positive\":1,\"samples_seen\":{}}}",
481            u64::MAX
482        );
483        let mut p: Precision = serde_json::from_str(&json).unwrap();
484        let result = p.update(true, true);
485        assert!(result.is_err(), "expected overflow");
486        assert_eq!(p.samples_seen(), u64::MAX);
487        assert_eq!(p.true_positive, 1);
488        assert_eq!(p.false_positive, 1);
489    }
490
491    #[test]
492    #[cfg(feature = "serde")]
493    fn precision_serde_rejects_missing_samples_seen() {
494        // Old state without samples_seen must be rejected: true-negative
495        // count cannot be reconstructed from TP/FP alone.
496        let json = "{\"true_positive\":1,\"false_positive\":1}";
497        assert!(serde_json::from_str::<Precision>(json).is_err());
498    }
499
500    #[test]
501    #[cfg(feature = "serde")]
502    fn precision_serde_rejects_inconsistent_samples_seen() {
503        // samples_seen < tp + fp is internally inconsistent.
504        let json = "{\"true_positive\":5,\"false_positive\":5,\"samples_seen\":3}";
505        assert!(serde_json::from_str::<Precision>(json).is_err());
506    }
507
508    #[test]
509    #[cfg(feature = "serde")]
510    fn recall_serde_rejects_missing_samples_seen() {
511        let json = "{\"true_positive\":1,\"false_negative\":1}";
512        assert!(serde_json::from_str::<Recall>(json).is_err());
513    }
514
515    #[test]
516    #[cfg(feature = "serde")]
517    fn f1_serde_rejects_missing_samples_seen() {
518        let json = "{\"true_positive\":1,\"false_positive\":1,\"false_negative\":1}";
519        assert!(serde_json::from_str::<F1Score>(json).is_err());
520    }
521
522    #[test]
523    #[cfg(feature = "serde")]
524    fn metric_serde_roundtrip_preserves_samples_seen() {
525        let mut p = Precision::default();
526        for _ in 0..10 {
527            p.update(true, true).unwrap();
528        }
529        let json = serde_json::to_string(&p).unwrap();
530        let restored: Precision = serde_json::from_str(&json).unwrap();
531        assert_eq!(restored.samples_seen(), 10);
532        assert_eq!(restored.true_positive, 10);
533    }
534
535    #[test]
536    fn reset_clears_samples_seen() {
537        let mut p = Precision::default();
538        p.update(true, true).unwrap();
539        p.update(false, false).unwrap();
540        assert_eq!(p.samples_seen(), 2);
541        p.reset();
542        assert_eq!(p.samples_seen(), 0);
543        assert_eq!(p.true_positive, 0);
544        assert_eq!(p.false_positive, 0);
545    }
546}