use crate::error::{DatarustError, Result};
use crate::linalg::cholesky;
use crate::matrix::Matrix;
use crate::stats;
use crate::traits::{Estimator, Predictor, Regressor};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum LinearSolver {
#[default]
Cholesky,
Svd,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct LinearRegression {
fit_intercept: bool,
solver: LinearSolver,
coef_: Vec<f64>,
intercept_: f64,
n_features_in_: usize,
fitted: bool,
}
impl Default for LinearRegression {
fn default() -> Self {
Self::new()
}
}
impl LinearRegression {
pub fn new() -> Self {
Self {
fit_intercept: true,
solver: LinearSolver::Cholesky,
coef_: Vec::new(),
intercept_: 0.0,
n_features_in_: 0,
fitted: false,
}
}
pub fn with_fit_intercept(mut self, b: bool) -> Self {
self.fit_intercept = b;
self
}
pub fn with_solver(mut self, s: LinearSolver) -> Self {
self.solver = s;
self
}
pub fn coef(&self) -> &[f64] {
&self.coef_
}
pub fn intercept(&self) -> f64 {
self.intercept_
}
pub fn n_features_in(&self) -> usize {
self.n_features_in_
}
pub fn score(&self, x: &Matrix, y: &[f64]) -> Result<f64> {
let pred = Predictor::predict(self, x)?;
crate::metrics::regression::r2_score(y, &pred)
}
fn solve_normal(
a_flat: Vec<f64>,
c: Vec<f64>,
p: usize,
solver: LinearSolver,
) -> Result<Vec<f64>> {
match solver {
LinearSolver::Cholesky => cholesky::solve_spd_system(&a_flat, p, &c),
LinearSolver::Svd => solve_via_eig_pinv(&a_flat, &c, p),
}
}
}
pub(crate) fn solve_via_eig_pinv(a: &[f64], b: &[f64], p: usize) -> Result<Vec<f64>> {
let mut a_buf = a.to_vec();
let (vals, vecs) = crate::decomposition::jacobi::eigh_flat(&mut a_buf, p)
.ok_or_else(|| DatarustError::Singular("eigendecomposition failed".into()))?;
let max_val = vals.iter().cloned().fold(0.0_f64, f64::max);
let tol = (max_val.max(1.0) * 1e-10).max(f64::MIN_POSITIVE);
let mut x = vec![0.0_f64; p];
for k in 0..p {
let lambda = vals[k];
if lambda.abs() <= tol {
continue;
}
let vk = &vecs[k * p..(k + 1) * p];
let mut dot = 0.0;
for i in 0..p {
dot += vk[i] * b[i];
}
let scale = dot / lambda;
for i in 0..p {
x[i] += scale * vk[i];
}
}
Ok(x)
}
impl Estimator for LinearRegression {}
impl Predictor for LinearRegression {
fn fit(&mut self, x: &Matrix, y: &[f64]) -> Result<()> {
let n = x.nrows();
let p = x.ncols();
if n == 0 {
return Err(DatarustError::EmptyInput("X has no rows".into()));
}
if p == 0 {
return Err(DatarustError::EmptyInput("X has no columns".into()));
}
if y.len() != n {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} targets", n),
actual: format!("{} targets", y.len()),
});
}
x.validate_finite()?;
super::validate_finite_targets(y)?;
let x_slice = x.as_slice();
let (design, y_work, x_mean, y_mean) = if self.fit_intercept {
let x_mean = stats::column_mean_flat(x_slice, n, p);
let y_mean = y.iter().sum::<f64>() / n as f64;
let mut xc = vec![0.0; n * p];
for i in 0..n {
for j in 0..p {
xc[i * p + j] = x_slice[i * p + j] - x_mean[j];
}
}
let yc: Vec<f64> = y.iter().map(|&v| v - y_mean).collect();
(xc, yc, x_mean, y_mean)
} else {
(x_slice.to_vec(), y.to_vec(), Vec::new(), 0.0)
};
let design_mat = Matrix::from_flat(n, p, design)?;
let xt = design_mat.transpose();
let xtx = xt.matmul(&design_mat)?; let y_col = Matrix::from_flat(n, 1, y_work)?;
let xty_mat = xt.matmul(&y_col)?; let xty = xty_mat.as_slice().to_vec();
let beta = Self::solve_normal(xtx.as_slice().to_vec(), xty, p, self.solver)?;
let intercept = if self.fit_intercept {
y_mean
- x_mean
.iter()
.zip(beta.iter())
.map(|(m, &bj)| m * bj)
.sum::<f64>()
} else {
0.0
};
self.coef_ = beta;
self.intercept_ = intercept;
self.n_features_in_ = p;
self.fitted = true;
Ok(())
}
fn predict(&self, x: &Matrix) -> Result<Vec<f64>> {
if !self.fitted {
return Err(DatarustError::NotFitted("LinearRegression".into()));
}
if self.coef_.len() != self.n_features_in_
|| !self.intercept_.is_finite()
|| self.coef_.iter().any(|v| !v.is_finite())
{
return Err(DatarustError::InvalidInput(
"LinearRegression has inconsistent fitted state".into(),
));
}
if x.ncols() != self.n_features_in_ {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} features", self.n_features_in_),
actual: format!("{} features", x.ncols()),
});
}
x.validate_finite()?;
let p = self.n_features_in_;
let beta = &self.coef_;
let intercept = self.intercept_;
let n = x.nrows();
let src = x.as_slice();
let mut out = vec![intercept; n];
for i in 0..n {
let row = &src[i * p..(i + 1) * p];
let mut s = intercept;
for j in 0..p {
s += beta[j] * row[j];
}
out[i] = s;
}
Ok(out)
}
fn is_fitted(&self) -> bool {
self.fitted
}
}
impl Regressor for LinearRegression {
fn name(&self) -> &'static str {
"LinearRegression"
}
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: &[f64], b: &[f64], tol: f64) -> bool {
a.len() == b.len() && a.iter().zip(b.iter()).all(|(x, y)| (x - y).abs() <= tol)
}
#[test]
fn fit_perfect_line_with_intercept() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0]]).unwrap();
let y = vec![3.0, 5.0, 7.0, 9.0];
let mut m = LinearRegression::new();
m.fit(&x, &y).unwrap();
assert!((m.coef()[0] - 2.0).abs() < 1e-9, "coef={}", m.coef()[0]);
assert!(
(m.intercept() - 1.0).abs() < 1e-9,
"intercept={}",
m.intercept()
);
let pred = m.predict(&x).unwrap();
assert!(approx(&pred, &y, 1e-9));
}
#[test]
fn fit_multivariate_known_coef() {
let rows: Vec<Vec<f64>> = (0..50)
.map(|i| {
let i = i as f64;
vec![i.sin(), i.cos(), (i + 1.0).ln()]
})
.collect();
let x = Matrix::new(rows.clone()).unwrap();
let y: Vec<f64> = rows
.iter()
.map(|r| 2.0 * r[0] - 3.5 * r[1] + 5.0 * r[2] + 7.0)
.collect();
let mut m = LinearRegression::new();
m.fit(&x, &y).unwrap();
assert!((m.coef()[0] - 2.0).abs() < 1e-6, "coef0={}", m.coef()[0]);
assert!((m.coef()[1] - (-3.5)).abs() < 1e-6, "coef1={}", m.coef()[1]);
assert!((m.coef()[2] - 5.0).abs() < 1e-6, "coef2={}", m.coef()[2]);
assert!(
(m.intercept() - 7.0).abs() < 1e-6,
"intercept={}",
m.intercept()
);
}
#[test]
fn fit_intercept_false() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![3.0]]).unwrap();
let y = vec![3.0, 6.0, 9.0];
let mut m = LinearRegression::new().with_fit_intercept(false);
m.fit(&x, &y).unwrap();
assert!((m.coef()[0] - 3.0).abs() < 1e-9);
assert!(m.intercept().abs() < 1e-12);
}
#[test]
fn cholesky_and_svd_agree_full_rank() {
let rows: Vec<Vec<f64>> = (0..30)
.map(|i| {
let i = i as f64;
vec![i.sin(), (i + 7.0).ln(), (i * 0.3).exp()]
})
.collect();
let x = Matrix::new(rows.clone()).unwrap();
let y: Vec<f64> = rows
.iter()
.map(|r| 1.5 * r[0] - 2.0 * r[1] + 0.3 * r[2] + 4.0)
.collect();
let mut m_chol = LinearRegression::new();
m_chol.fit(&x, &y).unwrap();
let mut m_svd = LinearRegression::new().with_solver(LinearSolver::Svd);
m_svd.fit(&x, &y).unwrap();
for i in 0..3 {
assert!(
(m_chol.coef()[i] - m_svd.coef()[i]).abs() < 1e-6,
"solver disagreement at {i}: chol={} svd={}",
m_chol.coef()[i],
m_svd.coef()[i]
);
}
assert!((m_chol.intercept() - m_svd.intercept()).abs() < 1e-6);
}
#[test]
fn svd_handles_rank_deficiency() {
let x = Matrix::new(vec![vec![1.0, 1.0], vec![2.0, 2.0], vec![3.0, 3.0]]).unwrap();
let y = vec![2.0, 4.0, 6.0];
let mut m_chol = LinearRegression::new();
let chol_res = m_chol.fit(&x, &y);
assert!(matches!(chol_res, Err(DatarustError::Singular(_))));
let mut m_svd = LinearRegression::new().with_solver(LinearSolver::Svd);
m_svd.fit(&x, &y).unwrap();
let pred = m_svd.predict(&x).unwrap();
assert!(approx(&pred, &y, 1e-6));
}
#[test]
fn predict_before_fit_errors() {
let m = LinearRegression::new();
let x = Matrix::new(vec![vec![1.0]]).unwrap();
let err = m.predict(&x).unwrap_err();
assert!(matches!(err, DatarustError::NotFitted(_)));
}
#[test]
fn predict_shape_mismatch() {
let x = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
let mut m = LinearRegression::new();
m.fit(&x, &[3.0, 5.0]).unwrap();
let bad = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
let err = m.predict(&bad).unwrap_err();
assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
#[test]
fn fit_shape_mismatch_y() {
let x = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
let mut m = LinearRegression::new();
let err = m.fit(&x, &[1.0]).unwrap_err(); assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
#[test]
fn fit_predict_convenience() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![3.0]]).unwrap();
let y = vec![2.0, 4.0, 6.0];
let mut m = LinearRegression::new();
let pred = m.fit_predict(&x, &y).unwrap();
assert!(approx(&pred, &y, 1e-9));
}
#[test]
fn n_features_in() {
let rows: Vec<Vec<f64>> = (0..10)
.map(|i| {
let i = i as f64;
vec![i.sin(), i.cos(), (i + 1.0).ln()]
})
.collect();
let x = Matrix::new(rows).unwrap();
let y: Vec<f64> = (0..10).map(|i| i as f64).collect();
let mut m = LinearRegression::new();
m.fit(&x, &y).unwrap();
assert_eq!(m.n_features_in(), 3);
}
#[test]
fn fit_rejects_y_length_mismatch() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![3.0]]).unwrap();
let mut m = LinearRegression::new();
let err = m.fit(&x, &[1.0, 2.0]).unwrap_err(); assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
#[test]
fn predict_returns_n_rows() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0]]).unwrap();
let y = vec![2.0, 4.0, 6.0, 8.0];
let mut m = LinearRegression::new().with_fit_intercept(false);
m.fit(&x, &y).unwrap();
let pred = m.predict(&x).unwrap();
assert_eq!(pred.len(), x.nrows());
}
#[test]
fn constant_target() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![5.0]]).unwrap();
let y = vec![7.0, 7.0, 7.0];
let mut m = LinearRegression::new();
m.fit(&x, &y).unwrap();
assert!(
(m.intercept() - 7.0).abs() < 1e-9,
"intercept={}",
m.intercept()
);
assert!(m.coef()[0].abs() < 1e-9);
}
#[test]
fn fit_new_data_predicts_correctly() {
let x = Matrix::new(vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0]]).unwrap();
let y = vec![2.0, 4.0, 6.0, 8.0]; let mut m = LinearRegression::new().with_fit_intercept(false);
m.fit(&x, &y).unwrap();
let new = Matrix::new(vec![vec![10.0]]).unwrap();
let pred = m.predict(&new).unwrap();
assert!((pred[0] - 20.0).abs() < 1e-9);
}
}