use gam_linalg::faer_ndarray::FaerEigh;
use ndarray::{Array2, 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)
}
#[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");
}
}