use crate::{dif2::split_2, *};
use aligned_vec::{avec, ABox, CACHELINE_ALIGN};
#[cfg(feature = "std")]
use core::time::Duration;
#[cfg(feature = "std")]
use dyn_stack::PodBuffer;
use dyn_stack::{PodStack, StackReq};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum FftAlgo {
Dif2,
Dit2,
Dif4,
Dit4,
Dif8,
Dit8,
Dif16,
Dit16,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum Method {
UserProvided(FftAlgo),
#[cfg(feature = "std")]
#[cfg_attr(docsrs, doc(cfg(feature = "std")))]
Measure(Duration),
}
#[cfg(feature = "std")]
fn measure_n_runs(
n_runs: u128,
algo: FftAlgo,
buf: &mut [c64],
twiddles_init: &[c64],
twiddles: &[c64],
stack: &mut PodStack,
) -> Duration {
let n = buf.len();
let (scratch, _) = stack.make_aligned_raw::<c64>(n, CACHELINE_ALIGN);
let [fwd, _] = get_fn_ptr(algo, n);
use crate::time::Instant;
let now = Instant::now();
for _ in 0..n_runs {
fwd(buf, scratch, twiddles, twiddles_init);
}
now.elapsed()
}
#[cfg(feature = "std")]
fn duration_div_f64(duration: Duration, n: f64) -> Duration {
Duration::from_secs_f64(duration.as_secs_f64() / n)
}
#[cfg(feature = "std")]
pub(crate) fn measure_fastest_scratch(n: usize) -> StackReq {
let align = CACHELINE_ALIGN;
StackReq::new_aligned::<c64>(2 * n, align) .and(StackReq::new_aligned::<c64>(n, align)) .and(StackReq::new_aligned::<c64>(n, align))
}
#[cfg(feature = "std")]
pub(crate) fn measure_fastest(
min_bench_duration_per_algo: Duration,
n: usize,
stack: &mut PodStack,
) -> (FftAlgo, Duration) {
const N_ALGOS: usize = 8;
const MIN_DURATION: Duration = if cfg!(target_arch = "wasm32") {
Duration::from_millis(10)
} else {
Duration::from_millis(1)
};
assert!(n.is_power_of_two());
let align = CACHELINE_ALIGN;
let f = |_| c64 { re: 0.0, im: 0.0 };
let (twiddles, stack) = stack.make_aligned_with::<c64>(2 * n, align, f);
let twiddles_init = &twiddles[..n];
let twiddles = &twiddles[n..];
let (buf, stack) = stack.make_aligned_with::<c64>(n, align, f);
{
drop(stack.make_aligned_with::<c64>(n, align, f));
}
let mut avg_durations = [Duration::ZERO; N_ALGOS];
let discriminant_to_algo = |i: usize| -> FftAlgo {
match i {
0 => FftAlgo::Dif2,
1 => FftAlgo::Dit2,
2 => FftAlgo::Dif4,
3 => FftAlgo::Dit4,
4 => FftAlgo::Dif8,
5 => FftAlgo::Dit8,
6 => FftAlgo::Dif16,
7 => FftAlgo::Dit16,
_ => unreachable!(),
}
};
for (i, avg) in (0..N_ALGOS).zip(&mut avg_durations) {
let algo = discriminant_to_algo(i);
let (init_n_runs, approx_duration) = {
let mut n_runs: u128 = 1;
loop {
let duration = measure_n_runs(n_runs, algo, buf, twiddles_init, twiddles, stack);
if duration < MIN_DURATION {
n_runs *= 2;
} else {
break (n_runs, duration_div_f64(duration, n_runs as f64));
}
}
};
let n_runs = (min_bench_duration_per_algo.as_secs_f64() / approx_duration.as_secs_f64())
.ceil() as u128;
*avg = if n_runs <= init_n_runs {
approx_duration
} else {
let duration = measure_n_runs(n_runs, algo, buf, twiddles_init, twiddles, stack);
duration_div_f64(duration, n_runs as f64)
};
}
let best_time = avg_durations.iter().min().unwrap();
let best_index = avg_durations
.iter()
.position(|elem| elem == best_time)
.unwrap();
(discriminant_to_algo(best_index), *best_time)
}
#[derive(Clone)]
pub struct Plan {
fwd: fn(&mut [c64], &mut [c64], &[c64], &[c64]),
inv: fn(&mut [c64], &mut [c64], &[c64], &[c64]),
twiddles: ABox<[c64]>,
twiddles_inv: ABox<[c64]>,
algo: FftAlgo,
}
impl core::fmt::Debug for Plan {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Plan")
.field("algo", &self.algo)
.field("fft_size", &self.fft_size())
.finish()
}
}
fn do_nothing(_: &mut [c64], _: &mut [c64], _: &[c64], _: &[c64]) {}
pub(crate) fn get_fn_ptr(
algo: FftAlgo,
n: usize,
) -> [fn(&mut [c64], &mut [c64], &[c64], &[c64]); 2] {
if n == 1 {
return [do_nothing; 2];
}
use FftAlgo::*;
match algo {
Dif2 => dif2::fft_impl_dispatch(n),
Dit2 => dit2::fft_impl_dispatch(n),
Dif4 => dif4::fft_impl_dispatch(n),
Dit4 => dit4::fft_impl_dispatch(n),
Dif8 => dif8::fft_impl_dispatch(n),
Dit8 => dit8::fft_impl_dispatch(n),
Dif16 => dif16::fft_impl_dispatch(n),
Dit16 => dit16::fft_impl_dispatch(n),
}
}
impl Plan {
#[cfg_attr(feature = "std", doc = " ```")]
#[cfg_attr(not(feature = "std"), doc = " ```ignore")]
pub fn new(n: usize, method: Method) -> Self {
assert!(n.is_power_of_two());
assert!(n.trailing_zeros() < 11);
let algo = match method {
Method::UserProvided(algo) => algo,
#[cfg(feature = "std")]
Method::Measure(duration) => {
let mut buf = PodBuffer::try_new(measure_fastest_scratch(n)).unwrap();
measure_fastest(duration, n, PodStack::new(&mut buf)).0
}
};
let [fwd, inv] = get_fn_ptr(algo, n);
let mut twiddles = avec![c64::default(); 2 * n].into_boxed_slice();
let mut twiddles_inv = avec![c64::default(); 2 * n].into_boxed_slice();
use FftAlgo::*;
let r = match algo {
Dif2 | Dit2 => 2,
Dif4 | Dit4 => 4,
Dif8 | Dit8 => 8,
Dif16 | Dit16 => 16,
};
fft_simd::init_wt(r, n, &mut twiddles, &mut twiddles_inv);
Self {
fwd,
inv,
twiddles,
algo,
twiddles_inv,
}
}
#[cfg_attr(feature = "std", doc = " ```")]
#[cfg_attr(not(feature = "std"), doc = " ```ignore")]
pub fn fft_size(&self) -> usize {
self.twiddles.len() / 2
}
pub fn algo(&self) -> FftAlgo {
self.algo
}
#[cfg_attr(feature = "std", doc = " ```")]
#[cfg_attr(not(feature = "std"), doc = " ```ignore")]
pub fn fft_scratch(&self) -> StackReq {
StackReq::new_aligned::<c64>(self.fft_size(), CACHELINE_ALIGN)
}
#[cfg_attr(feature = "std", doc = " ```")]
#[cfg_attr(not(feature = "std"), doc = " ```ignore")]
pub fn fwd(&self, buf: &mut [c64], stack: &mut PodStack) {
let n = self.fft_size();
let (scratch, _) = stack.make_aligned_raw::<c64>(n, CACHELINE_ALIGN);
let (w_init, w) = split_2(&self.twiddles);
(self.fwd)(buf, scratch, w_init, w)
}
#[cfg_attr(feature = "std", doc = " ```")]
#[cfg_attr(not(feature = "std"), doc = " ```ignore")]
pub fn inv(&self, buf: &mut [c64], stack: &mut PodStack) {
let n = self.fft_size();
let (scratch, _) = stack.make_aligned_raw::<c64>(n, CACHELINE_ALIGN);
let (w_init, w) = split_2(&self.twiddles_inv);
(self.inv)(buf, scratch, w_init, w)
}
}
#[cfg(test)]
mod tests {
use crate::{
c64, dif16, dif2, dif4, dif8, dit16, dit2, dit4, dit8,
fft_simd::{init_wt, FftSimd, Pod},
};
use num_complex::ComplexFloat;
use rand::random;
use rustfft::FftPlanner;
extern crate alloc;
use alloc::vec;
fn test_fft_simd<c64xN: Pod>(simd: impl FftSimd<c64xN>) {
for (r, fft) in [
(2, dif2::fft_impl(simd)),
(2, dit2::fft_impl(simd)),
(4, dif4::fft_impl(simd)),
(4, dit4::fft_impl(simd)),
(8, dif8::fft_impl(simd)),
(8, dit8::fft_impl(simd)),
(16, dif16::fft_impl(simd)),
(16, dit16::fft_impl(simd)),
] {
if simd.lane_count() > r {
continue;
}
for exp in 1..=10 {
let n: usize = 1 << exp;
if simd.lane_count() > 1 && simd.lane_count() * r > n {
continue;
}
let [fwd, inv] = fft.make_fn_ptr(n);
fn test_inner(
n: usize,
r: usize,
fwd: fn(&mut [c64], &mut [c64], &[c64], &[c64]),
inv: fn(&mut [c64], &mut [c64], &[c64], &[c64]),
) {
let mut scratch = vec![c64::default(); n];
let mut twiddles = vec![c64::default(); 2 * n];
let mut twiddles_inv = vec![c64::default(); 2 * n];
init_wt(r, n, &mut twiddles, &mut twiddles_inv);
let mut x = vec![c64::default(); n];
for z in &mut x {
*z = c64::new(random(), random());
}
let orig = x.clone();
fwd(&mut x, &mut scratch, &twiddles[..n], &twiddles[n..]);
{
let mut planner = FftPlanner::new();
let plan = planner.plan_fft_forward(n);
let mut y = orig.clone();
plan.process(&mut y);
for (z_expected, z_actual) in y.iter().zip(&x) {
assert!((*z_expected - *z_actual).abs() < 1e-12);
}
}
inv(&mut x, &mut scratch, &twiddles_inv[..n], &twiddles_inv[n..]);
for z in &mut x {
*z /= n as f64;
}
for (z_expected, z_actual) in orig.iter().zip(&x) {
assert!((*z_expected - *z_actual).abs() < 1e-14);
}
}
test_inner(n, r, fwd, inv);
}
}
}
#[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
#[test]
fn test_fft() {
test_fft_simd(crate::fft_simd::Scalar);
#[cfg(any(target_arch = "x86_64", target_arch = "x86"))]
{
if let Some(simd) = pulp::x86::V3::try_new() {
test_fft_simd(simd);
}
#[cfg(feature = "avx512")]
if let Some(simd) = pulp::x86::V4::try_new() {
test_fft_simd(simd);
}
}
}
}