use super::fft::{FftFloat, FftPlanner};
pub struct RfftPlanner<T: FftFloat> {
pub(crate) n: usize,
pub(crate) fft_n2: FftPlanner<T>,
pub(crate) post_twiddle_re: Vec<T>,
pub(crate) post_twiddle_im: Vec<T>,
pub(crate) scratch_re: Vec<T>,
pub(crate) scratch_im: Vec<T>,
}
impl<T: FftFloat> RfftPlanner<T> {
pub fn new(n: usize) -> Self {
assert!(n > 0, "RFFT size must be positive");
assert!(
n.is_power_of_two(),
"RFFT size must be a power of two, got {n}"
);
let n_half = n / 2;
let fft_n2 = FftPlanner::new(n_half);
let tau = T::tau();
let n_t = T::from_usize(n);
let mut post_twiddle_re = Vec::with_capacity(n_half + 1);
let mut post_twiddle_im = Vec::with_capacity(n_half + 1);
for k in 0..=n_half {
let angle = tau * T::from_usize(k) * n_t.recip();
post_twiddle_re.push(angle.cos());
post_twiddle_im.push(-angle.sin());
}
let scratch_re = vec![T::from_usize(0); n_half];
let scratch_im = vec![T::from_usize(0); n_half];
Self {
n,
fft_n2,
post_twiddle_re,
post_twiddle_im,
scratch_re,
scratch_im,
}
}
#[inline]
pub fn len(&self) -> usize {
self.n
}
#[inline]
pub fn is_empty(&self) -> bool {
self.n == 0
}
pub fn process_forward(&mut self, input: &[T], out_re: &mut [T], out_im: &mut [T]) {
let n = self.n;
let n_half = n / 2;
let expected = n_half + 1;
debug_assert_eq!(input.len(), n, "input length mismatch");
debug_assert_eq!(out_re.len(), expected, "out_re length mismatch");
debug_assert_eq!(out_im.len(), expected, "out_im length mismatch");
for i in 0..n_half {
self.scratch_re[i] = input[2 * i];
self.scratch_im[i] = input[2 * i + 1];
}
self.fft_n2
.process(&mut self.scratch_re, &mut self.scratch_im);
out_re[0] = self.scratch_re[0] + self.scratch_im[0];
out_im[0] = T::from_usize(0);
if n_half > 1 {
let half = T::from_usize(2).recip();
for k in 1..n_half {
let nk = n_half - k;
let re_k = self.scratch_re[k];
let im_k = self.scratch_im[k];
let re_nk = self.scratch_re[nk];
let im_nk = self.scratch_im[nk];
let even_re = half * (re_k + re_nk);
let even_im = half * (im_k - im_nk);
let odd_re = half * (re_k - re_nk);
let odd_im = half * (im_k + im_nk);
let w_re = self.post_twiddle_re[k];
let w_im = self.post_twiddle_im[k];
out_re[k] = even_re + odd_re.mul_add(w_im, odd_im * w_re);
out_im[k] = even_im - odd_re.mul_add(w_re, -odd_im * w_im);
}
}
out_re[n_half] = self.scratch_re[0] - self.scratch_im[0];
out_im[n_half] = T::from_usize(0);
}
pub fn process_inverse(&self, in_re: &mut [T], in_im: &mut [T], out: &mut [T]) {
let n = self.n;
let n_half = n / 2;
let expected = n_half + 1;
debug_assert_eq!(in_re.len(), expected, "in_re length mismatch");
debug_assert_eq!(in_im.len(), expected, "in_im length mismatch");
debug_assert_eq!(out.len(), n, "output length mismatch");
let two = T::from_usize(2);
let half = two.recip();
{
let x0_re = in_re[0];
let xn2_re = in_re[n_half];
in_re[0] = (x0_re + xn2_re) * half;
in_im[0] = (x0_re - xn2_re) * half;
}
let quarter_n = n / 4;
for k in 1..quarter_n {
let nk = n_half - k;
let xk_re = in_re[k];
let xk_im = in_im[k];
let xnk_re = in_re[nk];
let xnk_im = in_im[nk];
let even_re = (xk_re + xnk_re) * half;
let even_im = (xk_im - xnk_im) * half;
let diff_re = (xk_re - xnk_re) * half;
let sum_im = (xk_im + xnk_im) * half;
let w_re = self.post_twiddle_re[k];
let w_im = self.post_twiddle_im[k];
let odd_re = w_im.mul_add(diff_re, -w_re * sum_im);
let odd_im = w_re.mul_add(diff_re, w_im * sum_im);
in_re[k] = even_re + odd_re;
in_im[k] = even_im + odd_im;
in_re[nk] = even_re - odd_re;
in_im[nk] = -even_im + odd_im;
}
if quarter_n > 0 && n_half.is_multiple_of(2) {
let k_mid = quarter_n;
let x_re = in_re[k_mid];
let x_im = in_im[k_mid];
in_re[k_mid] = x_re;
in_im[k_mid] = -x_im;
}
self.fft_n2
.process_inverse(&mut in_re[..n_half], &mut in_im[..n_half]);
for k in 0..n_half {
out[2 * k] = in_re[k];
out[2 * k + 1] = in_im[k];
}
}
}