use std::sync::Arc;
use std::cmp::max;
use num_integer::Integer;
use num_complex::Complex;
use strength_reduce::StrengthReducedUsize;
use transpose;
use crate::common::FFTnum;
use crate::array_utils;
use crate::{Length, IsInverse, Fft};
pub struct GoodThomasAlgorithm<T> {
width: usize,
width_size_fft: Arc<dyn Fft<T>>,
height: usize,
height_size_fft: Arc<dyn Fft<T>>,
input_x_stride: usize,
input_y_stride: usize,
inplace_scratch_len: usize,
outofplace_scratch_len: usize,
len: StrengthReducedUsize,
inverse: bool,
}
impl<T: FFTnum> GoodThomasAlgorithm<T> {
pub fn new(width_fft: Arc<dyn Fft<T>>, height_fft: Arc<dyn Fft<T>>) -> Self {
assert_eq!(
width_fft.is_inverse(), height_fft.is_inverse(),
"width_fft and height_fft must both be inverse, or neither. got width inverse={}, height inverse={}",
width_fft.is_inverse(), height_fft.is_inverse());
let width = width_fft.len();
let height = height_fft.len();
let is_inverse = width_fft.is_inverse();
let gcd_data = i64::extended_gcd(&(width as i64), &(height as i64));
assert!(gcd_data.gcd == 1,
"Invalid input width and height to Good-Thomas Algorithm: ({},{}): Inputs must be coprime",
width,
height);
let width_inverse = if gcd_data.x >= 0 { gcd_data.x } else { gcd_data.x + height as i64 } as usize;
let height_inverse = if gcd_data.y >= 0 { gcd_data.y } else { gcd_data.y + width as i64 } as usize;
let len = width * height;
let width_inplace_scratch = height_fft.get_inplace_scratch_len();
let height_inplace_scratch = width_fft.get_inplace_scratch_len();
let height_outofplace_scratch = width_fft.get_out_of_place_scratch_len();
let outofplace_scratch = max(height_inplace_scratch, width_inplace_scratch);
let inplace_extra = max(if width_inplace_scratch > len { width_inplace_scratch } else { 0 }, height_outofplace_scratch);
Self {
width: width,
width_size_fft: width_fft,
height: height,
height_size_fft: height_fft,
input_x_stride: height_inverse as usize * height,
input_y_stride: width_inverse as usize * width,
inplace_scratch_len: len + inplace_extra,
outofplace_scratch_len: if outofplace_scratch > len { outofplace_scratch } else { 0 },
len: StrengthReducedUsize::new(width * height),
inverse: is_inverse,
}
}
fn perform_fft_inplace(&self, buffer: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
let (scratch, inner_scratch) = scratch.split_at_mut(self.len());
for (y, row) in scratch.chunks_mut(self.width).enumerate() {
let input_base = y * self.input_y_stride;
for (x, output_cell) in row.iter_mut().enumerate() {
let input_index = (input_base + x * self.input_x_stride) % self.len;
*output_cell = buffer[input_index];
}
}
let width_scratch = if inner_scratch.len() > buffer.len() { &mut inner_scratch[..] } else { &mut buffer[..] };
self.width_size_fft.process_inplace_multi(scratch, width_scratch);
transpose::transpose(scratch, buffer, self.width, self.height);
self.height_size_fft.process_multi(buffer, scratch, inner_scratch);
for (x, row) in scratch.chunks(self.height).enumerate() {
let output_base = x * self.height;
for (y, input_cell) in row.iter().enumerate() {
let output_index = (output_base + y * self.width) % self.len;
buffer[output_index] = *input_cell;
}
}
}
fn perform_fft_out_of_place(&self, input: &mut [Complex<T>], output: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
for (y, row) in output.chunks_mut(self.width).enumerate() {
let input_base = y * self.input_y_stride;
for (x, output_cell) in row.iter_mut().enumerate() {
let input_index = (input_base + x * self.input_x_stride) % self.len;
*output_cell = input[input_index];
}
}
let width_scratch = if scratch.len() > input.len() { &mut scratch[..] } else { &mut input[..] };
self.width_size_fft.process_inplace_multi(output, width_scratch);
transpose::transpose(output, input, self.width, self.height);
let height_scratch = if scratch.len() > output.len() { &mut scratch[..] } else { &mut output[..] };
self.height_size_fft.process_inplace_multi(input, height_scratch);
for (x, row) in input.chunks(self.height).enumerate() {
let output_base = x * self.height;
for (y, input_cell) in row.iter().enumerate() {
let output_index = (output_base + y * self.width) % self.len;
output[output_index] = *input_cell;
}
}
}
}
boilerplate_fft!(GoodThomasAlgorithm,
|this: &GoodThomasAlgorithm<_>| this.len.get(),
|this: &GoodThomasAlgorithm<_>| this.inplace_scratch_len,
|this: &GoodThomasAlgorithm<_>| this.outofplace_scratch_len
);
pub struct GoodThomasAlgorithmSmall<T> {
width: usize,
width_size_fft: Arc<dyn Fft<T>>,
height: usize,
height_size_fft: Arc<dyn Fft<T>>,
input_output_map: Box<[usize]>,
inverse: bool,
}
impl<T: FFTnum> GoodThomasAlgorithmSmall<T> {
pub fn new(width_fft: Arc<dyn Fft<T>>, height_fft: Arc<dyn Fft<T>>) -> Self {
assert_eq!(
width_fft.is_inverse(), height_fft.is_inverse(),
"n1_fft and height_fft must both be inverse, or neither. got width inverse={}, height inverse={}",
width_fft.is_inverse(), height_fft.is_inverse());
let width = width_fft.len();
let height = height_fft.len();
let len = width * height;
let gcd_data = i64::extended_gcd(&(width as i64), &(height as i64));
assert!(gcd_data.gcd == 1,
"Invalid input width and height to Good-Thomas Algorithm: ({},{}): Inputs must be coprime",
width,
height);
let width_inverse = if gcd_data.x >= 0 { gcd_data.x } else { gcd_data.x + height as i64 } as usize;
let height_inverse = if gcd_data.y >= 0 { gcd_data.y } else { gcd_data.y + width as i64 } as usize;
let input_iter = (0..len)
.map(|i| (i % width, i / width))
.map(|(x, y)| (x * height + y * width) % len);
let output_iter = (0..len)
.map(|i| (i % height, i / height))
.map(|(y, x)| (x * height * height_inverse as usize + y * width * width_inverse as usize) % len);
let input_output_map: Vec<usize> = input_iter.chain(output_iter).collect();
Self {
inverse: width_fft.is_inverse(),
width: width,
width_size_fft: width_fft,
height: height,
height_size_fft: height_fft,
input_output_map: input_output_map.into_boxed_slice(),
}
}
fn perform_fft_out_of_place(&self, input: &mut [Complex<T>], output: &mut [Complex<T>], _scratch: &mut [Complex<T>]) {
assert_eq!(self.len(), input.len());
assert_eq!(self.len(), output.len());
let (input_map, output_map) = self.input_output_map.split_at(self.len());
for (output_element, &input_index) in output.iter_mut().zip(input_map.iter()) {
*output_element = input[input_index];
}
self.width_size_fft.process_inplace_multi(output, input);
unsafe { array_utils::transpose_small(self.width, self.height, output, input) };
self.height_size_fft.process_inplace_multi(input, output);
for (input_element, &output_index) in input.iter().zip(output_map.iter()) {
output[output_index] = *input_element;
}
}
fn perform_fft_inplace(&self, buffer: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
assert_eq!(self.len(), buffer.len());
assert_eq!(self.len(), scratch.len());
let (input_map, output_map) = self.input_output_map.split_at(self.len());
for (output_element, &input_index) in scratch.iter_mut().zip(input_map.iter()) {
*output_element = buffer[input_index];
}
self.width_size_fft.process_inplace_multi(scratch, buffer);
unsafe { array_utils::transpose_small(self.width, self.height, scratch, buffer) };
self.height_size_fft.process_multi(buffer, scratch, &mut []);
for (input_element, &output_index) in scratch.iter().zip(output_map.iter()) {
buffer[output_index] = *input_element;
}
}
}
boilerplate_fft!(GoodThomasAlgorithmSmall,
|this: &GoodThomasAlgorithmSmall<_>| this.width * this.height,
|this: &GoodThomasAlgorithmSmall<_>| this.len(),
|_| 0
);
#[cfg(test)]
mod unit_tests {
use super::*;
use std::sync::Arc;
use crate::test_utils::check_fft_algorithm;
use crate::algorithm::DFT;
use num_integer::gcd;
#[test]
fn test_good_thomas() {
for width in 1..12 {
for height in 1..12 {
if gcd(width, height) == 1 {
test_good_thomas_with_lengths(width, height, false);
test_good_thomas_with_lengths(width, height, true);
}
}
}
}
#[test]
fn test_good_thomas_small() {
let butterfly_sizes = [2,3,4,5,6,7,8,16];
for width in &butterfly_sizes {
for height in &butterfly_sizes {
if gcd(*width, *height) == 1 {
test_good_thomas_small_with_lengths(*width, *height, false);
test_good_thomas_small_with_lengths(*width, *height, true);
}
}
}
}
fn test_good_thomas_with_lengths(width: usize, height: usize, inverse: bool) {
let width_fft = Arc::new(DFT::new(width, inverse)) as Arc<dyn Fft<f32>>;
let height_fft = Arc::new(DFT::new(height, inverse)) as Arc<dyn Fft<f32>>;
let fft = GoodThomasAlgorithm::new(width_fft, height_fft);
check_fft_algorithm(&fft, width * height, inverse);
}
fn test_good_thomas_small_with_lengths(width: usize, height: usize, inverse: bool) {
let width_fft = Arc::new(DFT::new(width, inverse)) as Arc<dyn Fft<f32>>;
let height_fft = Arc::new(DFT::new(height, inverse)) as Arc<dyn Fft<f32>>;
let fft = GoodThomasAlgorithmSmall::new(width_fft, height_fft);
check_fft_algorithm(&fft, width * height, inverse);
}
}