use crate::error::{DatarustError, Result};
pub(crate) fn cholesky_decompose(a: &[f64], n: usize) -> Result<Vec<f64>> {
if n == 0 {
return Err(DatarustError::EmptyInput("cholesky of 0×0 matrix".into()));
}
if a.len() != n * n {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} elements ({}×{})", n * n, n, n),
actual: format!("{} elements", a.len()),
});
}
let mut l = vec![0.0_f64; n * n];
let diag_max = (0..n)
.map(|i| a[i * n + i])
.fold(0.0_f64, f64::max)
.max(0.0);
let tol = (diag_max * 1e-14).max(f64::MIN_POSITIVE);
for i in 0..n {
for j in 0..=i {
let mut sum = a[i * n + j];
for k in 0..j {
sum -= l[i * n + k] * l[j * n + k];
}
if i == j {
if sum <= tol {
return Err(DatarustError::Singular(format!(
"matrix is not positive-definite (pivot {:.3e} <= tol {:.3e} at index {})",
sum, tol, i
)));
}
l[i * n + j] = sum.sqrt();
} else {
let diag = l[j * n + j];
l[i * n + j] = sum / diag;
}
}
}
Ok(l)
}
pub(crate) fn solve_spd(l: &[f64], n: usize, b: &[f64]) -> Result<Vec<f64>> {
if l.len() != n * n {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} elements ({}×{})", n * n, n, n),
actual: format!("{} elements", l.len()),
});
}
if b.len() != n {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} elements", n),
actual: format!("{} elements", b.len()),
});
}
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 * n + k] * y[k];
}
let diag = l[i * n + i];
if diag == 0.0 {
return Err(DatarustError::Singular(format!(
"zero diagonal in Cholesky factor at index {i}"
)));
}
y[i] = sum / diag;
}
let mut x = vec![0.0_f64; n];
for ii in 0..n {
let i = n - 1 - ii;
let mut sum = y[i];
for k in (i + 1)..n {
sum -= l[k * n + i] * x[k];
}
let diag = l[i * n + i];
x[i] = sum / diag;
}
Ok(x)
}
pub(crate) fn solve_spd_system(a: &[f64], n: usize, b: &[f64]) -> Result<Vec<f64>> {
let l = cholesky_decompose(a, n)?;
solve_spd(&l, n, b)
}
#[cfg(test)]
mod tests {
use super::*;
fn reconstruct(l: &[f64], n: usize) -> Vec<f64> {
let mut a = vec![0.0; n * n];
for i in 0..n {
for j in 0..n {
let mut s = 0.0;
for k in 0..n {
s += l[i * n + k] * l[j * n + k];
}
a[i * n + j] = s;
}
}
a
}
#[test]
fn cholesky_identity() {
let n = 3;
let a = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0];
let l = cholesky_decompose(&a, n).unwrap();
let recon = reconstruct(&l, n);
for x in 0..n * n {
assert!((recon[x] - a[x]).abs() < 1e-12);
}
assert!((l[0] - 1.0).abs() < 1e-12);
}
#[test]
fn cholesky_known_pd() {
let a = vec![4.0, 12.0, -16.0, 12.0, 37.0, -43.0, -16.0, -43.0, 98.0];
let n = 3;
let l = cholesky_decompose(&a, n).unwrap();
let recon = reconstruct(&l, n);
for x in 0..n * n {
assert!((recon[x] - a[x]).abs() < 1e-9, "mismatch at {x}");
}
for i in 0..n {
for j in (i + 1)..n {
assert!(l[i * n + j].abs() < 1e-12);
}
}
}
#[test]
fn cholesky_non_pd_returns_singular() {
let a = vec![1.0, 2.0, 2.0, 1.0];
let err = cholesky_decompose(&a, 2).unwrap_err();
assert!(matches!(err, DatarustError::Singular(_)));
}
#[test]
fn cholesky_zero_diagonal_singular() {
let a = vec![0.0, 0.0, 0.0, 1.0];
let err = cholesky_decompose(&a, 2).unwrap_err();
assert!(matches!(err, DatarustError::Singular(_)));
}
#[test]
fn solve_spd_basic() {
let a = vec![4.0, 2.0, 2.0, 3.0];
let b = vec![10.0, 11.0]; let l = cholesky_decompose(&a, 2).unwrap();
let x = solve_spd(&l, 2, &b).unwrap();
assert!((x[0] - 1.0).abs() < 1e-9);
assert!((x[1] - 3.0).abs() < 1e-9);
}
#[test]
fn solve_spd_round_trip() {
let n = 4;
let m_flat = vec![
1.0, 2.0, 0.0, 1.0, 3.0, 1.0, 4.0, 2.0, 0.0, 5.0, 1.0, 3.0, 2.0, 1.0, 3.0, 1.0,
];
let mut a = vec![0.0; n * n];
for i in 0..n {
for j in 0..n {
let mut s = 0.0;
for k in 0..n {
s += m_flat[k * n + i] * m_flat[k * n + j];
}
a[i * n + j] = s + if i == j { n as f64 } else { 0.0 };
}
}
let true_x = [1.5, -2.0, 0.5, 3.0];
let b: Vec<f64> = (0..n)
.map(|i| (0..n).map(|j| a[i * n + j] * true_x[j]).sum())
.collect();
let x = solve_spd_system(&a, n, &b).unwrap();
for i in 0..n {
assert!((x[i] - true_x[i]).abs() < 1e-8, "mismatch at {i}");
}
}
#[test]
fn solve_spd_shape_mismatch() {
let l = vec![1.0; 4]; let b = vec![1.0, 2.0, 3.0]; let err = solve_spd(&l, 2, &b).unwrap_err();
assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
#[test]
fn solve_spd_zero_diagonal() {
let l = vec![0.0, 0.0, 1.0, 0.0];
let b = vec![1.0, 2.0];
let err = solve_spd(&l, 2, &b).unwrap_err();
assert!(matches!(err, DatarustError::Singular(_)));
}
#[test]
fn cholesky_empty_rejected() {
let err = cholesky_decompose(&[], 0).unwrap_err();
assert!(matches!(err, DatarustError::EmptyInput(_)));
}
}