use super::RegressionError;
const PIVOT_EPS: f64 = 1e-12;
pub(super) fn solve(mut a: Vec<Vec<f64>>, mut b: Vec<f64>) -> Result<Vec<f64>, RegressionError> {
let n = a.len();
for col in 0..n {
let mut pivot_row = col;
let mut best = pivot_magnitude(&a, col, col);
for r in (col + 1)..n {
let mag = pivot_magnitude(&a, r, col);
if mag > best {
best = mag;
pivot_row = r;
}
}
if best < PIVOT_EPS {
return Err(RegressionError::Singular);
}
a.swap(col, pivot_row);
b.swap(col, pivot_row);
let pivot = a
.get(col)
.and_then(|row| row.get(col))
.copied()
.unwrap_or(0.0);
let pivot_a = a.get(col).cloned().unwrap_or_default();
let pivot_b = at(&b, col);
for r in (col + 1)..n {
let factor = a
.get(r)
.and_then(|row| row.get(col))
.copied()
.unwrap_or(0.0)
/ pivot;
if let Some(target_row) = a.get_mut(r) {
for k in col..n {
let pivot_val = pivot_a.get(k).copied().unwrap_or(0.0);
if let Some(target) = target_row.get_mut(k) {
*target -= factor * pivot_val;
}
}
}
if let Some(bv) = b.get_mut(r) {
*bv -= factor * pivot_b;
}
}
}
back_substitute(&a, &b)
}
fn back_substitute(a: &[Vec<f64>], b: &[f64]) -> Result<Vec<f64>, RegressionError> {
let n = a.len();
let mut beta = vec![0.0_f64; n];
for row in (0..n).rev() {
let mut acc = at(b, row);
for k in (row + 1)..n {
let a_rk = a.get(row).and_then(|r| r.get(k)).copied().unwrap_or(0.0);
let beta_k = beta.get(k).copied().unwrap_or(0.0);
acc -= a_rk * beta_k;
}
let pivot = a.get(row).and_then(|r| r.get(row)).copied().unwrap_or(0.0);
if pivot.abs() < PIVOT_EPS {
return Err(RegressionError::Singular);
}
if let Some(beta_row) = beta.get_mut(row) {
*beta_row = acc / pivot;
}
}
Ok(beta)
}
fn pivot_magnitude(a: &[Vec<f64>], r: usize, c: usize) -> f64 {
a.get(r)
.and_then(|row| row.get(c))
.copied()
.unwrap_or(0.0)
.abs()
}
fn at(v: &[f64], i: usize) -> f64 {
v.get(i).copied().unwrap_or(0.0)
}
pub(super) fn count_to_f64(n: usize) -> f64 {
let wide = u64::try_from(n).unwrap_or(u64::MAX);
let hi = u32::try_from(wide >> 32).unwrap_or(0);
let lo = u32::try_from(wide & 0xFFFF_FFFF).unwrap_or(0);
f64::from(hi).mul_add(4_294_967_296.0, f64::from(lo))
}
pub(super) fn mean(values: &[f64]) -> f64 {
let n = count_to_f64(values.len());
if n > 0.0 {
values.iter().sum::<f64>() / n
} else {
0.0
}
}
pub(super) fn column_means(x: &[Vec<f64>], n_cols: usize) -> Vec<f64> {
let n = count_to_f64(x.len());
let mut sums = vec![0.0_f64; n_cols];
for row in x {
for (s, &v) in sums.iter_mut().zip(row) {
*s += v;
}
}
if n > 0.0 {
for s in &mut sums {
*s /= n;
}
}
sums
}
pub(super) fn centered(row: &[f64], col_means: &[f64], j: usize) -> f64 {
let v = row.get(j).copied().unwrap_or(0.0);
let m = col_means.get(j).copied().unwrap_or(0.0);
v - m
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn solves_two_by_two_system() -> Result<(), RegressionError> {
let a = vec![vec![2.0, 1.0], vec![1.0, 3.0]];
let b = vec![5.0, 10.0];
let beta = solve(a, b)?;
let x = beta.first().copied().unwrap_or(f64::NAN);
let y = beta.get(1).copied().unwrap_or(f64::NAN);
assert!((x - 1.0).abs() < 1e-12, "x was {x}");
assert!((y - 3.0).abs() < 1e-12, "y was {y}");
Ok(())
}
#[test]
fn singular_matrix_is_reported() {
let a = vec![vec![1.0, 2.0], vec![2.0, 4.0]];
let b = vec![1.0, 2.0];
assert_eq!(solve(a, b), Err(RegressionError::Singular));
}
#[test]
fn count_widens_exactly() {
assert!((count_to_f64(4096) - 4096.0).abs() < 1e-12);
}
}