Skip to main content

qualia_core_db/solvers/learning/regression/
lasso.rs

1//! Lasso regression (ISL ch 6.2.2) — L1-penalized least squares by cyclic
2//! coordinate descent with soft-thresholding.
3//!
4//! Minimise `½‖y − Xβ‖² + λ‖β‖₁` (intercept not penalized; handled by centering).
5//! Unlike ridge, the L1 penalty drives some coefficients **exactly to zero**
6//! (variable selection). Coordinate descent updates one `βⱼ` at a time via the
7//! soft-threshold operator, keeping a running residual for O(np) per sweep.
8//! Scalar fit-loop → CPU (the per-coordinate dot is `Reduction`-class).
9
10use crate::solvers::learning::LearningError;
11use crate::solvers::statistics::descriptive::mean;
12
13/// A fitted lasso model.
14#[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    /// Number of non-zero coefficients (the selected variables).
39    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
55/// Fit lasso with penalty `lambda ≥ 0` by coordinate descent. `lambda = 0` reduces
56/// to OLS (up to the iteration tolerance). Fails closed on shape mismatch / `n < 2`.
57pub 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    // Centre predictors and response (so the intercept drops out of the penalty).
74    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    // Column squared norms zⱼ = Σ Xcᵢⱼ².
90    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    // Running residual r = yc − Xc·β  (β starts at 0 → r = yc).
101    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; // constant predictor contributes nothing
111            }
112            // ρⱼ = Xcⱼ·r + zⱼ·βⱼ  (add back coordinate j's own contribution).
113            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                // r ← r − Xcⱼ·Δβ.
121                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        // Intercept falls back to the response mean.
172        assert!((m.intercept - 6.0).abs() < 1e-9);
173    }
174
175    #[test]
176    fn selects_the_relevant_predictor() {
177        // x1 drives y; x2 is pure noise → lasso should keep x1, zero x2.
178        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]; // ≈ 2·x1
180        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}