use crate::kernel::{Complex, Float};
use crate::prelude::*;
pub struct AliasStage<T: Float> {
pub b_count: usize,
pub coeff0: Vec<Complex<T>>,
pub coeff1: Vec<Complex<T>>,
}
#[inline]
fn abs_f64<T: Float>(c: Complex<T>) -> f64 {
let re = c.re.to_f64().unwrap_or(0.0);
let im = c.im.to_f64().unwrap_or(0.0);
libm::sqrt(re * re + im * im)
}
fn subtract_frequency<T: Float>(
stages: &mut [AliasStage<T>],
n: usize,
f: usize,
coeff: Complex<T>,
) {
let theta = <T as Float>::TWO_PI * T::from_usize(f) / T::from_usize(n);
let (sin_t, cos_t) = Float::sin_cos(theta);
let w = Complex::new(cos_t, sin_t);
let coeff_w = coeff * w;
for stage in stages.iter_mut() {
let b = f % stage.b_count;
stage.coeff0[b] = stage.coeff0[b] - coeff;
stage.coeff1[b] = stage.coeff1[b] - coeff_w;
}
}
pub fn certified_peel<T: Float>(
stages: &mut [AliasStage<T>],
n: usize,
max_recover: usize,
abs_threshold: f64,
) -> Option<Vec<(usize, Complex<T>)>> {
if stages.is_empty() || n == 0 {
return None;
}
let mut peak = 0.0_f64;
for stage in stages.iter() {
for c in stage.coeff0.iter() {
let m = abs_f64(*c);
if m > peak {
peak = m;
}
}
}
if peak <= abs_threshold {
return Some(Vec::new());
}
let occ_floor = (1e-9 * peak).max(abs_threshold);
let phase_tol = 1e-6_f64;
let cert_tol = 1e-6 * peak;
let two_pi = core::f64::consts::PI * 2.0;
let mut recovered: Vec<(usize, Complex<T>)> = Vec::new();
let iter_cap = max_recover.saturating_mul(4).saturating_add(8);
let mut iterations = 0usize;
'peel: loop {
iterations += 1;
if iterations > iter_cap || recovered.len() >= max_recover {
break;
}
for si in 0..stages.len() {
let b_count = stages[si].b_count;
for b in 0..b_count {
let c0 = stages[si].coeff0[b];
let m0 = abs_f64(c0);
if m0 <= occ_floor {
continue;
}
let c1 = stages[si].coeff1[b];
let a0 = libm::atan2(c0.im.to_f64().unwrap_or(0.0), c0.re.to_f64().unwrap_or(0.0));
let a1 = libm::atan2(c1.im.to_f64().unwrap_or(0.0), c1.re.to_f64().unwrap_or(0.0));
let dtheta = (a1 - a0).rem_euclid(two_pi);
let f_est = (libm::round(dtheta * (n as f64) / two_pi) as usize) % n;
if f_est % b_count != b {
continue;
}
let m1 = abs_f64(c1);
if (m1 - m0).abs() > phase_tol * m0 {
continue;
}
if recovered.iter().any(|(ff, _)| *ff == f_est) {
continue;
}
subtract_frequency(stages, n, f_est, c0);
recovered.push((f_est, c0));
continue 'peel;
}
}
break;
}
let mut residual = 0.0_f64;
for stage in stages.iter() {
for c in stage.coeff0.iter().chain(stage.coeff1.iter()) {
let m = abs_f64(*c);
if m > residual {
residual = m;
}
}
}
if residual <= cert_tol {
Some(recovered)
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
fn stage_from_spectrum(spectrum: &[Complex<f64>], b_count: usize) -> AliasStage<f64> {
let n = spectrum.len();
let two_pi = core::f64::consts::PI * 2.0;
let mut coeff0 = vec![Complex::new(0.0, 0.0); b_count];
let mut coeff1 = vec![Complex::new(0.0, 0.0); b_count];
for (f, &x) in spectrum.iter().enumerate() {
let b = f % b_count;
let theta = two_pi * (f as f64) / (n as f64);
let w = Complex::new(theta.cos(), theta.sin());
coeff0[b] = coeff0[b] + x;
coeff1[b] = coeff1[b] + x * w;
}
AliasStage {
b_count,
coeff0,
coeff1,
}
}
#[test]
fn test_single_tone_recovered_exactly() {
let n = 256;
let mut spectrum = vec![Complex::new(0.0, 0.0); n];
spectrum[37] = Complex::new(3.0, -1.5);
let mut stages = vec![stage_from_spectrum(&spectrum, 16)];
let recovered =
certified_peel(&mut stages, n, n, 1e-10).expect("clean single tone must certify");
assert_eq!(recovered.len(), 1);
assert_eq!(recovered[0].0, 37);
assert!((recovered[0].1.re - 3.0).abs() < 1e-9);
assert!((recovered[0].1.im - (-1.5)).abs() < 1e-9);
}
#[test]
fn test_multiple_isolated_tones() {
let n = 512;
let planted = [
(11usize, Complex::new(1.0, 0.0)),
(200, Complex::new(-2.0, 0.5)),
(301, Complex::new(0.25, 0.75)),
];
let mut spectrum = vec![Complex::new(0.0, 0.0); n];
for &(f, v) in &planted {
spectrum[f] = v;
}
let mut stages = vec![
stage_from_spectrum(&spectrum, 32),
stage_from_spectrum(&spectrum, 64),
];
let recovered =
certified_peel(&mut stages, n, n, 1e-10).expect("isolated tones must certify");
assert_eq!(recovered.len(), 3);
for &(f, v) in &planted {
let found = recovered
.iter()
.find(|(rf, _)| *rf == f)
.expect("planted frequency must be recovered");
assert!((found.1.re - v.re).abs() < 1e-9);
assert!((found.1.im - v.im).abs() < 1e-9);
}
}
#[test]
fn test_permanent_collision_fails_certification() {
let n = 256;
let mut spectrum = vec![Complex::new(0.0, 0.0); n];
spectrum[10] = Complex::new(1.0, 0.0);
spectrum[10 + 128] = Complex::new(0.5, 0.3);
let mut stages = vec![
stage_from_spectrum(&spectrum, 16),
stage_from_spectrum(&spectrum, 32),
];
let result = certified_peel(&mut stages, n, n, 1e-10);
assert!(result.is_none(), "an unresolved collision must not certify");
}
#[test]
fn test_silent_signal_certifies_empty() {
let n = 64;
let spectrum = vec![Complex::new(0.0, 0.0); n];
let mut stages = vec![stage_from_spectrum(&spectrum, 16)];
let recovered =
certified_peel(&mut stages, n, n, 1e-10).expect("silent signal certifies as empty");
assert!(recovered.is_empty());
}
#[test]
fn test_empty_stages_returns_none() {
let mut stages: Vec<AliasStage<f64>> = Vec::new();
assert!(certified_peel(&mut stages, 64, 64, 1e-10).is_none());
}
#[test]
fn test_sub_threshold_signal_ignored() {
let n = 64;
let mut spectrum = vec![Complex::new(0.0, 0.0); n];
spectrum[5] = Complex::new(1e-12, 0.0);
let mut stages = vec![stage_from_spectrum(&spectrum, 16)];
let recovered =
certified_peel(&mut stages, n, n, 1e-10).expect("sub-threshold certifies empty");
assert!(recovered.is_empty());
}
}