Skip to main content

qualia_core_db/solvers/learning/regression/
ridge.rs

1//! Ridge regression (ISL ch 6.2.1, PRML ch 3) — L2-penalized least squares.
2//!
3//! Minimise `‖y − Xβ‖² + λ‖β‖²` (the intercept is **not** penalized). Centering y
4//! and the predictors removes the intercept from the penalized solve, leaving
5//! `(XcᵀXc + λI)β = Xcᵀyc`, solved with `linear_algebra::cholesky` (the penalty
6//! makes the system positive-definite even for collinear predictors — ridge's whole
7//! point). Kernel-class `DenseLinear`, dispatch-ready.
8
9use crate::solvers::learning::LearningError;
10use crate::solvers::linear_algebra::cholesky::{cholesky_factor, cholesky_solve};
11use crate::solvers::linear_algebra::gemm::{gemm, matvec, Transpose};
12use crate::solvers::statistics::descriptive::mean;
13
14/// A fitted ridge model: slope coefficients (predictor-aligned) plus an
15/// un-penalized intercept.
16#[derive(Debug, Clone)]
17pub struct RidgeModel {
18    pub coefficients: Vec<f64>,
19    pub intercept: f64,
20    pub lambda: f64,
21}
22
23impl RidgeModel {
24    pub fn predict_row(&self, x_row: &[f64]) -> f64 {
25        self.intercept
26            + self
27                .coefficients
28                .iter()
29                .zip(x_row)
30                .map(|(b, x)| b * x)
31                .sum::<f64>()
32    }
33    pub fn predict(&self, x: &[f64], n: usize, p: usize) -> Vec<f64> {
34        (0..n)
35            .map(|i| self.predict_row(&x[i * p..(i + 1) * p]))
36            .collect()
37    }
38}
39
40/// Fit ridge regression with penalty `lambda ≥ 0`. `lambda = 0` reproduces OLS.
41/// Fails closed on shape mismatch / `n < 2`.
42pub fn fit(
43    x: &[f64],
44    y: &[f64],
45    n: usize,
46    p: usize,
47    lambda: f64,
48) -> Result<RidgeModel, LearningError> {
49    if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
50        return Err(LearningError::InvalidDimension);
51    }
52    if n < 2 || lambda < 0.0 {
53        return Err(LearningError::InsufficientData);
54    }
55
56    // Column means and centred predictors; centred response.
57    let mut xbar = vec![0.0; p];
58    let mut colbuf = vec![0.0; n];
59    for j in 0..p {
60        for i in 0..n {
61            colbuf[i] = x[i * p + j];
62        }
63        xbar[j] = mean(&colbuf).ok_or(LearningError::InsufficientData)?;
64    }
65    let ybar = mean(y).ok_or(LearningError::InsufficientData)?;
66    let mut xc = vec![0.0; n * p];
67    let mut yc = vec![0.0; n];
68    for i in 0..n {
69        for j in 0..p {
70            xc[i * p + j] = x[i * p + j] - xbar[j];
71        }
72        yc[i] = y[i] - ybar;
73    }
74
75    // A = XcᵀXc + λI, b = Xcᵀyc.
76    let mut a = vec![0.0; p * p];
77    gemm(
78        Transpose::Yes,
79        Transpose::No,
80        p,
81        p,
82        n,
83        1.0,
84        &xc,
85        &xc,
86        0.0,
87        &mut a,
88    )?;
89    for j in 0..p {
90        a[j * p + j] += lambda;
91    }
92    let mut b = vec![0.0; p];
93    matvec(Transpose::Yes, p, n, &xc, &yc, &mut b)?;
94
95    let mut l = vec![0.0; p * p];
96    cholesky_factor(p, &a, &mut l).map_err(|_| LearningError::Singular)?;
97    let mut coefficients = vec![0.0; p];
98    cholesky_solve(p, &l, &b, &mut coefficients)?;
99
100    // Recover the un-penalized intercept.
101    let intercept = ybar
102        - coefficients
103            .iter()
104            .zip(xbar.iter())
105            .map(|(b, m)| b * m)
106            .sum::<f64>();
107
108    Ok(RidgeModel {
109        coefficients,
110        intercept,
111        lambda,
112    })
113}
114
115#[cfg(test)]
116mod tests {
117    use super::*;
118    use crate::solvers::learning::regression::linear;
119
120    #[test]
121    fn lambda_zero_matches_ols() {
122        let x = [1.0, 2.0, 2.0, 1.0, 3.0, 0.0, 4.0, 5.0, 5.0, 4.0];
123        let y = [3.0, 5.0, 4.0, 9.0, 13.0];
124        let ridge = fit(&x, &y, 5, 2, 0.0).unwrap();
125        let ols = linear::fit(&x, &y, 5, 2, true).unwrap();
126        assert!((ridge.intercept - ols.coefficients[0]).abs() < 1e-7);
127        assert!((ridge.coefficients[0] - ols.coefficients[1]).abs() < 1e-7);
128        assert!((ridge.coefficients[1] - ols.coefficients[2]).abs() < 1e-7);
129    }
130
131    #[test]
132    fn penalty_shrinks_coefficients() {
133        let x = [1.0, 2.0, 2.0, 1.0, 3.0, 0.0, 4.0, 5.0, 5.0, 4.0];
134        let y = [3.0, 5.0, 4.0, 9.0, 13.0];
135        let small = fit(&x, &y, 5, 2, 0.0).unwrap();
136        let big = fit(&x, &y, 5, 2, 50.0).unwrap();
137        let norm = |m: &RidgeModel| m.coefficients.iter().map(|c| c * c).sum::<f64>().sqrt();
138        assert!(
139            norm(&big) < norm(&small),
140            "ridge must shrink: {} !< {}",
141            norm(&big),
142            norm(&small)
143        );
144    }
145
146    #[test]
147    fn stable_on_collinear_where_ols_fails() {
148        // x2 = 2·x1 (collinear) — OLS is singular, ridge regularizes and solves.
149        let x = [1.0, 2.0, 2.0, 4.0, 3.0, 6.0, 4.0, 8.0, 5.0, 10.0];
150        let y = [1.0, 2.0, 3.0, 4.0, 5.0];
151        assert_eq!(
152            linear::fit(&x, &y, 5, 2, true).unwrap_err(),
153            LearningError::Singular
154        );
155        let ridge = fit(&x, &y, 5, 2, 1.0).unwrap();
156        assert!(ridge.coefficients.iter().all(|c| c.is_finite()));
157    }
158}