qualia_core_db/solvers/learning/metrics/
regression.rs1use crate::solvers::statistics::descriptive::mean;
5
6pub fn mse(y_true: &[f64], y_pred: &[f64]) -> Option<f64> {
8 let n = y_true.len();
9 if n == 0 || n != y_pred.len() {
10 return None;
11 }
12 let s: f64 = y_true
13 .iter()
14 .zip(y_pred)
15 .map(|(y, p)| (y - p) * (y - p))
16 .sum();
17 Some(s / n as f64)
18}
19
20pub fn rmse(y_true: &[f64], y_pred: &[f64]) -> Option<f64> {
22 mse(y_true, y_pred).map(f64::sqrt)
23}
24
25pub fn mae(y_true: &[f64], y_pred: &[f64]) -> Option<f64> {
27 let n = y_true.len();
28 if n == 0 || n != y_pred.len() {
29 return None;
30 }
31 let s: f64 = y_true.iter().zip(y_pred).map(|(y, p)| (y - p).abs()).sum();
32 Some(s / n as f64)
33}
34
35pub fn r2_score(y_true: &[f64], y_pred: &[f64]) -> Option<f64> {
38 let n = y_true.len();
39 if n == 0 || n != y_pred.len() {
40 return None;
41 }
42 let ybar = mean(y_true)?;
43 let mut sse = 0.0;
44 let mut sst = 0.0;
45 for (y, p) in y_true.iter().zip(y_pred) {
46 sse += (y - p) * (y - p);
47 sst += (y - ybar) * (y - ybar);
48 }
49 if sst == 0.0 {
50 return None;
51 }
52 Some(1.0 - sse / sst)
53}
54
55#[cfg(test)]
56mod tests {
57 use super::*;
58
59 #[test]
60 fn perfect_prediction() {
61 let y = [1.0, 2.0, 3.0, 4.0];
62 assert_eq!(mse(&y, &y), Some(0.0));
63 assert_eq!(rmse(&y, &y), Some(0.0));
64 assert_eq!(mae(&y, &y), Some(0.0));
65 assert!((r2_score(&y, &y).unwrap() - 1.0).abs() < 1e-12);
66 }
67
68 #[test]
69 fn known_values() {
70 let y = [3.0, -0.5, 2.0, 7.0];
71 let p = [2.5, 0.0, 2.0, 8.0];
72 assert!((mse(&y, &p).unwrap() - 0.375).abs() < 1e-12);
74 assert!((mae(&y, &p).unwrap() - 0.5).abs() < 1e-12);
75 assert!((r2_score(&y, &p).unwrap() - 0.948_608).abs() < 1e-4);
77 }
78
79 #[test]
80 fn mean_predictor_is_zero_r2() {
81 let y = [1.0, 2.0, 3.0, 4.0, 5.0];
82 let p = [3.0; 5]; assert!(r2_score(&y, &p).unwrap().abs() < 1e-12);
84 }
85
86 #[test]
87 fn guards() {
88 assert_eq!(mse(&[], &[]), None);
89 assert_eq!(mse(&[1.0], &[1.0, 2.0]), None);
90 assert_eq!(r2_score(&[5.0, 5.0], &[1.0, 2.0]), None); }
92}