1use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
4use crate::loss::log_loss::BinaryLogLoss;
5use crate::traits::Metric;
6
7#[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#[derive(Debug, Clone, Default)]
55#[cfg_attr(feature = "serde", derive(serde::Serialize))]
56pub struct Precision {
57 true_positive: u64,
58 false_positive: u64,
59 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 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#[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#[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#[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(); m.update(false, true).unwrap(); m.update(true, false).unwrap(); 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(); m.update(false, true).unwrap(); m.update(true, false).unwrap(); 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(); m.update(false, true).unwrap(); m.update(true, false).unwrap(); 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 #[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 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 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 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 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}