#![forbid(unsafe_code)]
#![allow(clippy::needless_range_loop)]
pub const NEAR_ZERO_PIVOT_REL: f64 = 1e-13;
#[derive(Debug, Clone, PartialEq)]
pub enum LdltError {
ZeroPivot(usize),
NearZeroPivot {
column: usize,
pivot: f64,
scale: f64,
suggested_shift: f64,
},
InvalidInput(&'static str),
SizeMismatch {
expected: usize,
got: usize,
},
}
#[derive(Debug, Clone)]
pub struct SparseLdlt {
n: usize,
lp: Vec<usize>, li: Vec<usize>, lx: Vec<f64>, d: Vec<f64>, order: Vec<usize>,
shift: f64,
}
fn diagonal_scale(n: usize, col_ptr: &[usize], row_idx: &[usize], values: &[f64]) -> f64 {
let mut scale = 0.0f64;
for k in 0..n {
let mut dk = 0.0f64;
for p in col_ptr[k]..col_ptr[k + 1] {
if row_idx[p] == k {
dk += values[p];
}
}
scale = scale.max(dk.abs());
}
scale
}
#[allow(clippy::type_complexity)]
fn with_diagonal_shift(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
shift: f64,
) -> (Vec<usize>, Vec<usize>, Vec<f64>) {
let mut cp = Vec::with_capacity(n + 1);
let mut ri = Vec::with_capacity(row_idx.len() + n);
let mut vx = Vec::with_capacity(values.len() + n);
cp.push(0usize);
for k in 0..n {
for p in col_ptr[k]..col_ptr[k + 1] {
ri.push(row_idx[p]);
vx.push(values[p]);
}
ri.push(k);
vx.push(shift);
cp.push(ri.len());
}
(cp, ri, vx)
}
impl SparseLdlt {
pub fn factor(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
) -> Result<Self, LdltError> {
Self::factor_inner(n, col_ptr, row_idx, values, None)
}
pub fn factor_reporting_collapse(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
) -> Result<(Self, Vec<usize>), LdltError> {
let mut collapsed = Vec::new();
let f = Self::factor_inner(n, col_ptr, row_idx, values, Some(&mut collapsed))?;
Ok((f, collapsed))
}
fn factor_inner(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
mut collapsed: Option<&mut Vec<usize>>,
) -> Result<Self, LdltError> {
if col_ptr.len() != n + 1 {
return Err(LdltError::InvalidInput("col_ptr length must be n + 1"));
}
if row_idx.len() != values.len() {
return Err(LdltError::InvalidInput("row_idx and values length mismatch"));
}
if col_ptr[n] != row_idx.len() {
return Err(LdltError::InvalidInput("col_ptr[n] must equal the nonzero count"));
}
for k in 0..n {
if col_ptr[k] > col_ptr[k + 1] {
return Err(LdltError::InvalidInput("col_ptr must be non-decreasing"));
}
}
for &r in row_idx {
if r >= n {
return Err(LdltError::InvalidInput("row index out of range"));
}
}
for &v in values {
if !v.is_finite() {
return Err(LdltError::InvalidInput(
"values contain a non-finite entry (NaN or infinity)",
));
}
}
let ap = col_ptr;
let ai = row_idx;
let ax = values;
let scale = diagonal_scale(n, ap, ai, ax);
let mut parent = vec![usize::MAX; n];
let mut flag = vec![usize::MAX; n];
let mut lnz = vec![0usize; n];
for k in 0..n {
flag[k] = k;
for p in ap[k]..ap[k + 1] {
let mut i = ai[p];
if i < k {
while flag[i] != k {
if parent[i] == usize::MAX {
parent[i] = k;
}
lnz[i] += 1;
flag[i] = k;
i = parent[i];
}
}
}
}
let mut lp = vec![0usize; n + 1];
for k in 0..n {
lp[k + 1] = lp[k] + lnz[k];
}
let mut li = vec![0usize; lp[n]];
let mut lx = vec![0.0f64; lp[n]];
let mut d = vec![0.0f64; n];
let mut y = vec![0.0f64; n]; let mut pattern = vec![0usize; n];
let mut fill = vec![0usize; n]; for f in flag.iter_mut() {
*f = usize::MAX;
}
for k in 0..n {
let mut top = n;
flag[k] = k;
y[k] = 0.0;
for p in ap[k]..ap[k + 1] {
let i = ai[p];
if i <= k {
y[i] += ax[p];
let mut len = 0usize;
let mut ii = i;
while flag[ii] != k {
pattern[len] = ii;
len += 1;
flag[ii] = k;
ii = parent[ii];
}
while len > 0 {
len -= 1;
top -= 1;
pattern[top] = pattern[len];
}
}
}
d[k] = y[k];
y[k] = 0.0;
for idx in top..n {
let i = pattern[idx];
let yi = y[i];
y[i] = 0.0;
let start = lp[i];
let used = fill[i];
for p in start..start + used {
y[li[p]] -= lx[p] * yi;
}
let l_ki = yi / d[i];
d[k] -= l_ki * yi;
let slot = start + used;
li[slot] = k;
lx[slot] = l_ki;
fill[i] = used + 1;
}
if d[k] == 0.0 {
return Err(LdltError::ZeroPivot(k));
}
if scale > 0.0 && d[k].abs() < NEAR_ZERO_PIVOT_REL * scale {
if let Some(list) = collapsed.as_deref_mut() {
list.push(k);
continue;
}
return Err(LdltError::NearZeroPivot {
column: k,
pivot: d[k],
scale,
suggested_shift: NEAR_ZERO_PIVOT_REL.sqrt() * scale,
});
}
}
Ok(SparseLdlt { n, lp, li, lx, d, order: (0..n).collect(), shift: 0.0 })
}
pub fn factor_perm(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
order: &[usize],
) -> Result<Self, LdltError> {
if order.len() != n {
return Err(LdltError::InvalidInput("order length must be n"));
}
let mut pos = vec![usize::MAX; n]; for (new, &old) in order.iter().enumerate() {
if old >= n || pos[old] != usize::MAX {
return Err(LdltError::InvalidInput(
"order must be a permutation of 0..n",
));
}
pos[old] = new;
}
let mut entries: Vec<(usize, f64)> = Vec::with_capacity(values.len());
let mut pcp = vec![0usize; n + 1];
for k in 0..n {
let old_k = order[k];
for p in col_ptr[old_k]..col_ptr[old_k + 1] {
entries.push((pos[row_idx[p]], values[p]));
}
entries[pcp[k]..].sort_unstable_by_key(|e| e.0);
let mut w = pcp[k];
let mut r = pcp[k];
while r < entries.len() {
let (row, mut val) = entries[r];
r += 1;
while r < entries.len() && entries[r].0 == row {
val += entries[r].1;
r += 1;
}
entries[w] = (row, val);
w += 1;
}
entries.truncate(w);
pcp[k + 1] = entries.len();
}
let pri: Vec<usize> = entries.iter().map(|e| e.0).collect();
let pv: Vec<f64> = entries.iter().map(|e| e.1).collect();
let mut f = Self::factor(n, &pcp, &pri, &pv)?;
f.order = order.to_vec();
Ok(f)
}
pub fn factor_shifted(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
) -> Result<Self, LdltError> {
Self::shifted_retry(n, col_ptr, row_idx, values, None)
}
pub fn factor_perm_shifted(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
order: &[usize],
) -> Result<Self, LdltError> {
Self::shifted_retry(n, col_ptr, row_idx, values, Some(order))
}
fn shifted_retry(
n: usize,
col_ptr: &[usize],
row_idx: &[usize],
values: &[f64],
order: Option<&[usize]>,
) -> Result<Self, LdltError> {
let attempt = |cp: &[usize], ri: &[usize], vx: &[f64]| match order {
Some(o) => Self::factor_perm(n, cp, ri, vx, o),
None => Self::factor(n, cp, ri, vx),
};
let mut last = match attempt(col_ptr, row_idx, values) {
Ok(f) => return Ok(f),
Err(e) => e,
};
let mut shift = match last {
LdltError::NearZeroPivot {
suggested_shift, ..
} => suggested_shift,
LdltError::ZeroPivot(_) => {
NEAR_ZERO_PIVOT_REL.sqrt() * diagonal_scale(n, col_ptr, row_idx, values)
}
other => return Err(other),
};
if shift <= 0.0 {
return Err(last);
}
for _ in 0..8 {
let (cp, ri, vx) = with_diagonal_shift(n, col_ptr, row_idx, values, shift);
match attempt(&cp, &ri, &vx) {
Ok(mut f) => {
f.shift = shift;
return Ok(f);
}
Err(e) => last = e,
}
shift *= 8.0;
}
Err(last)
}
pub fn shift(&self) -> f64 {
self.shift
}
pub fn dim(&self) -> usize {
self.n
}
pub fn d(&self) -> &[f64] {
&self.d
}
pub fn nnz(&self) -> usize {
self.lp[self.n]
}
pub fn flops(&self) -> u64 {
let mut f = 0u64;
for j in 0..self.n {
let c = (self.lp[j + 1] - self.lp[j]) as u64;
f += c * c + 3 * c;
}
f
}
pub fn solve(&self, b: &[f64]) -> Result<Vec<f64>, LdltError> {
if b.len() != self.n {
return Err(LdltError::SizeMismatch { expected: self.n, got: b.len() });
}
let identity = self.order.len() == self.n && self.order.iter().enumerate().all(|(k, &o)| o == k);
let mut x = if identity {
b.to_vec()
} else {
self.order.iter().map(|&o| b[o]).collect()
};
for j in 0..self.n {
let xj = x[j];
for p in self.lp[j]..self.lp[j + 1] {
x[self.li[p]] -= self.lx[p] * xj;
}
}
for j in 0..self.n {
x[j] /= self.d[j];
}
for j in (0..self.n).rev() {
let mut acc = x[j];
for p in self.lp[j]..self.lp[j + 1] {
acc -= self.lx[p] * x[self.li[p]];
}
x[j] = acc;
}
if identity {
Ok(x)
} else {
let mut out = vec![0.0f64; self.n];
for (k, &o) in self.order.iter().enumerate() {
out[o] = x[k];
}
Ok(out)
}
}
}
pub fn amd(n: usize, col_ptr: &[usize], row_idx: &[usize]) -> Vec<usize> {
let mut adj: Vec<Vec<usize>> = vec![Vec::new(); n];
for k in 0..n {
for p in col_ptr[k]..col_ptr[k + 1] {
let i = row_idx[p];
if i < n && i != k {
adj[k].push(i);
adj[i].push(k);
}
}
}
for a in adj.iter_mut() {
a.sort_unstable();
a.dedup();
}
let mut elem_vars: Vec<Vec<usize>> = Vec::new();
let mut elem_alive: Vec<bool> = Vec::new();
let mut elems_of: Vec<Vec<usize>> = vec![Vec::new(); n];
let mut alive = vec![true; n];
let mut deg: Vec<usize> = adj.iter().map(Vec::len).collect();
let mut flag = vec![usize::MAX; n]; let mut next_stamp = 0usize; let mut order = Vec::with_capacity(n);
let mut heap: std::collections::BinaryHeap<std::cmp::Reverse<(usize, usize)>> =
(0..n).map(|u| std::cmp::Reverse((deg[u], u))).collect();
for _step in 0..n {
let i = loop {
match heap.pop() {
Some(std::cmp::Reverse((d, u))) => {
if alive[u] && deg[u] == d {
break u;
}
}
None => break usize::MAX,
}
};
if i == usize::MAX {
break;
}
alive[i] = false;
order.push(i);
next_stamp += 1;
let stamp = next_stamp;
let mut nb: Vec<usize> = Vec::with_capacity(deg[i] + 1);
for &a in &adj[i] {
if a < n && alive[a] && flag[a] != stamp {
flag[a] = stamp;
nb.push(a);
}
}
for &e in &elems_of[i] {
if !elem_alive[e] {
continue;
}
for &x in &elem_vars[e] {
if x < n && alive[x] && flag[x] != stamp {
flag[x] = stamp;
nb.push(x);
}
}
}
for &e in &elems_of[i] {
elem_alive[e] = false;
}
if nb.is_empty() {
continue;
}
let elem_id = elem_vars.len();
elem_vars.push(nb.clone());
elem_alive.push(true);
for &j in &nb {
elems_of[j].retain(|&e| elem_alive[e]);
elems_of[j].push(elem_id);
}
for &j in &nb {
next_stamp += 1;
let estamp = next_stamp;
let mut count = 0usize;
flag[j] = estamp;
let scan = |xs: &[usize], flag: &mut Vec<usize>, count: &mut usize| {
for &x in xs {
if x < n && alive[x] && flag[x] != estamp {
flag[x] = estamp;
*count += 1;
}
}
};
scan(&adj[j], &mut flag, &mut count);
for &e in &elems_of[j] {
scan(&elem_vars[e], &mut flag, &mut count);
}
deg[j] = count;
heap.push(std::cmp::Reverse((count, j)));
}
}
if order.len() < n {
for u in 0..n {
if alive[u] {
order.push(u);
}
}
}
order
}
#[cfg(test)]
mod tests {
use super::*;
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
}
}
#[allow(clippy::type_complexity)]
fn random_symmetric(
n: usize,
density: f64,
diag_shift: f64,
seed: u64,
) -> (Vec<usize>, Vec<usize>, Vec<f64>, Vec<Vec<f64>>) {
let mut rng = Rng(seed);
let mut dense = vec![vec![0.0f64; n]; n];
for i in 0..n {
for j in (i + 1)..n {
if (rng.next_f64() + 1.0) / 2.0 < density {
let v = rng.next_f64();
dense[i][j] = v;
dense[j][i] = v;
}
}
dense[i][i] = rng.next_f64() + diag_shift;
}
let mut col_ptr = vec![0usize];
let mut row_idx = Vec::new();
let mut values = Vec::new();
for j in 0..n {
for i in 0..n {
if dense[i][j] != 0.0 {
row_idx.push(i);
values.push(dense[i][j]);
}
}
col_ptr.push(row_idx.len());
}
(col_ptr, row_idx, values, dense)
}
fn residual_inf(dense: &[Vec<f64>], x: &[f64], b: &[f64]) -> f64 {
let n = b.len();
(0..n)
.map(|i| {
let ax: f64 = (0..n).map(|j| dense[i][j] * x[j]).sum();
(ax - b[i]).abs()
})
.fold(0.0, f64::max)
}
fn negative_eigs(mat: &[Vec<f64>]) -> usize {
let n = mat.len();
let mut a = mat.to_vec();
for _sweep in 0..100 {
let mut off = 0.0;
for p in 0..n {
for q in (p + 1)..n {
off += a[p][q] * a[p][q];
}
}
if off < 1e-20 {
break;
}
for p in 0..n {
for q in (p + 1)..n {
if a[p][q].abs() < 1e-18 {
continue;
}
let theta = (a[q][q] - a[p][p]) / (2.0 * a[p][q]);
let t = theta.signum() / (theta.abs() + (theta * theta + 1.0).sqrt());
let c = 1.0 / (t * t + 1.0).sqrt();
let s = t * c;
for k in 0..n {
let akp = a[k][p];
let akq = a[k][q];
a[k][p] = c * akp - s * akq;
a[k][q] = s * akp + c * akq;
}
for k in 0..n {
let apk = a[p][k];
let aqk = a[q][k];
a[p][k] = c * apk - s * aqk;
a[q][k] = s * apk + c * aqk;
}
}
}
}
(0..n).filter(|&i| a[i][i] < -1e-9).count()
}
#[test]
fn spd_solves_accurately_with_no_negative_pivots() {
for seed in 0..25u64 {
let n = 6 + (seed as usize % 18);
let (cp, ri, v, dense) = random_symmetric(n, 0.4, n as f64 + 2.0, seed * 7 + 1);
let mut rng = Rng(seed * 13 + 3);
let b: Vec<f64> = (0..n).map(|_| rng.next_f64()).collect();
let f = SparseLdlt::factor(n, &cp, &ri, &v).expect("SPD factor");
let x = f.solve(&b).unwrap();
assert!(residual_inf(&dense, &x, &b) < 1e-9, "seed {seed}: residual too large");
assert_eq!(f.d().iter().filter(|&&d| d < 0.0).count(), 0);
}
}
#[test]
fn indefinite_solves_and_inertia_is_correct() {
let mut indefinite = 0;
for seed in 0..60u64 {
let n = 4 + (seed as usize % 10);
let (cp, ri, v, dense) = random_symmetric(n, 0.35, 0.5, seed * 5 + 9);
let mut rng = Rng(seed * 17 + 2);
let b: Vec<f64> = (0..n).map(|_| rng.next_f64()).collect();
let f = match SparseLdlt::factor(n, &cp, &ri, &v) {
Ok(f) => f,
Err(_) => continue, };
let x = f.solve(&b).unwrap();
assert!(residual_inf(&dense, &x, &b) < 1e-7, "seed {seed}: residual too large");
let neg = f.d().iter().filter(|&&d| d < 0.0).count();
assert_eq!(neg, negative_eigs(&dense), "seed {seed}: inertia mismatch");
if neg > 0 {
indefinite += 1;
}
}
assert!(indefinite >= 5, "expected several indefinite cases, got {indefinite}");
}
#[test]
fn rejects_malformed_input() {
assert!(matches!(SparseLdlt::factor(2, &[0, 1], &[0], &[1.0]), Err(LdltError::InvalidInput(_))));
}
#[test]
fn rejects_non_finite_values() {
let cp: &[usize] = &[0, 1, 2];
let ri: &[usize] = &[0, 1];
assert!(matches!(
SparseLdlt::factor(2, cp, ri, &[f64::NAN, 1.0]),
Err(LdltError::InvalidInput(_))
));
assert!(matches!(
SparseLdlt::factor(2, cp, ri, &[1.0, f64::INFINITY]),
Err(LdltError::InvalidInput(_))
));
}
#[test]
fn solve_rejects_wrong_rhs_length() {
let f = SparseLdlt::factor(3, &[0, 2, 5, 7], &[0, 1, 0, 1, 2, 1, 2],
&[2.0, 1.0, 1.0, -3.0, 1.0, 1.0, 2.0]).unwrap();
assert_eq!(
f.solve(&[1.0, 2.0]),
Err(LdltError::SizeMismatch { expected: 3, got: 2 })
);
}
#[test]
fn near_zero_pivot_is_reported_not_returned() {
let cp: &[usize] = &[0, 2, 4];
let ri: &[usize] = &[0, 1, 0, 1];
let v: &[f64] = &[1e-18, 1.0, 1.0, 1.0];
match SparseLdlt::factor(2, cp, ri, v) {
Err(LdltError::NearZeroPivot {
column,
pivot,
scale,
suggested_shift,
}) => {
assert_eq!(column, 0);
assert_eq!(pivot, 1e-18);
assert_eq!(scale, 1.0);
assert!(suggested_shift > NEAR_ZERO_PIVOT_REL * scale);
}
other => panic!("expected NearZeroPivot, got {other:?}"),
}
assert!(matches!(
SparseLdlt::factor(1, &[0, 1], &[0], &[0.0]),
Err(LdltError::ZeroPivot(0))
));
match SparseLdlt::factor_perm(2, cp, ri, v, &[0, 1]) {
Err(LdltError::NearZeroPivot { column, .. }) => assert_eq!(column, 0),
other => panic!("expected NearZeroPivot from factor_perm, got {other:?}"),
}
}
#[test]
fn factor_shifted_recovers_and_reports_the_shift() {
let cp: &[usize] = &[0, 2, 4];
let ri: &[usize] = &[0, 1, 0, 1];
let v: &[f64] = &[1e-18, 1.0, 1.0, 1.0];
let f = SparseLdlt::factor_shifted(2, cp, ri, v).expect("shifted factor");
let sh = f.shift();
assert!(sh > 0.0, "shift was {sh}");
let b = [1.0, 2.0];
let x = f.solve(&b).unwrap();
let a = [[1e-18 + sh, 1.0], [1.0, 1.0 + sh]];
for i in 0..2 {
let ax = a[i][0] * x[0] + a[i][1] * x[1];
assert!((ax - b[i]).abs() < 1e-9, "row {i}: {ax} vs {}", b[i]);
}
let g = SparseLdlt::factor_perm_shifted(2, cp, ri, v, &[0, 1]).expect("shifted perm");
assert!(g.shift() > 0.0);
let h = SparseLdlt::factor_shifted(2, cp, ri, &[3.0, 1.0, 1.0, 2.0]).unwrap();
assert_eq!(h.shift(), 0.0);
let plain = SparseLdlt::factor(2, cp, ri, &[3.0, 1.0, 1.0, 2.0]).unwrap();
assert_eq!(plain.shift(), 0.0);
}
#[test]
fn golden_known_answer() {
let f = SparseLdlt::factor(3, &[0, 2, 5, 7], &[0, 1, 0, 1, 2, 1, 2],
&[2.0, 1.0, 1.0, -3.0, 1.0, 1.0, 2.0]).unwrap();
assert_eq!(f.nnz(), 2);
assert_eq!(f.dim(), 3);
let d = f.d();
assert_eq!(d[0], 2.0);
assert_eq!(d[1], -3.5);
assert!((d[2] - 16.0 / 7.0).abs() < 1e-15, "d2 = {} (want 16/7)", d[2]);
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, "x[{i}] = {} (want {})", x[i], want[i]);
}
}
}