pub(crate) fn dft_in_place(re: &mut [f64], im: &mut [f64], inverse: bool) -> Result<(), String> {
if re.len() != im.len() {
return Err(format!(
"discrete Fourier transform: real and imaginary parts must have equal length, got {} and {}",
re.len(),
im.len()
));
}
let n = re.len();
if n <= 1 {
return Ok(());
}
if n.is_power_of_two() {
fft_radix2_in_place(re, im, inverse);
} else {
bluestein_in_place(re, im, inverse)?;
}
Ok(())
}
fn fft_radix2_in_place(re: &mut [f64], im: &mut [f64], inverse: bool) {
let n = re.len();
assert!(
n.is_power_of_two(),
"fft_radix2_in_place requires a power-of-two length, got {n}; the caller \
must route other lengths to bluestein_in_place"
);
let mut j = 0usize;
for i in 1..n {
let mut bit = n >> 1;
while j & bit != 0 {
j ^= bit;
bit >>= 1;
}
j |= bit;
if i < j {
re.swap(i, j);
im.swap(i, j);
}
}
let sign = if inverse { 1.0_f64 } else { -1.0_f64 };
let mut len = 2usize;
while len <= n {
let half = len / 2;
let step = sign * std::f64::consts::TAU / len as f64;
let mut base = 0usize;
while base < n {
for k in 0..half {
let angle = step * k as f64;
let (wi, wr) = angle.sin_cos();
let lo = base + k;
let hi = lo + half;
let vr = re[hi] * wr - im[hi] * wi;
let vi = re[hi] * wi + im[hi] * wr;
let ur = re[lo];
let ui = im[lo];
re[lo] = ur + vr;
im[lo] = ui + vi;
re[hi] = ur - vr;
im[hi] = ui - vi;
}
base += len;
}
len <<= 1;
}
}
fn bluestein_in_place(re: &mut [f64], im: &mut [f64], inverse: bool) -> Result<(), String> {
let n = re.len();
let conv_len = (2 * n - 1)
.checked_next_power_of_two()
.ok_or_else(|| format!("discrete Fourier transform: length {n} overflows the chirp-z convolution length"))?;
let sign = if inverse { 1.0_f64 } else { -1.0_f64 };
let two_n = 2 * n;
let chirp = |m: usize| -> (f64, f64) {
let residue = (m % two_n) * (m % two_n) % two_n;
let angle = sign * std::f64::consts::PI * residue as f64 / n as f64;
angle.sin_cos()
};
let mut a_re = vec![0.0_f64; conv_len];
let mut a_im = vec![0.0_f64; conv_len];
for t in 0..n {
let (s, c) = chirp(t);
a_re[t] = re[t] * c - im[t] * s;
a_im[t] = re[t] * s + im[t] * c;
}
let mut b_re = vec![0.0_f64; conv_len];
let mut b_im = vec![0.0_f64; conv_len];
for m in 0..n {
let (s, c) = chirp(m);
b_re[m] = c;
b_im[m] = -s;
if m > 0 {
b_re[conv_len - m] = c;
b_im[conv_len - m] = -s;
}
}
fft_radix2_in_place(&mut a_re, &mut a_im, false);
fft_radix2_in_place(&mut b_re, &mut b_im, false);
for idx in 0..conv_len {
let pr = a_re[idx] * b_re[idx] - a_im[idx] * b_im[idx];
let pi = a_re[idx] * b_im[idx] + a_im[idx] * b_re[idx];
a_re[idx] = pr;
a_im[idx] = pi;
}
fft_radix2_in_place(&mut a_re, &mut a_im, true);
let scale = 1.0 / conv_len as f64;
for k in 0..n {
let (s, c) = chirp(k);
let cr = a_re[k] * scale;
let ci = a_im[k] * scale;
re[k] = cr * c - ci * s;
im[k] = cr * s + ci * c;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn naive_dft(re: &[f64], im: &[f64], inverse: bool) -> (Vec<f64>, Vec<f64>) {
let n = re.len();
let sign = if inverse { 1.0_f64 } else { -1.0_f64 };
let mut out_re = vec![0.0_f64; n];
let mut out_im = vec![0.0_f64; n];
for k in 0..n {
for t in 0..n {
let angle = sign * std::f64::consts::TAU * (k as f64) * (t as f64) / (n as f64);
let (s, c) = angle.sin_cos();
out_re[k] += re[t] * c - im[t] * s;
out_im[k] += re[t] * s + im[t] * c;
}
}
(out_re, out_im)
}
fn sequence(n: usize, seed: u64) -> (Vec<f64>, Vec<f64>) {
let mut state = seed;
let mut next = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 11) as f64) / ((1u64 << 53) as f64) - 0.5
};
let re: Vec<f64> = (0..n).map(|_| next()).collect();
let im: Vec<f64> = (0..n).map(|_| next()).collect();
(re, im)
}
#[test]
fn fast_transform_matches_the_naive_dft_at_every_length_class() {
for &n in &[1usize, 2, 3, 4, 7, 8, 12, 13, 16, 31, 64, 96, 101, 128, 200] {
for inverse in [false, true] {
let (re0, im0) = sequence(n, 0xD1CE + n as u64);
let (want_re, want_im) = naive_dft(&re0, &im0, inverse);
let mut re = re0.clone();
let mut im = im0.clone();
dft_in_place(&mut re, &mut im, inverse).expect("equal-length halves");
let scale = want_re
.iter()
.chain(want_im.iter())
.fold(1.0_f64, |acc, v| acc.max(v.abs()));
let tol = 1.0e-12 * scale * (n as f64);
for k in 0..n {
assert!(
(re[k] - want_re[k]).abs() <= tol && (im[k] - want_im[k]).abs() <= tol,
"n={n} inverse={inverse} bin {k}: fast=({}, {}) naive=({}, {}) tol={tol:e}",
re[k],
im[k],
want_re[k],
want_im[k]
);
}
}
}
}
#[test]
fn forward_then_inverse_is_multiplication_by_the_length() {
for &n in &[5usize, 16, 33, 97, 256, 300] {
let (re0, im0) = sequence(n, 0xBEEF + n as u64);
let mut re = re0.clone();
let mut im = im0.clone();
dft_in_place(&mut re, &mut im, false).expect("forward");
dft_in_place(&mut re, &mut im, true).expect("inverse");
for k in 0..n {
let want_re = re0[k] * n as f64;
let want_im = im0[k] * n as f64;
let tol = 1.0e-10 * (n as f64);
assert!(
(re[k] - want_re).abs() <= tol && (im[k] - want_im).abs() <= tol,
"n={n} index {k}: round trip gave ({}, {}), expected ({want_re}, {want_im})",
re[k],
im[k]
);
}
}
}
#[test]
fn mismatched_halves_are_refused() {
let mut re = vec![0.0_f64; 4];
let mut im = vec![0.0_f64; 3];
assert!(dft_in_place(&mut re, &mut im, false).is_err());
}
}