use faer::Side;
use faer::prelude::*;
use gam_linalg::faer_ndarray::FaerEigh;
use ndarray::{Array1, Array2, ArrayView2};
use opt::{LmConfig, LmOutcome, LmState, lm_step};
use std::f64::consts::TAU;
fn optimal_hard_threshold_coefficient(beta: f64) -> f64 {
(2.0 * (beta + 1.0) + 8.0 * beta / ((beta + 1.0) + (beta * beta + 14.0 * beta + 1.0).sqrt()))
.sqrt()
}
fn singular_value_order(
singular_values: &[f64],
rows: usize,
cols: usize,
sigma: f64,
max_order: usize,
) -> usize {
let sigma_1 = singular_values.first().copied().unwrap_or(0.0);
let n_big = rows.max(cols) as f64;
let numerical_floor = sigma_1 * n_big * f64::EPSILON;
let threshold = if sigma > 0.0 {
let beta = rows.min(cols) as f64 / n_big;
let sigma_entry = std::f64::consts::SQRT_2 * sigma;
(optimal_hard_threshold_coefficient(beta) * sigma_entry * n_big.sqrt()).max(numerical_floor)
} else {
numerical_floor
};
let threshold_order = singular_values
.iter()
.filter(|&&s| s > threshold)
.count()
.min(max_order);
let dramatic_gap = (f32::EPSILON as f64).sqrt();
let mut gap_order = None;
for k in 1..max_order.min(singular_values.len()) {
if singular_values[k - 1] <= 0.0 {
break;
}
if singular_values[k] / singular_values[k - 1] < dramatic_gap {
gap_order = Some(k);
break;
}
}
gap_order.unwrap_or(threshold_order)
}
#[derive(Clone, Debug, PartialEq)]
pub struct SphereSpike {
pub direction: Vec<f64>,
pub amplitude: f64,
}
#[derive(Clone, Debug)]
pub struct SphereRecovery {
pub spikes: Vec<SphereSpike>,
pub model_order: usize,
pub residual: f64,
pub eigenvalues: Vec<f64>,
}
pub fn sphere_lift(spikes: &[SphereSpike], d: usize) -> Result<Array2<f64>, String> {
if d == 0 {
return Err("sphere_lift: ambient dimension must be positive".into());
}
let mut m = Array2::<f64>::zeros((d, d));
for (idx, spike) in spikes.iter().enumerate() {
if spike.direction.len() != d {
return Err(format!(
"sphere_lift: spike {idx} direction has length {}, expected {d}",
spike.direction.len()
));
}
let a = spike.amplitude;
for i in 0..d {
for j in 0..d {
m[[i, j]] += a * spike.direction[i] * spike.direction[j];
}
}
}
Ok(m)
}
pub fn canonicalize_direction(v: &mut [f64]) {
let mut pivot = 0usize;
let mut best = 0.0_f64;
for (i, &x) in v.iter().enumerate() {
if x.abs() > best {
best = x.abs();
pivot = i;
}
}
if v.get(pivot).is_some_and(|&x| x < 0.0) {
for x in v.iter_mut() {
*x = -*x;
}
}
}
fn sphere_eigenvalue_floor(lambda_max: f64, d: usize, sigma: f64) -> f64 {
const EDGE_SAFETY: f64 = 2.0;
let numerical_floor = lambda_max.abs() * (d as f64) * f64::EPSILON;
if sigma > 0.0 {
(EDGE_SAFETY * 2.0 * sigma * (d as f64).sqrt()).max(numerical_floor)
} else {
numerical_floor
}
}
pub fn recover_sphere_spikes(
m_hat: ArrayView2<'_, f64>,
sigma: f64,
) -> Result<SphereRecovery, String> {
let d = m_hat.nrows();
if d == 0 || m_hat.ncols() != d {
return Err(format!(
"recover_sphere_spikes: lift must be square and non-empty; got {:?}",
m_hat.dim()
));
}
if m_hat.iter().any(|x| !x.is_finite()) {
return Err("recover_sphere_spikes: lift has non-finite entries".into());
}
let mut sym = Array2::<f64>::zeros((d, d));
for i in 0..d {
for j in 0..d {
sym[[i, j]] = 0.5 * (m_hat[[i, j]] + m_hat[[j, i]]);
}
}
let (evals, evecs) = sym
.eigh(Side::Lower)
.map_err(|e| format!("recover_sphere_spikes: eigendecomposition failed: {e:?}"))?;
let order_desc: Vec<usize> = {
let mut idx: Vec<usize> = (0..d).collect();
idx.sort_by(|&a, &b| evals[b].total_cmp(&evals[a]));
idx
};
let eigenvalues: Vec<f64> = order_desc.iter().map(|&i| evals[i]).collect();
let lambda_max = eigenvalues.first().copied().unwrap_or(0.0);
let floor = sphere_eigenvalue_floor(lambda_max, d, sigma);
let model_order = eigenvalues
.iter()
.filter(|&&lam| lam > floor)
.count()
.min(d);
let mut spikes: Vec<SphereSpike> = Vec::with_capacity(model_order);
for k in 0..model_order {
let col = order_desc[k];
let mut direction: Vec<f64> = (0..d).map(|r| evecs[[r, col]]).collect();
let norm = direction.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm > 0.0 {
for x in direction.iter_mut() {
*x /= norm;
}
}
canonicalize_direction(&mut direction);
spikes.push(SphereSpike {
direction,
amplitude: eigenvalues[k],
});
}
let mut residual_sq = 0.0;
for i in 0..d {
for j in 0..d {
let mut fit = 0.0;
for spike in &spikes {
fit += spike.amplitude * spike.direction[i] * spike.direction[j];
}
let diff = m_hat[[i, j]] - fit;
residual_sq += diff * diff;
}
}
Ok(SphereRecovery {
spikes,
model_order,
residual: residual_sq.sqrt(),
eigenvalues,
})
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct TorusSpike {
pub theta: f64,
pub phi: f64,
pub amplitude: f64,
}
#[derive(Clone, Debug)]
pub struct TorusRecovery {
pub spikes: Vec<TorusSpike>,
pub model_order: usize,
pub residual: f64,
pub enhanced_singular_values: Vec<f64>,
}
pub fn torus_lift(spikes: &[TorusSpike], h1: usize, h2: usize) -> Result<Mat<c64>, String> {
if h1 == 0 || h2 == 0 {
return Err("torus_lift: both harmonic counts must be positive".into());
}
let grid = Mat::<c64>::from_fn(h1, h2, |a, b| {
let mut acc = c64::new(0.0, 0.0);
for spike in spikes {
let phase = TAU * (a as f64 * spike.theta + b as f64 * spike.phi);
acc += c64::new(spike.amplitude, 0.0) * c64::new(phase.cos(), phase.sin());
}
acc
});
Ok(grid)
}
fn matmul(a: &Mat<c64>, b: &Mat<c64>) -> Mat<c64> {
let (n, k, p) = (a.nrows(), a.ncols(), b.ncols());
Mat::<c64>::from_fn(n, p, |i, j| {
let mut acc = c64::new(0.0, 0.0);
for t in 0..k {
acc += a[(i, t)] * b[(t, j)];
}
acc
})
}
pub fn recover_torus_spikes(grid: &Mat<c64>, sigma: f64) -> Result<TorusRecovery, String> {
let h1 = grid.nrows();
let h2 = grid.ncols();
if h1 < 2 || h2 < 2 {
return Err(format!(
"recover_torus_spikes: need at least a 2×2 grid; got {h1}×{h2}"
));
}
for a in 0..h1 {
for b in 0..h2 {
if !grid[(a, b)].re.is_finite() || !grid[(a, b)].im.is_finite() {
return Err(format!("recover_torus_spikes: grid[{a},{b}] is not finite"));
}
}
}
let p = (h1 / 2 + 1).clamp(2, h1);
let k = (h2 / 2 + 1).clamp(2, h2);
let rows = k * p;
let cols = (h2 - k + 1) * (h1 - p + 1);
let max_order = rows.min(cols).min((p - 1) * k).min((k - 1) * p);
if max_order == 0 {
return Err(format!(
"recover_torus_spikes: grid {h1}×{h2} too small to resolve any spike"
));
}
let xe = Mat::<c64>::from_fn(rows, cols, |r, c| {
let (bk, pp) = (r / p, r % p);
let (bl, qq) = (c / (h1 - p + 1), c % (h1 - p + 1));
grid[(pp + qq, bk + bl)]
});
let svd = xe
.thin_svd()
.map_err(|e| format!("recover_torus_spikes: enhanced SVD failed: {e:?}"))?;
let singular_values: Vec<f64> = svd.S().column_vector().iter().map(|c| c.re).collect();
let model_order = singular_value_order(&singular_values, rows, cols, sigma, max_order);
if model_order == 0 {
let mut residual_sq = 0.0;
for a in 0..h1 {
for b in 0..h2 {
residual_sq += grid[(a, b)].norm_sqr();
}
}
return Ok(TorusRecovery {
spikes: Vec::new(),
model_order: 0,
residual: residual_sq.sqrt(),
enhanced_singular_values: singular_values,
});
}
let u = svd.U();
let es = Mat::<c64>::from_fn(rows, model_order, |r, c| u[(r, c)]);
let z_rows = (p - 1) * k;
let es_zup = Mat::<c64>::from_fn(z_rows, model_order, |r, c| {
let (bk, pp) = (r / (p - 1), r % (p - 1));
es[(bk * p + pp, c)]
});
let es_zdn = Mat::<c64>::from_fn(z_rows, model_order, |r, c| {
let (bk, pp) = (r / (p - 1), r % (p - 1));
es[(bk * p + pp + 1, c)]
});
let psi_z = es_zup.qr().solve_lstsq(&es_zdn);
let w_rows = (k - 1) * p;
let es_wup = Mat::<c64>::from_fn(w_rows, model_order, |r, c| es[(r, c)]);
let es_wdn = Mat::<c64>::from_fn(w_rows, model_order, |r, c| es[(r + p, c)]);
let psi_w = es_wup.qr().solve_lstsq(&es_wdn);
let gamma = c64::new(1.0, 0.618_033_988_749_895);
let psi_comb = Mat::<c64>::from_fn(model_order, model_order, |i, j| {
psi_z[(i, j)] + gamma * psi_w[(i, j)]
});
let eig = psi_comb
.eigen()
.map_err(|e| format!("recover_torus_spikes: joint eigenproblem failed: {e:?}"))?;
let e_vecs = eig.U().to_owned();
let bz = matmul(&psi_z, &e_vecs);
let dz = e_vecs.qr().solve_lstsq(&bz);
let bw = matmul(&psi_w, &e_vecs);
let dw = e_vecs.qr().solve_lstsq(&bw);
let mut z_phasors: Vec<c64> = Vec::with_capacity(model_order);
let mut w_phasors: Vec<c64> = Vec::with_capacity(model_order);
for j in 0..model_order {
let z = dz[(j, j)];
let w = dw[(j, j)];
let zn = z.norm();
let wn = w.norm();
z_phasors.push(if zn > 0.0 { z / zn } else { c64::new(1.0, 0.0) });
w_phasors.push(if wn > 0.0 { w / wn } else { c64::new(1.0, 0.0) });
}
let vander = Mat::<c64>::from_fn(h1 * h2, model_order, |r, j| {
let (a, b) = (r / h2, r % h2);
z_phasors[j].powu(a as u32) * w_phasors[j].powu(b as u32)
});
let rhs = Mat::<c64>::from_fn(h1 * h2, 1, |r, _| {
let (a, b) = (r / h2, r % h2);
grid[(a, b)]
});
let amps = vander.qr().solve_lstsq(&rhs);
let mut spikes: Vec<TorusSpike> = (0..model_order)
.map(|j| {
let theta = {
let t = z_phasors[j].arg() / TAU;
if t < 0.0 { t + 1.0 } else { t }
};
let phi = {
let t = w_phasors[j].arg() / TAU;
if t < 0.0 { t + 1.0 } else { t }
};
TorusSpike {
theta,
phi,
amplitude: amps[(j, 0)].re,
}
})
.collect();
spikes.sort_by(|a, b| a.theta.total_cmp(&b.theta).then(a.phi.total_cmp(&b.phi)));
let mut residual_sq = 0.0;
for a in 0..h1 {
for b in 0..h2 {
let mut fit = c64::new(0.0, 0.0);
for spike in &spikes {
let zp = c64::new((TAU * spike.theta).cos(), (TAU * spike.theta).sin());
let wp = c64::new((TAU * spike.phi).cos(), (TAU * spike.phi).sin());
fit += c64::new(spike.amplitude, 0.0) * zp.powu(a as u32) * wp.powu(b as u32);
}
residual_sq += (grid[(a, b)] - fit).norm_sqr();
}
}
Ok(TorusRecovery {
spikes,
model_order,
residual: residual_sq.sqrt(),
enhanced_singular_values: singular_values,
})
}
#[derive(Clone, Debug)]
pub struct PolishOptions {
pub max_iters: usize,
pub initial_damping: f64,
pub residual_tol: f64,
}
impl Default for PolishOptions {
fn default() -> Self {
Self {
max_iters: 64,
initial_damping: 1e-6,
residual_tol: 1e-12,
}
}
}
#[derive(Clone, Debug)]
pub struct PolishState {
pub amplitudes: Vec<f64>,
pub coords: Vec<Vec<f64>>,
}
#[derive(Clone, Debug)]
pub struct PolishResult {
pub state: PolishState,
pub residual: f64,
pub iterations: usize,
pub converged: bool,
}
pub fn polish_spikes<F>(
observed: &[f64],
phi: F,
init: PolishState,
opts: &PolishOptions,
) -> Result<PolishResult, String>
where
F: Fn(&[f64]) -> (Array1<f64>, Array2<f64>),
{
let big_d = observed.len();
let m = init.amplitudes.len();
if m == 0 || init.coords.len() != m {
return Err(format!(
"polish_spikes: amplitudes ({}) and coords ({}) must have equal positive length",
m,
init.coords.len()
));
}
if big_d == 0 {
return Err("polish_spikes: observed code must be non-empty".into());
}
let d = init.coords[0].len();
if d == 0 || init.coords.iter().any(|c| c.len() != d) {
return Err("polish_spikes: all latent coords must share one positive dimension".into());
}
if observed.iter().any(|x| !x.is_finite()) {
return Err("polish_spikes: observed code has non-finite entries".into());
}
let n_params = m * (d + 1);
let eval =
|state: &PolishState| -> Result<(Vec<f64>, f64, Vec<(Array1<f64>, Array2<f64>)>), String> {
let mut fit = vec![0.0_f64; big_d];
let mut cache = Vec::with_capacity(m);
for j in 0..m {
let (phi_j, jac_j) = phi(&state.coords[j]);
if phi_j.len() != big_d || jac_j.dim() != (big_d, d) {
return Err(format!(
"polish_spikes: phi returned Φ len {} / J {:?}, expected {big_d} / ({big_d}, {d})",
phi_j.len(),
jac_j.dim()
));
}
let a = state.amplitudes[j];
for i in 0..big_d {
fit[i] += a * phi_j[i];
}
cache.push((phi_j, jac_j));
}
let mut r = vec![0.0_f64; big_d];
let mut nsq = 0.0;
for i in 0..big_d {
r[i] = observed[i] - fit[i];
nsq += r[i] * r[i];
}
Ok((r, nsq.sqrt(), cache))
};
let mut state = init;
let (mut r, mut res_norm, mut cache) = eval(&state)?;
let lm_cfg = LmConfig::quartic(opts.initial_damping, 1e-12, 24);
let mut lm = LmState::new(&lm_cfg);
let mut converged = false;
let mut iterations = 0;
for _ in 0..opts.max_iters {
iterations += 1;
let mut jdata = vec![0.0_f64; big_d * n_params];
for j in 0..m {
let (phi_j, jac_j) = &cache[j];
let a = state.amplitudes[j];
let base = j * (d + 1);
for i in 0..big_d {
jdata[i * n_params + base] = phi_j[i];
for c in 0..d {
jdata[i * n_params + base + 1 + c] = a * jac_j[[i, c]];
}
}
}
let outcome = lm_step(
&mut lm,
&lm_cfg,
|mu: f64| {
let sqrt_mu = mu.sqrt();
let aug = Mat::<f64>::from_fn(big_d + n_params, n_params, |row, col| {
if row < big_d {
jdata[row * n_params + col]
} else if row - big_d == col {
sqrt_mu
} else {
0.0
}
});
let rhs = Mat::<f64>::from_fn(big_d + n_params, 1, |row, _| {
if row < big_d { r[row] } else { 0.0 }
});
let delta = aug.qr().solve_lstsq(&rhs);
let mut trial = state.clone();
for j in 0..m {
let base = j * (d + 1);
trial.amplitudes[j] += delta[(base, 0)];
for c in 0..d {
trial.coords[j][c] += delta[(base + 1 + c, 0)];
}
}
let (r_new, res_new, cache_new) = eval(&trial)?;
Ok::<_, String>(Some((trial, r_new, res_new, cache_new)))
},
|&(_, _, res_new, _): &_| res_new < res_norm,
)?;
match outcome {
LmOutcome::Accepted { candidate, .. } => {
let (trial, r_new, res_new, cache_new) = candidate;
let improvement = res_norm - res_new;
state = trial;
r = r_new;
cache = cache_new;
res_norm = res_new;
if improvement < opts.residual_tol {
converged = true;
}
}
LmOutcome::Exhausted { .. } => {
converged = true;
}
}
if converged {
break;
}
}
Ok(PolishResult {
state,
residual: res_norm,
iterations,
converged,
})
}
#[cfg(test)]
mod tests {
use super::*;
use rand::RngExt as _;
use rand::SeedableRng;
use rand::rngs::StdRng;
fn gaussian(rng: &mut StdRng) -> f64 {
let u1 = rng.random::<f64>().max(1e-16);
let u2 = rng.random::<f64>();
(-2.0 * u1.ln()).sqrt() * (TAU * u2).cos()
}
fn basis_dir(i: usize, d: usize) -> Vec<f64> {
let mut v = vec![0.0; d];
v[i] = 1.0;
v
}
fn dir_dist(a: &[f64], b: &[f64]) -> f64 {
let dot: f64 = a.iter().zip(b).map(|(x, y)| x * y).sum();
(1.0 - dot.abs()).max(0.0)
}
#[test]
fn sphere_noiseless_roundtrip() {
let d = 4;
let planted = vec![
SphereSpike {
direction: basis_dir(0, d),
amplitude: 2.0,
},
SphereSpike {
direction: basis_dir(1, d),
amplitude: 1.3,
},
SphereSpike {
direction: basis_dir(2, d),
amplitude: 0.7,
},
];
let m = sphere_lift(&planted, d).expect("lift");
let rec = recover_sphere_spikes(m.view(), 0.0).expect("recover");
assert_eq!(rec.model_order, 3, "order from clean spectrum");
for (r, p) in rec.spikes.iter().zip(planted.iter()) {
assert!(dir_dist(&r.direction, &p.direction) < 1e-8, "direction");
assert!((r.amplitude - p.amplitude).abs() < 1e-8, "amplitude");
}
assert!(rec.residual < 1e-8, "residual {:.3e}", rec.residual);
}
#[test]
fn sphere_antipodal_gauge_deterministic() {
let d = 3;
let raw = vec![0.6, -0.8, 0.0];
let neg = vec![-0.6, 0.8, 0.0];
let m_pos = sphere_lift(
&[SphereSpike {
direction: raw.clone(),
amplitude: 1.5,
}],
d,
)
.unwrap();
let m_neg = sphere_lift(
&[SphereSpike {
direction: neg,
amplitude: 1.5,
}],
d,
)
.unwrap();
let rec_pos = recover_sphere_spikes(m_pos.view(), 0.0).unwrap();
let rec_neg = recover_sphere_spikes(m_neg.view(), 0.0).unwrap();
assert_eq!(rec_pos.spikes.len(), 1);
assert_eq!(rec_neg.spikes.len(), 1);
let mut canon = raw.clone();
canonicalize_direction(&mut canon);
for k in 0..d {
assert!((rec_pos.spikes[0].direction[k] - canon[k]).abs() < 1e-10);
assert!((rec_neg.spikes[0].direction[k] - canon[k]).abs() < 1e-10);
}
}
#[test]
fn sphere_multiplicity_detection() {
let d = 5;
for m in 1..=3 {
let amps = [2.5, 1.7, 1.0];
let planted: Vec<SphereSpike> = (0..m)
.map(|i| SphereSpike {
direction: basis_dir(i, d),
amplitude: amps[i],
})
.collect();
let mut lift = sphere_lift(&planted, d).unwrap();
let mut rng = StdRng::seed_from_u64(100 + m as u64);
let sigma = 1e-3;
for i in 0..d {
for j in i..d {
let e = sigma * gaussian(&mut rng);
lift[[i, j]] += e;
if i != j {
lift[[j, i]] += e;
}
}
}
let rec = recover_sphere_spikes(lift.view(), sigma).unwrap();
assert_eq!(rec.model_order, m, "multiplicity m={m}");
}
}
#[test]
fn sphere_noise_recovery() {
let d = 3;
let sigma = 0.05;
let planted = vec![
SphereSpike {
direction: basis_dir(0, d),
amplitude: 3.0,
},
SphereSpike {
direction: basis_dir(1, d),
amplitude: 1.5,
},
];
let mut lift = sphere_lift(&planted, d).unwrap();
let mut rng = StdRng::seed_from_u64(2024);
for i in 0..d {
for j in i..d {
let e = sigma * gaussian(&mut rng);
lift[[i, j]] += e;
if i != j {
lift[[j, i]] += e;
}
}
}
let rec = recover_sphere_spikes(lift.view(), sigma).unwrap();
assert_eq!(rec.model_order, 2, "order under moderate noise");
for (r, p) in rec.spikes.iter().zip(planted.iter()) {
assert!(
dir_dist(&r.direction, &p.direction) < 0.05,
"direction {:.3e}",
dir_dist(&r.direction, &p.direction)
);
assert!(
(r.amplitude - p.amplitude).abs() < 0.5,
"amplitude {:.3e}",
(r.amplitude - p.amplitude).abs()
);
}
}
fn torus_dist(a: f64, b: f64) -> f64 {
let d = (a - b).abs();
d.min(1.0 - d)
}
fn torus_match_error(rec: &[TorusSpike], planted: &[TorusSpike]) -> (f64, f64) {
assert_eq!(rec.len(), planted.len(), "spike-count mismatch");
let mut max_pos = 0.0_f64;
let mut max_amp = 0.0_f64;
let mut used = vec![false; planted.len()];
for r in rec {
let mut best = usize::MAX;
let mut best_d = f64::INFINITY;
for (j, p) in planted.iter().enumerate() {
if used[j] {
continue;
}
let dd = torus_dist(r.theta, p.theta).max(torus_dist(r.phi, p.phi));
if dd < best_d {
best_d = dd;
best = j;
}
}
used[best] = true;
max_pos = max_pos.max(best_d);
max_amp = max_amp.max((r.amplitude - planted[best].amplitude).abs());
}
(max_pos, max_amp)
}
#[test]
fn torus_noiseless_roundtrip_m1_m2() {
let (h1, h2) = (6, 6);
for planted in [
vec![TorusSpike {
theta: 0.23,
phi: 0.61,
amplitude: 1.4,
}],
vec![
TorusSpike {
theta: 0.15,
phi: 0.72,
amplitude: 1.0,
},
TorusSpike {
theta: 0.63,
phi: 0.28,
amplitude: 0.8,
},
],
] {
let grid = torus_lift(&planted, h1, h2).unwrap();
let rec = recover_torus_spikes(&grid, 0.0).unwrap();
assert_eq!(rec.model_order, planted.len(), "order");
let (pos_err, amp_err) = torus_match_error(&rec.spikes, &planted);
assert!(pos_err < 1e-8, "position error {pos_err:.3e}");
assert!(amp_err < 1e-8, "amplitude error {amp_err:.3e}");
assert!(rec.residual < 1e-8, "residual {:.3e}", rec.residual);
}
}
#[test]
fn torus_multiplicity_detection() {
let (h1, h2) = (6, 6);
let candidates = [
vec![TorusSpike {
theta: 0.30,
phi: 0.40,
amplitude: 1.2,
}],
vec![
TorusSpike {
theta: 0.12,
phi: 0.70,
amplitude: 1.1,
},
TorusSpike {
theta: 0.55,
phi: 0.20,
amplitude: 0.9,
},
],
vec![
TorusSpike {
theta: 0.10,
phi: 0.15,
amplitude: 1.3,
},
TorusSpike {
theta: 0.45,
phi: 0.62,
amplitude: 1.0,
},
TorusSpike {
theta: 0.80,
phi: 0.35,
amplitude: 0.8,
},
],
];
for planted in candidates {
let grid = torus_lift(&planted, h1, h2).unwrap();
let rec = recover_torus_spikes(&grid, 0.0).unwrap();
assert_eq!(rec.model_order, planted.len(), "multiplicity");
let (pos_err, _) = torus_match_error(&rec.spikes, &planted);
assert!(pos_err < 1e-7, "position error {pos_err:.3e}");
}
}
#[test]
fn torus_adversarial_pairing() {
let (h1, h2) = (7, 7);
let planted = vec![
TorusSpike {
theta: 0.20,
phi: 0.75,
amplitude: 1.0,
},
TorusSpike {
theta: 0.60,
phi: 0.25,
amplitude: 0.9,
},
];
let grid = torus_lift(&planted, h1, h2).unwrap();
let rec = recover_torus_spikes(&grid, 0.0).unwrap();
assert_eq!(rec.model_order, 2, "order");
let (pos_err, _) = torus_match_error(&rec.spikes, &planted);
assert!(
pos_err < 1e-7,
"correct-pairing position error {pos_err:.3e}"
);
let ghost = [
TorusSpike {
theta: 0.20,
phi: 0.25,
amplitude: 1.0,
},
TorusSpike {
theta: 0.60,
phi: 0.75,
amplitude: 0.9,
},
];
let (ghost_err, _) = torus_match_error(&rec.spikes, &ghost);
assert!(
ghost_err > 0.3,
"recovery must not be the ghost pairing (err {ghost_err:.3e})"
);
}
#[test]
fn torus_noise_recovery() {
let (h1, h2) = (6, 6);
let sigma = 0.05;
let planted = vec![
TorusSpike {
theta: 0.18,
phi: 0.66,
amplitude: 1.2,
},
TorusSpike {
theta: 0.62,
phi: 0.24,
amplitude: 1.0,
},
];
let mut grid = torus_lift(&planted, h1, h2).unwrap();
let mut rng = StdRng::seed_from_u64(7);
for a in 0..h1 {
for b in 0..h2 {
let re = grid[(a, b)].re + sigma * gaussian(&mut rng);
let im = grid[(a, b)].im + sigma * gaussian(&mut rng);
grid[(a, b)] = c64::new(re, im);
}
}
let rec = recover_torus_spikes(&grid, sigma).unwrap();
assert_eq!(rec.model_order, 2, "order under noise");
let (pos_err, amp_err) = torus_match_error(&rec.spikes, &planted);
assert!(pos_err < 0.03, "position error {pos_err:.3e}");
assert!(amp_err < 0.3, "amplitude error {amp_err:.3e}");
}
fn circle_phi(h_max: usize) -> impl Fn(&[f64]) -> (Array1<f64>, Array2<f64>) {
move |t: &[f64]| {
let t0 = t[0];
let d = 2 * h_max;
let mut phi = Array1::<f64>::zeros(d);
let mut jac = Array2::<f64>::zeros((d, 1));
for h in 1..=h_max {
let w = TAU * h as f64;
let (s, c) = (w * t0).sin_cos();
phi[2 * (h - 1)] = c;
phi[2 * (h - 1) + 1] = s;
jac[[2 * (h - 1), 0]] = -w * s;
jac[[2 * (h - 1) + 1, 0]] = w * c;
}
(phi, jac)
}
}
fn circle_code(spikes: &[(f64, f64)], h_max: usize) -> Vec<f64> {
let d = 2 * h_max;
let mut z = vec![0.0; d];
for &(t, a) in spikes {
for h in 1..=h_max {
let w = TAU * h as f64;
z[2 * (h - 1)] += a * (w * t).cos();
z[2 * (h - 1) + 1] += a * (w * t).sin();
}
}
z
}
#[test]
fn polish_converges_from_perturbed_seed() {
let h_max = 6;
let true_spikes = [(0.22, 1.1), (0.61, 0.8)];
let z = circle_code(&true_spikes, h_max);
let init = PolishState {
amplitudes: vec![1.1 + 0.12, 0.8 - 0.1],
coords: vec![vec![0.22 + 0.02], vec![0.61 - 0.025]],
};
let (_, seed_res, _) = {
let phi = circle_phi(h_max);
let mut fit = vec![0.0; z.len()];
for j in 0..init.amplitudes.len() {
let (p, _) = phi(&init.coords[j]);
for i in 0..z.len() {
fit[i] += init.amplitudes[j] * p[i];
}
}
let nsq: f64 = z.iter().zip(&fit).map(|(a, b)| (a - b) * (a - b)).sum();
((), nsq.sqrt(), ())
};
let res =
polish_spikes(&z, circle_phi(h_max), init, &PolishOptions::default()).expect("polish");
assert!(res.converged, "should converge");
assert!(res.residual < 1e-6, "residual {:.3e}", res.residual);
assert!(res.residual < seed_res, "polish must reduce the residual");
}
#[test]
fn polish_rejects_shape_mismatch() {
let h_max = 4;
let z = circle_code(&[(0.3, 1.0)], h_max);
let bad = |_t: &[f64]| (Array1::<f64>::zeros(3), Array2::<f64>::zeros((3, 1)));
let init = PolishState {
amplitudes: vec![1.0],
coords: vec![vec![0.3]],
};
assert!(polish_spikes(&z, bad, init, &PolishOptions::default()).is_err());
}
}