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;
}
}