use std::sync::Arc;
use num_integer::div_ceil;
use num_complex::Complex;
use num_traits::Zero;
use crate::{Length, IsInverse, Fft};
use super::{AvxNum, avx_vector::{AvxVector, AvxVector128, AvxVector256, AvxArray, AvxArrayMut}};
use super::CommonSimdData;
pub struct BluesteinsAvx<T: AvxNum> {
inner_fft_multiplier: Box<[T::VectorType]>,
common_data: CommonSimdData<T, T::VectorType>,
}
boilerplate_avx_fft_commondata!(BluesteinsAvx);
impl<T: AvxNum> BluesteinsAvx<T> {
fn compute_bluesteins_twiddle(index: usize, len: usize, inverse: bool) -> Complex<T> {
let index_float = index as f64;
let index_squared = index_float * index_float;
T::generate_twiddle_factor_floatindex(index_squared, len*2, !inverse)
}
#[inline(always)]
unsafe fn mul_complex_conjugated<V: AvxVector>(left: V, right: V) -> V {
let (left_real, left_imag) = V::duplicate_complex_components(left);
let right_shuffled = V::swap_complex_components(right);
let output_right = V::mul(left_imag, right_shuffled);
V::fmsubadd(left_real, right, output_right)
}
#[inline]
pub fn new(len: usize, inner_fft: Arc<dyn Fft<T>>) -> Result<Self, ()> {
let has_avx = is_x86_feature_detected!("avx");
let has_fma = is_x86_feature_detected!("fma");
if has_avx && has_fma {
Ok(unsafe { Self::new_with_avx(len, inner_fft) })
} else {
Err(())
}
}
#[target_feature(enable = "avx")]
unsafe fn new_with_avx(len: usize, inner_fft: Arc<dyn Fft<T>>) -> Self {
let inner_fft_len = inner_fft.len();
assert!(len * 2 - 1 <= inner_fft_len, "Bluestein's algorithm requires inner_fft.len() >= self.len() * 2 - 1. Expected >= {}, got {}", len * 2 - 1, inner_fft_len);
assert_eq!(inner_fft_len % T::VectorType::COMPLEX_PER_VECTOR, 0, "BluesteinsAvx requires its inner_fft.len() to be a multiple of {} (IE the number of complex numbers in a single vector) inner_fft.len() = {}", T::VectorType::COMPLEX_PER_VECTOR, inner_fft_len);
let inner_len_float = T::from_usize(inner_fft_len).unwrap();
let inverse = inner_fft.is_inverse();
let mut inner_fft_input = vec![Complex::zero(); inner_fft_len];
for i in 0..len {
inner_fft_input[i] = Self::compute_bluesteins_twiddle(i, len, inverse) / inner_len_float;
}
for i in 1..len {
inner_fft_input[inner_fft_len - i] = inner_fft_input[i];
}
let mut inner_fft_scratch = vec![Complex::zero(); inner_fft.get_inplace_scratch_len()];
inner_fft.process_inplace_with_scratch(&mut inner_fft_input, &mut inner_fft_scratch);
let conjugation_mask = AvxVector256::broadcast_complex_elements(Complex::new(T::zero(), -T::zero()));
let inner_fft_multiplier = inner_fft_input.chunks_exact(T::VectorType::COMPLEX_PER_VECTOR).map(|chunk| {
let chunk_vector = chunk.load_complex(0);
AvxVector::xor(chunk_vector, conjugation_mask) }).collect::<Vec<_>>().into_boxed_slice();
let chunk_count = div_ceil(len, T::VectorType::COMPLEX_PER_VECTOR);
let twiddles : Vec<_> = (0..chunk_count).map(|x| {
let mut twiddle_chunk = [Complex::zero();4];
for i in 0..T::VectorType::COMPLEX_PER_VECTOR {
twiddle_chunk[i] = Self::compute_bluesteins_twiddle(x*T::VectorType::COMPLEX_PER_VECTOR+i, len, !inverse);
}
twiddle_chunk.load_complex(0)
}).collect();
let required_scratch = inner_fft_input.len() + inner_fft_scratch.len();
Self {
inner_fft_multiplier,
common_data: CommonSimdData {
inner_fft,
twiddles: twiddles.into_boxed_slice(),
len,
inplace_scratch_len: required_scratch,
outofplace_scratch_len: required_scratch,
inverse,
}
}
}
#[target_feature(enable = "avx", enable = "fma")]
unsafe fn prepare_bluesteins(&self, input: &[Complex<T>], inner_fft_buffer: &mut [Complex<T>]) {
let chunk_count = self.common_data.twiddles.len() - 1;
let remainder = self.len() - chunk_count * T::VectorType::COMPLEX_PER_VECTOR;
for (i, twiddle) in self.common_data.twiddles[..chunk_count].iter().enumerate() {
let index = i * T::VectorType::COMPLEX_PER_VECTOR;
let input_vector = input.load_complex(index);
let product_vector = AvxVector::mul_complex(input_vector, *twiddle);
inner_fft_buffer.store_complex(product_vector, index);
}
{
let remainder_twiddle = self.common_data.twiddles[chunk_count];
let remainder_index = chunk_count * T::VectorType::COMPLEX_PER_VECTOR;
let remainder_data = match remainder {
1 => input.load_partial1_complex(remainder_index).zero_extend(),
2 => if T::VectorType::COMPLEX_PER_VECTOR == 2 {
input.load_complex(remainder_index)
} else {
input.load_partial2_complex(remainder_index).zero_extend()
},
3 => input.load_partial3_complex(remainder_index),
4 => input.load_complex(remainder_index),
_ => unreachable!(),
};
let twiddled_remainder = AvxVector::mul_complex(remainder_twiddle, remainder_data);
inner_fft_buffer.store_complex(twiddled_remainder, remainder_index);
}
let zerofill_start = chunk_count + 1;
for i in zerofill_start..(inner_fft_buffer.len()/T::VectorType::COMPLEX_PER_VECTOR) {
let index = i * T::VectorType::COMPLEX_PER_VECTOR;
inner_fft_buffer.store_complex(AvxVector::zero(), index);
}
}
#[target_feature(enable = "avx", enable = "fma")]
unsafe fn finalize_bluesteins(&self, inner_fft_buffer: &[Complex<T>], output: &mut [Complex<T>]) {
let chunk_count = self.common_data.twiddles.len() - 1;
let remainder = self.len() - chunk_count * T::VectorType::COMPLEX_PER_VECTOR;
for (i, twiddle) in self.common_data.twiddles[..chunk_count].iter().enumerate() {
let index = i * T::VectorType::COMPLEX_PER_VECTOR;
let inner_vector = inner_fft_buffer.load_complex(index);
let product_vector = Self::mul_complex_conjugated(inner_vector, *twiddle);
output.store_complex(product_vector, index);
}
{
let remainder_twiddle = self.common_data.twiddles[chunk_count];
let remainder_index = chunk_count * T::VectorType::COMPLEX_PER_VECTOR;
let inner_vector = inner_fft_buffer.load_complex(remainder_index);
let product_vector = Self::mul_complex_conjugated(inner_vector, remainder_twiddle);
match remainder {
1 => output.store_partial1_complex(product_vector.lo(), remainder_index),
2 => if T::VectorType::COMPLEX_PER_VECTOR == 2 {
output.store_complex(product_vector, remainder_index)
} else {
output.store_partial2_complex(product_vector.lo(), remainder_index)
},
3 => output.store_partial3_complex(product_vector, remainder_index),
4 => output.store_complex(product_vector, remainder_index),
_ => unreachable!(),
};
}
}
#[target_feature(enable = "avx", enable = "fma")]
unsafe fn pairwise_complex_multiply_conjugated(buffer: &mut [Complex<T>], multiplier: &[T::VectorType]) {
for (i, right) in multiplier.iter().enumerate() {
let left = buffer.load_complex(i*T::VectorType::COMPLEX_PER_VECTOR);
let product = Self::mul_complex_conjugated(left, *right);
buffer.store_complex(product, i*T::VectorType::COMPLEX_PER_VECTOR);
}
}
fn perform_fft_inplace(&self, buffer: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
let (inner_input, inner_scratch) = scratch.split_at_mut(self.inner_fft_multiplier.len()*T::VectorType::COMPLEX_PER_VECTOR);
unsafe { self.prepare_bluesteins(buffer, inner_input) };
self.common_data.inner_fft.process_inplace_with_scratch(inner_input, inner_scratch);
unsafe { Self::pairwise_complex_multiply_conjugated(inner_input, &self.inner_fft_multiplier) };
self.common_data.inner_fft.process_inplace_with_scratch(inner_input, inner_scratch);
unsafe { self.finalize_bluesteins(inner_input, buffer) };
}
fn perform_fft_out_of_place(&self, input: &[Complex<T>], output: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
let (inner_input, inner_scratch) = scratch.split_at_mut(self.inner_fft_multiplier.len()*T::VectorType::COMPLEX_PER_VECTOR);
unsafe { self.prepare_bluesteins(input, inner_input) };
self.common_data.inner_fft.process_inplace_with_scratch(inner_input, inner_scratch);
unsafe { Self::pairwise_complex_multiply_conjugated(inner_input, &self.inner_fft_multiplier) };
self.common_data.inner_fft.process_inplace_with_scratch(inner_input, inner_scratch);
unsafe { self.finalize_bluesteins(inner_input, output) };
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
use std::sync::Arc;
use crate::test_utils::check_fft_algorithm;
use crate::algorithm::DFT;
#[test]
fn test_bluesteins_avx_f32() {
for len in 2..16 {
let minimum_inner : usize = len * 2 - 1;
let remainder = minimum_inner % 4;
let next_multiple_of_4 = minimum_inner - remainder + 4;
let maximum_inner = minimum_inner.checked_next_power_of_two().unwrap() + 1;
for inner_len in (next_multiple_of_4..maximum_inner).step_by(4) {
test_bluesteins_avx_with_length::<f32>(len, inner_len, false);
test_bluesteins_avx_with_length::<f32>(len, inner_len, true);
}
}
}
#[test]
fn test_bluesteins_avx_f64() {
for len in 2..16 {
let minimum_inner : usize = len * 2 - 1;
let remainder = minimum_inner % 2;
let next_multiple_of_2 = minimum_inner + remainder;
let maximum_inner = minimum_inner.checked_next_power_of_two().unwrap() + 1;
for inner_len in (next_multiple_of_2..maximum_inner).step_by(2) {
test_bluesteins_avx_with_length::<f64>(len, inner_len, false);
test_bluesteins_avx_with_length::<f64>(len, inner_len, true);
}
}
}
fn test_bluesteins_avx_with_length<T: AvxNum + num_traits::Float>(len: usize, inner_len: usize, inverse: bool) {
let inner_fft = Arc::new(DFT::new(inner_len, inverse));
let fft : BluesteinsAvx<T> = BluesteinsAvx::new(len, inner_fft).expect("Can't run test because this machine doesn't have the required instruction sets");
check_fft_algorithm(&fft, len, inverse);
}
}