use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use crate::error::{RegressionError, Result};
#[derive(Debug, Clone)]
pub struct LassoFit {
x: Array2<f64>,
y: Array1<f64>,
lambda: f64,
coefficients: Array1<f64>,
fitted: Array1<f64>,
residuals: Array1<f64>,
rss: f64,
intercept_col: Option<usize>,
n_nonzero: usize,
iterations: usize,
n: usize,
p: usize,
}
impl LassoFit {
pub fn new(x: Array2<f64>, y: Array1<f64>, lambda: f64) -> Result<Self> {
Self::with_options(x, y, lambda, 1e-7, 10_000)
}
pub fn with_options(
x: Array2<f64>,
y: Array1<f64>,
lambda: f64,
tol: f64,
max_iter: usize,
) -> Result<Self> {
if x.nrows() == 0 || x.ncols() == 0 {
return Err(RegressionError::EmptyInput { what: "X" });
}
if y.len() != x.nrows() {
return Err(RegressionError::ShapeMismatch {
what: "y length vs X rows",
expected: x.nrows(),
got: y.len(),
});
}
if lambda < 0.0 || lambda.is_nan() {
return Err(RegressionError::InvalidParameter {
msg: format!("lasso lambda must be >= 0, got {lambda}"),
});
}
let n = x.nrows();
let p = x.ncols();
let intercept_col = detect_constant_column(&x);
let has_intercept = intercept_col.is_some();
let pred: Vec<usize> = (0..p).filter(|&j| Some(j) != intercept_col).collect();
let q = pred.len();
let nf = n as f64;
let y_mean = if has_intercept { y.sum() / nf } else { 0.0 };
let mut means = vec![0.0; q];
let mut sds = vec![1.0; q];
for (k, &j) in pred.iter().enumerate() {
let col = x.column(j);
let m = if has_intercept { col.sum() / nf } else { 0.0 };
means[k] = m;
let sd = (col.iter().map(|v| (v - m).powi(2)).sum::<f64>() / nf).sqrt();
sds[k] = if sd > 0.0 { sd } else { 1.0 };
}
let mut z = vec![0.0f64; n * q];
for i in 0..n {
for (k, &j) in pred.iter().enumerate() {
z[i * q + k] = (x[(i, j)] - means[k]) / sds[k];
}
}
let yc: Vec<f64> = (0..n).map(|i| y[i] - y_mean).collect();
let mut beta = vec![0.0f64; q];
let mut r = yc.clone();
let mut iterations = 0usize;
let mut converged = false;
while iterations < max_iter {
iterations += 1;
let mut max_delta = 0.0f64;
for k in 0..q {
let mut zr = 0.0;
for i in 0..n {
zr += z[i * q + k] * r[i];
}
let rho = zr / nf + beta[k];
let new = soft_threshold(rho, lambda);
let delta = new - beta[k];
if delta != 0.0 {
for i in 0..n {
r[i] -= z[i * q + k] * delta;
}
beta[k] = new;
max_delta = max_delta.max(delta.abs());
}
}
if max_delta < tol {
converged = true;
break;
}
}
if !converged {
return Err(RegressionError::NotConverged {
iterations,
msg: "lasso coordinate descent did not reach tolerance".into(),
});
}
let mut coefficients = Array1::<f64>::zeros(p);
let mut slopes_orig = vec![0.0; q];
for (k, &j) in pred.iter().enumerate() {
let b = beta[k] / sds[k];
slopes_orig[k] = b;
coefficients[j] = b;
}
if let Some(c) = intercept_col {
coefficients[c] = y_mean - (0..q).map(|k| means[k] * slopes_orig[k]).sum::<f64>();
}
let fitted = x.dot(&coefficients);
let residuals = &y - &fitted;
let rss: f64 = residuals.iter().map(|e| e * e).sum();
let n_nonzero = slopes_orig.iter().filter(|b| b.abs() > 1e-12).count();
Ok(Self {
x,
y,
lambda,
coefficients,
fitted,
residuals,
rss,
intercept_col,
n_nonzero,
iterations,
n,
p,
})
}
pub fn lambda(&self) -> f64 {
self.lambda
}
pub fn n_observations(&self) -> usize {
self.n
}
pub fn n_parameters(&self) -> usize {
self.p
}
pub fn has_intercept(&self) -> bool {
self.intercept_col.is_some()
}
pub fn iterations(&self) -> usize {
self.iterations
}
pub fn design_matrix(&self) -> ArrayView2<'_, f64> {
self.x.view()
}
pub fn coefficients(&self) -> ArrayView1<'_, f64> {
self.coefficients.view()
}
pub fn fitted_values(&self) -> ArrayView1<'_, f64> {
self.fitted.view()
}
pub fn residuals(&self) -> ArrayView1<'_, f64> {
self.residuals.view()
}
pub fn residual_sum_of_squares(&self) -> f64 {
self.rss
}
pub fn response(&self) -> ArrayView1<'_, f64> {
self.y.view()
}
pub fn active_set(&self) -> Vec<usize> {
(0..self.p)
.filter(|&j| Some(j) != self.intercept_col && self.coefficients[j].abs() > 1e-12)
.collect()
}
pub fn n_nonzero(&self) -> usize {
self.n_nonzero
}
pub fn effective_df(&self) -> f64 {
self.n_nonzero as f64 + if self.has_intercept() { 1.0 } else { 0.0 }
}
pub fn log_likelihood(&self) -> f64 {
let n = self.n as f64;
-0.5 * n * ((2.0 * std::f64::consts::PI).ln() + 1.0 + (self.rss / n).ln())
}
pub fn aic(&self) -> f64 {
-2.0 * self.log_likelihood() + 2.0 * self.effective_df()
}
pub fn bic(&self) -> f64 {
-2.0 * self.log_likelihood() + (self.n as f64).ln() * self.effective_df()
}
}
fn soft_threshold(a: f64, lambda: f64) -> f64 {
if a > lambda {
a - lambda
} else if a < -lambda {
a + lambda
} else {
0.0
}
}
fn detect_constant_column(x: &Array2<f64>) -> Option<usize> {
for (j, col) in x.columns().into_iter().enumerate() {
let first = col[0];
let scale = first.abs().max(1.0);
if col.iter().all(|&v| (v - first).abs() <= 1e-12 * scale) {
return Some(j);
}
}
None
}