#![allow(clippy::needless_range_loop)]
use sparse_ldlt::{LdltError, SparseLdlt};
struct Rng(u64);
impl Rng {
fn next_f64(&mut self) -> f64 {
self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
((self.0 >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0
}
fn below(&mut self, n: usize) -> usize {
((self.next_f64() + 1.0) / 2.0 * n as f64) as usize % n
}
}
fn matvec(n: usize, cp: &[usize], ri: &[usize], v: &[f64], x: &[f64]) -> Vec<f64> {
let mut y = vec![0.0f64; n];
for j in 0..n {
for p in cp[j]..cp[j + 1] {
let i = ri[p];
y[i] += v[p] * x[j];
if i != j {
y[j] += v[p] * x[i];
}
}
}
y
}
fn residual_inf(n: usize, cp: &[usize], ri: &[usize], v: &[f64], x: &[f64], b: &[f64]) -> f64 {
matvec(n, cp, ri, v, x)
.iter()
.zip(b)
.map(|(ax, bi)| (ax - bi).abs())
.fold(0.0f64, f64::max)
}
fn solves_original_system(
n: usize,
cp: &[usize],
ri: &[usize],
v: &[f64],
seed: u64,
) -> Result<(), String> {
let mut rng = Rng(seed);
let b: Vec<f64> = (0..n).map(|_| rng.next_f64()).collect();
let f = match SparseLdlt::factor(n, cp, ri, v) {
Ok(f) => f,
Err(LdltError::ZeroPivot(_) | LdltError::NearZeroPivot { .. }) => return Ok(()), Err(e) => return Err(format!("unexpected error {e:?}")),
};
let x = match f.solve(&b) {
Ok(x) => x,
Err(e) => return Err(format!("solve failed on a factored system: {e:?}")),
};
let res = residual_inf(n, cp, ri, v, &x, &b);
if res > 1e-6 {
return Err(format!("residual {res}"));
}
Ok(())
}
#[test]
fn duplicates_are_summed_and_still_solve() {
for seed in 0..20u64 {
let n = 3 + (seed as usize % 9);
let mut rng = Rng(seed * 3 + 1);
let mut cp = vec![0usize];
let mut ri = Vec::new();
let mut v = Vec::new();
for j in 0..n {
for i in 0..=j {
if rng.below(3) == 0 || i == j {
let val = rng.next_f64() + if i == j { n as f64 } else { 0.0 };
ri.push(i);
v.push(val);
ri.push(i); v.push(val);
}
}
cp.push(ri.len());
}
solves_original_system(n, &cp, &ri, &v, seed)
.unwrap_or_else(|e| panic!("seed {seed}: duplicated entries: {e}"));
}
}
#[test]
fn explicit_zeros_are_harmless() {
let cp: &[usize] = &[0, 3, 7, 10];
let ri: &[usize] = &[0, 0, 1, 0, 0, 1, 2, 1, 1, 2];
let v: &[f64] = &[2.0, 0.0, 1.0, 1.0, 0.0, -3.0, 1.0, 1.0, 0.0, 2.0];
let f = SparseLdlt::factor(3, cp, ri, v).expect("explicit zeros must factor");
let x = f.solve(&[1.0, 2.0, 3.0]).unwrap();
let want = [0.5, 0.0, 1.5]; for i in 0..3 {
assert!((x[i] - want[i]).abs() < 1e-14);
}
}
#[test]
fn lower_triangle_only_is_accepted() {
let cp: &[usize] = &[0, 1, 2, 3];
let ri: &[usize] = &[0, 1, 2];
let v: &[f64] = &[4.0, 5.0, 6.0];
let f = SparseLdlt::factor(3, cp, ri, v).unwrap();
let x = f.solve(&[8.0, 10.0, 12.0]).unwrap();
assert!((x[0] - 2.0).abs() < 1e-14 && (x[1] - 2.0).abs() < 1e-14 && (x[2] - 2.0).abs() < 1e-14);
}
#[test]
fn unordered_rows_within_columns_are_accepted() {
for perm_seed in 0..10u64 {
let mut rng = Rng(perm_seed * 11 + 2);
let mut cols: Vec<Vec<(usize, f64)>> = vec![
vec![(0, 2.0), (1, 1.0)],
vec![(0, 1.0), (1, -3.0), (2, 1.0)],
vec![(1, 1.0), (2, 2.0)],
];
for col in &mut cols {
for k in (1..col.len()).rev() {
let s = rng.below(k + 1);
col.swap(k, s);
}
}
let mut cp = vec![0usize];
let mut ri = Vec::new();
let mut v = Vec::new();
for col in &cols {
for (r, val) in col {
ri.push(*r);
v.push(*val);
}
cp.push(ri.len());
}
let f = SparseLdlt::factor(3, &cp, &ri, &v).expect("shuffled rows must factor");
let x = f.solve(&[1.0, 2.0, 3.0]).unwrap();
let want = [0.5, 0.0, 1.5];
for i in 0..3 {
assert!((x[i] - want[i]).abs() < 1e-12, "perm {perm_seed}: x[{i}] drifted");
}
}
}
#[test]
fn degenerate_shapes_never_panic() {
assert!(SparseLdlt::factor(0, &[0], &[], &[]).is_ok());
assert!(matches!(
SparseLdlt::factor(0, &[], &[], &[]),
Err(LdltError::InvalidInput(_))
));
let f = SparseLdlt::factor(1, &[0, 1], &[0], &[2.0]).unwrap();
assert_eq!(f.solve(&[4.0]).unwrap(), vec![2.0]);
assert!(matches!(
SparseLdlt::factor(1, &[0, 1], &[0], &[0.0]),
Err(LdltError::ZeroPivot(0))
));
let r = SparseLdlt::factor(3, &[0, 2, 2, 3], &[0, 1, 2], &[1.0, 1.0, 1.0]);
assert!(matches!(r, Err(LdltError::ZeroPivot(1))), "got {r:?}");
let r = SparseLdlt::factor(2, &[0, 1, 1], &[0], &[1.0]);
assert!(matches!(r, Err(LdltError::ZeroPivot(1))), "got {r:?}");
}
#[test]
fn malformed_arrays_error_never_panic() {
assert!(matches!(
SparseLdlt::factor(2, &[0, 1], &[0, 1], &[1.0, 1.0]),
Err(LdltError::InvalidInput(_))
));
assert!(matches!(
SparseLdlt::factor(2, &[0, 1, 2], &[0, 1], &[1.0]),
Err(LdltError::InvalidInput(_))
));
assert!(matches!(
SparseLdlt::factor(2, &[0, 1, 5], &[0, 1], &[1.0, 1.0]),
Err(LdltError::InvalidInput(_))
));
assert!(matches!(
SparseLdlt::factor(2, &[0, 2, 1], &[0, 0], &[1.0, 1.0]),
Err(LdltError::InvalidInput(_))
));
assert!(matches!(
SparseLdlt::factor(2, &[0, 1, 2], &[0, 2], &[1.0, 1.0]),
Err(LdltError::InvalidInput(_))
));
assert!(SparseLdlt::factor(0, &[0], &[], &[]).is_ok());
}
#[test]
fn random_valid_shape_garbage_never_panics() {
for seed in 0..300u64 {
let n = 1 + (seed as usize % 24);
let mut rng = Rng(seed * 7919 + 13);
let nnz_cap = n * (n + 1) / 2;
let nnz = 1 + rng.below(nnz_cap.min(64));
let mut cp = vec![0usize; n + 1];
let mut ri = Vec::new();
let mut v = Vec::new();
for _ in 0..nnz {
let j = rng.below(n);
let i = rng.below(j + 1); ri.push(i);
let val = match seed % 5 {
0 => rng.next_f64() * 1e12,
1 => rng.next_f64() * 1e-6,
_ => rng.next_f64(),
};
v.push(val);
cp[j + 1] += 1;
}
for j in 1..=n {
cp[j] += cp[j - 1];
}
let factored = SparseLdlt::factor(n, &cp, &ri, &v);
match factored {
Ok(f) => {
assert!(
f.d().iter().all(|d| d.is_finite()),
"seed {seed}: non-finite pivot with finite input"
);
let b: Vec<f64> = (0..n).map(|_| rng.next_f64()).collect();
let x = f.solve(&b).unwrap();
assert!(x.iter().all(|xi| xi.is_finite()), "seed {seed}: non-finite solve");
}
Err(LdltError::ZeroPivot(_) | LdltError::NearZeroPivot { .. }) => {} Err(e) => panic!("seed {seed}: unexpected {e:?}"),
}
}
}
#[test]
fn full_symmetric_storage_matches_upper_only() {
let upper = SparseLdlt::factor(2, &[0, 1, 3], &[0, 0, 1], &[3.0, 1.0, 2.0]).unwrap();
let full = SparseLdlt::factor(2, &[0, 2, 4], &[0, 1, 0, 1], &[3.0, 1.0, 1.0, 2.0]).unwrap();
let b = [5.0, 7.0];
let x1 = upper.solve(&b).unwrap();
let x2 = full.solve(&b).unwrap();
for i in 0..2 {
assert!((x1[i] - x2[i]).abs() < 1e-13);
}
}