use std::sync::Arc;
use std::arch::x86_64::*;
use std::convert::TryInto;
use num_integer::{Integer, div_ceil};
use num_complex::Complex;
use num_traits::Zero;
use strength_reduce::StrengthReducedUsize;
use primal_check::miller_rabin;
use crate::math_utils;
use crate::{Length, IsInverse, Fft};
use super::{AvxNum, avx_vector::{AvxVector, AvxVector256, AvxVector128, AvxArray, AvxArrayMut}};
use super::avx_vector;
#[derive(Clone)]
struct VectorizedMultiplyMod {
b: __m256i,
divisor: __m256i,
intermediate: __m256i,
}
impl VectorizedMultiplyMod {
#[target_feature(enable = "avx")]
unsafe fn new(b: u32, divisor: u32) -> Self {
assert!(divisor.leading_zeros() > 0, "divisor must be less than {}, got {}", 1 << 31, divisor);
let b = b % divisor;
let intermediate = ((b as i64) << 32) / divisor as i64;
Self {
b: _mm256_set1_epi64x(b as i64),
divisor: _mm256_set1_epi64x(divisor as i64),
intermediate: _mm256_set1_epi64x(intermediate),
}
}
#[allow(unused)]
#[target_feature(enable = "avx2")]
unsafe fn mul_rem(&self, a: __m256i) -> __m256i {
let masked_divisor = _mm256_blend_epi32(self.divisor, _mm256_setzero_si256(), 0xAA);
let quotient = _mm256_srli_epi64(_mm256_mul_epu32(a, self.intermediate), 32);
let numerator = _mm256_mul_epu32(a, self.b);
let quotient_product = _mm256_mul_epu32(quotient, masked_divisor);
let remainder = _mm256_sub_epi64(numerator, quotient_product);
let casted_remainder = _mm256_castsi256_pd(remainder);
let subtracted_remainder = _mm256_castsi256_pd(_mm256_sub_epi64(remainder, masked_divisor));
let wrapped_remainder = _mm256_castpd_si256(_mm256_blendv_pd(subtracted_remainder, casted_remainder, subtracted_remainder));
wrapped_remainder
}
}
pub struct RadersAvx2<T: AvxNum> {
input_index_multiplier: VectorizedMultiplyMod,
input_index_init: __m256i,
output_index_mapping: Box<[__m128i]>,
twiddles: Box<[T::VectorType]>,
inner_fft: Arc<dyn Fft<T>>,
len: usize,
inplace_scratch_len: usize,
outofplace_scratch_len: usize,
inverse: bool,
}
impl<T: AvxNum> RadersAvx2<T> {
#[inline]
pub fn new(inner_fft: Arc<dyn Fft<T>>) -> Result<Self, ()> {
let has_avx = is_x86_feature_detected!("avx");
let has_avx2 = is_x86_feature_detected!("avx2");
let has_fma = is_x86_feature_detected!("fma");
if has_avx && has_avx2 && has_fma {
Ok(unsafe { Self::new_with_avx(inner_fft) })
} else {
Err(())
}
}
#[target_feature(enable = "avx")]
unsafe fn new_with_avx(inner_fft: Arc<dyn Fft<T>>) -> Self {
let inner_fft_len = inner_fft.len();
let len = inner_fft_len + 1;
assert!(miller_rabin(len as u64), "For raders algorithm, inner_fft.len() + 1 must be prime. Expected prime number, got {} + 1 = {}", inner_fft_len, len);
let inverse = inner_fft.is_inverse();
let reduced_len = StrengthReducedUsize::new(len);
let primitive_root = math_utils::primitive_root(len as u64).unwrap() as usize;
let gcd_data = i64::extended_gcd(&(primitive_root as i64), &(len as i64));
let primitive_root_inverse = if gcd_data.x >= 0 { gcd_data.x } else { gcd_data.x + len as i64 } as usize;
let unity_scale = T::from_f64(1f64 / inner_fft_len as f64).unwrap();
let mut inner_fft_input = vec![Complex::zero(); inner_fft_len];
let mut twiddle_input = 1;
for input_cell in &mut inner_fft_input {
let twiddle = T::generate_twiddle_factor(twiddle_input, len, inverse);
*input_cell = twiddle * unity_scale;
twiddle_input = (twiddle_input * primitive_root_inverse) % reduced_len;
}
let required_inner_scratch = inner_fft.get_inplace_scratch_len();
let extra_inner_scratch = if required_inner_scratch <= inner_fft_len { 0 } else { required_inner_scratch };
let mut inner_fft_scratch = vec![Zero::zero(); required_inner_scratch];
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 : Box<[_]> = inner_fft_input.chunks(T::VectorType::COMPLEX_PER_VECTOR).map(|chunk| {
let chunk_vector = match chunk.len() {
1 => chunk.load_partial1_complex(0).zero_extend(),
2 => if chunk.len() == T::VectorType::COMPLEX_PER_VECTOR { chunk.load_complex(0) } else {chunk.load_partial2_complex(0).zero_extend()},
3 => chunk.load_partial3_complex(0),
4 => chunk.load_complex(0),
_ => unreachable!()
};
AvxVector::xor(chunk_vector, conjugation_mask) }).collect();
const NUM_POWERS : usize = 5;
let mut root_powers = [0; NUM_POWERS];
let mut current_power = 1;
for i in 0..NUM_POWERS {
root_powers[i] = current_power;
current_power = (current_power * primitive_root) % reduced_len;
}
let (input_index_multiplier, input_index_init) = if T::VectorType::COMPLEX_PER_VECTOR == 4 {
(VectorizedMultiplyMod::new(root_powers[4] as u32, len as u32), _mm256_loadu_si256(root_powers.as_ptr().add(1) as *const __m256i))
} else {
let duplicated_powers = [root_powers[1],root_powers[1],root_powers[2],root_powers[2],];
(VectorizedMultiplyMod::new(root_powers[2] as u32, len as u32), _mm256_loadu_si256(duplicated_powers.as_ptr() as *const __m256i))
};
let mapping_size = 1 + div_ceil(len, T::VectorType::COMPLEX_PER_VECTOR) * T::VectorType::COMPLEX_PER_VECTOR;
let mut output_mapping_inverse: Vec<i32> = vec![0; mapping_size];
let mut output_index = 1;
for i in 1..len {
output_index = (output_index * primitive_root_inverse) % reduced_len;
output_mapping_inverse[output_index] = i.try_into().unwrap();
}
let output_index_mapping = if T::VectorType::COMPLEX_PER_VECTOR == 4 {
(&output_mapping_inverse[1..]).chunks_exact(T::VectorType::COMPLEX_PER_VECTOR).map(|chunk| _mm_loadu_si128(chunk.as_ptr() as *const __m128i)).collect::<Box<[__m128i]>>()
} else {
(&output_mapping_inverse[1..]).chunks_exact(T::VectorType::COMPLEX_PER_VECTOR).map(|chunk| {
let duplicated_indexes = [chunk[0], chunk[0], chunk[1], chunk[1]];
_mm_loadu_si128(duplicated_indexes.as_ptr() as *const __m128i)
}).collect::<Box<[__m128i]>>()
};
Self {
input_index_multiplier,
input_index_init,
output_index_mapping,
inner_fft: inner_fft,
twiddles: inner_fft_multiplier,
len,
inplace_scratch_len: len + extra_inner_scratch,
outofplace_scratch_len: extra_inner_scratch,
inverse,
}
}
#[target_feature(enable = "avx2", enable = "avx", enable = "fma")]
unsafe fn prepare_raders(&self, input: &[Complex<T>], output: &mut [Complex<T>]) -> (Complex<T>, Complex<T>) {
let mut vector_sum = T::VectorType::zero();
let mut indexes = self.input_index_init;
let first_element = input[0];
let index_multiplier = self.input_index_multiplier.clone();
let mut chunks_iter = (&mut output[1..]).chunks_exact_mut(T::VectorType::COMPLEX_PER_VECTOR);
for chunk in chunks_iter.by_ref() {
let gathered_elements = T::VectorType::gather_complex_avx2_index64(input.as_ptr(), indexes);
indexes = index_multiplier.mul_rem(indexes);
vector_sum = AvxVector::add(vector_sum, gathered_elements);
chunk.store_complex(gathered_elements, 0);
}
let output_remainder = chunks_iter.into_remainder();
if output_remainder.len() == 2 {
let half_data = AvxVector128::gather64_complex_avx2(input.as_ptr(), _mm256_castsi256_si128(indexes));
vector_sum = AvxVector::add(vector_sum, AvxVector128::zero_extend(half_data));
output_remainder.store_partial2_complex(half_data, 0);
}
(first_element, vector_sum.hadd_complex() + first_element)
}
#[target_feature(enable = "avx2", enable = "avx", enable = "fma")]
unsafe fn finalize_raders(&self, input: &[Complex<T>], output: &mut [Complex<T>], first_input: Complex<T>) {
let output_add = AvxVector256::broadcast_complex_elements(first_input);
let conjugation_mask = AvxVector256::broadcast_complex_elements(Complex::new(T::zero(), -T::zero()));
let mut chunks_iter = (&mut output[1..]).chunks_exact_mut(T::VectorType::COMPLEX_PER_VECTOR);
for (i, chunk) in chunks_iter.by_ref().enumerate() {
let index_chunk = *self.output_index_mapping.get_unchecked(i);
let gathered_elements = T::VectorType::gather_complex_avx2_index32(input.as_ptr(), index_chunk);
let conjugated_elements = AvxVector::xor(gathered_elements, conjugation_mask);
let added_elements = AvxVector::add(output_add, conjugated_elements);
chunk.store_complex(added_elements, 0);
}
let output_remainder = chunks_iter.into_remainder();
if output_remainder.len() == 2 {
let index_chunk = *self.output_index_mapping.get_unchecked(self.output_index_mapping.len() - 1);
let half_data = AvxVector128::gather32_complex_avx2(input.as_ptr(), index_chunk);
let conjugated_elements = AvxVector::xor(half_data, conjugation_mask.lo());
let added_elements = AvxVector::add(output_add.lo(), conjugated_elements);
output_remainder.store_partial2_complex(added_elements, 0);
}
}
fn perform_fft_out_of_place(&self, input: &mut [Complex<T>], output: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
let (first_input, first_output) = unsafe { self.prepare_raders(input, output) };
let inner_input = &mut input[1..];
let inner_output = &mut output[1..];
let inner_scratch = if scratch.len() > 0 { &mut scratch[..] } else { &mut inner_input[..] };
self.inner_fft.process_inplace_with_scratch(inner_output, inner_scratch);
unsafe { avx_vector::pairwise_complex_mul_conjugated(&mut inner_output[..], &mut inner_input[..], &self.twiddles) };
let inner_scratch = if scratch.len() > 0 { scratch } else { &mut inner_output[..] };
self.inner_fft.process_inplace_with_scratch(inner_input, inner_scratch);
output[0] = first_output;
unsafe { self.finalize_raders(input, output, first_input); }
}
fn perform_fft_inplace(&self, buffer: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
let (scratch, extra_scratch) = scratch.split_at_mut(self.len());
let (first_input, first_output) = unsafe { self.prepare_raders(buffer, scratch) };
let truncated_scratch = &mut scratch[1..];
let inner_scratch = if extra_scratch.len() > 0 { extra_scratch } else { &mut buffer[..] };
self.inner_fft.process_inplace_with_scratch(truncated_scratch, inner_scratch);
unsafe { avx_vector::pairwise_complex_mul_assign_conjugated(truncated_scratch, &self.twiddles) };
self.inner_fft.process_inplace_with_scratch(truncated_scratch, inner_scratch);
buffer[0] = first_output;
unsafe { self.finalize_raders(scratch, buffer, first_input); }
}
}
boilerplate_avx_fft!(RadersAvx2,
|this: &RadersAvx2<_>| this.len,
|this: &RadersAvx2<_>| this.inplace_scratch_len,
|this: &RadersAvx2<_>| this.outofplace_scratch_len
);
#[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_raders_avx_f32() {
for len in 3..100 {
if miller_rabin(len as u64) {
test_raders_with_length::<f32>(len, false);
test_raders_with_length::<f32>(len, true);
}
}
}
#[test]
fn test_raders_avx_f64() {
for len in 3..100 {
if miller_rabin(len as u64) {
test_raders_with_length::<f64>(len, false);
test_raders_with_length::<f64>(len, true);
}
}
}
fn test_raders_with_length<T: AvxNum + num_traits::Float>(len: usize, inverse: bool) {
let inner_fft = Arc::new(DFT::new(len - 1, inverse));
let fft = RadersAvx2::new(inner_fft).unwrap();
check_fft_algorithm::<T>(&fft, len, inverse);
}
}