use faer::prelude::*;
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Spike {
pub t: f64,
pub amplitude: f64,
}
#[derive(Clone, Debug)]
pub struct SpikeRecovery {
pub spikes: Vec<Spike>,
pub model_order: usize,
pub residual: f64,
pub hankel_singular_values: Vec<f64>,
}
pub fn separation_limit(n_harmonics: usize) -> f64 {
const CFG_CONSTANT: f64 = 2.0;
if n_harmonics == 0 {
return f64::INFINITY;
}
CFG_CONSTANT / (n_harmonics as f64)
}
pub fn recover_spikes(fourier_coeffs: &[(f64, f64)], sigma: f64) -> Result<SpikeRecovery, String> {
let n = fourier_coeffs.len();
if n < 2 {
return Err(format!(
"super-resolution needs at least 2 harmonics to resolve a spike; got {n}"
));
}
for (h, &(c, s)) in fourier_coeffs.iter().enumerate() {
if !c.is_finite() || !s.is_finite() {
return Err(format!("fourier_coeffs[{h}] = ({c}, {s}) is not finite"));
}
}
let samples: Vec<c64> = fourier_coeffs
.iter()
.map(|&(c, s)| c64::new(c, s))
.collect();
let l = n / 2;
let rows = n - l;
let cols = l + 1;
let max_order = l;
let hankel = Mat::<c64>::from_fn(rows, cols, |i, k| samples[i + k]);
let svd = hankel
.thin_svd()
.map_err(|e| format!("Hankel SVD failed to converge: {e:?}"))?;
let singular_values: Vec<f64> = svd.S().column_vector().iter().map(|c| c.re).collect();
let threshold = order_threshold(&singular_values, rows, cols, sigma);
let threshold_order = singular_values
.iter()
.filter(|&&s| s > threshold)
.count()
.min(max_order);
let dramatic_gap = (f32::EPSILON as f64).sqrt();
let gap_order = {
let mut 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 {
order = Some(k);
break;
}
}
order
};
let model_order = gap_order.unwrap_or(threshold_order);
if model_order == 0 {
let residual = samples.iter().map(|y| y.norm_sqr()).sum::<f64>().sqrt();
return Ok(SpikeRecovery {
spikes: Vec::new(),
model_order: 0,
residual,
hankel_singular_values: singular_values,
});
}
let v = svd.V();
let v_m = v.submatrix(0, 0, cols, model_order);
let v1 = v_m.submatrix(0, 0, l, model_order);
let v2 = v_m.submatrix(1, 0, l, model_order);
let phi = v1.qr().solve_lstsq(v2);
let roots = phi
.eigenvalues()
.map_err(|e| format!("matrix-pencil eigenproblem failed: {e:?}"))?;
let unit_phasors: Vec<c64> = roots
.iter()
.map(|z| {
let norm = z.norm();
if norm > 0.0 {
*z / norm
} else {
c64::new(1.0, 0.0)
}
})
.collect();
let solve_branch = |phasors: &[c64]| -> (Vec<Spike>, f64) {
let positions: Vec<f64> = phasors
.iter()
.map(|z| {
let t = z.arg() / std::f64::consts::TAU;
if t < 0.0 { t + 1.0 } else { t }
})
.collect();
let vander = Mat::<c64>::from_fn(n, model_order, |k, j| phasors[j].powu((k + 1) as u32));
let rhs = Mat::<c64>::from_fn(n, 1, |k, _| samples[k]);
let amps = vander.qr().solve_lstsq(&rhs);
let mut spikes: Vec<Spike> = (0..model_order)
.map(|j| Spike {
t: positions[j],
amplitude: amps[(j, 0)].re,
})
.collect();
spikes.sort_by(|a, b| a.t.total_cmp(&b.t));
let mut residual_sq = 0.0;
for (k, y) in samples.iter().enumerate() {
let mut fit = c64::new(0.0, 0.0);
for spike in &spikes {
let phasor = c64::new(
(std::f64::consts::TAU * spike.t).cos(),
(std::f64::consts::TAU * spike.t).sin(),
);
fit += phasor.powu((k + 1) as u32) * spike.amplitude;
}
residual_sq += (y - fit).norm_sqr();
}
(spikes, residual_sq)
};
let conj_phasors: Vec<c64> = unit_phasors.iter().map(|z| z.conj()).collect();
let (spikes_direct, residual_direct) = solve_branch(&unit_phasors);
let (spikes_reflected, residual_reflected) = solve_branch(&conj_phasors);
let (spikes, residual_sq) = if residual_reflected < residual_direct {
(spikes_reflected, residual_reflected)
} else {
(spikes_direct, residual_direct)
};
Ok(SpikeRecovery {
spikes,
model_order,
residual: residual_sq.sqrt(),
hankel_singular_values: singular_values,
})
}
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 order_threshold(singular_values: &[f64], rows: usize, cols: usize, sigma: f64) -> f64 {
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;
if sigma > 0.0 {
let beta = rows.min(cols) as f64 / n_big;
let sigma_entry = std::f64::consts::SQRT_2 * sigma;
let gd = optimal_hard_threshold_coefficient(beta) * sigma_entry * n_big.sqrt();
gd.max(numerical_floor)
} else {
numerical_floor
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::RngExt as _;
use rand::SeedableRng;
use rand::rngs::StdRng;
fn coeffs_from_spikes(spikes: &[(f64, f64)], n_harmonics: usize) -> Vec<(f64, f64)> {
(1..=n_harmonics)
.map(|h| {
let mut c = 0.0;
let mut s = 0.0;
for &(t, a) in spikes {
let phase = std::f64::consts::TAU * (h as f64) * t;
c += a * phase.cos();
s += a * phase.sin();
}
(c, s)
})
.collect()
}
fn add_noise(coeffs: &mut [(f64, f64)], sigma: f64, seed: u64) {
let mut rng = StdRng::seed_from_u64(seed);
for coeff in coeffs.iter_mut() {
coeff.0 += sigma * gaussian(&mut rng);
coeff.1 += sigma * gaussian(&mut rng);
}
}
fn gaussian(rng: &mut StdRng) -> f64 {
let u1 = rng.random::<f64>().max(1e-16);
let u2 = rng.random::<f64>();
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn circle_dist(a: f64, b: f64) -> f64 {
let d = (a - b).abs();
d.min(1.0 - d)
}
fn match_error(recovered: &[Spike], planted: &[(f64, f64)]) -> (f64, f64) {
assert_eq!(recovered.len(), planted.len(), "spike-count mismatch");
let mut planted_sorted = planted.to_vec();
planted_sorted.sort_by(|a, b| a.0.total_cmp(&b.0));
let mut max_t = 0.0_f64;
let mut max_a = 0.0_f64;
for (rec, &(t, a)) in recovered.iter().zip(planted_sorted.iter()) {
max_t = max_t.max(circle_dist(rec.t, t));
max_a = max_a.max((rec.amplitude - a).abs());
}
(max_t, max_a)
}
#[test]
fn exact_recovery_two_spikes_no_noise() {
let h = 8;
let planted = [(0.20, 1.0), (0.45, 0.7)];
let coeffs = coeffs_from_spikes(&planted, h);
let rec = recover_spikes(&coeffs, 0.0).expect("recovery");
assert_eq!(rec.model_order, 2, "order from clean singular values");
let (t_err, a_err) = match_error(&rec.spikes, &planted);
assert!(t_err < 1e-9, "position error {t_err:.3e}");
assert!(a_err < 1e-9, "amplitude error {a_err:.3e}");
assert!(rec.residual < 1e-9, "residual {:.3e}", rec.residual);
}
#[test]
fn separation_law_brackets_the_limit() {
let h = 8;
let sigma = 1e-3;
let base = 1.0 / (h as f64);
let mut errors = Vec::new();
for factor in [0.5_f64, 1.0, 2.0] {
let sep = factor * base;
let planted = [(0.30, 1.0), (0.30 + sep, 1.0)];
let mut coeffs = coeffs_from_spikes(&planted, h);
add_noise(&mut coeffs, sigma, 0xC0FFEE + (factor * 1000.0) as u64);
let rec = recover_spikes(&coeffs, sigma).expect("recovery");
let err = if rec.model_order == 2 {
match_error(&rec.spikes, &planted).0
} else {
sep.max(base)
};
errors.push(err);
}
let (err_half, _err_one, err_two) = (errors[0], errors[1], errors[2]);
assert!(
err_two < 0.02,
"position error at 2/H should be small, got {err_two:.3e}"
);
assert!(
err_half > err_two,
"0.5/H error {err_half:.3e} should exceed 2/H error {err_two:.3e}"
);
}
#[test]
fn model_order_selection_from_singular_values() {
let h = 8;
let sigma = 1e-2;
let planted1 = [(0.35, 1.0)];
let mut c1 = coeffs_from_spikes(&planted1, h);
add_noise(&mut c1, sigma, 11);
let rec1 = recover_spikes(&c1, sigma).expect("recovery m=1");
assert_eq!(rec1.model_order, 1, "should select order 1");
let planted3 = [(0.10, 1.0), (0.40, 0.9), (0.75, 1.1)];
let mut c3 = coeffs_from_spikes(&planted3, h);
add_noise(&mut c3, sigma, 22);
let rec3 = recover_spikes(&c3, sigma).expect("recovery m=3");
assert_eq!(rec3.model_order, 3, "should select order 3");
let (t_err, _) = match_error(&rec3.spikes, &planted3);
assert!(t_err < 0.05, "order-3 position error {t_err:.3e}");
}
#[test]
fn noise_robustness_two_spikes() {
let h = 8;
let sigma = 0.05;
let planted = [(0.20, 1.0), (0.65, 1.0)];
let mut coeffs = coeffs_from_spikes(&planted, h);
add_noise(&mut coeffs, sigma, 777);
let rec = recover_spikes(&coeffs, sigma).expect("recovery");
assert_eq!(rec.model_order, 2, "order under moderate noise");
let (t_err, a_err) = match_error(&rec.spikes, &planted);
assert!(
t_err < 10.0 * sigma / (h as f64),
"position error {t_err:.3e}"
);
assert!(a_err < 5.0 * sigma, "amplitude error {a_err:.3e}");
}
#[test]
fn separation_limit_is_two_over_h() {
assert!((separation_limit(8) - 0.25).abs() < 1e-15);
assert!((separation_limit(16) - 0.125).abs() < 1e-15);
assert!(separation_limit(0).is_infinite());
}
#[test]
fn too_few_harmonics_errors() {
assert!(recover_spikes(&[(1.0, 0.0)], 0.0).is_err());
assert!(recover_spikes(&[], 0.0).is_err());
}
}