const MAX_SWEEPS: usize = 100;
const TOL: f64 = 1e-12;
pub fn eigh(matrix: &[Vec<f64>]) -> Option<(Vec<f64>, Vec<Vec<f64>>)> {
let n = matrix.len();
if n == 0 || matrix[0].len() != n {
return None;
}
let mut flat = Vec::with_capacity(n * n);
for row in matrix {
flat.extend_from_slice(row);
}
let (vals, vecs_flat) = eigh_flat(&mut flat, n)?;
let vecs: Vec<Vec<f64>> = (0..n)
.map(|k| vecs_flat[k * n..(k + 1) * n].to_vec())
.collect();
Some((vals, vecs))
}
pub(crate) fn eigh_flat(a: &mut [f64], n: usize) -> Option<(Vec<f64>, Vec<f64>)> {
if n == 0 || a.len() != n * n {
return None;
}
if n == 1 {
return Some((vec![a[0]], vec![1.0]));
}
let mut v = vec![0.0; n * n];
for i in 0..n {
v[i * n + i] = 1.0;
}
for _ in 0..MAX_SWEEPS {
let off = off_diagonal_norm_flat(a, n);
if off < TOL {
break;
}
for p in 0..n {
for q in (p + 1)..n {
let apq = a[p * n + q];
if apq.abs() < 1e-300 {
continue;
}
let app = a[p * n + p];
let aqq = a[q * n + q];
let theta = (aqq - app) / (2.0 * apq);
let t = theta.signum() / (theta.abs() + (theta * theta + 1.0).sqrt());
let c = 1.0 / (t * t + 1.0).sqrt();
let s = t * c;
rotate_flat(a, &mut v, n, p, q, c, s);
a[p * n + q] = 0.0;
a[q * n + p] = 0.0;
}
}
}
let mut idx: Vec<usize> = (0..n).collect();
idx.sort_by(|&i, &j| a[j * n + j].total_cmp(&a[i * n + i]));
let vals: Vec<f64> = idx.iter().map(|&i| a[i * n + i]).collect();
let mut vecs = vec![0.0; n * n];
for (k, &i) in idx.iter().enumerate() {
for x in 0..n {
vecs[k * n + x] = v[x * n + i];
}
}
Some((vals, vecs))
}
#[inline]
fn off_diagonal_norm_flat(a: &[f64], n: usize) -> f64 {
let mut sum = 0.0;
for i in 0..n {
let base = i * n;
for j in (i + 1)..n {
let v = a[base + j];
sum += v * v;
}
}
sum.sqrt()
}
#[allow(clippy::needless_range_loop)]
fn rotate_flat(a: &mut [f64], v: &mut [f64], n: usize, p: usize, q: usize, c: f64, s: f64) {
for r in 0..n {
if r == p || r == q {
continue;
}
let arp = a[r * n + p];
let arq = a[r * n + q];
let new_rp = c * arp - s * arq;
a[r * n + p] = new_rp;
a[p * n + r] = new_rp;
let new_rq = s * arp + c * arq;
a[r * n + q] = new_rq;
a[q * n + r] = new_rq;
}
let app = a[p * n + p];
let aqq = a[q * n + q];
let apq = a[p * n + q];
a[p * n + p] = c * c * app - 2.0 * s * c * apq + s * s * aqq;
a[q * n + q] = s * s * app + 2.0 * s * c * apq + c * c * aqq;
a[p * n + q] = 0.0;
a[q * n + p] = 0.0;
for r in 0..n {
let vrp = v[r * n + p];
let vrq = v[r * n + q];
v[r * n + p] = c * vrp - s * vrq;
v[r * n + q] = s * vrp + c * vrq;
}
}
pub fn covariance(x_centered: &[Vec<f64>], ddof: usize) -> Vec<Vec<f64>> {
crate::stats::covariance_centered(x_centered, ddof)
}
pub(crate) fn eigh_topk_flat(
matrix: &[f64],
n: usize,
k: usize,
iters: usize,
) -> Option<(Vec<f64>, Vec<f64>)> {
if n == 0 || matrix.len() != n * n {
return None;
}
let k = k.min(n);
if k == 0 {
return Some((vec![], vec![]));
}
if k >= n.saturating_sub(1) || n <= 3 {
let mut buf = matrix.to_vec();
let (vals, vecs) = eigh_flat(&mut buf, n)?;
return Some((
vals.into_iter().take(k).collect(),
vecs.into_iter().take(k * n).collect(),
));
}
let mut a = matrix.to_vec();
let mut out_vals = Vec::with_capacity(k);
let mut out_vecs = vec![0.0; k * n];
let mut v = vec![0.0; n];
let mut w = vec![0.0; n];
for j in 0..k {
v.iter_mut().enumerate().for_each(|(i, x)| {
*x = ((i as f64 * 0.5).sin() + 1.0) / n as f64;
});
let mut lambda = 0.0;
for _ in 0..iters {
sym_matvec(&a, n, &v, &mut w);
lambda = (0..n).map(|i| w[i] * v[i]).sum();
let norm = (0..n).map(|i| w[i] * w[i]).sum::<f64>().sqrt();
if norm < 1e-300 {
break;
}
let inv = 1.0 / norm;
v.iter_mut().zip(w.iter()).for_each(|(x, &y)| *x = y * inv);
}
out_vals.push(lambda);
out_vecs[j * n..(j + 1) * n].copy_from_slice(&v);
for r in 0..n {
let vr = v[r];
let base = r * n;
for c in 0..n {
a[base + c] -= lambda * vr * v[c];
}
}
}
let mut idx: Vec<usize> = (0..k).collect();
idx.sort_by(|&i, &j| out_vals[j].total_cmp(&out_vals[i]));
let mut sorted_vals = vec![0.0; k];
let mut sorted_vecs = vec![0.0; k * n];
for (new_k, &old) in idx.iter().enumerate() {
sorted_vals[new_k] = out_vals[old];
sorted_vecs[new_k * n..(new_k + 1) * n].copy_from_slice(&out_vecs[old * n..(old + 1) * n]);
}
Some((sorted_vals, sorted_vecs))
}
#[inline]
fn sym_matvec(a: &[f64], n: usize, x: &[f64], y: &mut [f64]) {
for (r, out) in y.iter_mut().enumerate().take(n) {
let base = r * n;
let mut s = 0.0;
for c in 0..n {
s += a[base + c] * x[c];
}
*out = s;
}
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn eigh_diagonal_matrix_returns_diagonal() {
let m = vec![vec![3.0, 0.0], vec![0.0, 1.0]];
let (vals, vecs) = eigh(&m).unwrap();
assert!(approx(vals[0], 3.0, 1e-9));
assert!(approx(vals[1], 1.0, 1e-9));
assert!(approx(vecs[0][0].abs(), 1.0, 1e-9));
assert!(approx(vecs[1][1].abs(), 1.0, 1e-9));
}
#[test]
fn eigh_known_2x2_matches_hand_calculation() {
let m = vec![vec![2.0, 1.0], vec![1.0, 2.0]];
let (vals, vecs) = eigh(&m).unwrap();
assert!(approx(vals[0], 3.0, 1e-9));
assert!(approx(vals[1], 1.0, 1e-9));
for v in &vecs {
let nrm: f64 = v.iter().map(|x| x * x).sum::<f64>().sqrt();
assert!(approx(nrm, 1.0, 1e-9));
}
}
#[test]
fn eigh_identity_matrix_all_ones() {
let n = 4;
let m: Vec<Vec<f64>> = (0..n)
.map(|i| (0..n).map(|j| if i == j { 1.0 } else { 0.0 }).collect())
.collect();
let (vals, _vecs) = eigh(&m).unwrap();
assert_eq!(vals.len(), n);
for v in &vals {
assert!(approx(*v, 1.0, 1e-9));
}
}
#[test]
fn eigh_eigenvectors_orthonormal() {
let m = vec![
vec![4.0, 1.0, 2.0],
vec![1.0, 3.0, 0.5],
vec![2.0, 0.5, 2.0],
];
let (vals, vecs) = eigh(&m).unwrap();
let n = vals.len();
for i in 0..n {
for j in 0..n {
let dot: f64 = (0..n).map(|k| vecs[k][i] * vecs[k][j]).sum();
let expected = if i == j { 1.0 } else { 0.0 };
assert!(approx(dot, expected, 1e-9), "VᵀV[{}][{}]={}", i, j, dot);
}
}
}
#[test]
fn eigh_reconstructs_original_matrix() {
let m = vec![
vec![4.0, 1.0, 0.5],
vec![1.0, 3.0, 1.5],
vec![0.5, 1.5, 2.0],
];
let (vals, vecs) = eigh(&m).unwrap();
let n = vals.len();
for i in 0..n {
for j in 0..n {
let a_ij: f64 = (0..n).map(|k| vals[k] * vecs[k][i] * vecs[k][j]).sum();
assert!(
approx(a_ij, m[i][j], 1e-9),
"A[{}][{}]={} want {}",
i,
j,
a_ij,
m[i][j]
);
}
}
}
#[test]
fn eigh_descending_order() {
let m = vec![
vec![1.0, 0.0, 0.0],
vec![0.0, 5.0, 0.0],
vec![0.0, 0.0, 2.0],
];
let (vals, _) = eigh(&m).unwrap();
assert!(vals[0] >= vals[1]);
assert!(vals[1] >= vals[2]);
assert!(approx(vals[0], 5.0, 1e-9));
assert!(approx(vals[2], 1.0, 1e-9));
}
#[test]
fn eigh_empty_returns_none() {
let m: Vec<Vec<f64>> = vec![];
assert!(eigh(&m).is_none());
}
#[test]
fn eigh_nonsquare_returns_none() {
let m = vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]];
assert!(eigh(&m).is_none());
}
#[test]
fn eigh_flat_single_element() {
let mut a = vec![7.0];
let (vals, vecs) = eigh_flat(&mut a, 1).unwrap();
assert!(approx(vals[0], 7.0, 1e-12));
assert!(approx(vecs[0], 1.0, 1e-12));
}
#[test]
fn covariance_centered_matches_definition() {
let x = vec![vec![1.0, 2.0], vec![-1.0, -2.0]];
let cov = covariance(&x, 0);
assert!(approx(cov[0][0], 1.0, 1e-9));
assert!(approx(cov[1][1], 4.0, 1e-9));
assert!(approx(cov[0][1], 2.0, 1e-9));
assert!(approx(cov[1][0], 2.0, 1e-9));
}
#[test]
fn eigh_topk_validates_inputs_and_handles_empty_request() {
assert!(eigh_topk_flat(&[], 0, 1, 10).is_none());
assert!(eigh_topk_flat(&[1.0, 0.0, 0.0], 2, 1, 10).is_none());
let matrix = [4.0, 0.0, 0.0, 1.0];
let (values, vectors) = eigh_topk_flat(&matrix, 2, 0, 10).unwrap();
assert!(values.is_empty());
assert!(vectors.is_empty());
}
#[test]
fn eigh_topk_matches_known_diagonal_spectrum_in_both_paths() {
let small = [4.0, 0.0, 0.0, 1.0];
let (values, vectors) = eigh_topk_flat(&small, 2, 1, 10).unwrap();
assert_eq!(values, vec![4.0]);
assert_eq!(vectors.len(), 2);
let large = [
9.0, 0.0, 0.0, 0.0, 0.0, 5.0, 0.0, 0.0, 0.0, 0.0, 2.0, 0.0, 0.0, 0.0, 0.0, 1.0,
];
let (values, vectors) = eigh_topk_flat(&large, 4, 2, 200).unwrap();
assert!(approx(values[0], 9.0, 1e-9));
assert!(approx(values[1], 5.0, 1e-9));
assert_eq!(vectors.len(), 8);
for vector in vectors.chunks_exact(4) {
let norm = vector.iter().map(|x| x * x).sum::<f64>().sqrt();
assert!(approx(norm, 1.0, 1e-9));
}
}
}