use super::MathError;
const PIVOT_TOLERANCE: f64 = 1e-14;
#[allow(clippy::needless_range_loop)]
pub fn solve(a: &mut [Vec<f64>], b: &mut [f64]) -> Result<(), MathError> {
let n = a.len();
if n == 0 || b.len() != n {
return Err(MathError::DimensionMismatch);
}
for row in a.iter() {
if row.len() != n {
return Err(MathError::DimensionMismatch);
}
}
for k in 0..n {
let mut pivot_row = k;
let mut pivot_val = a[k][k].abs();
for i in (k + 1)..n {
let v = a[i][k].abs();
if v > pivot_val {
pivot_val = v;
pivot_row = i;
}
}
if pivot_val < PIVOT_TOLERANCE {
return Err(MathError::Singular);
}
if pivot_row != k {
a.swap(k, pivot_row);
b.swap(k, pivot_row);
}
let pivot = a[k][k];
for i in (k + 1)..n {
let factor = a[i][k] / pivot;
if factor == 0.0 {
continue;
}
a[i][k] = 0.0;
for j in (k + 1)..n {
a[i][j] -= factor * a[k][j];
}
b[i] -= factor * b[k];
}
}
for i in (0..n).rev() {
let mut sum = b[i];
for j in (i + 1)..n {
sum -= a[i][j] * b[j];
}
let pivot = a[i][i];
if pivot.abs() < PIVOT_TOLERANCE {
return Err(MathError::Singular);
}
b[i] = sum / pivot;
}
Ok(())
}
#[allow(clippy::needless_range_loop, clippy::many_single_char_names)]
pub fn solve_spd(a: &[Vec<f64>], b: &[f64]) -> Result<Vec<f64>, MathError> {
let n = a.len();
if n == 0 || b.len() != n {
return Err(MathError::DimensionMismatch);
}
for row in a {
if row.len() != n {
return Err(MathError::DimensionMismatch);
}
}
let mut l = vec![vec![0.0_f64; n]; n];
for i in 0..n {
for j in 0..=i {
let mut sum = a[i][j];
for k in 0..j {
sum -= l[i][k] * l[j][k];
}
if i == j {
if sum <= 0.0 || !sum.is_finite() {
return Err(MathError::NotSpd);
}
l[i][j] = sum.sqrt();
} else {
let pivot = l[j][j];
if pivot.abs() < PIVOT_TOLERANCE {
return Err(MathError::NotSpd);
}
l[i][j] = sum / pivot;
}
}
}
let mut y = vec![0.0_f64; n];
for i in 0..n {
let mut sum = b[i];
for k in 0..i {
sum -= l[i][k] * y[k];
}
y[i] = sum / l[i][i];
}
let mut x = vec![0.0_f64; n];
for i in (0..n).rev() {
let mut sum = y[i];
for k in (i + 1)..n {
sum -= l[k][i] * x[k];
}
x[i] = sum / l[i][i];
}
Ok(x)
}
#[cfg(test)]
mod tests {
use super::*;
fn matvec(a: &[Vec<f64>], x: &[f64]) -> Vec<f64> {
a.iter()
.map(|row| row.iter().zip(x.iter()).map(|(&aij, &xj)| aij * xj).sum())
.collect()
}
#[test]
fn solve_identity_2x2() {
let mut a = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
let mut b = vec![3.0, 5.0];
solve(&mut a, &mut b).unwrap();
assert!((b[0] - 3.0).abs() < 1e-15);
assert!((b[1] - 5.0).abs() < 1e-15);
}
#[test]
fn solve_3x3_known_inverse() {
let a_orig = [
vec![1.0, 2.0, 3.0],
vec![2.0, 5.0, 3.0],
vec![1.0, 0.0, 8.0],
];
let mut a: Vec<Vec<f64>> = a_orig.to_vec();
let mut b = matvec(&a_orig, &[1.0, 2.0, 3.0]);
solve(&mut a, &mut b).unwrap();
assert!((b[0] - 1.0).abs() < 1e-12);
assert!((b[1] - 2.0).abs() < 1e-12);
assert!((b[2] - 3.0).abs() < 1e-12);
}
#[test]
fn solve_requires_pivoting() {
let mut a = vec![vec![0.0, 1.0], vec![1.0, 0.0]];
let mut b = vec![2.0, 3.0];
solve(&mut a, &mut b).unwrap();
assert!((b[0] - 3.0).abs() < 1e-15);
assert!((b[1] - 2.0).abs() < 1e-15);
}
#[test]
fn solve_singular_matrix_rejected() {
let mut a = vec![vec![1.0, 2.0], vec![2.0, 4.0]];
let mut b = vec![1.0, 2.0];
let err = solve(&mut a, &mut b).unwrap_err();
assert!(matches!(err, MathError::Singular));
}
#[test]
fn solve_dimension_mismatch_rejected() {
let mut a = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
let mut b = vec![1.0];
let err = solve(&mut a, &mut b).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
let mut a = vec![vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]];
let mut b = vec![1.0, 2.0];
let err = solve(&mut a, &mut b).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
}
#[test]
fn solve_empty_rejected() {
let mut a: Vec<Vec<f64>> = vec![];
let mut b: Vec<f64> = vec![];
let err = solve(&mut a, &mut b).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
}
#[test]
fn solve_random_5x5() {
let a_orig: Vec<Vec<f64>> = vec![
vec![10.0, 1.0, 0.0, 2.0, -1.0],
vec![1.0, 12.0, -2.0, 1.0, 0.0],
vec![0.0, -1.0, 9.0, 0.0, 2.0],
vec![2.0, 0.0, 1.0, 8.0, -2.0],
vec![-1.0, 0.0, 1.0, -2.0, 11.0],
];
let x_true = vec![1.0, -2.0, 3.0, -4.0, 5.0];
let b_true = matvec(&a_orig, &x_true);
let mut a = a_orig.clone();
let mut b = b_true.clone();
solve(&mut a, &mut b).unwrap();
for (got, exp) in b.iter().zip(x_true.iter()) {
assert!((got - exp).abs() < 1e-10);
}
}
#[test]
fn solve_spd_identity_3x3() {
let a = vec![
vec![1.0, 0.0, 0.0],
vec![0.0, 1.0, 0.0],
vec![0.0, 0.0, 1.0],
];
let b = vec![2.0, 3.0, 5.0];
let x = solve_spd(&a, &b).unwrap();
assert!((x[0] - 2.0).abs() < 1e-15);
assert!((x[1] - 3.0).abs() < 1e-15);
assert!((x[2] - 5.0).abs() < 1e-15);
}
#[test]
fn solve_spd_standard_3x3() {
let a = vec![
vec![4.0, 12.0, -16.0],
vec![12.0, 37.0, -43.0],
vec![-16.0, -43.0, 98.0],
];
let x_true = vec![1.0, 1.0, 1.0];
let b = matvec(&a, &x_true);
let x = solve_spd(&a, &b).unwrap();
for (got, exp) in x.iter().zip(x_true.iter()) {
assert!((got - exp).abs() < 1e-10);
}
}
#[test]
fn solve_spd_rejects_non_spd() {
let a = vec![vec![1.0, 2.0], vec![2.0, 1.0]];
let b = vec![1.0, 1.0];
let err = solve_spd(&a, &b).unwrap_err();
assert!(matches!(err, MathError::NotSpd));
}
#[test]
fn solve_spd_rejects_zero_diagonal() {
let a = vec![vec![0.0, 0.0], vec![0.0, 1.0]];
let b = vec![0.0, 1.0];
let err = solve_spd(&a, &b).unwrap_err();
assert!(matches!(err, MathError::NotSpd));
}
#[test]
fn solve_spd_dimension_mismatch() {
let a = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
let b = vec![1.0];
let err = solve_spd(&a, &b).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
let a = vec![vec![1.0, 0.0]];
let b = vec![1.0];
let err = solve_spd(&a, &b).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
}
#[test]
fn solve_spd_empty_rejected() {
let a: Vec<Vec<f64>> = vec![];
let b: Vec<f64> = vec![];
let err = solve_spd(&a, &b).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
}
#[test]
fn solve_spd_cross_check_with_solve() {
let a = vec![
vec![4.0, 1.0, 0.0],
vec![1.0, 3.0, 1.0],
vec![0.0, 1.0, 2.0],
];
let x_true = vec![1.0, -1.0, 2.0];
let b = matvec(&a, &x_true);
let x_spd = solve_spd(&a, &b).unwrap();
let mut a_mut = a.clone();
let mut b_mut = b.clone();
solve(&mut a_mut, &mut b_mut).unwrap();
for (s, g) in x_spd.iter().zip(b_mut.iter()) {
assert!((s - g).abs() < 1e-10);
}
}
}