qualia_core_db/solvers/learning/regression/
lasso.rs1use crate::solvers::learning::LearningError;
11use crate::solvers::statistics::descriptive::mean;
12
13#[derive(Debug, Clone)]
15pub struct LassoModel {
16 pub coefficients: Vec<f64>,
17 pub intercept: f64,
18 pub lambda: f64,
19 pub n_iter: usize,
20 pub converged: bool,
21}
22
23impl LassoModel {
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 pub fn n_selected(&self) -> usize {
40 self.coefficients.iter().filter(|&&c| c != 0.0).count()
41 }
42}
43
44#[inline]
45fn soft_threshold(z: f64, gamma: f64) -> f64 {
46 if z > gamma {
47 z - gamma
48 } else if z < -gamma {
49 z + gamma
50 } else {
51 0.0
52 }
53}
54
55pub fn fit(
58 x: &[f64],
59 y: &[f64],
60 n: usize,
61 p: usize,
62 lambda: f64,
63 max_iter: usize,
64 tol: f64,
65) -> Result<LassoModel, LearningError> {
66 if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
67 return Err(LearningError::InvalidDimension);
68 }
69 if n < 2 || lambda < 0.0 {
70 return Err(LearningError::InsufficientData);
71 }
72
73 let mut xbar = vec![0.0; p];
75 let mut colbuf = vec![0.0; n];
76 for j in 0..p {
77 for i in 0..n {
78 colbuf[i] = x[i * p + j];
79 }
80 xbar[j] = mean(&colbuf).ok_or(LearningError::InsufficientData)?;
81 }
82 let ybar = mean(y).ok_or(LearningError::InsufficientData)?;
83 let mut xc = vec![0.0; n * p];
84 for i in 0..n {
85 for j in 0..p {
86 xc[i * p + j] = x[i * p + j] - xbar[j];
87 }
88 }
89 let mut znorm = vec![0.0; p];
91 for j in 0..p {
92 let mut s = 0.0;
93 for i in 0..n {
94 s += xc[i * p + j] * xc[i * p + j];
95 }
96 znorm[j] = s;
97 }
98
99 let mut beta = vec![0.0; p];
100 let mut r: Vec<f64> = y.iter().map(|&yi| yi - ybar).collect();
102
103 let mut converged = false;
104 let mut iters = 0;
105 for it in 1..=max_iter.max(1) {
106 iters = it;
107 let mut max_delta = 0.0_f64;
108 for j in 0..p {
109 if znorm[j] == 0.0 {
110 continue; }
112 let mut rho = znorm[j] * beta[j];
114 for i in 0..n {
115 rho += xc[i * p + j] * r[i];
116 }
117 let new = soft_threshold(rho, lambda) / znorm[j];
118 let delta = new - beta[j];
119 if delta != 0.0 {
120 for i in 0..n {
122 r[i] -= xc[i * p + j] * delta;
123 }
124 beta[j] = new;
125 max_delta = max_delta.max(delta.abs());
126 }
127 }
128 if max_delta < tol {
129 converged = true;
130 break;
131 }
132 }
133
134 let intercept = ybar
135 - beta
136 .iter()
137 .zip(xbar.iter())
138 .map(|(b, m)| b * m)
139 .sum::<f64>();
140 Ok(LassoModel {
141 coefficients: beta,
142 intercept,
143 lambda,
144 n_iter: iters,
145 converged,
146 })
147}
148
149#[cfg(test)]
150mod tests {
151 use super::*;
152 use crate::solvers::learning::regression::linear;
153
154 #[test]
155 fn lambda_zero_approaches_ols() {
156 let x = [1.0, 5.0, 2.0, 4.0, 3.0, 3.0, 4.0, 2.0, 5.0, 1.0];
157 let y = [3.0, 4.0, 6.0, 8.0, 11.0];
158 let lasso = fit(&x, &y, 5, 2, 0.0, 5000, 1e-10).unwrap();
159 let ols = linear::fit(&x, &y, 5, 2, true).unwrap();
160 assert!((lasso.coefficients[0] - ols.coefficients[1]).abs() < 1e-4);
161 assert!((lasso.coefficients[1] - ols.coefficients[2]).abs() < 1e-4);
162 assert!((lasso.intercept - ols.coefficients[0]).abs() < 1e-4);
163 }
164
165 #[test]
166 fn large_penalty_zeros_all_coefficients() {
167 let x = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
168 let y = [2.0, 4.0, 6.0, 8.0, 10.0];
169 let m = fit(&x, &y, 5, 2, 1e6, 1000, 1e-9).unwrap();
170 assert_eq!(m.n_selected(), 0);
171 assert!((m.intercept - 6.0).abs() < 1e-9);
173 }
174
175 #[test]
176 fn selects_the_relevant_predictor() {
177 let x = [1.0, 0.3, 2.0, -0.1, 3.0, 0.2, 4.0, -0.3, 5.0, 0.1, 6.0, 0.0];
179 let y = [2.0, 4.1, 5.9, 8.0, 10.1, 12.0]; let m = fit(&x, &y, 6, 2, 1.0, 5000, 1e-10).unwrap();
181 assert!(m.converged);
182 assert!(
183 m.coefficients[0].abs() > 0.5,
184 "x1 should be selected: {}",
185 m.coefficients[0]
186 );
187 assert_eq!(m.coefficients[1], 0.0, "x2 (noise) should be zeroed");
188 }
189}