use crate::linalg::{symmetric_eigen, EigenError};
use crate::smooth::Bounds;
const DIM: usize = 3;
const EIGVAL_REL_TOL: f64 = 1e-12;
#[derive(Debug, Clone, PartialEq)]
pub struct Embedding {
pub coords: Vec<[f64; DIM]>,
pub eigenvalues: [f64; DIM],
pub fit3: f64,
pub negative_share: f64,
pub degenerate_axes: usize,
pub negative_centroid_sq: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EmbedError {
Eigen(EigenError),
BadShape,
}
impl From<EigenError> for EmbedError {
fn from(e: EigenError) -> Self {
Self::Eigen(e)
}
}
#[must_use]
pub fn reference_distances(b: &Bounds) -> Vec<f64> {
let n = b.len();
let mut d = vec![0.0; n * n];
for i in 0..n {
for j in (i + 1)..n {
let u = b.upper(i, j);
d[i * n + j] = u;
d[j * n + i] = u;
}
}
d
}
#[must_use]
pub fn metric_matrix(dist: &[f64], n: usize) -> (Vec<f64>, usize) {
assert_eq!(dist.len(), n * n, "距离表不是 {n}×{n}");
if n == 0 {
return (Vec::new(), 0);
}
let nf = n as f64;
let mut sum_sq = 0.0;
for i in 0..n {
for j in (i + 1)..n {
sum_sq += dist[i * n + j] * dist[i * n + j];
}
}
sum_sq /= nf * nf;
let mut sq0 = vec![0.0; n];
let mut negative = 0;
for i in 0..n {
let row: f64 = (0..n).map(|j| dist[i * n + j] * dist[i * n + j]).sum();
sq0[i] = row / nf - sum_sq;
if sq0[i] < 0.0 {
negative += 1;
}
}
let mut t = vec![0.0; n * n];
for i in 0..n {
for j in 0..n {
t[i * n + j] = 0.5 * (sq0[i] + sq0[j] - dist[i * n + j] * dist[i * n + j]);
}
}
(t, negative)
}
pub fn embed(dist: &[f64], n: usize) -> Result<Embedding, EmbedError> {
if dist.len() != n * n {
return Err(EmbedError::BadShape);
}
let (t, negative_centroid_sq) = metric_matrix(dist, n);
let eig = symmetric_eigen(&t, n)?;
let mut abs_mass = 0.0;
let mut neg_mass = 0.0;
for &v in &eig.values {
abs_mass += v.abs();
if v <= 0.0 {
neg_mass += -v;
}
}
let scale = eig.values.first().map_or(0.0, |v| v.abs());
let cut = EIGVAL_REL_TOL * scale;
let mut eigenvalues = [0.0; DIM];
let mut degenerate_axes = 0;
for (k, lam) in eigenvalues.iter_mut().enumerate() {
let v = eig.values.get(k).copied().unwrap_or(0.0);
if v > cut {
*lam = v;
} else {
degenerate_axes += 1;
}
}
let mut coords = vec![[0.0; DIM]; n];
for (k, &lam) in eigenvalues.iter().enumerate() {
let s = lam.sqrt();
let Some(vk) = (k < eig.len()).then(|| eig.vector(k)) else {
continue;
};
for (i, c) in coords.iter_mut().enumerate() {
c[k] = s * vk[i];
}
}
Ok(Embedding {
coords,
eigenvalues,
fit3: if abs_mass > 0.0 {
eigenvalues.iter().sum::<f64>() / abs_mass
} else {
0.0
},
negative_share: if abs_mass > 0.0 {
neg_mass / abs_mass
} else {
0.0
},
degenerate_axes,
negative_centroid_sq,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn lcg(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
((*state >> 11) as f64) / ((1u64 << 53) as f64) * 2.0 - 1.0
}
fn random_points(n: usize, seed: u64) -> Vec<[f64; 3]> {
let mut st = seed;
(0..n)
.map(|_| [lcg(&mut st) * 6.0, lcg(&mut st) * 6.0, lcg(&mut st) * 6.0])
.collect()
}
fn exact_distances(p: &[[f64; 3]]) -> Vec<f64> {
let n = p.len();
let mut d = vec![0.0; n * n];
for i in 0..n {
for j in 0..n {
d[i * n + j] = (0..3)
.map(|k| (p[i][k] - p[j][k]).powi(2))
.sum::<f64>()
.sqrt();
}
}
d
}
fn max_dist_dev(a: &[f64], b: &[f64]) -> f64 {
a.iter().zip(b).fold(0.0_f64, |w, (x, y)| {
crate::linalg::max_nan_wins(w, (x - y).abs())
})
}
#[test]
fn 最大偏差不许把_nan_洗成零() {
let a = [1.0, 2.0, 3.0];
assert!((max_dist_dev(&a, &[1.0, 2.0, 3.5]) - 0.5).abs() < 1e-12);
assert_eq!(max_dist_dev(&a, &a), 0.0, "同一张表偏差是 0");
assert!(
max_dist_dev(&a, &[1.0, f64::NAN, 3.0]).is_nan(),
"右边带 NaN"
);
assert!(
max_dist_dev(&[f64::NAN, 2.0, 3.0], &a).is_nan(),
"左边带 NaN"
);
}
#[test]
fn 真实三维点集精确回嵌() {
for n in [4usize, 5, 9, 17, 33] {
for seed in [1u64, 0xf00d, 0xdead_beef] {
let pts = random_points(n, seed);
let d = exact_distances(&pts);
let e = embed(&d, n).unwrap();
let back = exact_distances(&e.coords);
let dev = max_dist_dev(&d, &back);
assert!(dev < 1e-12, "n={n} seed={seed} 距离偏差 {dev:.3e}");
assert_eq!(e.degenerate_axes, 0, "n={n} seed={seed} 不该有退化轴");
assert!(e.fit3 > 1.0 - 1e-12, "n={n} seed={seed} fit3={}", e.fit3);
assert!(
e.negative_share < 1e-12,
"n={n} 负份额 {}",
e.negative_share
);
assert_eq!(e.negative_centroid_sq, 0);
}
}
}
#[test]
fn 嵌出来的坐标以质心为原点() {
let pts = random_points(11, 7);
let e = embed(&exact_distances(&pts), 11).unwrap();
for k in 0..3 {
let c: f64 = e.coords.iter().map(|p| p[k]).sum::<f64>() / 11.0;
assert!(c.abs() < 1e-12, "第 {k} 轴质心 {c:.3e} 不在原点");
}
}
#[test]
fn 平面点集只用掉两根轴() {
let mut st = 3u64;
let pts: Vec<[f64; 3]> = (0..8)
.map(|_| [lcg(&mut st) * 4.0, lcg(&mut st) * 4.0, 0.0])
.collect();
let d = exact_distances(&pts);
let e = embed(&d, 8).unwrap();
assert_eq!(e.degenerate_axes, 1, "特征值 {:?}", e.eigenvalues);
assert!(max_dist_dev(&d, &exact_distances(&e.coords)) < 1e-9);
}
#[test]
fn 接近平面但不是平面的结构不许被压平() {
let mut st = 5u64;
let pts: Vec<[f64; 3]> = (0..10)
.map(|_| [lcg(&mut st) * 4.0, lcg(&mut st) * 4.0, lcg(&mut st) * 0.02])
.collect();
let d = exact_distances(&pts);
let e = embed(&d, 10).unwrap();
let ratio = e.eigenvalues[2] / e.eigenvalues[0];
assert!(
(1e-12..1e-4).contains(&ratio),
"构造失效:λ₃/λ₁ = {ratio:.3e} 没落在化学相关的那一段里"
);
assert_eq!(e.degenerate_axes, 0, "第三维是真的,不许判成退化");
let dev = max_dist_dev(&d, &exact_distances(&e.coords));
assert!(dev < 1e-11, "距离偏差 {dev:.3e} —— 结构被压平了");
}
#[test]
fn 摆不进三维时降级而不是失败() {
let n = 5;
let mut d = vec![0.0; n * n];
for i in 0..n {
for j in 0..n {
if i != j {
d[i * n + j] = 1.0;
}
}
}
let e = embed(&d, n).expect("必须给出坐标,不能失败");
assert_eq!(e.coords.len(), n);
assert!(e.fit3 < 1.0, "fit3={} 应当小于 1", e.fit3);
assert!(e.negative_share < 1e-12, "负份额 {}", e.negative_share);
assert_eq!(e.degenerate_axes, 0, "四维单形的前三个特征值都是正的");
}
#[test]
fn 自相矛盾的距离表会露出负特征值() {
let n = 3;
let mut d = vec![0.0; n * n];
let set = |d: &mut Vec<f64>, i: usize, j: usize, v: f64| {
d[i * n + j] = v;
d[j * n + i] = v;
};
set(&mut d, 0, 1, 1.0);
set(&mut d, 1, 2, 1.0);
set(&mut d, 0, 2, 10.0);
let e = embed(&d, n).expect("坏表也要给坐标");
assert!(e.negative_share > 0.01, "负份额只有 {}", e.negative_share);
assert_eq!(e.degenerate_axes, 2, "特征值 {:?}", e.eigenvalues);
assert!(
e.fit3 < 0.85,
"fit3 = {:.6},这张表两根轴都塌了还报高分",
e.fit3
);
for (i, c) in e.coords.iter().enumerate() {
assert!(
c.iter().all(|x| x.is_finite()),
"第 {i} 个原子的坐标不是有限数:{c:?}"
);
}
}
#[test]
fn 参考距离表整张取上限() {
let mut b = Bounds::new(3, 1.0, 5.0);
b.set_upper(0, 1, 2.5);
b.set_lower(0, 1, 2.0);
let d = reference_distances(&b);
let at = |i: usize, j: usize| d[i * 3 + j];
assert_eq!(at(0, 1), 2.5);
assert_eq!(at(1, 0), 2.5, "必须对称");
assert_eq!(at(0, 0), 0.0, "对角必须是 0");
assert_eq!(at(0, 2), 5.0);
}
#[test]
fn 空表与退化尺寸() {
let e = embed(&[], 0).unwrap();
assert!(e.coords.is_empty());
let e = embed(&[0.0], 1).unwrap();
assert_eq!(e.coords, vec![[0.0; 3]]);
assert_eq!(e.degenerate_axes, 3);
let d = vec![0.0, 1.5, 1.5, 0.0];
let e = embed(&d, 2).unwrap();
assert_eq!(e.degenerate_axes, 2);
let back = exact_distances(&e.coords);
assert!((back[1] - 1.5).abs() < 1e-12, "两原子距离 {}", back[1]);
}
#[test]
fn 形状不对要报错() {
assert_eq!(embed(&[1.0, 2.0], 3), Err(EmbedError::BadShape));
}
#[test]
fn 平移旋转不改变结果的距离() {
let pts = random_points(9, 99);
let d1 = exact_distances(&pts);
let (c, s) = (0.6_f64, 0.8_f64);
let moved: Vec<[f64; 3]> = pts
.iter()
.map(|p| {
[
c * p[0] - s * p[1] + 100.0,
s * p[0] + c * p[1] - 50.0,
p[2] + 7.0,
]
})
.collect();
let d2 = exact_distances(&moved);
assert!(max_dist_dev(&d1, &d2) < 1e-12, "构造有问题:距离本该不变");
let e1 = embed(&d1, 9).unwrap();
let e2 = embed(&d2, 9).unwrap();
let dev = max_dist_dev(&exact_distances(&e1.coords), &exact_distances(&e2.coords));
assert!(dev < 1e-12, "偏差 {dev:.3e}");
}
}