pub fn eigh(a: &[f64], n: usize) -> (Vec<f64>, Vec<f64>) {
let mut am = a.to_vec();
let mut w = vec![0f64; n];
let _info = crate::blas::dsyevd(&mut am, &mut w, n);
(w, am)
}
pub fn eigh_batch(mats: &[Vec<f64>], n: usize) -> Vec<(Vec<f64>, Vec<f64>)> {
use rayon::prelude::*;
mats.par_iter().map(|a| eigh(a, n)).collect()
}
pub fn matrix_fn(a: &[f64], n: usize, f: impl Fn(f64) -> f64) -> Vec<f64> {
let (w, v) = eigh(a, n);
let fl: Vec<f64> = w.iter().map(|&l| f(l)).collect();
let mut out = vec![0f64; n * n];
for k in 0..n {
let fk = fl[k];
for i in 0..n {
let vik = fk * v[k * n + i];
for j in 0..n {
out[i * n + j] += vik * v[k * n + j];
}
}
}
out
}
pub fn logm(a: &[f64], n: usize) -> Vec<f64> {
matrix_fn(a, n, |l| l.max(1e-12).ln())
}
pub fn expm(a: &[f64], n: usize) -> Vec<f64> {
matrix_fn(a, n, |l| l.exp())
}
pub fn sqrtm(a: &[f64], n: usize) -> Vec<f64> {
matrix_fn(a, n, |l| l.max(0.0).sqrt())
}
pub fn invsqrtm(a: &[f64], n: usize) -> Vec<f64> {
matrix_fn(a, n, |l| 1.0 / l.max(1e-12).sqrt())
}
fn sqrt_invsqrt(a: &[f64], n: usize) -> (Vec<f64>, Vec<f64>) {
let (w, v) = eigh(a, n);
let sh: Vec<f64> = w.iter().map(|&l| l.max(0.0).sqrt()).collect();
let ish: Vec<f64> = w.iter().map(|&l| 1.0 / l.max(1e-12).sqrt()).collect();
(reconstruct(&sh, &v, n), reconstruct(&ish, &v, n))
}
fn matrix_fn_batch(mats: &[Vec<f64>], n: usize, f: impl Fn(f64) -> f64 + Sync) -> Vec<Vec<f64>> {
use rayon::prelude::*;
mats.par_iter().map(|a| matrix_fn(a, n, &f)).collect()
}
pub fn logm_batch(covs: &[Vec<f64>], n: usize) -> Vec<Vec<f64>> {
matrix_fn_batch(covs, n, |l| l.max(1e-12).ln())
}
pub fn expm_batch(mats: &[Vec<f64>], n: usize) -> Vec<Vec<f64>> {
matrix_fn_batch(mats, n, |l| l.exp())
}
pub fn sqrtm_batch(covs: &[Vec<f64>], n: usize) -> Vec<Vec<f64>> {
matrix_fn_batch(covs, n, |l| l.max(0.0).sqrt())
}
pub fn invsqrtm_batch(covs: &[Vec<f64>], n: usize) -> Vec<Vec<f64>> {
matrix_fn_batch(covs, n, |l| 1.0 / l.max(1e-12).sqrt())
}
fn matmul(a: &[f64], b: &[f64], n: usize) -> Vec<f64> {
let mut o = vec![0f64; n * n];
for i in 0..n {
for k in 0..n {
let aik = a[i * n + k];
if aik == 0.0 {
continue;
}
for j in 0..n {
o[i * n + j] += aik * b[k * n + j];
}
}
}
o
}
pub fn airm_dist2(a: &[f64], b: &[f64], n: usize) -> f64 {
let w = invsqrtm(a, n);
let m = matmul(&matmul(&w, b, n), &w, n);
let (evals, _) = eigh(&m, n);
evals.iter().map(|&l| l.max(1e-12).ln().powi(2)).sum()
}
pub fn karcher_mean(covs: &[Vec<f64>], n: usize, iters: usize, tol: f64) -> Vec<f64> {
let k = covs.len().max(1) as f64;
let mut m = vec![0f64; n * n];
for c in covs {
for i in 0..n * n {
m[i] += c[i] / k;
}
}
for _ in 0..iters {
let msqrt = sqrtm(&m, n);
let minv = invsqrtm(&m, n);
let mut sbar = vec![0f64; n * n];
for c in covs {
let wcw = matmul(&matmul(&minv, c, n), &minv, n);
let l = logm(&wcw, n);
for i in 0..n * n {
sbar[i] += l[i] / k;
}
}
let e = expm(&sbar, n);
m = matmul(&matmul(&msqrt, &e, n), &msqrt, n);
let norm: f64 = sbar.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm < tol {
break;
}
}
m
}
pub fn karcher_mean_weighted(
covs: &[Vec<f64>],
weights: &[f64],
n: usize,
iters: usize,
tol: f64,
) -> Vec<f64> {
assert_eq!(
covs.len(),
weights.len(),
"karcher_mean_weighted: {} covs but {} weights",
covs.len(),
weights.len()
);
let wsum: f64 = weights.iter().sum();
let wbar: Vec<f64> = if wsum.abs() < 1e-300 {
let u = 1.0 / covs.len().max(1) as f64;
vec![u; covs.len()]
} else {
weights.iter().map(|&w| w / wsum).collect()
};
let mut m = vec![0f64; n * n];
for (c, &w) in covs.iter().zip(&wbar) {
for i in 0..n * n {
m[i] += w * c[i];
}
}
for _ in 0..iters {
let msqrt = sqrtm(&m, n);
let minv = invsqrtm(&m, n);
let mut sbar = vec![0f64; n * n];
for (c, &w) in covs.iter().zip(&wbar) {
let wcw = matmul(&matmul(&minv, c, n), &minv, n);
let l = logm(&wcw, n);
for i in 0..n * n {
sbar[i] += w * l[i];
}
}
let e = expm(&sbar, n);
m = matmul(&matmul(&msqrt, &e, n), &msqrt, n);
let norm: f64 = sbar.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm < tol {
break;
}
}
m
}
fn transpose(a: &[f64], n: usize) -> Vec<f64> {
let mut t = vec![0f64; n * n];
for i in 0..n {
for j in 0..n {
t[j * n + i] = a[i * n + j];
}
}
t
}
fn symmetrize(a: &[f64], n: usize) -> Vec<f64> {
let mut o = vec![0f64; n * n];
for i in 0..n {
for j in 0..n {
o[i * n + j] = 0.5 * (a[i * n + j] + a[j * n + i]);
}
}
o
}
pub fn bimap(w: &[f64], x: &[f64], m: usize, n: usize) -> Vec<f64> {
let mut t = vec![0f64; m * n];
crate::blas::dgemm(w, x, &mut t, m, n, n);
let mut wt = vec![0f64; n * m];
for i in 0..m {
for j in 0..n {
wt[j * m + i] = w[i * n + j];
}
}
let mut y = vec![0f64; m * m];
crate::blas::dgemm(&t, &wt, &mut y, m, n, m);
y
}
pub fn reeig(x: &[f64], n: usize, eps: f64) -> Vec<f64> {
matrix_fn(x, n, |l| l.max(eps))
}
pub fn logeig(x: &[f64], n: usize, eps: f64) -> Vec<f64> {
matrix_fn(x, n, |l| l.max(eps).ln())
}
pub fn spectral_backward(
x: &[f64],
dy: &[f64],
n: usize,
f: impl Fn(f64) -> f64,
fprime: impl Fn(f64) -> f64,
) -> Vec<f64> {
let (lam, v) = eigh(x, n);
spectral_backward_precomputed(&lam, &v, dy, n, f, fprime)
}
fn mm(a: &[f64], b: &[f64], n: usize) -> Vec<f64> {
let mut o = vec![0f64; n * n];
crate::blas::dgemm(a, b, &mut o, n, n, n);
o
}
fn reconstruct(fk: &[f64], v: &[f64], n: usize) -> Vec<f64> {
let mut sv = vec![0f64; n * n];
for k in 0..n {
for i in 0..n {
sv[k * n + i] = fk[k] * v[k * n + i];
}
}
mm(&transpose(&sv, n), v, n)
}
pub fn spectral_forward_packed(x: &[f64], n: usize, f: impl Fn(f64) -> f64) -> Vec<f64> {
let (lam, v) = eigh(x, n);
let fk: Vec<f64> = lam.iter().map(|&l| f(l)).collect();
let y = reconstruct(&fk, &v, n);
let mut out = vec![0f64; 2 * n * n + n];
out[..n * n].copy_from_slice(&y);
out[n * n..n * n + n].copy_from_slice(&lam);
out[n * n + n..].copy_from_slice(&v);
out
}
pub fn spectral_backward_precomputed(
lam: &[f64],
v: &[f64],
dy: &[f64],
n: usize,
f: impl Fn(f64) -> f64,
fprime: impl Fn(f64) -> f64,
) -> Vec<f64> {
let g = symmetrize(dy, n);
let vt = transpose(v, n);
let c = mm(&mm(v, &g, n), &vt, n);
let fl: Vec<f64> = lam.iter().map(|&l| f(l)).collect();
let scale = lam.iter().fold(0f64, |acc, &l| acc.max(l.abs())).max(1.0);
let tol = scale * 1e-9;
let mut mmat = vec![0f64; n * n];
for a in 0..n {
for b in 0..n {
let d = lam[a] - lam[b];
let p = if d.abs() > tol {
(fl[a] - fl[b]) / d
} else {
fprime(lam[a])
};
mmat[a * n + b] = p * c[a * n + b];
}
}
let dx = mm(&mm(&vt, &mmat, n), v, n);
symmetrize(&dx, n)
}
pub fn reeig_backward(x: &[f64], dy: &[f64], n: usize, eps: f64) -> Vec<f64> {
let (lam, v) = eigh(x, n);
reeig_backward_precomputed(&lam, &v, dy, n, eps)
}
pub fn reeig_backward_precomputed(
lam: &[f64],
v: &[f64],
dy: &[f64],
n: usize,
eps: f64,
) -> Vec<f64> {
spectral_backward_precomputed(
lam,
v,
dy,
n,
|l| l.max(eps),
|l| {
if l > eps { 1.0 } else { 0.0 }
},
)
}
pub fn logeig_backward(x: &[f64], dy: &[f64], n: usize, eps: f64) -> Vec<f64> {
let (lam, v) = eigh(x, n);
logeig_backward_precomputed(&lam, &v, dy, n, eps)
}
pub fn logeig_backward_precomputed(
lam: &[f64],
v: &[f64],
dy: &[f64],
n: usize,
eps: f64,
) -> Vec<f64> {
spectral_backward_precomputed(lam, v, dy, n, |l| l.max(eps).ln(), |l| 1.0 / l.max(eps))
}
pub fn spd_bn_transport(
x: &[f64],
mean: &[f64],
g: &[f64],
batch: usize,
n: usize,
eps: f64,
) -> Vec<f64> {
use rayon::prelude::*;
let ms = matrix_fn(mean, n, |l| 1.0 / l.max(eps).sqrt()); let gs = matrix_fn(g, n, |l| l.max(eps).sqrt()); let slices: Vec<Vec<f64>> = (0..batch)
.into_par_iter()
.map(|bi| {
let xi = &x[bi * n * n..(bi + 1) * n * n];
let ci = mm(&mm(&ms, xi, n), &ms, n); mm(&mm(&gs, &ci, n), &gs, n) })
.collect();
slices.concat()
}
pub fn spd_bn_backward_x(
mean: &[f64],
g: &[f64],
dy: &[f64],
batch: usize,
n: usize,
eps: f64,
) -> Vec<f64> {
use rayon::prelude::*;
let ms = matrix_fn(mean, n, |l| 1.0 / l.max(eps).sqrt());
let gs = matrix_fn(g, n, |l| l.max(eps).sqrt());
let q = mm(&ms, &gs, n); let qt = transpose(&q, n); let slices: Vec<Vec<f64>> = (0..batch)
.into_par_iter()
.map(|bi| {
let gi = &dy[bi * n * n..(bi + 1) * n * n];
symmetrize(&mm(&mm(&q, gi, n), &qt, n), n) })
.collect();
slices.concat()
}
pub fn spd_bn_backward_g(
x: &[f64],
mean: &[f64],
g: &[f64],
dy: &[f64],
batch: usize,
n: usize,
eps: f64,
) -> Vec<f64> {
use rayon::prelude::*;
let ms = matrix_fn(mean, n, |l| 1.0 / l.max(eps).sqrt());
let gs = matrix_fn(g, n, |l| l.max(eps).sqrt());
let dgs = (0..batch)
.into_par_iter()
.map(|bi| {
let xi = &x[bi * n * n..(bi + 1) * n * n];
let gi = &dy[bi * n * n..(bi + 1) * n * n];
let ci = mm(&mm(&ms, xi, n), &ms, n); let t1 = mm(&mm(gi, &gs, n), &ci, n); let t2 = mm(&mm(&ci, &gs, n), gi, n); let mut s = t1;
for k in 0..n * n {
s[k] += t2[k];
}
s
})
.reduce(
|| vec![0f64; n * n],
|mut a, b| {
for k in 0..n * n {
a[k] += b[k];
}
a
},
);
spectral_backward(
g,
&dgs,
n,
|l| l.max(eps).sqrt(),
|l| 0.5 / l.max(eps).sqrt(),
)
}
pub fn geodesic_interp(a: &[f64], b: &[f64], t: f64, n: usize) -> Vec<f64> {
let (asqrt, ainv) = sqrt_invsqrt(a, n);
let m = matmul(&matmul(&ainv, b, n), &ainv, n);
let mt = matrix_fn(&m, n, |l| l.max(1e-12).powf(t));
matmul(&matmul(&asqrt, &mt, n), &asqrt, n)
}
pub fn log_map(base: &[f64], x: &[f64], n: usize) -> Vec<f64> {
let (phalf, pinv) = sqrt_invsqrt(base, n);
let m = matmul(&matmul(&pinv, x, n), &pinv, n); let lm = logm(&m, n);
matmul(&matmul(&phalf, &lm, n), &phalf, n)
}
pub fn exp_map(base: &[f64], v: &[f64], n: usize) -> Vec<f64> {
let (phalf, pinv) = sqrt_invsqrt(base, n);
let m = matmul(&matmul(&pinv, v, n), &pinv, n); let em = expm(&m, n);
matmul(&matmul(&phalf, &em, n), &phalf, n)
}
pub fn parallel_transport(from: &[f64], to: &[f64], v: &[f64], n: usize) -> Vec<f64> {
let (phalf, pinv) = sqrt_invsqrt(from, n);
let wq = matmul(&matmul(&pinv, to, n), &pinv, n); let wq_half = sqrtm(&wq, n);
let e = matmul(&matmul(&phalf, &wq_half, n), &pinv, n); let et = transpose(&e, n);
let evet = matmul(&matmul(&e, v, n), &et, n);
symmetrize(&evet, n) }
fn add(a: &[f64], b: &[f64]) -> Vec<f64> {
a.iter().zip(b).map(|(x, y)| x + y).collect()
}
fn sqrt_ff(l: f64) -> f64 {
l.max(0.0).sqrt()
}
fn sqrt_fp(l: f64) -> f64 {
0.5 / l.max(1e-12).sqrt()
}
fn invsqrt_ff(l: f64) -> f64 {
1.0 / l.max(1e-12).sqrt()
}
fn invsqrt_fp(l: f64) -> f64 {
-0.5 * l.max(1e-12).powf(-1.5)
}
fn log_ff(l: f64) -> f64 {
l.max(1e-12).ln()
}
fn log_fp(l: f64) -> f64 {
1.0 / l.max(1e-12)
}
fn reconstruct_fn(lam: &[f64], v: &[f64], n: usize, f: impl Fn(f64) -> f64) -> Vec<f64> {
let fk: Vec<f64> = lam.iter().map(|&l| f(l)).collect();
reconstruct(&fk, v, n)
}
pub fn log_map_backward(base: &[f64], x: &[f64], dy: &[f64], n: usize) -> (Vec<f64>, Vec<f64>) {
let (lam_b, v_b) = eigh(base, n);
let a = reconstruct_fn(&lam_b, &v_b, n, sqrt_ff); let w = reconstruct_fn(&lam_b, &v_b, n, invsqrt_ff); let m = matmul(&matmul(&w, x, n), &w, n); let (lam_m, v_m) = eigh(&m, n);
let g = symmetrize(dy, n);
let lbar = matmul(&matmul(&a, &g, n), &a, n); let mbar = spectral_backward_precomputed(&lam_m, &v_m, &lbar, n, log_ff, log_fp);
let d_x = symmetrize(&matmul(&matmul(&w, &mbar, n), &w, n), n); let l = reconstruct_fn(&lam_m, &v_m, n, log_ff); let gal = matmul(&matmul(&g, &a, n), &l, n); let abar = add(&gal, &transpose(&gal, n)); let mwx = matmul(&matmul(&mbar, &w, n), x, n); let wbar = add(&mwx, &transpose(&mwx, n)); let d_base_a = spectral_backward_precomputed(&lam_b, &v_b, &abar, n, sqrt_ff, sqrt_fp);
let d_base_w = spectral_backward_precomputed(&lam_b, &v_b, &wbar, n, invsqrt_ff, invsqrt_fp);
let d_base = symmetrize(&add(&d_base_a, &d_base_w), n);
(d_base, d_x)
}
pub fn exp_map_backward(base: &[f64], v: &[f64], dy: &[f64], n: usize) -> (Vec<f64>, Vec<f64>) {
let (lam_b, v_b) = eigh(base, n);
let a = reconstruct_fn(&lam_b, &v_b, n, sqrt_ff);
let w = reconstruct_fn(&lam_b, &v_b, n, invsqrt_ff);
let m = matmul(&matmul(&w, v, n), &w, n); let (lam_m, v_m) = eigh(&m, n);
let g = symmetrize(dy, n);
let ebar = matmul(&matmul(&a, &g, n), &a, n); let mbar = spectral_backward_precomputed(&lam_m, &v_m, &ebar, n, f64::exp, f64::exp);
let d_v = symmetrize(&matmul(&matmul(&w, &mbar, n), &w, n), n); let e = reconstruct_fn(&lam_m, &v_m, n, f64::exp); let gae = matmul(&matmul(&g, &a, n), &e, n); let abar = add(&gae, &transpose(&gae, n));
let mwv = matmul(&matmul(&mbar, &w, n), v, n); let wbar = add(&mwv, &transpose(&mwv, n));
let d_base_a = spectral_backward_precomputed(&lam_b, &v_b, &abar, n, sqrt_ff, sqrt_fp);
let d_base_w = spectral_backward_precomputed(&lam_b, &v_b, &wbar, n, invsqrt_ff, invsqrt_fp);
let d_base = symmetrize(&add(&d_base_a, &d_base_w), n);
(d_base, d_v)
}
pub fn parallel_transport_backward(
from: &[f64],
to: &[f64],
v: &[f64],
dy: &[f64],
n: usize,
) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let (lam_f, v_f) = eigh(from, n);
let a = reconstruct_fn(&lam_f, &v_f, n, sqrt_ff); let w = reconstruct_fn(&lam_f, &v_f, n, invsqrt_ff); let mq = matmul(&matmul(&w, to, n), &w, n); let (lam_mq, v_mq) = eigh(&mq, n);
let s = reconstruct_fn(&lam_mq, &v_mq, n, sqrt_ff); let e = matmul(&matmul(&a, &s, n), &w, n); let et = transpose(&e, n);
let g = symmetrize(dy, n);
let d_v = symmetrize(&matmul(&matmul(&et, &g, n), &e, n), n);
let gev = matmul(&matmul(&g, &e, n), v, n);
let ebar: Vec<f64> = gev.iter().map(|x| 2.0 * x).collect();
let abar = matmul(&matmul(&ebar, &w, n), &s, n); let sbar = matmul(&matmul(&a, &ebar, n), &w, n); let wbar1 = matmul(&matmul(&s, &a, n), &ebar, n); let mqbar = spectral_backward_precomputed(&lam_mq, &v_mq, &sbar, n, sqrt_ff, sqrt_fp);
let d_to = symmetrize(&matmul(&matmul(&w, &mqbar, n), &w, n), n); let mqwq = matmul(&matmul(&mqbar, &w, n), to, n); let wbar = add(&wbar1, &add(&mqwq, &transpose(&mqwq, n))); let d_from_a = spectral_backward_precomputed(&lam_f, &v_f, &abar, n, sqrt_ff, sqrt_fp);
let d_from_w = spectral_backward_precomputed(&lam_f, &v_f, &wbar, n, invsqrt_ff, invsqrt_fp);
let d_from = symmetrize(&add(&d_from_a, &d_from_w), n);
(d_from, d_to, d_v)
}
pub fn matrix_fn_batch_backward(
x: &[f64],
dy: &[f64],
n: usize,
f: impl Fn(f64) -> f64 + Sync,
fprime: impl Fn(f64) -> f64 + Sync,
) -> Vec<f64> {
use rayon::prelude::*;
let batch = x.len() / (n * n);
(0..batch)
.into_par_iter()
.flat_map(|bi| {
let xi = &x[bi * n * n..(bi + 1) * n * n];
let dyi = &dy[bi * n * n..(bi + 1) * n * n];
spectral_backward(xi, dyi, n, &f, &fprime)
})
.collect()
}
pub fn eigh_packed(a: &[f64], n: usize) -> Vec<f64> {
let (w, vrow) = eigh(a, n); let u = transpose(&vrow, n); let mut out = vec![0f64; n + n * n];
out[..n].copy_from_slice(&w);
out[n..].copy_from_slice(&u);
out
}
pub fn eigh_backward_packed(fwd: &[f64], bar: &[f64], n: usize) -> Vec<f64> {
let lam = &fwd[..n];
let u = &fwd[n..];
let lbar = &bar[..n];
let ubar = &bar[n..];
let ut = transpose(u, n);
let c = mm(&ut, ubar, n); let scale = lam.iter().fold(0f64, |a, &l| a.max(l.abs())).max(1.0);
let tol = scale * 1e-9;
let mut mid = vec![0f64; n * n];
for i in 0..n {
for j in 0..n {
mid[i * n + j] = if i == j {
lbar[i]
} else {
let d = lam[j] - lam[i];
if d.abs() > tol { c[i * n + j] / d } else { 0.0 }
};
}
}
let abar = mm(&mm(u, &mid, n), &ut, n); symmetrize(&abar, n)
}
pub fn eigh_batch_packed(mats: &[Vec<f64>], n: usize) -> Vec<Vec<f64>> {
use rayon::prelude::*;
mats.par_iter().map(|a| eigh_packed(a, n)).collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn ident(n: usize) -> Vec<f64> {
let mut a = vec![0f64; n * n];
for i in 0..n {
a[i * n + i] = 1.0;
}
a
}
#[test]
fn logm_expm_roundtrip() {
let n = 4;
let mut a = vec![0f64; n * n];
for i in 0..n {
for j in 0..n {
a[i * n + j] = if i == j { 2.0 + i as f64 } else { 0.3 };
}
}
let back = expm(&logm(&a, n), n);
let err: f64 = a
.iter()
.zip(&back)
.map(|(x, y)| (x - y).abs())
.fold(0.0, f64::max);
assert!(err < 1e-6, "logm/expm roundtrip err {err}");
}
#[test]
fn sqrtm_squares_back() {
let n = 3;
let a = vec![4.0, 0.0, 0.0, 0.0, 9.0, 0.0, 0.0, 0.0, 16.0];
let s = sqrtm(&a, n);
assert!(
(s[0] - 2.0).abs() < 1e-9 && (s[4] - 3.0).abs() < 1e-9 && (s[8] - 4.0).abs() < 1e-9
);
}
#[test]
fn airm_dist_zero_to_self_and_invariant() {
let n = 3;
let a = vec![2.0, 0.1, 0.0, 0.1, 3.0, 0.2, 0.0, 0.2, 1.5];
assert!(airm_dist2(&a, &a, n) < 1e-9, "distance to self must be 0");
let k = 5.0;
let ka: Vec<f64> = a.iter().map(|x| x * k).collect();
let ki: Vec<f64> = ident(n).iter().map(|x| x * k).collect();
let d1 = airm_dist2(&a, &ident(n), n);
let d2 = airm_dist2(&ka, &ki, n);
assert!(
(d1 - d2).abs() < 1e-6,
"AIRM not affine-invariant: {d1} vs {d2}"
);
}
#[test]
fn karcher_mean_of_identicals_is_the_matrix() {
let n = 3;
let a = vec![2.0, 0.1, 0.0, 0.1, 3.0, 0.2, 0.0, 0.2, 1.5];
let m = karcher_mean(&[a.clone(), a.clone(), a.clone()], n, 20, 1e-10);
let err: f64 = a
.iter()
.zip(&m)
.map(|(x, y)| (x - y).abs())
.fold(0.0, f64::max);
assert!(err < 1e-6, "Karcher mean of identicals err {err}");
}
fn sym(n: usize, seed: f64) -> Vec<f64> {
let mut a = vec![0f64; n * n];
for i in 0..n {
for j in i..n {
let v = ((i as f64 * 3.0 + j as f64 * 1.7 + seed).sin()) * 0.5;
a[i * n + j] = v;
a[j * n + i] = v;
}
}
a
}
fn spd(n: usize, seed: f64) -> Vec<f64> {
let m = sym(n, seed);
let mut a = matmul(&m, &m, n); for i in 0..n {
a[i * n + i] += (n + 1) as f64;
}
a
}
fn frob(a: &[f64], b: &[f64]) -> f64 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
#[test]
fn bimap_matches_manual() {
let (m, n) = (2, 3);
let w = vec![1.0, 0.5, -0.25, 2.0, -1.0, 0.75];
let x = spd(n, 0.3);
let y = bimap(&w, &x, m, n);
let mut wx = vec![0f64; m * n];
for i in 0..m {
for j in 0..n {
let mut s = 0.0;
for k in 0..n {
s += w[i * n + k] * x[k * n + j];
}
wx[i * n + j] = s;
}
}
let mut expect = vec![0f64; m * m];
for i in 0..m {
for j in 0..m {
let mut s = 0.0;
for k in 0..n {
s += wx[i * n + k] * w[j * n + k];
}
expect[i * m + j] = s;
}
}
let err = y
.iter()
.zip(&expect)
.map(|(a, b)| (a - b).abs())
.fold(0.0, f64::max);
assert!(err < 1e-10, "bimap err {err}");
}
#[test]
fn reeig_floors_spectrum() {
let n = 4;
let mut x = vec![0f64; n * n];
for i in 0..n {
x[i * n + i] = [0.01, 0.5, 2.0, 5.0][i];
}
let eps = 0.1;
let y = reeig(&x, n, eps);
let (evals, _) = eigh(&y, n);
for &l in &evals {
assert!(l >= eps - 1e-9, "reeig eigenvalue {l} below eps {eps}");
}
assert!((y[0] - eps).abs() < 1e-9, "smallest not floored: {}", y[0]);
assert!((y[n + 1] - 0.5).abs() < 1e-9);
}
#[test]
fn logeig_expm_roundtrip() {
let n = 4;
let x = spd(n, 1.1);
let back = expm(&logeig(&x, n, 1e-12), n);
let err = x
.iter()
.zip(&back)
.map(|(a, b)| (a - b).abs())
.fold(0.0, f64::max);
assert!(err < 1e-8, "expm(logeig) roundtrip err {err}");
}
fn fd_spectral_check(
layer: impl Fn(&[f64], usize) -> Vec<f64>,
backward: impl Fn(&[f64], &[f64], usize) -> Vec<f64>,
n: usize,
seed: f64,
) {
let x = spd(n, seed);
let c = sym(n, seed + 10.0); let v = sym(n, seed + 20.0); let grad = backward(&x, &c, n);
let analytic = frob(&grad, &v);
let h = 1e-6;
let xp: Vec<f64> = x.iter().zip(&v).map(|(a, b)| a + h * b).collect();
let xm: Vec<f64> = x.iter().zip(&v).map(|(a, b)| a - h * b).collect();
let fd = (frob(&c, &layer(&xp, n)) - frob(&c, &layer(&xm, n))) / (2.0 * h);
let err = (analytic - fd).abs() / (fd.abs() + 1e-6);
assert!(
err < 1e-4,
"spectral backward FD mismatch: analytic {analytic}, fd {fd}, rel {err}"
);
}
#[test]
fn logeig_backward_matches_fd() {
fd_spectral_check(
|x, n| logeig(x, n, 1e-12),
|x, dy, n| logeig_backward(x, dy, n, 1e-12),
5,
0.7,
);
}
#[test]
fn reeig_backward_matches_fd() {
let eps = 3.0;
fd_spectral_check(
move |x, n| reeig(x, n, eps),
move |x, dy, n| reeig_backward(x, dy, n, eps),
5,
0.42,
);
}
#[test]
fn spd_bn_transport_normalizes() {
let n = 3;
let batch = 4;
let mut x = vec![0f64; batch * n * n];
let mut slices = Vec::new();
for bi in 0..batch {
let s = spd(n, bi as f64 * 0.9 + 0.2);
x[bi * n * n..(bi + 1) * n * n].copy_from_slice(&s);
slices.push(s);
}
let mean = karcher_mean(&slices, n, 50, 1e-12);
let mut g = vec![0f64; n * n];
for i in 0..n {
g[i * n + i] = 1.0;
}
let y = spd_bn_transport(&x, &mean, &g, batch, n, 1e-12);
let yslices: Vec<Vec<f64>> = (0..batch)
.map(|bi| y[bi * n * n..(bi + 1) * n * n].to_vec())
.collect();
let ymean = karcher_mean(&yslices, n, 50, 1e-12);
let mut ident = vec![0f64; n * n];
for i in 0..n {
ident[i * n + i] = 1.0;
}
let err = ymean
.iter()
.zip(&ident)
.map(|(a, b)| (a - b).abs())
.fold(0.0, f64::max);
assert!(err < 1e-4, "SPD-BN did not normalize to I: err {err}");
}
#[test]
fn spd_bn_backward_x_matches_fd() {
let n = 3;
let batch = 2;
let mean = spd(n, 5.0);
let g = spd(n, 6.0);
let mut x = vec![0f64; batch * n * n];
for bi in 0..batch {
let s = spd(n, bi as f64 + 0.3);
x[bi * n * n..(bi + 1) * n * n].copy_from_slice(&s);
}
let mut c = vec![0f64; batch * n * n];
for bi in 0..batch {
let s = sym(n, bi as f64 + 7.0);
c[bi * n * n..(bi + 1) * n * n].copy_from_slice(&s);
}
let eps = 1e-12;
let grad = spd_bn_backward_x(&mean, &g, &c, batch, n, eps);
let mut v = vec![0f64; batch * n * n];
for bi in 0..batch {
let s = sym(n, bi as f64 + 8.0);
v[bi * n * n..(bi + 1) * n * n].copy_from_slice(&s);
}
let analytic = frob(&grad, &v);
let h = 1e-6;
let xp: Vec<f64> = x.iter().zip(&v).map(|(a, b)| a + h * b).collect();
let xm: Vec<f64> = x.iter().zip(&v).map(|(a, b)| a - h * b).collect();
let fd = (frob(&c, &spd_bn_transport(&xp, &mean, &g, batch, n, eps))
- frob(&c, &spd_bn_transport(&xm, &mean, &g, batch, n, eps)))
/ (2.0 * h);
let err = (analytic - fd).abs() / (fd.abs() + 1e-6);
assert!(err < 1e-4, "SPD-BN dX FD mismatch: {analytic} vs {fd}");
}
#[test]
fn spd_bn_backward_g_matches_fd() {
let n = 3;
let batch = 2;
let mean = spd(n, 5.0);
let g = spd(n, 6.0);
let mut x = vec![0f64; batch * n * n];
for bi in 0..batch {
let s = spd(n, bi as f64 + 0.3);
x[bi * n * n..(bi + 1) * n * n].copy_from_slice(&s);
}
let mut c = vec![0f64; batch * n * n];
for bi in 0..batch {
let s = sym(n, bi as f64 + 7.0);
c[bi * n * n..(bi + 1) * n * n].copy_from_slice(&s);
}
let eps = 1e-12;
let grad = spd_bn_backward_g(&x, &mean, &g, &c, batch, n, eps);
let vg = sym(n, 9.0); let analytic = frob(&grad, &vg);
let h = 1e-6;
let gp: Vec<f64> = g.iter().zip(&vg).map(|(a, b)| a + h * b).collect();
let gm: Vec<f64> = g.iter().zip(&vg).map(|(a, b)| a - h * b).collect();
let fd = (frob(&c, &spd_bn_transport(&x, &mean, &gp, batch, n, eps))
- frob(&c, &spd_bn_transport(&x, &mean, &gm, batch, n, eps)))
/ (2.0 * h);
let err = (analytic - fd).abs() / (fd.abs() + 1e-6);
assert!(err < 1e-4, "SPD-BN dG FD mismatch: {analytic} vs {fd}");
}
#[test]
fn spectral_backward_handles_degenerate_eigenvalues() {
let n = 3;
let vv = [1.0, 0.5, -0.3];
let mut x = vec![0f64; n * n];
for i in 0..n {
for j in 0..n {
let d = if i == j { 2.0 } else { 0.0 };
x[i * n + j] = d + vv[i] * vv[j];
}
}
let c = sym(n, 3.0);
let dir = sym(n, 4.0);
let grad = logeig_backward(&x, &c, n, 1e-12);
let analytic = frob(&grad, &dir);
let h = 1e-6;
let xp: Vec<f64> = x.iter().zip(&dir).map(|(a, b)| a + h * b).collect();
let xm: Vec<f64> = x.iter().zip(&dir).map(|(a, b)| a - h * b).collect();
let fd = (frob(&c, &logeig(&xp, n, 1e-12)) - frob(&c, &logeig(&xm, n, 1e-12))) / (2.0 * h);
let err = (analytic - fd).abs() / (fd.abs() + 1e-6);
assert!(
err < 1e-4,
"degenerate spectral backward FD: analytic {analytic} vs fd {fd}"
);
}
#[test]
fn geodesic_interp_endpoints_and_midpoint() {
let n = 3;
let a = spd(n, 1.0);
let b = spd(n, 2.0);
let maxdiff = |p: &[f64], q: &[f64]| {
p.iter()
.zip(q)
.map(|(x, y)| (x - y).abs())
.fold(0.0, f64::max)
};
assert!(maxdiff(&geodesic_interp(&a, &b, 0.0, n), &a) < 1e-8);
assert!(maxdiff(&geodesic_interp(&a, &b, 1.0, n), &b) < 1e-8);
let mid = geodesic_interp(&a, &b, 0.5, n);
let (ev, _) = eigh(&mid, n);
assert!(
ev.iter().all(|&l| l > 1e-9),
"geodesic midpoint not SPD: {ev:?}"
);
}
fn maxerr(a: &[f64], b: &[f64]) -> f64 {
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs())
.fold(0.0, f64::max)
}
fn asymmetry(a: &[f64], n: usize) -> f64 {
let mut m = 0.0f64;
for i in 0..n {
for j in 0..n {
m = m.max((a[i * n + j] - a[j * n + i]).abs());
}
}
m
}
fn airm_inner(p: &[f64], u: &[f64], v: &[f64], n: usize) -> f64 {
let pinv = matrix_fn(p, n, |l| 1.0 / l.max(1e-12));
let m = matmul(&matmul(&pinv, u, n), &pinv, n); let mut tr = 0.0;
for i in 0..n {
for j in 0..n {
tr += m[i * n + j] * v[j * n + i];
}
}
tr
}
#[test]
fn karcher_mean_weighted_uniform_matches_unweighted() {
let n = 3;
let covs = vec![spd(n, 0.2), spd(n, 1.3), spd(n, 2.7), spd(n, 3.1)];
let unw = karcher_mean(&covs, n, 60, 1e-13);
let wtd = karcher_mean_weighted(&covs, &[1.0; 4], n, 60, 1e-13);
assert!(maxerr(&unw, &wtd) < 1e-9, "uniform weighted != unweighted");
let scaled = karcher_mean_weighted(&covs, &[7.5; 4], n, 60, 1e-13);
assert!(maxerr(&wtd, &scaled) < 1e-12, "weight scale changed result");
}
#[test]
fn karcher_mean_weighted_dominant_weight_selects_matrix() {
let n = 3;
let covs = vec![spd(n, 0.5), spd(n, 1.5), spd(n, 2.5)];
let m = karcher_mean_weighted(&covs, &[1e-9, 1.0, 1e-9], n, 80, 1e-14);
assert!(
maxerr(&covs[1], &m) < 1e-6,
"dominant-weight barycentre should ≈ that cov"
);
}
#[test]
fn karcher_mean_weighted_is_stationary() {
let n = 4;
let covs = vec![
spd(n, 0.3),
spd(n, 1.1),
spd(n, 2.2),
spd(n, 3.9),
spd(n, 5.0),
];
let w = [0.4, 0.1, 0.25, 0.05, 0.2];
let m = karcher_mean_weighted(&covs, &w, n, 200, 1e-15);
let wsum: f64 = w.iter().sum();
let mut g = vec![0f64; n * n];
for (c, &wi) in covs.iter().zip(&w) {
let l = log_map(&m, c, n);
for k in 0..n * n {
g[k] += (wi / wsum) * l[k];
}
}
let gnorm: f64 = g.iter().map(|x| x * x).sum::<f64>().sqrt();
assert!(
gnorm < 1e-6,
"weighted barycentre not stationary: |g|={gnorm}"
);
}
#[test]
fn log_exp_map_roundtrip_and_identity_base() {
let n = 4;
let base = spd(n, 0.7);
let x = spd(n, 2.4);
let v = log_map(&base, &x, n);
let back = exp_map(&base, &v, n);
assert!(maxerr(&x, &back) < 1e-8, "exp∘log roundtrip");
assert!(asymmetry(&v, n) < 1e-10, "log_map output not symmetric");
let i = ident(n);
assert!(
maxerr(&log_map(&i, &x, n), &logm(&x, n)) < 1e-9,
"log_map(I, ·) != logm"
);
let s = sym(n, 3.0);
assert!(
maxerr(&exp_map(&i, &s, n), &expm(&s, n)) < 1e-9,
"exp_map(I, ·) != expm"
);
}
#[test]
fn exp_log_maps_reproduce_geodesic() {
let n = 3;
let a = spd(n, 1.0);
let b = spd(n, 2.6);
let v = log_map(&a, &b, n);
for &t in &[0.0, 0.25, 0.5, 0.9, 1.0] {
let tv: Vec<f64> = v.iter().map(|x| x * t).collect();
let via_exp = exp_map(&a, &tv, n);
let gi = geodesic_interp(&a, &b, t, n);
assert!(
maxerr(&via_exp, &gi) < 1e-8,
"exp_map(t·log) != geodesic_interp at t={t}"
);
}
}
#[test]
fn parallel_transport_is_airm_isometry() {
let n = 4;
let p = spd(n, 0.9);
let q = spd(n, 3.3);
let v = sym(n, 1.2);
let w = sym(n, 2.4);
let gv = parallel_transport(&p, &q, &v, n);
let gw = parallel_transport(&p, &q, &w, n);
let before = airm_inner(&p, &v, &w, n);
let after = airm_inner(&q, &gv, &gw, n);
let rel = (before - after).abs() / (before.abs() + 1e-9);
assert!(
rel < 1e-7,
"transport not an isometry: ⟨v,w⟩_P={before} vs ⟨Γv,Γw⟩_Q={after}"
);
assert!(asymmetry(&gv, n) < 1e-9, "transported vector not symmetric");
let same = parallel_transport(&p, &p, &v, n);
assert!(maxerr(&same, &v) < 1e-8, "transport P→P must be identity");
}
#[test]
fn batched_matrix_fns_match_scalar() {
let n = 3;
let covs: Vec<Vec<f64>> = (0..6).map(|i| spd(n, i as f64 * 0.6 + 0.1)).collect();
let syms: Vec<Vec<f64>> = (0..6).map(|i| sym(n, i as f64 * 0.4 + 0.2)).collect();
let lb = logm_batch(&covs, n);
let sb = sqrtm_batch(&covs, n);
let ib = invsqrtm_batch(&covs, n);
for (i, c) in covs.iter().enumerate() {
assert!(maxerr(&lb[i], &logm(c, n)) < 1e-12, "logm_batch[{i}]");
assert!(maxerr(&sb[i], &sqrtm(c, n)) < 1e-12, "sqrtm_batch[{i}]");
assert!(
maxerr(&ib[i], &invsqrtm(c, n)) < 1e-12,
"invsqrtm_batch[{i}]"
);
}
let eb = expm_batch(&syms, n);
for (i, s) in syms.iter().enumerate() {
assert!(maxerr(&eb[i], &expm(s, n)) < 1e-12, "expm_batch[{i}]");
}
assert!(logm_batch(&[], n).is_empty());
}
fn fd_dir_check(
grad: &[f64],
dir: &[f64],
c: &[f64],
forward: impl Fn(&[f64]) -> Vec<f64>,
at: &[f64],
name: &str,
) {
let h = 1e-6;
let xp: Vec<f64> = at.iter().zip(dir).map(|(a, b)| a + h * b).collect();
let xm: Vec<f64> = at.iter().zip(dir).map(|(a, b)| a - h * b).collect();
let fd = (frob(c, &forward(&xp)) - frob(c, &forward(&xm))) / (2.0 * h);
let an = frob(grad, dir);
assert!(
(an - fd).abs() / (fd.abs() + 1e-6) < 1e-4,
"{name}: analytic {an} vs fd {fd}"
);
}
#[test]
fn log_map_backward_matches_fd() {
let n = 3;
let base = spd(n, 0.7);
let x = spd(n, 2.1);
let c = sym(n, 3.0); let (d_base, d_x) = log_map_backward(&base, &x, &c, n);
fd_dir_check(
&d_x,
&sym(n, 4.0),
&c,
|x| log_map(&base, x, n),
&x,
"log_map d_x",
);
fd_dir_check(
&d_base,
&sym(n, 5.0),
&c,
|b| log_map(b, &x, n),
&base,
"log_map d_base",
);
}
#[test]
fn exp_map_backward_matches_fd() {
let n = 3;
let base = spd(n, 0.8);
let v = sym(n, 1.4); let c = sym(n, 3.2);
let (d_base, d_v) = exp_map_backward(&base, &v, &c, n);
fd_dir_check(
&d_v,
&sym(n, 4.1),
&c,
|v| exp_map(&base, v, n),
&v,
"exp_map d_v",
);
fd_dir_check(
&d_base,
&sym(n, 5.1),
&c,
|b| exp_map(b, &v, n),
&base,
"exp_map d_base",
);
}
#[test]
fn parallel_transport_backward_matches_fd() {
let n = 3;
let from = spd(n, 0.6);
let to = spd(n, 1.9);
let v = sym(n, 2.5);
let c = sym(n, 3.3);
let (d_from, d_to, d_v) = parallel_transport_backward(&from, &to, &v, &c, n);
fd_dir_check(
&d_from,
&sym(n, 4.0),
&c,
|p| parallel_transport(p, &to, &v, n),
&from,
"transport d_from",
);
fd_dir_check(
&d_to,
&sym(n, 5.0),
&c,
|q| parallel_transport(&from, q, &v, n),
&to,
"transport d_to",
);
fd_dir_check(
&d_v,
&sym(n, 6.0),
&c,
|vv| parallel_transport(&from, &to, vv, n),
&v,
"transport d_v",
);
}
#[test]
fn matrix_fn_batch_backward_matches_fd() {
let n = 3;
let batch = 3;
let mut x = vec![0.0; batch * n * n];
let mut c = vec![0.0; batch * n * n];
for bi in 0..batch {
x[bi * n * n..(bi + 1) * n * n].copy_from_slice(&spd(n, bi as f64 * 0.5 + 0.2));
c[bi * n * n..(bi + 1) * n * n].copy_from_slice(&sym(n, bi as f64 + 3.0));
}
let grad =
matrix_fn_batch_backward(&x, &c, n, |l| l.max(1e-12).ln(), |l| 1.0 / l.max(1e-12));
let mut dir = vec![0.0; batch * n * n];
for bi in 0..batch {
dir[bi * n * n..(bi + 1) * n * n].copy_from_slice(&sym(n, bi as f64 + 7.0));
}
let logm_flat = |xx: &[f64]| {
let covs: Vec<Vec<f64>> = (0..batch)
.map(|bi| xx[bi * n * n..(bi + 1) * n * n].to_vec())
.collect();
logm_batch(&covs, n).concat()
};
fd_dir_check(
&grad,
&dir,
&c,
logm_flat,
&x,
"matrix_fn_batch(logm) backward",
);
}
fn spd_sep(n: usize) -> Vec<f64> {
let mut a = vec![0f64; n * n];
for i in 0..n {
for j in i..n {
let v = if i == j {
(i as f64 + 1.0) * 2.0
} else {
0.05 * ((i + j) as f64).cos()
};
a[i * n + j] = v;
a[j * n + i] = v;
}
}
a
}
#[test]
fn eigh_packed_reconstructs() {
let n = 4;
let a = spd_sep(n);
let fwd = eigh_packed(&a, n);
let (lam, u) = (&fwd[..n], &fwd[n..]);
let mut recon = vec![0f64; n * n];
for i in 0..n {
for j in 0..n {
let mut s = 0.0;
for k in 0..n {
s += u[i * n + k] * lam[k] * u[j * n + k];
}
recon[i * n + j] = s;
}
}
assert!(maxerr(&recon, &a) < 1e-9, "eigh_packed A = U Λ Uᵀ failed");
assert!(lam.windows(2).all(|w| w[0] <= w[1] + 1e-12));
}
#[test]
fn eigh_backward_matches_fd() {
let n = 4;
let a = spd_sep(n);
let fwd = eigh_packed(&a, n);
let u_base = fwd[n..].to_vec();
let lbar: Vec<f64> = (0..n).map(|i| ((i as f64 + 1.0) * 0.7).sin()).collect();
let ubar = sym(n, 5.0);
let mut bar = vec![0f64; n + n * n];
bar[..n].copy_from_slice(&lbar);
bar[n..].copy_from_slice(&ubar);
let abar = eigh_backward_packed(&fwd, &bar, n);
let loss = |mat: &[f64]| -> f64 {
let fw = eigh_packed(mat, n);
let (l, mut uu) = (fw[..n].to_vec(), fw[n..].to_vec());
for j in 0..n {
let mut dot = 0.0;
for i in 0..n {
dot += uu[i * n + j] * u_base[i * n + j];
}
if dot < 0.0 {
for i in 0..n {
uu[i * n + j] = -uu[i * n + j];
}
}
}
let mut s = 0.0;
for i in 0..n {
s += lbar[i] * l[i];
}
for k in 0..n * n {
s += ubar[k] * uu[k];
}
s
};
let dir = sym(n, 4.0);
let h = 1e-6;
let ap: Vec<f64> = a.iter().zip(&dir).map(|(x, y)| x + h * y).collect();
let am: Vec<f64> = a.iter().zip(&dir).map(|(x, y)| x - h * y).collect();
let fd = (loss(&ap) - loss(&am)) / (2.0 * h);
let an = frob(&abar, &dir);
assert!(
(an - fd).abs() / (fd.abs() + 1e-6) < 1e-4,
"eigh backward FD: analytic {an} vs fd {fd}"
);
}
#[test]
fn eigh_batch_packed_matches_scalar() {
let n = 3;
let mats: Vec<Vec<f64>> = (0..4).map(|i| spd(n, i as f64 * 0.4 + 0.2)).collect();
let batched = eigh_batch_packed(&mats, n);
for (i, m) in mats.iter().enumerate() {
assert!(
maxerr(&batched[i], &eigh_packed(m, n)) < 1e-12,
"eigh_batch[{i}]"
);
}
}
}