use crate::error::{CsError, CsResult};
use crate::handle::LcgRng;
#[inline]
fn cmul(ar: f64, ai: f64, br: f64, bi: f64) -> (f64, f64) {
(ar * br - ai * bi, ar * bi + ai * br)
}
fn dft(x: &[f64], n: usize) -> Vec<f64> {
let mut out = vec![0.0_f64; 2 * n];
let two_pi = std::f64::consts::TAU;
for k in 0..n {
let mut sr = 0.0_f64;
let mut si = 0.0_f64;
for j in 0..n {
let angle = -two_pi * (j as f64) * (k as f64) / (n as f64);
let (wr, wi) = (angle.cos(), angle.sin());
let xr = x[2 * j];
let xi = x[2 * j + 1];
let (pr, pi) = cmul(xr, xi, wr, wi);
sr += pr;
si += pi;
}
out[2 * k] = sr;
out[2 * k + 1] = si;
}
out
}
fn idft(x: &[f64], n: usize) -> Vec<f64> {
let mut out = vec![0.0_f64; 2 * n];
let two_pi = std::f64::consts::TAU;
let inv_n = 1.0 / (n as f64);
for k in 0..n {
let mut sr = 0.0_f64;
let mut si = 0.0_f64;
for j in 0..n {
let angle = two_pi * (j as f64) * (k as f64) / (n as f64);
let (wr, wi) = (angle.cos(), angle.sin());
let xr = x[2 * j];
let xi = x[2 * j + 1];
let (pr, pi) = cmul(xr, xi, wr, wi);
sr += pr;
si += pi;
}
out[2 * k] = sr * inv_n;
out[2 * k + 1] = si * inv_n;
}
out
}
#[derive(Debug, Clone)]
pub struct CodedDiffraction {
pub n: usize,
pub n_masks: usize,
masks: Vec<f64>,
}
#[derive(Debug, Clone, Copy)]
pub enum MaskKind {
Octanary,
UniformPhase,
Rademacher,
}
impl CodedDiffraction {
pub fn new(n: usize, n_masks: usize, kind: MaskKind, rng: &mut LcgRng) -> CsResult<Self> {
if n == 0 || n_masks == 0 {
return Err(CsError::InvalidParameter(
"coded diffraction: n and n_masks must be > 0".into(),
));
}
let mut masks = vec![0.0_f64; n_masks * 2 * n];
for l in 0..n_masks {
for j in 0..n {
let (re, im) = sample_mask_entry(kind, rng);
masks[l * 2 * n + 2 * j] = re;
masks[l * 2 * n + 2 * j + 1] = im;
}
}
Ok(Self { n, n_masks, masks })
}
#[must_use]
pub fn n_measurements(&self) -> usize {
self.n_masks * self.n
}
fn mask(&self, l: usize) -> &[f64] {
&self.masks[l * 2 * self.n..(l + 1) * 2 * self.n]
}
fn apply(&self, l: usize, x: &[f64]) -> Vec<f64> {
let d = self.mask(l);
let mut mod_x = vec![0.0_f64; 2 * self.n];
for j in 0..self.n {
let (pr, pi) = cmul(d[2 * j], d[2 * j + 1], x[2 * j], x[2 * j + 1]);
mod_x[2 * j] = pr;
mod_x[2 * j + 1] = pi;
}
dft(&mod_x, self.n)
}
fn apply_adjoint(&self, l: usize, z: &[f64]) -> Vec<f64> {
let d = self.mask(l);
let inv = idft(z, self.n);
let mut out = vec![0.0_f64; 2 * self.n];
for j in 0..self.n {
let (pr, pi) = cmul(d[2 * j], -d[2 * j + 1], inv[2 * j], inv[2 * j + 1]);
out[2 * j] = pr;
out[2 * j + 1] = pi;
}
out
}
pub fn forward(&self, x: &[f64]) -> CsResult<Vec<f64>> {
if x.len() != 2 * self.n {
return Err(CsError::DimensionMismatch {
a: x.len(),
b: 2 * self.n,
});
}
let mut y = vec![0.0_f64; self.n_measurements()];
for l in 0..self.n_masks {
let ax = self.apply(l, x);
for k in 0..self.n {
let re = ax[2 * k];
let im = ax[2 * k + 1];
y[l * self.n + k] = re * re + im * im;
}
}
Ok(y)
}
pub fn wirtinger_flow(
&self,
y: &[f64],
cfg: &WirtingerConfig,
rng: &mut LcgRng,
) -> CsResult<Vec<f64>> {
let m = self.n_measurements();
if y.len() != m {
return Err(CsError::DimensionMismatch { a: y.len(), b: m });
}
let sum_y: f64 = y.iter().sum();
let lambda_sq = (self.n as f64) * sum_y / (m as f64);
let lambda = lambda_sq.max(1e-30).sqrt();
let mut z = random_unit_complex(self.n, rng);
for _ in 0..cfg.power_iters {
let yz = self.apply_y(y, &z, m);
let nrm = cnorm(&yz);
if nrm < 1e-300 {
return Err(CsError::NumericalInstability(
"WF spectral init: degenerate leading eigenvector".into(),
));
}
for v in z.iter_mut() {
*v = 0.0;
}
for j in 0..2 * self.n {
z[j] = yz[j] / nrm;
}
}
for v in z.iter_mut() {
*v *= lambda;
}
let step = cfg.step_size / lambda_sq.max(1e-30);
for _ in 0..cfg.max_iter {
let grad = self.wf_gradient(y, &z, m);
for j in 0..2 * self.n {
z[j] -= step * grad[j];
}
}
Ok(z)
}
fn apply_y(&self, y: &[f64], z: &[f64], m: usize) -> Vec<f64> {
let inv_m = 1.0 / (m as f64);
let mut acc = vec![0.0_f64; 2 * self.n];
for l in 0..self.n_masks {
let az = self.apply(l, z); let mut w = vec![0.0_f64; 2 * self.n];
for k in 0..self.n {
let yr = y[l * self.n + k];
w[2 * k] = yr * az[2 * k];
w[2 * k + 1] = yr * az[2 * k + 1];
}
let contrib = self.apply_adjoint(l, &w);
for j in 0..2 * self.n {
acc[j] += contrib[j];
}
}
for v in acc.iter_mut() {
*v *= inv_m;
}
acc
}
fn wf_gradient(&self, y: &[f64], z: &[f64], m: usize) -> Vec<f64> {
let inv_m = 1.0 / (m as f64);
let mut grad = vec![0.0_f64; 2 * self.n];
for l in 0..self.n_masks {
let az = self.apply(l, z);
let mut resid = vec![0.0_f64; 2 * self.n];
for k in 0..self.n {
let re = az[2 * k];
let im = az[2 * k + 1];
let mag_sq = re * re + im * im;
let factor = mag_sq - y[l * self.n + k]; resid[2 * k] = factor * re;
resid[2 * k + 1] = factor * im;
}
let contrib = self.apply_adjoint(l, &resid);
for j in 0..2 * self.n {
grad[j] += contrib[j];
}
}
for v in grad.iter_mut() {
*v *= inv_m;
}
grad
}
}
#[derive(Debug, Clone)]
pub struct WirtingerConfig {
pub power_iters: usize,
pub max_iter: usize,
pub step_size: f64,
}
impl Default for WirtingerConfig {
fn default() -> Self {
Self {
power_iters: 50,
max_iter: 400,
step_size: 0.2,
}
}
}
fn cnorm(x: &[f64]) -> f64 {
x.iter().map(|v| v * v).sum::<f64>().sqrt()
}
fn random_unit_complex(n: usize, rng: &mut LcgRng) -> Vec<f64> {
let mut z = vec![0.0_f64; 2 * n];
for v in z.iter_mut() {
*v = rng.next_normal();
}
let nrm = cnorm(&z).max(1e-300);
for v in z.iter_mut() {
*v /= nrm;
}
z
}
fn sample_mask_entry(kind: MaskKind, rng: &mut LcgRng) -> (f64, f64) {
match kind {
MaskKind::Octanary => {
let phase = rng.next_usize(4);
let (pr, pi) = match phase {
0 => (1.0, 0.0),
1 => (-1.0, 0.0),
2 => (0.0, 1.0),
_ => (0.0, -1.0),
};
let b2 = if rng.next_f64() < 0.8 {
std::f64::consts::FRAC_1_SQRT_2
} else {
3.0_f64.sqrt()
};
(pr * b2, pi * b2)
}
MaskKind::UniformPhase => {
let theta = std::f64::consts::TAU * rng.next_f64();
(theta.cos(), theta.sin())
}
MaskKind::Rademacher => {
if rng.next_bool() {
(1.0, 0.0)
} else {
(-1.0, 0.0)
}
}
}
}
pub fn phase_aligned_error(x_hat: &[f64], x: &[f64]) -> CsResult<f64> {
if x_hat.len() != x.len() {
return Err(CsError::DimensionMismatch {
a: x_hat.len(),
b: x.len(),
});
}
let mut ipr = 0.0_f64;
let mut ipi = 0.0_f64;
let n = x.len() / 2;
for k in 0..n {
let (ar, ai) = (x_hat[2 * k], x_hat[2 * k + 1]);
let (br, bi) = (x[2 * k], x[2 * k + 1]);
ipr += ar * br + ai * bi;
ipi += ar * bi - ai * br;
}
let mag = (ipr * ipr + ipi * ipi).sqrt();
let (cr, ci) = if mag > 1e-300 {
(ipr / mag, ipi / mag)
} else {
(1.0, 0.0)
};
let mut err_sq = 0.0_f64;
let mut x_sq = 0.0_f64;
for k in 0..n {
let (ar, ai) = (x_hat[2 * k], x_hat[2 * k + 1]);
let (rr, ri) = cmul(ar, ai, cr, ci);
let dr = rr - x[2 * k];
let di = ri - x[2 * k + 1];
err_sq += dr * dr + di * di;
x_sq += x[2 * k] * x[2 * k] + x[2 * k + 1] * x[2 * k + 1];
}
Ok((err_sq / x_sq.max(1e-300)).sqrt())
}
#[cfg(test)]
mod tests {
use super::*;
fn complex_signal(reals: &[f64], imags: &[f64]) -> Vec<f64> {
let mut x = Vec::with_capacity(reals.len() * 2);
for (r, i) in reals.iter().zip(imags.iter()) {
x.push(*r);
x.push(*i);
}
x
}
#[test]
fn dft_idft_round_trip() {
let n = 6;
let x = complex_signal(
&[1.0, -2.0, 3.0, 0.5, -1.0, 2.0],
&[0.0, 1.0, -1.0, 0.0, 2.0, -0.5],
);
let back = idft(&dft(&x, n), n);
for k in 0..2 * n {
assert!(
(back[k] - x[k]).abs() < 1e-9,
"k={k}: {} vs {}",
back[k],
x[k]
);
}
}
#[test]
fn constructor_rejects_zero_dims() {
let mut rng = LcgRng::new(1);
assert!(CodedDiffraction::new(0, 4, MaskKind::Octanary, &mut rng).is_err());
assert!(CodedDiffraction::new(8, 0, MaskKind::Octanary, &mut rng).is_err());
}
#[test]
fn forward_shapes() {
let mut rng = LcgRng::new(2);
let cdp = CodedDiffraction::new(8, 5, MaskKind::UniformPhase, &mut rng).expect("ok");
assert_eq!(cdp.n_measurements(), 40);
let x = vec![0.1_f64; 16];
let y = cdp.forward(&x).expect("ok");
assert_eq!(y.len(), 40);
assert!(y.iter().all(|&v| v >= 0.0), "intensities non-negative");
}
#[test]
fn forward_dimension_mismatch() {
let mut rng = LcgRng::new(3);
let cdp = CodedDiffraction::new(8, 4, MaskKind::Octanary, &mut rng).expect("ok");
assert!(matches!(
cdp.forward(&[0.0; 10]),
Err(CsError::DimensionMismatch { .. })
));
}
#[test]
fn forward_phase_invariant_intensity() {
let mut rng = LcgRng::new(4);
let cdp = CodedDiffraction::new(6, 4, MaskKind::UniformPhase, &mut rng).expect("ok");
let x = complex_signal(
&[1.0, 2.0, -1.0, 0.5, 0.0, 3.0],
&[0.0, -1.0, 1.0, 0.0, 2.0, 0.0],
);
let y0 = cdp.forward(&x).expect("ok");
let (cr, ci) = (
std::f64::consts::FRAC_PI_3.cos(),
std::f64::consts::FRAC_PI_3.sin(),
);
let mut x_rot = vec![0.0_f64; x.len()];
for k in 0..6 {
let (rr, ri) = cmul(x[2 * k], x[2 * k + 1], cr, ci);
x_rot[2 * k] = rr;
x_rot[2 * k + 1] = ri;
}
let y1 = cdp.forward(&x_rot).expect("ok");
for (a, b) in y0.iter().zip(y1.iter()) {
assert!((a - b).abs() < 1e-8, "{a} vs {b}");
}
}
#[test]
fn phase_aligned_error_zero_for_rotation() {
let x = complex_signal(&[1.0, -2.0, 0.5], &[0.5, 1.0, -1.0]);
let (cr, ci) = (0.6_f64, 0.8_f64); let mut x_rot = vec![0.0_f64; x.len()];
for k in 0..3 {
let (rr, ri) = cmul(x[2 * k], x[2 * k + 1], cr, ci);
x_rot[2 * k] = rr;
x_rot[2 * k + 1] = ri;
}
let err = phase_aligned_error(&x_rot, &x).expect("ok");
assert!(err < 1e-9, "err = {err}");
}
#[test]
fn phase_aligned_error_dim_mismatch() {
assert!(matches!(
phase_aligned_error(&[1.0, 0.0], &[1.0, 0.0, 0.0, 0.0]),
Err(CsError::DimensionMismatch { .. })
));
}
#[test]
fn wirtinger_flow_recovers_real_signal() {
let n = 5usize;
let mut rng = LcgRng::new(20);
let cdp = CodedDiffraction::new(n, 10, MaskKind::Octanary, &mut rng).expect("ok");
let x = complex_signal(&[1.0, -1.0, 0.5, 2.0, -0.5], &[0.0; 5]);
let y = cdp.forward(&x).expect("ok");
let cfg = WirtingerConfig {
power_iters: 120,
max_iter: 2500,
step_size: 0.15,
};
let mut rng2 = LcgRng::new(99);
let x_hat = cdp.wirtinger_flow(&y, &cfg, &mut rng2).expect("ok");
let err = phase_aligned_error(&x_hat, &x).expect("ok");
assert!(err < 0.15, "relative error = {err}");
}
#[test]
fn wirtinger_flow_recovers_complex_signal() {
let n = 4usize;
let mut rng = LcgRng::new(31);
let cdp = CodedDiffraction::new(n, 8, MaskKind::Octanary, &mut rng).expect("ok");
let x = complex_signal(&[1.0, 0.0, -1.0, 0.5], &[0.5, 1.0, 0.0, -1.0]);
let y = cdp.forward(&x).expect("ok");
let cfg = WirtingerConfig {
power_iters: 100,
max_iter: 2000,
step_size: 0.15,
};
let mut rng2 = LcgRng::new(7);
let x_hat = cdp.wirtinger_flow(&y, &cfg, &mut rng2).expect("ok");
let err = phase_aligned_error(&x_hat, &x).expect("ok");
assert!(err < 0.2, "relative error = {err}");
}
#[test]
fn wirtinger_flow_dimension_mismatch() {
let mut rng = LcgRng::new(5);
let cdp = CodedDiffraction::new(8, 4, MaskKind::Octanary, &mut rng).expect("ok");
let mut rng2 = LcgRng::new(6);
assert!(matches!(
cdp.wirtinger_flow(&[0.0; 10], &WirtingerConfig::default(), &mut rng2),
Err(CsError::DimensionMismatch { .. })
));
}
#[test]
fn recovered_signal_reproduces_measurements() {
let n = 4usize;
let mut rng = LcgRng::new(40);
let cdp = CodedDiffraction::new(n, 7, MaskKind::Octanary, &mut rng).expect("ok");
let x = complex_signal(&[2.0, -1.0, 0.0, 1.0], &[0.0, 0.5, -0.5, 0.0]);
let y = cdp.forward(&x).expect("ok");
let cfg = WirtingerConfig {
power_iters: 100,
max_iter: 2000,
step_size: 0.15,
};
let mut rng2 = LcgRng::new(8);
let x_hat = cdp.wirtinger_flow(&y, &cfg, &mut rng2).expect("ok");
let y_hat = cdp.forward(&x_hat).expect("ok");
let mut num = 0.0_f64;
let mut den = 0.0_f64;
for (a, b) in y_hat.iter().zip(y.iter()) {
num += (a - b) * (a - b);
den += b * b;
}
let rel = (num / den.max(1e-30)).sqrt();
assert!(rel < 0.2, "measurement reproduction error = {rel}");
}
#[test]
fn rademacher_masks_are_real() {
let mut rng = LcgRng::new(50);
let cdp = CodedDiffraction::new(6, 3, MaskKind::Rademacher, &mut rng).expect("ok");
for l in 0..3 {
let m = cdp.mask(l);
for j in 0..6 {
assert_eq!(m[2 * j + 1], 0.0, "imag part must be zero");
assert!((m[2 * j].abs() - 1.0).abs() < 1e-12, "must be ±1");
}
}
}
#[test]
fn zero_signal_zero_measurements() {
let mut rng = LcgRng::new(60);
let cdp = CodedDiffraction::new(5, 4, MaskKind::Octanary, &mut rng).expect("ok");
let y = cdp.forward(&vec![0.0_f64; 10]).expect("ok");
assert!(y.iter().all(|&v| v == 0.0));
}
}