use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use opt::{BacktrackConfig, armijo_roundoff_cushion, backtracking_line_search};
use crate::manifold::{
GeometryError, GeometryResult, RiemannianManifold, check_len, cholesky_spd, dot, flatten,
from_flat, inverse, jacobi_symmetric, spectral_map_spd, spectral_map_symmetric, sym,
tangent_basis_metric_orthonormal,
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SpdManifold {
n: usize,
}
impl SpdManifold {
const SYM_REL_TOL: f64 = 1.0e-9;
pub const fn new(n: usize) -> Self {
Self { n }
}
fn matrix(&self, point: ArrayView1<'_, f64>) -> GeometryResult<Array2<f64>> {
let raw = from_flat(point, self.n, self.n)?;
let mut max_abs = 0.0_f64;
let mut max_asym = 0.0_f64;
for i in 0..self.n {
for j in 0..self.n {
max_abs = max_abs.max(raw[[i, j]].abs());
max_asym = max_asym.max((raw[[i, j]] - raw[[j, i]]).abs());
}
}
if !max_asym.is_finite() || max_asym > Self::SYM_REL_TOL * max_abs.max(1.0) {
return Err(GeometryError::InvalidPoint(
"SPD point must be a symmetric matrix",
));
}
let p = sym(&raw);
cholesky_spd(&p)?;
Ok(p)
}
fn affine_inner(
&self,
p: &Array2<f64>,
u: &Array2<f64>,
v: &Array2<f64>,
) -> GeometryResult<f64> {
use gam_linalg::faer_ndarray::fast_ab;
let pinv = inverse(p)?;
let a = fast_ab(&fast_ab(&fast_ab(&pinv, u), &pinv), v);
let mut trace = 0.0;
for i in 0..self.n {
trace += a[[i, i]];
}
Ok(trace)
}
}
impl RiemannianManifold for SpdManifold {
fn dim(&self) -> usize {
self.n * (self.n + 1) / 2
}
fn ambient_dim(&self) -> usize {
self.n * self.n
}
fn tangent_basis(&self, point: ArrayView1<'_, f64>) -> GeometryResult<Array2<f64>> {
check_len("SPD point", point.len(), self.ambient_dim())?;
tangent_basis_metric_orthonormal(self, point, self.n, self.n)
}
fn exp_map(
&self,
point: ArrayView1<'_, f64>,
tangent_vec: ArrayView1<'_, f64>,
) -> GeometryResult<Array1<f64>> {
use gam_linalg::faer_ndarray::fast_ab;
let p = self.matrix(point)?;
let u = sym(&from_flat(tangent_vec, self.n, self.n)?);
let sqrt_p = spectral_map_spd(&p, |x| Ok(x.sqrt()))?;
let inv_sqrt_p = spectral_map_spd(&p, |x| Ok(1.0 / x.sqrt()))?;
let middle = fast_ab(&fast_ab(&inv_sqrt_p, &u), &inv_sqrt_p);
let exp_middle = spectral_map_symmetric(&middle, |x| Ok(x.exp()))?;
Ok(flatten(&sym(&fast_ab(
&fast_ab(&sqrt_p, &exp_middle),
&sqrt_p,
))))
}
fn log_map(
&self,
p_from: ArrayView1<'_, f64>,
p_to: ArrayView1<'_, f64>,
) -> GeometryResult<Array1<f64>> {
use gam_linalg::faer_ndarray::fast_ab;
let p = self.matrix(p_from)?;
let q = self.matrix(p_to)?;
let sqrt_p = spectral_map_spd(&p, |x| Ok(x.sqrt()))?;
let inv_sqrt_p = spectral_map_spd(&p, |x| Ok(1.0 / x.sqrt()))?;
let middle = fast_ab(&fast_ab(&inv_sqrt_p, &q), &inv_sqrt_p);
let log_middle = spectral_map_spd(&middle, |x| Ok(x.ln()))?;
Ok(flatten(&sym(&fast_ab(
&fast_ab(&sqrt_p, &log_middle),
&sqrt_p,
))))
}
fn parallel_transport(
&self,
point_along: ArrayView2<'_, f64>,
vec: ArrayView1<'_, f64>,
) -> GeometryResult<Array1<f64>> {
check_len("SPD transported vector", vec.len(), self.ambient_dim())?;
if point_along.nrows() < 2 {
return Ok(flatten(&sym(&from_flat(vec, self.n, self.n)?)));
}
let p = self.matrix(point_along.row(0))?;
let q = self.matrix(point_along.row(point_along.nrows() - 1))?;
use gam_linalg::faer_ndarray::{fast_ab, fast_abt};
let u = sym(&from_flat(vec, self.n, self.n)?);
let inv_sqrt_p = spectral_map_spd(&p, |x| Ok(1.0 / x.sqrt()))?;
let middle = fast_ab(&fast_ab(&inv_sqrt_p, &q), &inv_sqrt_p);
let e = spectral_map_spd(&middle, |x| Ok(x.sqrt()))?;
let sqrt_p = spectral_map_spd(&p, |x| Ok(x.sqrt()))?;
let a = fast_ab(&fast_ab(&sqrt_p, &e), &inv_sqrt_p);
Ok(flatten(&sym(&fast_abt(&fast_ab(&a, &u), &a))))
}
fn metric_tensor(&self, point: ArrayView1<'_, f64>) -> GeometryResult<Array2<f64>> {
let p = self.matrix(point)?;
let pinv = inverse(&p)?;
let ambient = self.ambient_dim();
let mut g = Array2::<f64>::zeros((ambient, ambient));
for i in 0..self.n {
for j in 0..self.n {
for k in 0..self.n {
for l in 0..self.n {
g[[i * self.n + j, k * self.n + l]] = pinv[[i, k]] * pinv[[l, j]];
}
}
}
}
Ok(g)
}
fn christoffel_symbols(&self, point: ArrayView1<'_, f64>) -> GeometryResult<Vec<Array2<f64>>> {
let p = self.matrix(point)?;
let pinv = inverse(&p)?;
let ambient = self.ambient_dim();
let mut gamma = (0..ambient)
.map(|_| Array2::<f64>::zeros((ambient, ambient)))
.collect::<Vec<_>>();
for a in 0..ambient {
let ai = a / self.n;
let aj = a % self.n;
for b in 0..ambient {
let bi = b / self.n;
let bj = b % self.n;
let mut u = Array2::<f64>::zeros((self.n, self.n));
let mut v = Array2::<f64>::zeros((self.n, self.n));
u[[ai, aj]] = 1.0;
v[[bi, bj]] = 1.0;
let c = -0.5 * (u.dot(&pinv).dot(&v) + v.dot(&pinv).dot(&u));
for r in 0..self.n {
for s in 0..self.n {
gamma[r * self.n + s][[a, b]] = c[[r, s]];
}
}
}
}
Ok(gamma)
}
fn sectional_curvature(
&self,
point: ArrayView1<'_, f64>,
tangent_pair: (ArrayView1<'_, f64>, ArrayView1<'_, f64>),
) -> GeometryResult<f64> {
let p = self.matrix(point)?;
let u = sym(&from_flat(tangent_pair.0, self.n, self.n)?);
let v = sym(&from_flat(tangent_pair.1, self.n, self.n)?);
use gam_linalg::faer_ndarray::fast_ab;
let inv_sqrt_p = spectral_map_spd(&p, |x| Ok(1.0 / x.sqrt()))?;
let a = fast_ab(&fast_ab(&inv_sqrt_p, &u), &inv_sqrt_p);
let b = fast_ab(&fast_ab(&inv_sqrt_p, &v), &inv_sqrt_p);
let comm = &fast_ab(&a, &b) - &fast_ab(&b, &a);
let comm_norm = dot(flatten(&comm).view(), flatten(&comm).view());
let uu = self.affine_inner(&p, &u, &u)?;
let vv = self.affine_inner(&p, &v, &v)?;
let uv = self.affine_inner(&p, &u, &v)?;
let denom = uu * vv - uv * uv;
if denom.abs() <= 1.0e-14 {
return Err(GeometryError::Singular(
"SPD sectional curvature plane is degenerate",
));
}
Ok(-0.25 * comm_norm / denom)
}
fn project_tangent(
&self,
point: ArrayView1<'_, f64>,
vec: ArrayView1<'_, f64>,
) -> GeometryResult<Array1<f64>> {
check_len("SPD projection point", point.len(), self.ambient_dim())?;
Ok(flatten(&sym(&from_flat(vec, self.n, self.n)?)))
}
fn riemannian_gradient(
&self,
point: ArrayView1<'_, f64>,
euclidean_grad: ArrayView1<'_, f64>,
) -> GeometryResult<Array1<f64>> {
use gam_linalg::faer_ndarray::fast_ab;
let p = self.matrix(point)?;
let e = sym(&from_flat(euclidean_grad, self.n, self.n)?);
let grad = fast_ab(&fast_ab(&p, &e), &p);
Ok(flatten(&sym(&grad)))
}
fn exp_map_vjp(
&self,
point: ArrayView1<'_, f64>,
tangent_vec: ArrayView1<'_, f64>,
grad_output: ArrayView1<'_, f64>,
) -> GeometryResult<(Array1<f64>, Array1<f64>)> {
use gam_linalg::faer_ndarray::fast_ab;
let m = self.ambient_dim();
check_len("SPD exp_map_vjp point", point.len(), m)?;
check_len("SPD exp_map_vjp tangent", tangent_vec.len(), m)?;
check_len("SPD exp_map_vjp grad", grad_output.len(), m)?;
let p = self.matrix(point)?;
let u = sym(&from_flat(tangent_vec, self.n, self.n)?);
let (p_evals, p_vecs) = jacobi_symmetric(&p)?;
for &lam in p_evals.iter() {
if !(lam.is_finite() && lam > 0.0) {
return Err(GeometryError::InvalidPoint(
"SPD eigenvalue is not positive",
));
}
}
let sqrt_p = spectral_reconstruct(&p_vecs, &p_evals, f64::sqrt);
let inv_sqrt_p = spectral_reconstruct(&p_vecs, &p_evals, |x| 1.0 / x.sqrt());
let middle = sym(&fast_ab(&fast_ab(&inv_sqrt_p, &u), &inv_sqrt_p));
let (m_evals, m_vecs) = jacobi_symmetric(&middle)?;
let exp_middle = spectral_reconstruct(&m_vecs, &m_evals, f64::exp);
let g_y = sym(&from_flat(grad_output, self.n, self.n)?);
let g_e = fast_ab(&fast_ab(&sqrt_p, &g_y), &sqrt_p);
let g_s = &fast_ab(&fast_ab(&g_y, &sqrt_p), &exp_middle)
+ &fast_ab(&fast_ab(&exp_middle, &sqrt_p), &g_y);
let g_m = daleckii_krein_pullback(&m_vecs, &m_evals, exp_divided_difference, &sym(&g_e));
let g_u = fast_ab(&fast_ab(&inv_sqrt_p, &g_m), &inv_sqrt_p);
let g_s_inv =
&fast_ab(&fast_ab(&g_m, &inv_sqrt_p), &u) + &fast_ab(&fast_ab(&u, &inv_sqrt_p), &g_m);
let g_p = &daleckii_krein_pullback(&p_vecs, &p_evals, sqrt_divided_difference, &sym(&g_s))
+ &daleckii_krein_pullback(
&p_vecs,
&p_evals,
inv_sqrt_divided_difference,
&sym(&g_s_inv),
);
Ok((flatten(&sym(&g_p)), flatten(&sym(&g_u))))
}
}
fn spectral_reconstruct(
vecs: &Array2<f64>,
evals: &Array1<f64>,
f: impl Fn(f64) -> f64,
) -> Array2<f64> {
use gam_linalg::faer_ndarray::{fast_ab, fast_abt};
let n = evals.len();
let mut diag = Array2::<f64>::zeros((n, n));
for i in 0..n {
diag[[i, i]] = f(evals[i]);
}
fast_abt(&fast_ab(vecs, &diag), vecs)
}
fn daleckii_krein_pullback(
vecs: &Array2<f64>,
evals: &Array1<f64>,
divided_difference: impl Fn(f64, f64) -> f64,
c: &Array2<f64>,
) -> Array2<f64> {
use gam_linalg::faer_ndarray::{fast_ab, fast_abt, fast_atb};
let n = evals.len();
let mut inner = fast_ab(&fast_atb(vecs, c), vecs);
for i in 0..n {
for j in 0..n {
inner[[i, j]] *= divided_difference(evals[i], evals[j]);
}
}
fast_abt(&fast_ab(vecs, &inner), vecs)
}
fn exp_divided_difference(a: f64, b: f64) -> f64 {
if a == b {
return a.exp();
}
let hi = a.max(b);
let gap = (a - b).abs();
hi.exp() * (-(-gap).exp_m1() / gap)
}
fn sqrt_divided_difference(a: f64, b: f64) -> f64 {
1.0 / (a.sqrt() + b.sqrt())
}
fn inv_sqrt_divided_difference(a: f64, b: f64) -> f64 {
let (sa, sb) = (a.sqrt(), b.sqrt());
let (lo, hi) = if sa <= sb { (sa, sb) } else { (sb, sa) };
-((1.0 / hi) / (hi + lo)) / lo
}
fn affine_sq_norm(
n: usize,
inv_sqrt_p: &Array2<f64>,
v: ArrayView1<'_, f64>,
) -> GeometryResult<f64> {
use gam_linalg::faer_ndarray::fast_ab;
let vm = sym(&from_flat(v, n, n)?);
let whitened = sym(&fast_ab(&fast_ab(inv_sqrt_p, &vm), inv_sqrt_p));
let mut squared_norm = 0.0_f64;
for &value in &whitened {
if !value.is_finite() {
return Err(GeometryError::Singular(
"SPD affine metric norm is non-finite",
));
}
squared_norm += value * value;
}
if !squared_norm.is_finite() {
return Err(GeometryError::Singular("SPD affine metric norm overflowed"));
}
Ok(squared_norm)
}
pub fn spd_frechet_mean(
n: usize,
points: ArrayView2<'_, f64>,
weights: Option<ArrayView1<'_, f64>>,
tol: f64,
max_iter: usize,
) -> GeometryResult<Array1<f64>> {
let ambient = n * n;
let (m, cols) = points.dim();
if m == 0 || cols != ambient {
return Err(GeometryError::InvalidPoint(
"SPD Fréchet mean: points must be M×n² with M ≥ 1",
));
}
if !(tol.is_finite() && tol > 0.0) {
return Err(GeometryError::InvalidPoint(
"SPD Fréchet mean tolerance must be finite and positive",
));
}
let spd = SpdManifold::new(n);
let w = crate::normalize_weights(m, weights)
.map_err(|_| GeometryError::InvalidPoint("SPD Fréchet mean: invalid weights"))?;
let samples: Vec<Array1<f64>> = (0..m).map(|i| points.row(i).to_owned()).collect();
let dispersion = |p: ArrayView1<'_, f64>| -> GeometryResult<f64> {
let pm = spd.matrix(p)?;
let inv_sqrt_p = spectral_map_spd(&pm, |x| Ok(1.0 / x.sqrt()))?;
let mut acc = 0.0_f64;
for (i, x) in samples.iter().enumerate() {
let lg = spd.log_map(p, x.view())?;
acc += w[i] * affine_sq_norm(n, &inv_sqrt_p, lg.view())?;
}
Ok(acc)
};
let mut p = Array1::<f64>::zeros(ambient);
for (i, x) in samples.iter().enumerate() {
p.scaled_add(w[i], x);
}
p = flatten(&sym(&from_flat(p.view(), n, n)?));
let mut f_cur = dispersion(p.view())?;
const ARMIJO_C1: f64 = opt::constants::ARMIJO_C1;
let stationarity = |point: ArrayView1<'_, f64>| -> GeometryResult<(Array1<f64>, f64)> {
let pm = spd.matrix(point)?;
let inv_sqrt_p = spectral_map_spd(&pm, |x| Ok(1.0 / x.sqrt()))?;
let mut xi = Array1::<f64>::zeros(ambient);
for (i, x) in samples.iter().enumerate() {
let lg = spd.log_map(point, x.view())?;
xi.scaled_add(w[i], &lg);
}
let residual = affine_sq_norm(n, &inv_sqrt_p, xi.view())?.sqrt();
Ok((xi, residual))
};
for iteration in 0..max_iter {
let (xi, grad_norm) = stationarity(p.view())?;
if grad_norm <= tol {
return Ok(p);
}
let pred = grad_norm * grad_norm; let f_tol = armijo_roundoff_cushion(f_cur);
let accepted = backtracking_line_search(
BacktrackConfig::default(),
|t| -> GeometryResult<Option<(f64, Array1<f64>)>> {
let step = &xi * t;
let cand = match spd.exp_map(p.view(), step.view()) {
Ok(candidate) => candidate,
Err(GeometryError::InvalidPoint(_) | GeometryError::Singular(_)) => {
return Ok(None);
}
Err(error) => return Err(error),
};
let f_cand = match dispersion(cand.view()) {
Ok(value) => value,
Err(GeometryError::InvalidPoint(_) | GeometryError::Singular(_)) => {
return Ok(None);
}
Err(error) => return Err(error),
};
Ok(Some((f_cand, cand)))
},
|t, f_cand| f_cand <= f_cur - 2.0 * ARMIJO_C1 * t * pred + f_tol,
)?;
match accepted {
Some(step) => {
f_cur = step.value;
p = step.payload;
}
None => {
return Err(GeometryError::NonConvergence {
context: "SPD Fréchet mean",
iterations: iteration + 1,
residual: grad_norm,
tolerance: tol,
});
}
}
}
let (_, residual) = stationarity(p.view())?;
if residual <= tol {
Ok(p)
} else {
Err(GeometryError::NonConvergence {
context: "SPD Fréchet mean",
iterations: max_iter,
residual,
tolerance: tol,
})
}
}
#[cfg(test)]
mod tangent_basis_tests {
use super::SpdManifold;
use crate::manifold::RiemannianManifold;
use ndarray::Array1;
#[test]
fn spd_riemannian_gradient_is_affine_metric_riesz_representative() {
let spd = SpdManifold::new(2);
let p = Array1::from(vec![2.0, 0.0, 0.0, 1.0]);
let differential = Array1::from(vec![1.0, 0.0, 0.0, 1.0]);
let tangent = Array1::from(vec![0.7, 0.2, 0.2, -0.3]);
let gradient = spd
.riemannian_gradient(p.view(), differential.view())
.expect("affine-invariant metric raise");
let metric = spd.metric_tensor(p.view()).expect("SPD metric tensor");
let lhs = gradient.dot(&metric.dot(&tangent));
let rhs = differential.dot(&tangent);
assert!(
(lhs - rhs).abs() <= 1.0e-12,
"Riesz identity failed: g_P(grad, xi)={lhs}, <E, xi>={rhs}"
);
let expected = Array1::from(vec![4.0, 0.0, 0.0, 1.0]);
for (got, want) in gradient.iter().zip(expected.iter()) {
assert!((got - want).abs() <= 1.0e-12);
}
let projected = spd
.project_tangent(p.view(), differential.view())
.expect("Euclidean tangent projection");
assert!(
(&gradient - &projected).dot(&(&gradient - &projected)) > 1.0,
"affine metric raise unexpectedly equals Euclidean projection"
);
}
#[test]
fn spd_tangent_basis_metric_orthonormal() {
let spd = SpdManifold::new(2);
let p = Array1::from(vec![2.0, 0.5, 0.5, 1.0]);
let q = spd.tangent_basis(p.view()).expect("tangent basis");
let w = spd.metric_tensor(p.view()).expect("metric tensor");
let d = spd.dim();
assert_eq!(q.ncols(), d, "basis must have dim() columns");
let wq = w.dot(&q);
let gram = q.t().dot(&wq);
for i in 0..d {
for j in 0..d {
let want = if i == j { 1.0 } else { 0.0 };
assert!(
(gram[[i, j]] - want).abs() <= 1.0e-10,
"QᵀWQ != I at ({i},{j}): got {}",
gram[[i, j]]
);
}
}
}
}
#[cfg(test)]
mod exp_map_vjp_tests {
use super::{SpdManifold, exp_divided_difference};
use crate::manifold::RiemannianManifold;
use ndarray::{Array1, Array2};
fn flat(m: &Array2<f64>) -> Array1<f64> {
Array1::from_iter(m.iter().copied())
}
fn spd_with_eigs(d: [f64; 3], theta: f64, phi: f64) -> Array2<f64> {
let (c1, s1) = (theta.cos(), theta.sin());
let (c2, s2) = (phi.cos(), phi.sin());
let g1 =
Array2::from_shape_vec((3, 3), vec![c1, -s1, 0.0, s1, c1, 0.0, 0.0, 0.0, 1.0]).unwrap();
let g2 =
Array2::from_shape_vec((3, 3), vec![1.0, 0.0, 0.0, 0.0, c2, -s2, 0.0, s2, c2]).unwrap();
let r = g1.dot(&g2);
let mut dm = Array2::<f64>::zeros((3, 3));
for i in 0..3 {
dm[[i, i]] = d[i];
}
r.dot(&dm).dot(&r.t())
}
fn assert_vjp_matches_fd(p: &Array2<f64>, t: &Array2<f64>, g: &Array2<f64>) {
let spd = SpdManifold::new(3);
let (pf, tf, gf) = (flat(p), flat(t), flat(g));
let (grad_p, grad_t) = spd
.exp_map_vjp(pf.view(), tf.view(), gf.view())
.expect("SPD exp_map_vjp");
let scalar = |pv: &Array1<f64>, tv: &Array1<f64>| -> f64 {
let y = spd.exp_map(pv.view(), tv.view()).expect("exp_map");
y.dot(&gf)
};
let eps = 1.0e-6;
for i in 0..3 {
for j in i..3 {
let mut h = Array2::<f64>::zeros((3, 3));
h[[i, j]] = 1.0;
h[[j, i]] = 1.0;
let hf = flat(&h);
let fd = (scalar(&(&pf + &(&hf * eps)), &tf) - scalar(&(&pf - &(&hf * eps)), &tf))
/ (2.0 * eps);
let analytic = grad_p.dot(&hf);
assert!(
(fd - analytic).abs() <= 1.0e-5 * (1.0 + fd.abs()),
"grad_point mismatch along sym e({i},{j}): fd {fd:.9e} vs vjp {analytic:.9e}"
);
}
}
for idx in 0..9 {
let mut hf = Array1::<f64>::zeros(9);
hf[idx] = 1.0;
let fd = (scalar(&pf, &(&tf + &(&hf * eps))) - scalar(&pf, &(&tf - &(&hf * eps))))
/ (2.0 * eps);
let analytic = grad_t.dot(&hf);
assert!(
(fd - analytic).abs() <= 1.0e-5 * (1.0 + fd.abs()),
"grad_tangent mismatch along e{idx}: fd {fd:.9e} vs vjp {analytic:.9e}"
);
}
}
#[test]
fn spd_exp_map_vjp_matches_fd_generic_spectrum() {
let p = spd_with_eigs([3.0, 1.2, 0.4], 0.7, 1.1);
let t =
Array2::from_shape_vec((3, 3), vec![0.3, -0.2, 0.5, 0.1, -0.4, 0.2, -0.3, 0.6, 0.1])
.unwrap();
let g =
Array2::from_shape_vec((3, 3), vec![1.0, 0.4, -0.3, 0.2, -0.8, 0.5, 0.7, -0.1, 0.9])
.unwrap();
assert_vjp_matches_fd(&p, &t, &g);
}
#[test]
fn spd_exp_map_vjp_matches_fd_clustered_point_spectrum() {
let p = spd_with_eigs([2.0, 2.0, 0.5], 0.9, 0.3);
let t =
Array2::from_shape_vec((3, 3), vec![0.2, 0.1, -0.3, 0.1, -0.1, 0.4, -0.3, 0.4, 0.3])
.unwrap();
let g =
Array2::from_shape_vec((3, 3), vec![0.5, -0.6, 0.2, -0.6, 0.3, 0.8, 0.2, 0.8, -0.4])
.unwrap();
assert_vjp_matches_fd(&p, &t, &g);
}
#[test]
fn spd_exp_map_vjp_matches_fd_degenerate_exp_spectrum() {
let p = spd_with_eigs([1.5, 0.8, 2.5], 0.4, 1.3);
let t = &p * 0.35;
let g =
Array2::from_shape_vec((3, 3), vec![0.9, 0.1, -0.2, 0.1, -0.5, 0.3, -0.2, 0.3, 0.6])
.unwrap();
assert_vjp_matches_fd(&p, &t, &g);
}
#[test]
fn spd_exp_map_vjp_zero_tangent_reduces_to_identity_pullback() {
let spd = SpdManifold::new(3);
let p = spd_with_eigs([2.0, 1.0, 0.5], 0.2, 0.8);
let g = Array2::from_shape_vec((3, 3), vec![1.0, 0.3, 0.0, 0.3, -0.7, 0.2, 0.0, 0.2, 0.4])
.unwrap();
let zeros = Array1::<f64>::zeros(9);
let (grad_p, grad_t) = spd
.exp_map_vjp(flat(&p).view(), zeros.view(), flat(&g).view())
.expect("VJP at zero tangent");
let gs = crate::manifold::sym(&g);
for (a, b) in grad_p.iter().zip(gs.iter()) {
assert!(
(a - b).abs() <= 1.0e-12,
"grad_point at T=0 must be sym(G): {a} vs {b}"
);
}
for (a, b) in grad_t.iter().zip(gs.iter()) {
assert!(
(a - b).abs() <= 1.0e-12,
"grad_tangent at T=0 must be sym(G): {a} vs {b}"
);
}
}
#[test]
fn exp_divided_difference_stays_finite_across_underflow_range() {
let got = exp_divided_difference(-1.0, -1500.0);
let expected = (-1.0_f64).exp() / 1499.0;
assert!(got.is_finite());
assert!((got - expected).abs() <= f64::EPSILON * expected);
}
}
#[cfg(test)]
mod frechet_mean_tests {
use super::{SpdManifold, affine_sq_norm, spd_frechet_mean};
use crate::manifold::{GeometryError, RiemannianManifold, spectral_map_spd};
use ndarray::{Array1, Array2};
fn diag_flat(d: &[f64]) -> Array1<f64> {
let n = d.len();
let mut m = Array2::<f64>::zeros((n, n));
for i in 0..n {
m[[i, i]] = d[i];
}
Array1::from_iter(m.iter().copied())
}
fn stack(rows: &[Array1<f64>]) -> Array2<f64> {
let m = rows.len();
let k = rows[0].len();
let mut s = Array2::<f64>::zeros((m, k));
for (i, r) in rows.iter().enumerate() {
for (j, &v) in r.iter().enumerate() {
s[[i, j]] = v;
}
}
s
}
fn residual(spd: &SpdManifold, p: &Array1<f64>, rows: &[Array1<f64>], w: &[f64]) -> f64 {
let k = p.len();
let mut xi = Array1::<f64>::zeros(k);
for (x, &wi) in rows.iter().zip(w) {
xi.scaled_add(wi, &spd.log_map(p.view(), x.view()).expect("log_map"));
}
let pm = spd.matrix(p.view()).expect("SPD mean");
let inv_sqrt_p =
spectral_map_spd(&pm, |value| Ok(1.0 / value.sqrt())).expect("inverse square root");
affine_sq_norm(spd.n, &inv_sqrt_p, xi.view())
.expect("affine norm")
.sqrt()
}
#[test]
fn spd_frechet_mean_matches_geometric_mean_on_commuting_extreme_magnitudes() {
let n = 3;
let diags = [
[1e6, 1e-6, 1.0],
[1e-6, 1.0, 1e6],
[1.0, 1e6, 1e-6],
[1e2, 1e-2, 1e2],
];
let rows: Vec<Array1<f64>> = diags.iter().map(|d| diag_flat(d)).collect();
let m = rows.len();
let mut want = [0.0_f64; 3];
for k in 0..n {
let mut s = 0.0;
for d in &diags {
s += d[k].ln();
}
want[k] = (s / m as f64).exp();
}
let p = spd_frechet_mean(n, stack(&rows).view(), None, 1e-12, 500)
.expect("frechet mean converges on commuting extreme-magnitude SPD");
let spd = SpdManifold::new(n);
for i in 0..n {
for j in 0..n {
let got = p[i * n + j];
let exp = if i == j { want[i] } else { 0.0 };
let scale = exp.abs().max(1.0);
assert!(
(got - exp).abs() <= 1e-7 * scale,
"commuting mean[{i},{j}] = {got:.6e}, want {exp:.6e}"
);
}
}
let w = vec![1.0 / m as f64; m];
let r = residual(&spd, &p, &rows, &w);
assert!(r < 1e-9, "commuting case residual {r:.3e} not at floor");
}
#[test]
fn spd_frechet_mean_weighted_matches_weighted_geometric_mean() {
let n = 2;
let diags = [[4.0, 0.25], [0.5, 16.0], [9.0, 1.0]];
let raw_w = [0.5, 0.3, 0.2];
let rows: Vec<Array1<f64>> = diags.iter().map(|d| diag_flat(d)).collect();
let mut want = [0.0_f64; 2];
for k in 0..n {
let mut s = 0.0;
for (d, &wi) in diags.iter().zip(&raw_w) {
s += wi * d[k].ln();
}
want[k] = s.exp();
}
let wv = Array1::from(raw_w.to_vec());
let p = spd_frechet_mean(n, stack(&rows).view(), Some(wv.view()), 1e-12, 500)
.expect("weighted frechet mean converges");
for k in 0..n {
let got = p[k * n + k];
assert!(
(got - want[k]).abs() <= 1e-9 * want[k].max(1.0),
"weighted mean diag[{k}] = {got:.9e}, want {want_k:.9e}",
want_k = want[k]
);
}
}
#[test]
fn spd_frechet_mean_converges_below_sqrt_eps_on_spread_non_commuting() {
let n = 2;
let angles = [0.0_f64, 0.6, 1.2, 1.9, 2.7];
let eig = [
(12.0_f64, 0.4_f64),
(0.5, 9.0),
(3.0, 0.2),
(0.3, 6.0),
(5.0, 0.7),
];
let mut rows: Vec<Array1<f64>> = Vec::new();
for (&th, &(a, b)) in angles.iter().zip(&eig) {
let (c, s) = (th.cos(), th.sin());
let m00 = c * c * a + s * s * b;
let m01 = c * s * (a - b);
let m11 = s * s * a + c * c * b;
rows.push(Array1::from(vec![m00, m01, m01, m11]));
}
let m = rows.len();
let tol = 1e-9;
let p = spd_frechet_mean(n, stack(&rows).view(), None, tol, 1000)
.expect("spread non-commuting frechet mean reaches its certificate");
let spd = SpdManifold::new(n);
let w = vec![1.0 / m as f64; m];
let r = residual(&spd, &p, &rows, &w);
assert!(
r <= tol,
"spread non-commuting residual {r:.3e} exceeds requested tolerance {tol:.3e}"
);
let disp = |q: &Array1<f64>| -> f64 {
rows.iter()
.map(|x| {
let lg = spd.log_map(q.view(), x.view()).expect("log_map");
let g = spd.metric_tensor(q.view()).expect("metric");
lg.dot(&g.dot(&lg)) / m as f64
})
.sum()
};
let v_mean = disp(&p);
for x in &rows {
assert!(
v_mean < disp(x),
"mean does not minimize dispersion: V(mean)={v_mean:.6e}"
);
}
}
#[test]
fn spd_frechet_mean_budget_shortfall_is_typed_non_convergence() {
let n = 2;
let rows = [
diag_flat(&[4.0, 0.25]),
Array1::from(vec![1.0, 0.5, 0.5, 3.0]),
diag_flat(&[0.3, 6.0]),
];
match spd_frechet_mean(n, stack(&rows).view(), None, 1e-14, 1) {
Err(GeometryError::NonConvergence {
context,
iterations,
residual,
tolerance,
}) => {
assert_eq!(context, "SPD Fréchet mean");
assert_eq!(iterations, 1);
assert!(residual.is_finite() && residual > tolerance);
assert_eq!(tolerance, 1e-14);
}
other => panic!("expected typed SPD Fréchet exhaustion, got {other:?}"),
}
}
#[test]
fn spd_frechet_mean_is_equivariant_at_uniformly_tiny_scale() {
let n = 2;
let rows = [
diag_flat(&[4.0, 0.25]),
Array1::from(vec![1.0, 0.4, 0.4, 2.5]),
diag_flat(&[0.6, 3.0]),
];
let unit_mean =
spd_frechet_mean(n, stack(&rows).view(), None, 1.0e-11, 500).expect("unit-scale mean");
let scale = 1.0e-16;
let tiny_rows: Vec<Array1<f64>> = rows.iter().map(|row| row * scale).collect();
let tiny_mean = spd_frechet_mean(n, stack(&tiny_rows).view(), None, 1.0e-11, 500)
.expect("uniformly tiny SPD data remain valid");
for (&tiny, &unit) in tiny_mean.iter().zip(&unit_mean) {
let expected = scale * unit;
assert!(
(tiny - expected).abs() <= 2.0e-10 * expected.abs().max(scale),
"scale equivariance failed: tiny mean {tiny:.6e}, expected {expected:.6e}"
);
}
let weights = vec![1.0 / tiny_rows.len() as f64; tiny_rows.len()];
let achieved = residual(&SpdManifold::new(n), &tiny_mean, &tiny_rows, &weights);
assert!(achieved <= 1.0e-11, "tiny-scale residual {achieved:.3e}");
}
#[test]
fn affine_stationarity_norm_rejects_non_finite_tangents() {
let inv_sqrt_p = Array2::eye(2);
let tangent = Array1::from(vec![f64::NAN, 0.0, 0.0, 1.0]);
assert!(affine_sq_norm(2, &inv_sqrt_p, tangent.view()).is_err());
}
}
#[cfg(test)]
mod parallel_transport_tests {
use super::SpdManifold;
use crate::manifold::{RiemannianManifold, from_flat, sym};
use ndarray::{Array1, Array2};
fn rotated_diag(theta: f64, a: f64, b: f64) -> Array1<f64> {
let (c, s) = (theta.cos(), theta.sin());
let m00 = c * c * a + s * s * b;
let m01 = c * s * (a - b);
let m11 = s * s * a + c * c * b;
Array1::from(vec![m00, m01, m01, m11])
}
fn fixture() -> (SpdManifold, Array1<f64>, Array1<f64>) {
let spd = SpdManifold::new(2);
let p = rotated_diag(0.3, 3.0, 0.5);
let q = rotated_diag(-0.5, 1.2, 4.0);
(spd, p, q)
}
fn path2(a: &Array1<f64>, b: &Array1<f64>) -> Array2<f64> {
let mut m = Array2::<f64>::zeros((2, a.len()));
for (col, &x) in a.iter().enumerate() {
m[[0, col]] = x;
}
for (col, &x) in b.iter().enumerate() {
m[[1, col]] = x;
}
m
}
#[test]
fn parallel_transport_preserves_affine_inner_product() {
let (spd, p, q) = fixture();
let path = path2(&p, &q);
let u = Array1::from(vec![1.0, 0.4, 0.4, -0.7]);
let v = Array1::from(vec![-0.3, 0.9, 0.9, 1.6]);
let tu = spd.parallel_transport(path.view(), u.view()).expect("Γ(U)");
let tv = spd.parallel_transport(path.view(), v.view()).expect("Γ(V)");
let pm = spd.matrix(p.view()).expect("P");
let qm = spd.matrix(q.view()).expect("Q");
let um = sym(&from_flat(u.view(), 2, 2).expect("U"));
let vm = sym(&from_flat(v.view(), 2, 2).expect("V"));
let tum = sym(&from_flat(tu.view(), 2, 2).expect("ΓU"));
let tvm = sym(&from_flat(tv.view(), 2, 2).expect("ΓV"));
let before = spd.affine_inner(&pm, &um, &vm).expect("⟨U,V⟩_P");
let after = spd.affine_inner(&qm, &tum, &tvm).expect("⟨ΓU,ΓV⟩_Q");
assert!(
(before - after).abs() <= 1e-10 * before.abs().max(1.0),
"parallel transport is not an isometry: ⟨U,V⟩_P={before:.12e}, ⟨ΓU,ΓV⟩_Q={after:.12e}"
);
}
#[test]
fn parallel_transport_matches_geodesic_velocity_identity() {
let (spd, p, q) = fixture();
let forward = path2(&p, &q);
let v_p_to_q = spd.log_map(p.view(), q.view()).expect("log_P(Q)");
let v_q_to_p = spd.log_map(q.view(), p.view()).expect("log_Q(P)");
let transported = spd
.parallel_transport(forward.view(), v_p_to_q.view())
.expect("Γ(log_P Q)");
for (i, (&t, &v)) in transported.iter().zip(v_q_to_p.iter()).enumerate() {
assert!(
(t + v).abs() <= 1e-9 * v.abs().max(1.0),
"component {i}: Γ(log_P Q)={t:.12e}, −log_Q P={:.12e}",
-v
);
}
}
#[test]
fn parallel_transport_round_trip_is_identity() {
let (spd, p, q) = fixture();
let forward = path2(&p, &q);
let backward = path2(&q, &p);
let u = Array1::from(vec![0.6, -0.2, -0.2, 1.1]);
let out = spd
.parallel_transport(forward.view(), u.view())
.expect("Γ_{P→Q}(U)");
let back = spd
.parallel_transport(backward.view(), out.view())
.expect("Γ_{Q→P}(Γ_{P→Q}(U))");
for (i, (&b, &orig)) in back.iter().zip(u.iter()).enumerate() {
assert!(
(b - orig).abs() <= 1e-9 * orig.abs().max(1.0),
"component {i}: round-trip {b:.12e} vs original {orig:.12e}"
);
}
}
}
#[cfg(test)]
mod christoffel_tests {
use super::SpdManifold;
use crate::manifold::{RiemannianManifold, flatten, from_flat};
use ndarray::{Array1, Array2};
fn symmetric_basis(n: usize) -> Vec<Array2<f64>> {
let mut basis = Vec::with_capacity(n * (n + 1) / 2);
for i in 0..n {
let mut m = Array2::<f64>::zeros((n, n));
m[[i, i]] = 1.0;
basis.push(m);
}
for i in 0..n {
for j in (i + 1)..n {
let mut m = Array2::<f64>::zeros((n, n));
m[[i, j]] = 1.0;
m[[j, i]] = 1.0;
basis.push(m);
}
}
basis
}
fn base_point(n: usize) -> Array2<f64> {
let mut p = Array2::<f64>::zeros((n, n));
for i in 0..n {
p[[i, i]] = 1.0 + i as f64;
}
for i in 0..n {
for j in (i + 1)..n {
let v = 0.05 * (i as f64 + 1.0) - 0.03 * (j as f64 + 1.0) + 0.1;
p[[i, j]] = v;
p[[j, i]] = v;
}
}
p
}
#[test]
fn christoffel_matches_fd_of_metric_on_symmetric_chart() {
let n = 3;
let m = SpdManifold::new(n);
let p0 = base_point(n);
let basis = symmetric_basis(n);
let basis_flat: Vec<Array1<f64>> = basis.iter().map(flatten).collect();
let d = basis.len();
assert_eq!(d, n * (n + 1) / 2);
let ambient = m.ambient_dim();
let point_at = |x: &[f64]| -> Array1<f64> {
let mut p = p0.clone();
for (a, &xa) in x.iter().enumerate() {
if xa != 0.0 {
p = &p + &(&basis[a] * xa);
}
}
flatten(&p)
};
let contract = |g: &Array2<f64>, b: usize, c: usize| -> f64 {
basis_flat[b].dot(&g.dot(&basis_flat[c]))
};
let x0 = vec![0.0_f64; d];
let h = 1e-6;
let mut dg = vec![vec![vec![0.0_f64; d]; d]; d]; for a in 0..d {
let mut xp = x0.clone();
xp[a] += h;
let mut xn = x0.clone();
xn[a] -= h;
let gp = m.metric_tensor(point_at(&xp).view()).expect("G(x+h e_a)");
let gn = m.metric_tensor(point_at(&xn).view()).expect("G(x-h e_a)");
for b in 0..d {
for c in 0..d {
dg[a][b][c] = (contract(&gp, b, c) - contract(&gn, b, c)) / (2.0 * h);
}
}
}
let point0 = point_at(&x0);
let gamma = m.christoffel_symbols(point0.view()).expect("Γ tensor");
let connection_matrix = |a: usize, b: usize| -> Array2<f64> {
let mut gamma_vec = Array1::<f64>::zeros(ambient);
for out in 0..ambient {
let mut acc = 0.0;
for p_idx in 0..ambient {
let coeff = basis_flat[a][p_idx];
if coeff == 0.0 {
continue;
}
for q_idx in 0..ambient {
acc += coeff * gamma[out][[p_idx, q_idx]] * basis_flat[b][q_idx];
}
}
gamma_vec[out] = acc;
}
from_flat(gamma_vec.view(), n, n).expect("Γ(E_a,E_b) as n×n")
};
for a in 0..d {
for b in 0..d {
let gamma_mat = connection_matrix(a, b);
for c in 0..d {
let lhs = m
.affine_inner(&p0, &gamma_mat, &basis[c])
.expect("⟨Γ(E_a,E_b), E_c⟩");
let rhs = 0.5 * (dg[a][b][c] + dg[b][a][c] - dg[c][a][b]);
assert!(
(lhs - rhs).abs() <= 1e-6 * rhs.abs().max(1.0),
"a={a} b={b} c={c}: ⟨Γ,E_c⟩_analytic={lhs:.10e} vs FD-of-metric={rhs:.10e}"
);
}
}
}
}
#[test]
fn sectional_curvature_vanishes_on_commuting_diagonal_plane() {
let m = SpdManifold::new(2);
let p = Array1::from(vec![2.0_f64, 0.0, 0.0, 3.0]); let u = Array1::from(vec![1.0_f64, 0.0, 0.0, 0.0]); let v = Array1::from(vec![0.0_f64, 0.0, 0.0, 1.0]); let k = m
.sectional_curvature(p.view(), (u.view(), v.view()))
.expect("sectional curvature on commuting plane");
assert!(
k.abs() <= 1e-12,
"expected flat commuting plane, got κ={k:.3e}"
);
}
#[test]
fn sectional_curvature_is_nonpositive_on_noncommuting_plane() {
let m = SpdManifold::new(2);
let p = Array1::from(vec![1.0_f64, 0.0, 0.0, 1.0]); let u = Array1::from(vec![1.0_f64, 0.0, 0.0, -1.0]); let v = Array1::from(vec![0.0_f64, 1.0, 1.0, 0.0]); let k = m
.sectional_curvature(p.view(), (u.view(), v.view()))
.expect("sectional curvature on non-commuting plane");
assert!(
k < -1e-6,
"expected strictly negative curvature, got κ={k:.3e}"
);
}
}