pub fn cholesky(a: &[f64], n: usize, jitter: f64) -> Vec<f64> {
debug_assert_eq!(a.len(), n * n);
let mut l = vec![0.0f64; n * n];
for i in 0..n {
for j in 0..=i {
let mut s = a[i * n + j];
for t in 0..j {
s -= l[i * n + t] * l[j * n + t];
}
if i == j {
l[i * n + i] = if s > jitter { s.sqrt() } else { jitter.sqrt() };
} else {
l[i * n + j] = s / l[j * n + j];
}
}
}
l
}
pub fn mahal2(l: &[f64], v: &[f64], n: usize) -> f64 {
debug_assert_eq!(l.len(), n * n);
debug_assert_eq!(v.len(), n);
let mut w = vec![0.0f64; n];
let mut d2 = 0.0;
for i in 0..n {
let mut s = v[i];
for t in 0..i {
s -= l[i * n + t] * w[t];
}
let wi = s / l[i * n + i];
w[i] = wi;
d2 += wi * wi;
}
d2
}
pub fn top_eig(s: &[f64], n: usize, iters: usize) -> (f64, Vec<f64>) {
debug_assert_eq!(s.len(), n * n);
let mut v: Vec<f64> = (0..n).map(|i| 1.0 + 1e-3 * i as f64).collect();
let mut norm = v.iter().map(|x| x * x).sum::<f64>().sqrt();
for x in &mut v {
*x /= norm;
}
let mut lam = 0.0;
for _ in 0..iters {
let w: Vec<f64> = (0..n)
.map(|i| (0..n).map(|j| s[i * n + j] * v[j]).sum::<f64>())
.collect();
norm = w.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm <= 0.0 {
return (0.0, v);
}
let w_normed: Vec<f64> = w.iter().map(|x| x / norm).collect();
lam = (0..n)
.map(|i| w_normed[i] * (0..n).map(|j| s[i * n + j] * w_normed[j]).sum::<f64>())
.sum::<f64>();
v = w_normed;
}
(lam.max(0.0), v)
}
pub fn top_factors(s: &[f64], n: usize, r: usize) -> Vec<(f64, Vec<f64>)> {
debug_assert_eq!(s.len(), n * n);
let mut work = s.to_vec();
let mean_diag = (0..n).map(|i| s[i * n + i]).sum::<f64>() / n as f64;
let cutoff = 0.01 * mean_diag;
let mut out = Vec::with_capacity(r);
for _ in 0..r {
let (lam, v) = top_eig(&work, n, 60);
if lam <= cutoff {
break;
}
for i in 0..n {
for j in 0..n {
work[i * n + j] -= lam * v[i] * v[j];
}
}
out.push((lam, v));
}
out
}
pub fn solve_sym(a: &[f64], b: &[f64], n: usize) -> Vec<f64> {
debug_assert_eq!(a.len(), n * n);
debug_assert_eq!(b.len(), n);
let l = cholesky(a, n, 1e-12);
let mut y = vec![0.0f64; n];
for i in 0..n {
let mut s = b[i];
for t in 0..i {
s -= l[i * n + t] * y[t];
}
y[i] = s / l[i * n + i];
}
let mut x = vec![0.0f64; n];
for i in (0..n).rev() {
let mut s = y[i];
for t in i + 1..n {
s -= l[t * n + i] * x[t];
}
x[i] = s / l[i * n + i];
}
x
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cholesky_of_identity_is_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(&a, n, 1e-12);
for i in 0..n {
for j in 0..n {
let expected = if i == j { 1.0 } else { 0.0 };
assert!((l[i * n + j] - expected).abs() < 1e-12);
}
}
}
#[test]
fn cholesky_reconstruction() {
let n = 3;
let a = vec![
4.0, 2.0, 1.0, 2.0, 5.0, 3.0, 1.0, 3.0, 6.0,
];
let l = cholesky(&a, n, 1e-12);
for i in 0..n {
for j in 0..n {
let mut sum = 0.0;
for t in 0..n {
sum += l[i * n + t] * l[j * n + t];
}
assert!(
(sum - a[i * n + j]).abs() < 1e-10,
"L·Lᵀ[{i},{j}] = {sum}, A = {}",
a[i * n + j]
);
}
}
}
#[test]
fn mahal2_of_standard_normal() {
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 v = vec![1.0, 2.0, 3.0];
let l = cholesky(&a, n, 1e-12);
let d2 = mahal2(&l, &v, n);
assert!((d2 - 14.0).abs() < 1e-10, "expected 14, got {d2}");
}
#[test]
fn top_eig_identity_returns_unit_eigenvalue() {
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 (lam, _) = top_eig(&a, n, 50);
assert!((lam - 1.0).abs() < 1e-6);
}
#[test]
fn top_eig_recovers_dominant_direction() {
let n = 3;
let a = vec![5.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let (lam, vec) = top_eig(&a, n, 100);
assert!((lam - 5.0).abs() < 1e-6);
assert!(vec[0].abs() > 0.99);
}
#[test]
fn top_factors_finds_all_significant_and_stops() {
let n = 3;
let a = vec![
5.0, 0.0, 0.0, 0.0, 3.0, 0.0, 0.0, 0.0, 0.01,
];
let facs = top_factors(&a, n, 3);
assert_eq!(facs.len(), 2);
assert!((facs[0].0 - 5.0).abs() < 1e-3);
assert!((facs[1].0 - 3.0).abs() < 1e-3);
}
#[test]
fn solve_sym_recovers_x() {
let a = vec![4.0, 1.0, 1.0, 3.0];
let b = vec![1.0, 2.0];
let x = solve_sym(&a, &b, 2);
assert!((x[0] - 1.0 / 11.0).abs() < 1e-12);
assert!((x[1] - 7.0 / 11.0).abs() < 1e-12);
}
}