1use crate::error::{
4 RillError, checked_finite_add, checked_increment, ensure_finite, ensure_finite_target,
5};
6use crate::traits::Metric;
7
8#[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 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#[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 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#[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 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#[derive(Debug, Clone, Default)]
157#[cfg_attr(feature = "serde", derive(serde::Serialize))]
158pub struct R2 {
159 ss_res: f64,
160 mean_truth: f64,
161 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 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 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 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 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 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 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(); m.update(5.0, 4.0).unwrap(); 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(); m.update(5.0, 4.0).unwrap(); 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 m.update(1.0, 2.0).unwrap();
364 m.update(3.0, 2.0).unwrap();
365 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 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 m.update(y, y).unwrap();
397 }
398 assert!((m.value().unwrap() - 1.0).abs() < 1e-9);
399
400 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 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 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 #[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 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 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 assert!(m.value().is_none());
517 }
518}