use nalgebra::{DMatrix, DVector};
use crate::error::{RegressionError, Result};
pub(crate) struct QrFit {
pub coef: DVector<f64>,
pub fitted: DVector<f64>,
pub leverage: DVector<f64>,
pub xtx_inv: DMatrix<f64>,
pub singular_values: Vec<f64>,
}
pub(crate) fn ols_via_qr(x: &DMatrix<f64>, y: &DVector<f64>) -> Result<QrFit> {
let n = x.nrows();
let p = x.ncols();
let qr = x.clone().qr();
let q = qr.q();
let r = qr.r();
let r_diag_min = (0..p)
.map(|i| r[(i, i)].abs())
.fold(f64::INFINITY, f64::min);
if r_diag_min <= f64::EPSILON * 100.0 * r[(0, 0)].abs().max(1.0) {
return Err(RegressionError::RankDeficient);
}
let r_inv = r.try_inverse().ok_or(RegressionError::RankDeficient)?;
let qty = q.transpose() * y;
let coef = &r_inv * qty;
let xtx_inv = &r_inv * r_inv.transpose();
let mut leverage = DVector::<f64>::zeros(n);
for i in 0..n {
let mut s = 0.0;
for j in 0..q.ncols() {
s += q[(i, j)] * q[(i, j)];
}
leverage[i] = s;
}
let fitted = x * &coef;
let singular_values = x
.clone()
.singular_values()
.iter()
.copied()
.collect::<Vec<_>>();
Ok(QrFit {
coef,
fitted,
leverage,
xtx_inv,
singular_values,
})
}
pub(crate) fn aux_r_squared(x: &DMatrix<f64>, y: &DVector<f64>) -> Option<f64> {
let n = y.len() as f64;
let mean = y.sum() / n;
let tss: f64 = y.iter().map(|v| (v - mean).powi(2)).sum();
if tss <= 0.0 {
return None;
}
let qr = x.clone().qr();
let r = qr.r();
if x.ncols() > x.nrows() {
return None;
}
let p = x.ncols();
let r_diag_min = (0..p)
.map(|i| r[(i, i)].abs())
.fold(f64::INFINITY, f64::min);
if r_diag_min <= f64::EPSILON * 100.0 * r[(0, 0)].abs().max(1.0) {
return Some(1.0);
}
let r_inv = match r.try_inverse() {
Some(inv) => inv,
None => return Some(1.0),
};
let coef = &r_inv * (qr.q().transpose() * y);
let fitted = x * &coef;
let rss: f64 = y
.iter()
.zip(fitted.iter())
.map(|(a, b)| (a - b).powi(2))
.sum();
Some((1.0 - rss / tss).clamp(0.0, 1.0))
}
pub(crate) fn dmatrix_from_rows(n: usize, p: usize, data: &[f64]) -> DMatrix<f64> {
DMatrix::from_row_iterator(n, p, data.iter().copied())
}
pub(crate) fn dvector_from_slice(data: &[f64]) -> DVector<f64> {
DVector::from_iterator(data.len(), data.iter().copied())
}