use num_complex::Complex;
use num_traits::Zero;
use crate::common::FFTnum;
use crate::algorithm::butterflies::{Butterfly2, Butterfly4, Butterfly8, Butterfly16};
use crate::{Length, IsInverse, Fft};
pub struct Radix4<T> {
twiddles: Box<[Complex<T>]>,
butterfly8: Butterfly8<T>,
butterfly16: Butterfly16<T>,
len: usize,
inverse: bool,
}
impl<T: FFTnum> Radix4<T> {
pub fn new(len: usize, inverse: bool) -> Self {
assert!(len.is_power_of_two(), "Radix4 algorithm requires a power-of-two input size. Got {}", len);
let num_bits = len.trailing_zeros();
let mut twiddle_stride = if num_bits%2 == 0 {
len / 64
} else {
len / 32
};
let mut twiddle_factors = Vec::with_capacity(len * 2);
while twiddle_stride > 0 {
let num_rows = len / (twiddle_stride * 4);
for i in 0..num_rows {
for k in 1..4 {
let twiddle = T::generate_twiddle_factor(i * k * twiddle_stride, len, inverse);
twiddle_factors.push(twiddle);
}
}
twiddle_stride >>= 2;
}
Radix4 {
twiddles: twiddle_factors.into_boxed_slice(),
butterfly8: Butterfly8::new(inverse),
butterfly16: Butterfly16::new(inverse),
len: len,
inverse: inverse,
}
}
fn perform_fft_out_of_place(&self, signal: &[Complex<T>], spectrum: &mut [Complex<T>], _scratch: &mut [Complex<T>]) {
match self.len() {
0|1 => spectrum.copy_from_slice(signal),
2 => {
spectrum.copy_from_slice(signal);
unsafe { Butterfly2::new(self.inverse).perform_fft_butterfly(spectrum) }
},
4 => {
spectrum.copy_from_slice(signal);
unsafe { Butterfly4::new(self.inverse).perform_fft_butterfly(spectrum) }
},
_ => {
prepare_radix4(signal.len(), signal, spectrum, 1);
let num_bits = signal.len().trailing_zeros();
let mut current_size = if num_bits % 2 == 0 {
self.butterfly16.process_inplace_multi(spectrum, &mut []);
64
} else {
self.butterfly8.process_inplace_multi(spectrum, &mut []);
32
};
let mut layer_twiddles: &[Complex<T>] = &self.twiddles;
while current_size <= signal.len() {
let num_rows = signal.len() / current_size;
for i in 0..num_rows {
unsafe {
butterfly_4(&mut spectrum[i * current_size..],
layer_twiddles,
current_size / 4,
self.inverse)
}
}
let twiddle_offset = (current_size * 3) / 4;
layer_twiddles = &layer_twiddles[twiddle_offset..];
current_size *= 4;
}
}
}
}
}
boilerplate_fft_oop!(Radix4, |this: &Radix4<_>| this.len);
fn prepare_radix4<T: FFTnum>(size: usize,
signal: &[Complex<T>],
spectrum: &mut [Complex<T>],
stride: usize) {
match size {
16 => unsafe {
for i in 0..16 {
*spectrum.get_unchecked_mut(i) = *signal.get_unchecked(i * stride);
}
},
8 => unsafe {
for i in 0..8 {
*spectrum.get_unchecked_mut(i) = *signal.get_unchecked(i * stride);
}
},
4 => unsafe {
for i in 0..4 {
*spectrum.get_unchecked_mut(i) = *signal.get_unchecked(i * stride);
}
},
2 => unsafe {
for i in 0..2 {
*spectrum.get_unchecked_mut(i) = *signal.get_unchecked(i * stride);
}
},
_ => {
for i in 0..4 {
prepare_radix4(size / 4,
&signal[i * stride..],
&mut spectrum[i * (size / 4)..],
stride * 4);
}
}
}
}
unsafe fn butterfly_4<T: FFTnum>(data: &mut [Complex<T>],
twiddles: &[Complex<T>],
num_ffts: usize,
inverse: bool)
{
let mut idx = 0usize;
let mut tw_idx = 0usize;
let mut scratch: [Complex<T>; 6] = [Zero::zero(); 6];
for _ in 0..num_ffts {
scratch[0] = data.get_unchecked(idx + 1 * num_ffts) * twiddles[tw_idx];
scratch[1] = data.get_unchecked(idx + 2 * num_ffts) * twiddles[tw_idx + 1];
scratch[2] = data.get_unchecked(idx + 3 * num_ffts) * twiddles[tw_idx + 2];
scratch[5] = data.get_unchecked(idx) - scratch[1];
*data.get_unchecked_mut(idx) = data.get_unchecked(idx) + scratch[1];
scratch[3] = scratch[0] + scratch[2];
scratch[4] = scratch[0] - scratch[2];
*data.get_unchecked_mut(idx + 2 * num_ffts) = data.get_unchecked(idx) - scratch[3];
*data.get_unchecked_mut(idx) = data.get_unchecked(idx) + scratch[3];
if inverse {
data.get_unchecked_mut(idx + num_ffts).re = scratch[5].re - scratch[4].im;
data.get_unchecked_mut(idx + num_ffts).im = scratch[5].im + scratch[4].re;
data.get_unchecked_mut(idx + 3 * num_ffts).re = scratch[5].re + scratch[4].im;
data.get_unchecked_mut(idx + 3 * num_ffts).im = scratch[5].im - scratch[4].re;
} else {
data.get_unchecked_mut(idx + num_ffts).re = scratch[5].re + scratch[4].im;
data.get_unchecked_mut(idx + num_ffts).im = scratch[5].im - scratch[4].re;
data.get_unchecked_mut(idx + 3 * num_ffts).re = scratch[5].re - scratch[4].im;
data.get_unchecked_mut(idx + 3 * num_ffts).im = scratch[5].im + scratch[4].re;
}
tw_idx += 3;
idx += 1;
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
use crate::test_utils::check_fft_algorithm;
#[test]
fn test_radix4() {
for pow in 0..8 {
let len = 1 << pow;
test_radix4_with_length(len, false);
test_radix4_with_length(len, true);
}
}
fn test_radix4_with_length(len: usize, inverse: bool) {
let fft = Radix4::new(len, inverse);
check_fft_algorithm::<f32>(&fft, len, inverse);
}
}