use super::SaeAtomBasisKind;
use gam_linalg::faer_ndarray::FaerEigh;
use ndarray::{Array2, Array3, ArrayView2};
use std::cmp::Ordering;
use std::collections::BinaryHeap;
use faer::Side;
const MDS_EIGENVALUE_FLOOR_FRAC: f64 = 1.0e-9;
fn intrinsic_seed_knn(n_points: usize, d: usize) -> usize {
let tangent_floor = 2 * d + 1;
let connectivity_floor = (n_points.max(2) as f64).log2().ceil() as usize;
tangent_floor.max(connectivity_floor).max(2)
}
fn intrinsic_landmark_count(n_points: usize, d: usize) -> usize {
const COVERAGE_MULTIPLIER: f64 = 4.0;
let coverage = (COVERAGE_MULTIPLIER * (n_points as f64).sqrt()).ceil() as usize;
let floor = 2 * (d + 1);
coverage.max(floor).min(n_points)
}
fn squared_distance(z: ArrayView2<'_, f64>, a: usize, b: usize) -> f64 {
let mut acc = 0.0;
for c in 0..z.ncols() {
let diff = z[[a, c]] - z[[b, c]];
acc += diff * diff;
}
acc
}
pub(crate) fn deterministic_knn_graph(z: ArrayView2<'_, f64>, k: usize) -> Vec<Vec<(usize, f64)>> {
let n = z.nrows();
let mut adj: Vec<Vec<(usize, f64)>> = vec![Vec::new(); n];
if n == 0 {
return adj;
}
let k = k.min(n.saturating_sub(1)).max(1);
let mut edges: std::collections::BTreeMap<(usize, usize), f64> =
std::collections::BTreeMap::new();
for i in 0..n {
let mut dists: Vec<(f64, usize)> = Vec::with_capacity(n - 1);
for j in 0..n {
if i != j {
dists.push((squared_distance(z, i, j), j));
}
}
dists.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
for &(dist2, j) in dists.iter().take(k) {
let key = (i.min(j), i.max(j));
edges.entry(key).or_insert_with(|| dist2.sqrt());
}
}
let mut parent: Vec<usize> = (0..n).collect();
fn find(parent: &mut [usize], mut x: usize) -> usize {
while parent[x] != x {
parent[x] = parent[parent[x]];
x = parent[x];
}
x
}
for &(a, b) in edges.keys() {
let ra = find(&mut parent, a);
let rb = find(&mut parent, b);
if ra != rb {
parent[ra.max(rb)] = ra.min(rb);
}
}
loop {
let mut best: Option<(f64, usize, usize)> = None;
for i in 0..n {
let ri = find(&mut parent, i);
for j in (i + 1)..n {
if find(&mut parent, j) == ri {
continue;
}
let d2 = squared_distance(z, i, j);
let better = match best {
None => true,
Some((bd, _, _)) => d2 < bd,
};
if better {
best = Some((d2, i, j));
}
}
}
match best {
None => break, Some((d2, i, j)) => {
edges.insert((i, j), d2.sqrt());
let ri = find(&mut parent, i);
let rj = find(&mut parent, j);
parent[ri.max(rj)] = ri.min(rj);
}
}
}
for (&(a, b), &w) in &edges {
adj[a].push((b, w));
adj[b].push((a, w));
}
adj
}
pub(crate) fn farthest_point_landmarks(z: ArrayView2<'_, f64>, count: usize) -> Vec<usize> {
let n = z.nrows();
if n == 0 {
return Vec::new();
}
let target = count.max(1).min(n);
let mut chosen: Vec<usize> = Vec::with_capacity(target);
chosen.push(0);
let mut nearest_sq: Vec<f64> = (0..n).map(|r| squared_distance(z, r, 0)).collect();
while chosen.len() < target {
let mut best = 0usize;
let mut best_d = -1.0;
for r in 0..n {
if nearest_sq[r] > best_d {
best_d = nearest_sq[r];
best = r;
}
}
if best_d <= 0.0 {
break; }
chosen.push(best);
for r in 0..n {
let dr = squared_distance(z, r, best);
if dr < nearest_sq[r] {
nearest_sq[r] = dr;
}
}
}
chosen
}
#[derive(PartialEq)]
struct DijkstraNode {
dist: f64,
node: usize,
}
impl Eq for DijkstraNode {}
impl Ord for DijkstraNode {
fn cmp(&self, other: &Self) -> Ordering {
other
.dist
.total_cmp(&self.dist)
.then_with(|| other.node.cmp(&self.node))
}
}
impl PartialOrd for DijkstraNode {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
fn dijkstra(adj: &[Vec<(usize, f64)>], source: usize) -> Vec<f64> {
let n = adj.len();
let mut dist = vec![f64::INFINITY; n];
dist[source] = 0.0;
let mut heap = BinaryHeap::new();
heap.push(DijkstraNode {
dist: 0.0,
node: source,
});
while let Some(DijkstraNode { dist: d, node }) = heap.pop() {
if d > dist[node] {
continue;
}
for &(nbr, w) in &adj[node] {
let nd = d + w;
if nd < dist[nbr] {
dist[nbr] = nd;
heap.push(DijkstraNode {
dist: nd,
node: nbr,
});
}
}
}
dist
}
pub(crate) fn landmark_geodesics(adj: &[Vec<(usize, f64)>], landmarks: &[usize]) -> Array2<f64> {
let n = adj.len();
let l = landmarks.len();
let mut out = Array2::<f64>::zeros((l, n));
for (li, &src) in landmarks.iter().enumerate() {
let d = dijkstra(adj, src);
for j in 0..n {
out[[li, j]] = d[j];
}
}
out
}
pub fn intrinsic_geodesic_embedding(
z: ArrayView2<'_, f64>,
d: usize,
) -> Result<Array2<f64>, String> {
let n = z.nrows();
if d == 0 {
return Ok(Array2::<f64>::zeros((n, 0)));
}
let mut out = Array2::<f64>::zeros((n, d));
if n == 0 || z.ncols() == 0 {
return Ok(out);
}
for ((row, col), &value) in z.indexed_iter() {
if !value.is_finite() {
return Err(format!(
"intrinsic_seed: Z must be finite; Z[{row}, {col}] = {value}"
));
}
}
if n == 1 {
return Ok(out);
}
if n == 2 {
let dist2 = squared_distance(z, 0, 1);
if !dist2.is_finite() {
return Err(
"intrinsic_seed: pairwise distance overflowed; rescale Z before seeding"
.to_string(),
);
}
let half_distance = 0.5 * dist2.sqrt();
out[[0, 0]] = -half_distance;
out[[1, 0]] = half_distance;
return Ok(out);
}
let k = intrinsic_seed_knn(n, d).min(n - 1);
let adj = deterministic_knn_graph(z, k);
let l_count = intrinsic_landmark_count(n, d);
let landmarks = farthest_point_landmarks(z, l_count);
let l = landmarks.len();
if l < 2 {
return Ok(out);
}
let geo = landmark_geodesics(&adj, &landmarks);
let mut d2 = Array2::<f64>::zeros((l, l));
for a in 0..l {
for b in 0..l {
let g = geo[[a, landmarks[b]]];
d2[[a, b]] = g * g;
}
}
for a in 0..l {
for b in (a + 1)..l {
let avg = 0.5 * (d2[[a, b]] + d2[[b, a]]);
d2[[a, b]] = avg;
d2[[b, a]] = avg;
}
}
let mut row_mean = vec![0.0_f64; l];
let mut grand = 0.0_f64;
for a in 0..l {
let mut s = 0.0;
for b in 0..l {
s += d2[[a, b]];
}
row_mean[a] = s / l as f64;
grand += s;
}
grand /= (l * l) as f64;
let mut b_mat = Array2::<f64>::zeros((l, l));
for a in 0..l {
for b in 0..l {
b_mat[[a, b]] = -0.5 * (d2[[a, b]] - row_mean[a] - row_mean[b] + grand);
}
}
for a in 0..l {
for b in (a + 1)..l {
let avg = 0.5 * (b_mat[[a, b]] + b_mat[[b, a]]);
b_mat[[a, b]] = avg;
b_mat[[b, a]] = avg;
}
}
let (evals, evecs) = b_mat
.eigh(Side::Lower)
.map_err(|err| format!("intrinsic_seed: MDS eigensolve failed: {err:?}"))?;
let leading = evals.iter().cloned().fold(0.0_f64, f64::max);
if !(leading > 0.0) {
return Ok(out); }
let floor = leading * MDS_EIGENVALUE_FLOOR_FRAC;
let mut order: Vec<usize> = (0..evals.len()).collect();
order.sort_by(|&i, &j| evals[j].total_cmp(&evals[i]).then_with(|| i.cmp(&j)));
let axes: Vec<usize> = order
.into_iter()
.filter(|&c| evals[c] > floor)
.take(d)
.collect();
for (col, &c) in axes.iter().enumerate() {
let inv = -0.5 / evals[c].sqrt();
for r in 0..n {
let mut acc = 0.0_f64;
for a in 0..l {
let g = geo[[a, r]];
let delta = g * g;
acc += evecs[[a, c]] * (delta - row_mean[a]);
}
out[[r, col]] = inv * acc;
}
}
Ok(out)
}
fn min_max_normalize_into(out: &mut Array3<f64>, atom_idx: usize, axis: usize, values: &[f64]) {
let (lo, hi) = values
.iter()
.fold((f64::INFINITY, f64::NEG_INFINITY), |(lo, hi), &v| {
(lo.min(v), hi.max(v))
});
let span = hi - lo;
if span > 0.0 && span.is_finite() {
for (row, &v) in values.iter().enumerate() {
out[[atom_idx, row, axis]] = (v - lo) / span - 0.5;
}
}
}
fn intrinsic_chart_embedding_axes(kind: &SaeAtomBasisKind, latent_dim: usize) -> usize {
match kind {
SaeAtomBasisKind::Periodic => 2 + latent_dim.max(1).saturating_sub(1),
SaeAtomBasisKind::Sphere => 3,
SaeAtomBasisKind::Torus => 2 * latent_dim.max(1),
_ => latent_dim.max(1),
}
}
pub fn sae_intrinsic_seed_initial_coords(
z: ArrayView2<'_, f64>,
basis_kinds: &[SaeAtomBasisKind],
atom_dim: &[usize],
) -> Result<Array3<f64>, String> {
let k_atoms = basis_kinds.len();
if atom_dim.len() != k_atoms {
return Err(format!(
"sae_intrinsic_seed_initial_coords: basis_kinds and atom_dim must have the same length; got {} and {}",
k_atoms,
atom_dim.len()
));
}
let (n_obs, _p) = z.dim();
let latent_d_max = atom_dim.iter().copied().max().unwrap_or(1).max(1);
let embedding_axes = basis_kinds
.iter()
.zip(atom_dim.iter().copied())
.map(|(kind, d)| intrinsic_chart_embedding_axes(kind, d))
.max()
.unwrap_or(1);
let mut out = Array3::<f64>::zeros((k_atoms, n_obs, latent_d_max));
if n_obs == 0 || z.ncols() == 0 || k_atoms == 0 {
return Ok(out);
}
let embed = intrinsic_geodesic_embedding(z, embedding_axes)?;
let two_pi = std::f64::consts::TAU;
for atom_idx in 0..k_atoms {
let d = atom_dim[atom_idx];
if d == 0 {
continue;
}
match &basis_kinds[atom_idx] {
SaeAtomBasisKind::Periodic => {
if embed.ncols() >= 2 {
for row in 0..n_obs {
let phase = embed[[row, 1]].atan2(embed[[row, 0]]) / two_pi;
out[[atom_idx, row, 0]] = phase - phase.floor();
}
} else {
let col: Vec<f64> = (0..n_obs).map(|r| embed[[r, 0]]).collect();
let (lo, hi) = col
.iter()
.fold((f64::INFINITY, f64::NEG_INFINITY), |(lo, hi), &v| {
(lo.min(v), hi.max(v))
});
let span = hi - lo;
if span > 0.0 {
for (row, &v) in col.iter().enumerate() {
out[[atom_idx, row, 0]] = (v - lo) / span;
}
}
}
for axis in 1..d {
if axis + 1 >= embed.ncols() {
break;
}
let vals: Vec<f64> = (0..n_obs).map(|r| embed[[r, axis + 1]]).collect();
min_max_normalize_into(&mut out, atom_idx, axis, &vals);
}
}
SaeAtomBasisKind::Sphere => {
for row in 0..n_obs {
let x = embed[[row, 0]];
let y = embed[[row, 1]];
let zz = embed[[row, 2]];
let norm = (x * x + y * y + zz * zz).sqrt().max(1.0e-24);
if d >= 1 {
out[[atom_idx, row, 0]] = (zz / norm).clamp(-1.0, 1.0).asin();
}
if d >= 2 {
out[[atom_idx, row, 1]] = y.atan2(x);
}
}
}
SaeAtomBasisKind::Torus => {
for axis in 0..d {
let (ca, cb) = (2 * axis, 2 * axis + 1);
if cb >= embed.ncols() {
break;
}
for row in 0..n_obs {
let phase = embed[[row, cb]].atan2(embed[[row, ca]]) / two_pi;
out[[atom_idx, row, axis]] = phase - phase.floor();
}
}
}
_ => {
for axis in 0..d {
if axis >= embed.ncols() {
break;
}
let vals: Vec<f64> = (0..n_obs).map(|r| embed[[r, axis]]).collect();
min_max_normalize_into(&mut out, atom_idx, axis, &vals);
}
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mds_recovers_known_configuration_up_to_rigid_motion() {
let n = 6usize;
let mut z = Array2::<f64>::zeros((n, 3));
for i in 0..2 {
for j in 0..3 {
let r = i * 3 + j;
z[[r, 0]] = i as f64;
z[[r, 1]] = j as f64;
z[[r, 2]] = 0.5; }
}
let embed = intrinsic_geodesic_embedding(z.view(), 2).unwrap();
let mut max_rel = 0.0_f64;
for a in 0..n {
for b in (a + 1)..n {
let amb = super::squared_distance(z.view(), a, b).sqrt();
let mut e2 = 0.0;
for c in 0..2 {
let diff = embed[[a, c]] - embed[[b, c]];
e2 += diff * diff;
}
let emb = e2.sqrt();
if amb > 1e-9 {
max_rel = max_rel.max(((emb - amb) / amb).abs());
}
}
}
assert!(
max_rel < 1e-6,
"MDS on a complete-graph (exact-Euclidean) configuration must reproduce \
its pairwise distances to rounding (max relative error {max_rel:.3e})"
);
}
#[test]
fn two_row_embedding_preserves_full_ambient_distance() {
let z = Array2::from_shape_vec((2, 3), vec![4.0, -2.0, 1.0, 4.0, 4.0, 9.0]).unwrap();
let embed = intrinsic_geodesic_embedding(z.view(), 1).unwrap();
assert_eq!(embed[[0, 0]], -5.0);
assert_eq!(embed[[1, 0]], 5.0);
assert_eq!(embed[[0, 0]] + embed[[1, 0]], 0.0);
}
#[test]
fn intrinsic_embedding_is_bit_identical_run_to_run() {
let n = 40usize;
let mut z = Array2::<f64>::zeros((n, 4));
for r in 0..n {
let t = r as f64 * 0.3;
z[[r, 0]] = t.sin();
z[[r, 1]] = t.cos();
z[[r, 2]] = (0.5 * t).sin();
z[[r, 3]] = 0.2 * t;
}
let a = intrinsic_geodesic_embedding(z.view(), 2).unwrap();
let b = intrinsic_geodesic_embedding(z.view(), 2).unwrap();
assert_eq!(a, b, "intrinsic embedding must be bit-identical run-to-run");
}
#[test]
fn intrinsic_seed_matches_pca_contract_shape_and_finite() {
let n = 30usize;
let p = 5usize;
let mut z = Array2::<f64>::zeros((n, p));
for r in 0..n {
for c in 0..p {
z[[r, c]] = ((r * 3 + c) as f64 * 0.21).sin() + 0.1 * (r as f64 - c as f64);
}
}
for (kind, d) in [
(SaeAtomBasisKind::Linear, 2usize),
(SaeAtomBasisKind::Periodic, 1usize),
(SaeAtomBasisKind::Sphere, 2usize),
(SaeAtomBasisKind::Torus, 2usize),
] {
let k = 3;
let kinds = vec![kind.clone(); k];
let dims = vec![d; k];
let seed = sae_intrinsic_seed_initial_coords(z.view(), &kinds, &dims).unwrap();
assert_eq!(seed.dim(), (k, n, d));
for v in seed.iter() {
assert!(
v.is_finite(),
"{kind:?}: non-finite intrinsic seed coord {v}"
);
}
}
}
#[test]
fn intrinsic_seed_allocates_every_chart_function_2240() {
assert_eq!(
intrinsic_chart_embedding_axes(&SaeAtomBasisKind::Periodic, 1),
2
);
assert_eq!(
intrinsic_chart_embedding_axes(&SaeAtomBasisKind::Periodic, 3),
4
);
assert_eq!(
intrinsic_chart_embedding_axes(&SaeAtomBasisKind::Sphere, 2),
3
);
assert_eq!(
intrinsic_chart_embedding_axes(&SaeAtomBasisKind::Torus, 2),
4
);
assert_eq!(
intrinsic_chart_embedding_axes(&SaeAtomBasisKind::Linear, 2),
2
);
let mut z = Array2::<f64>::zeros((10, 4));
for axis in 0..4 {
z[[2 * axis, axis]] = 1.0;
z[[2 * axis + 1, axis]] = -1.0;
}
z[[8, 0]] = 0.5;
z[[8, 1]] = -0.25;
z[[8, 2]] = 0.75;
z[[8, 3]] = 0.125;
z[[9, 0]] = -0.375;
z[[9, 1]] = 0.625;
z[[9, 2]] = 0.25;
z[[9, 3]] = -0.875;
let seed = sae_intrinsic_seed_initial_coords(
z.view(),
&[SaeAtomBasisKind::Sphere, SaeAtomBasisKind::Torus],
&[2, 2],
)
.unwrap();
for atom in 0..2 {
for axis in 0..2 {
let (lo, hi) = seed
.slice(ndarray::s![atom, .., axis])
.iter()
.fold((f64::INFINITY, f64::NEG_INFINITY), |(lo, hi), &v| {
(lo.min(v), hi.max(v))
});
assert!(
hi > lo,
"intrinsic chart {atom} axis {axis} must carry a genuine coordinate"
);
}
}
}
#[test]
fn intrinsic_seed_rejects_misaligned_atom_metadata() {
let z = Array2::<f64>::zeros((3, 2));
let error = sae_intrinsic_seed_initial_coords(
z.view(),
&[SaeAtomBasisKind::Linear, SaeAtomBasisKind::Sphere],
&[2],
)
.unwrap_err();
assert!(error.contains("same length"));
}
}