qualia_core_db/solvers/learning/regression/
ridge.rs1use 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#[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
40pub 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 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 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 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 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}