use std::sync::Arc;
use num_complex::Complex;
use num_traits::{FromPrimitive, Zero};
use common::{FFTnum, verify_length, verify_length_divisible};
use math_utils;
use twiddles;
use ::{Length, IsInverse, FFT};
pub struct RadersAlgorithm<T> {
inner_fft: Arc<FFT<T>>,
inner_fft_data: Box<[Complex<T>]>,
input_map: Box<[usize]>,
output_map: Box<[usize]>,
}
impl<T: FFTnum> RadersAlgorithm<T> {
pub fn new(len: usize, inner_fft: Arc<FFT<T>>) -> Self {
assert_eq!(len - 1, inner_fft.len(), "For raders algorithm, inner_fft.len() must be self.len() - 1. Expected {}, got {}", len - 1, inner_fft.len());
let inner_fft_len = len - 1;
let primitive_root = math_utils::primitive_root(len as u64).unwrap();
let root_inverse = math_utils::multiplicative_inverse(primitive_root, len as u64);
let unity_scale: T = FromPrimitive::from_f64(1f64 / inner_fft_len as f64).unwrap();
let mut inner_fft_input: Vec<Complex<T>> = (0..inner_fft_len)
.map(|i| math_utils::modular_exponent(root_inverse, i as u64, len as u64) as usize)
.map(|i| twiddles::single_twiddle(i, len, inner_fft.is_inverse()))
.map(|c| c * unity_scale)
.collect();
let mut inner_fft_output = vec![Zero::zero(); inner_fft_len];
inner_fft.process(&mut inner_fft_input, &mut inner_fft_output);
let input_map: Vec<usize> = (0..len-1).map(|i| math_utils::modular_exponent(primitive_root, (i + 1) as u64, len as u64) as usize - 1).collect();
let output_map: Vec<usize> = (0..len-1).map(|i| math_utils::modular_exponent(root_inverse, (i + 1) as u64, len as u64) as usize - 1).collect();
RadersAlgorithm {
inner_fft: inner_fft,
inner_fft_data: inner_fft_output.into_boxed_slice(),
input_map: input_map.into_boxed_slice(),
output_map: output_map.into_boxed_slice(),
}
}
fn perform_fft(&self, input: &mut [Complex<T>], output: &mut [Complex<T>]) {
let (first_input, input) = input.split_first_mut().unwrap();
let first_input_val = *first_input;
let (first_output, output) = output.split_first_mut().unwrap();
let input_sum: Complex<T> = input.iter().fold(Zero::zero(), |acc, &e| acc + e);
*first_output = first_input_val + input_sum;
for (&input_index, output_element) in self.input_map.iter().zip(output.iter_mut()) {
*output_element = input[input_index];
}
self.inner_fft.process(output, input);
for ((&input_cell, output_cell), &multiple) in input.iter().zip(output.iter_mut()).zip(self.inner_fft_data.iter()) {
*output_cell = (input_cell * multiple).conj();
}
self.inner_fft.process(output, input);
for element in input.iter_mut() {
*element = element.conj() + first_input_val;
}
for (&output_index, input_element) in self.output_map.iter().zip(input.iter()) {
output[output_index] = *input_element;
}
}
}
impl<T: FFTnum> FFT<T> for RadersAlgorithm<T> {
fn process(&self, input: &mut [Complex<T>], output: &mut [Complex<T>]) {
verify_length(input, output, self.len());
self.perform_fft(input, output);
}
fn process_multi(&self, input: &mut [Complex<T>], output: &mut [Complex<T>]) {
verify_length_divisible(input, output, self.len());
for (in_chunk, out_chunk) in input.chunks_mut(self.len()).zip(output.chunks_mut(self.len())) {
self.perform_fft(in_chunk, out_chunk);
}
}
}
impl<T> Length for RadersAlgorithm<T> {
#[inline(always)]
fn len(&self) -> usize {
self.inner_fft_data.len() + 1
}
}
impl<T> IsInverse for RadersAlgorithm<T> {
#[inline(always)]
fn is_inverse(&self) -> bool {
self.inner_fft.is_inverse()
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
use std::sync::Arc;
use test_utils::check_fft_algorithm;
use algorithm::DFT;
#[test]
fn test_raders() {
for &len in &[3,5,7,11,13] {
test_raders_with_length(len, false);
test_raders_with_length(len, true);
}
}
fn test_raders_with_length(len: usize, inverse: bool) {
let inner_fft = Arc::new(DFT::new(len - 1, inverse));
let fft = RadersAlgorithm::new(len, inner_fft);
check_fft_algorithm(&fft, len, inverse);
}
}