Skip to main content

qualia_core_db/solvers/learning/metrics/
regression.rs

1//! Regression metrics — error and fit measures over caller-owned prediction slices.
2//! Reuses `statistics::descriptive` for the mean (no re-implementation).
3
4use crate::solvers::statistics::descriptive::mean;
5
6/// Mean squared error `Σ(yᵢ−ŷᵢ)²/n`. `None` if lengths differ or are empty.
7pub 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
20/// Root mean squared error.
21pub fn rmse(y_true: &[f64], y_pred: &[f64]) -> Option<f64> {
22    mse(y_true, y_pred).map(f64::sqrt)
23}
24
25/// Mean absolute error.
26pub 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
35/// Coefficient of determination `R² = 1 − SSE/SST`. Can be negative for a model
36/// worse than the mean. `None` if lengths differ or `SST = 0` (constant target).
37pub 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        // errors: 0.5,-0.5,0,-1 → squared 0.25,0.25,0,1 → mse 0.375
73        assert!((mse(&y, &p).unwrap() - 0.375).abs() < 1e-12);
74        assert!((mae(&y, &p).unwrap() - 0.5).abs() < 1e-12);
75        // sklearn r2 for this pair ≈ 0.9486.
76        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]; // predicting the mean
83        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); // constant target
91    }
92}