diff --git a/benches/bench_rustfft.rs b/benches/bench_rustfft.rs index af9c6422..c00562a5 100644 --- a/benches/bench_rustfft.rs +++ b/benches/bench_rustfft.rs @@ -219,6 +219,30 @@ fn bench_good_thomas(b: &mut Bencher, width: usize, height: usize) { #[bench] fn good_thomas_2048_3(b: &mut Bencher) { bench_good_thomas(b, 2048, 3); } #[bench] fn good_thomas_2048_2187(b: &mut Bencher) { bench_good_thomas(b, 2048, 2187); } +/// Times just the FFT execution (not allocation and pre-calculation) +/// for a given length, specific to the Good-Thomas algorithm +fn bench_good_thomas_f64(b: &mut Bencher, width: usize, height: usize) { + + let mut planner = rustfft::FftPlanner::new(); + let width_fft = planner.plan_fft_forward(width); + let height_fft = planner.plan_fft_forward(height); + + let fft : Arc> = Arc::new(GoodThomasAlgorithm::new(width_fft, height_fft)); + + let mut buffer = vec![Complex::zero(); width * height]; + let mut scratch = vec![Complex::zero(); fft.get_inplace_scratch_len()]; + b.iter(|| {fft.process_with_scratch(&mut buffer, &mut scratch);} ); +} + +#[bench] fn good_thomas_64_0002_3(b: &mut Bencher) { bench_good_thomas_f64(b, 2, 3); } +#[bench] fn good_thomas_64_0003_4(b: &mut Bencher) { bench_good_thomas_f64(b, 3, 4); } +#[bench] fn good_thomas_64_0004_5(b: &mut Bencher) { bench_good_thomas_f64(b, 4, 5); } +#[bench] fn good_thomas_64_0007_32(b: &mut Bencher) { bench_good_thomas_f64(b, 7, 32); } +#[bench] fn good_thomas_64_0032_27(b: &mut Bencher) { bench_good_thomas_f64(b, 32, 27); } +#[bench] fn good_thomas_64_0256_243(b: &mut Bencher) { bench_good_thomas_f64(b, 256, 243); } +#[bench] fn good_thomas_64_2048_3(b: &mut Bencher) { bench_good_thomas_f64(b, 2048, 3); } +#[bench] fn good_thomas_64_2048_2187(b: &mut Bencher) { bench_good_thomas_f64(b, 2048, 2187); } + /// Times just the FFT setup (not execution) /// for a given length, specific to the Good-Thomas algorithm fn bench_good_thomas_setup(b: &mut Bencher, width: usize, height: usize) { @@ -310,6 +334,30 @@ fn bench_mixed_radix(b: &mut Bencher, width: usize, height: usize) { #[bench] fn mixed_radix_2048_3(b: &mut Bencher) { bench_mixed_radix(b, 2048, 3); } #[bench] fn mixed_radix_2048_2187(b: &mut Bencher) { bench_mixed_radix(b, 2048, 2187); } +/// Times just the FFT execution (not allocation and pre-calculation) +/// for a given length, specific to the Mixed-Radix algorithm +fn bench_mixed_radix_f64(b: &mut Bencher, width: usize, height: usize) { + + let mut planner = rustfft::FftPlanner::new(); + let width_fft = planner.plan_fft_forward(width); + let height_fft = planner.plan_fft_forward(height); + + let fft : Arc> = Arc::new(MixedRadix::new(width_fft, height_fft)); + + let mut buffer = vec![Complex{re: 0_f64, im: 0_f64}; fft.len()]; + let mut scratch = vec![Complex{re: 0_f64, im: 0_f64}; fft.get_inplace_scratch_len()]; + b.iter(|| {fft.process_with_scratch(&mut buffer, &mut scratch);} ); +} + +#[bench] fn mixed_radix_64_0002_3(b: &mut Bencher) { bench_mixed_radix_f64(b, 2, 3); } +#[bench] fn mixed_radix_64_0003_4(b: &mut Bencher) { bench_mixed_radix_f64(b, 3, 4); } +#[bench] fn mixed_radix_64_0004_5(b: &mut Bencher) { bench_mixed_radix_f64(b, 4, 5); } +#[bench] fn mixed_radix_64_0007_32(b: &mut Bencher) { bench_mixed_radix_f64(b, 7, 32); } +#[bench] fn mixed_radix_64_0032_27(b: &mut Bencher) { bench_mixed_radix_f64(b, 32, 27); } +#[bench] fn mixed_radix_64_0256_243(b: &mut Bencher) { bench_mixed_radix_f64(b, 256, 243); } +#[bench] fn mixed_radix_64_2048_3(b: &mut Bencher) { bench_mixed_radix_f64(b, 2048, 3); } +#[bench] fn mixed_radix_64_2048_2187(b: &mut Bencher) { bench_mixed_radix_f64(b, 2048, 2187); } + fn plan_butterfly_fft(len: usize) -> Arc> { match len { 2 => Arc::new(Butterfly2::new(FftDirection::Forward)), @@ -344,6 +392,40 @@ fn bench_mixed_radix_small(b: &mut Bencher, width: usize, height: usize) { #[bench] fn mixed_radix_small_0004_5(b: &mut Bencher) { bench_mixed_radix_small(b, 4, 5); } #[bench] fn mixed_radix_small_0007_32(b: &mut Bencher) { bench_mixed_radix_small(b, 7, 32); } +fn plan_butterfly_fft_f64(len: usize) -> Arc> { + match len { + 2 => Arc::new(Butterfly2::new(FftDirection::Forward)), + 3 => Arc::new(Butterfly3::new(FftDirection::Forward)), + 4 => Arc::new(Butterfly4::new(FftDirection::Forward)), + 5 => Arc::new(Butterfly5::new(FftDirection::Forward)), + 6 => Arc::new(Butterfly6::new(FftDirection::Forward)), + 7 => Arc::new(Butterfly7::new(FftDirection::Forward)), + 8 => Arc::new(Butterfly8::new(FftDirection::Forward)), + 16 => Arc::new(Butterfly16::new(FftDirection::Forward)), + 32 => Arc::new(Butterfly32::new(FftDirection::Forward)), + _ => panic!("Invalid butterfly size: {}", len), + } +} + +/// Times just the FFT execution (not allocation and pre-calculation) +/// for a given length, specific to the MixedRadixSmall algorithm +fn bench_mixed_radix_small_f64(b: &mut Bencher, width: usize, height: usize) { + + let width_fft = plan_butterfly_fft_f64(width); + let height_fft = plan_butterfly_fft_f64(height); + + let fft : Arc> = Arc::new(MixedRadixSmall::new(width_fft, height_fft)); + + let mut signal = vec![Complex{re: 0_f64, im: 0_f64}; width * height]; + let mut spectrum = signal.clone(); + b.iter(|| {fft.process_with_scratch(&mut signal, &mut spectrum);} ); +} + +#[bench] fn mixed_radix_small_64_0002_3(b: &mut Bencher) { bench_mixed_radix_small_f64(b, 2, 3); } +#[bench] fn mixed_radix_small_64_0003_4(b: &mut Bencher) { bench_mixed_radix_small_f64(b, 3, 4); } +#[bench] fn mixed_radix_small_64_0004_5(b: &mut Bencher) { bench_mixed_radix_small_f64(b, 4, 5); } +#[bench] fn mixed_radix_small_64_0007_32(b: &mut Bencher) { bench_mixed_radix_small_f64(b, 7, 32); } + /// Times just the FFT execution (not allocation and pre-calculation) /// for a given length, specific to the Mixed-Radix Double Butterfly algorithm fn bench_good_thomas_small(b: &mut Bencher, width: usize, height: usize) { @@ -363,6 +445,25 @@ fn bench_good_thomas_small(b: &mut Bencher, width: usize, height: usize) { #[bench] fn good_thomas_small_0004_5(b: &mut Bencher) { bench_good_thomas_small(b, 4, 5); } #[bench] fn good_thomas_small_0007_32(b: &mut Bencher) { bench_good_thomas_small(b, 7, 32); } +/// Times just the FFT execution (not allocation and pre-calculation) +/// for a given length, specific to the Mixed-Radix Double Butterfly algorithm +fn bench_good_thomas_small_f64(b: &mut Bencher, width: usize, height: usize) { + + let width_fft = plan_butterfly_fft_f64(width); + let height_fft = plan_butterfly_fft_f64(height); + + let fft : Arc> = Arc::new(GoodThomasAlgorithmSmall::new(width_fft, height_fft)); + + let mut signal = vec![Complex{re: 0_f64, im: 0_f64}; width * height]; + let mut spectrum = signal.clone(); + b.iter(|| {fft.process_with_scratch(&mut signal, &mut spectrum);} ); +} + +#[bench] fn good_thomas_small_64_0002_3(b: &mut Bencher) { bench_good_thomas_small_f64(b, 2, 3); } +#[bench] fn good_thomas_small_64_0003_4(b: &mut Bencher) { bench_good_thomas_small_f64(b, 3, 4); } +#[bench] fn good_thomas_small_64_0004_5(b: &mut Bencher) { bench_good_thomas_small_f64(b, 4, 5); } +#[bench] fn good_thomas_small_64_0007_32(b: &mut Bencher) { bench_good_thomas_small_f64(b, 7, 32); } + /// Times just the FFT execution (not allocation and pre-calculation) /// for a given length, specific to Rader's algorithm diff --git a/src/algorithm/mixed_radix.rs b/src/algorithm/mixed_radix.rs index 317d4cb5..a4630b11 100644 --- a/src/algorithm/mixed_radix.rs +++ b/src/algorithm/mixed_radix.rs @@ -324,13 +324,16 @@ impl MixedRadixSmall { // STEP 2: perform FFTs of size `height` self.height_size_fft.process_with_scratch(scratch, buffer); - // STEP 3: Apply twiddle factors - for (element, twiddle) in scratch.iter_mut().zip(self.twiddles.iter()) { - *element = *element * twiddle; - } - - // STEP 4: transpose again - unsafe { array_utils::transpose_small(self.height, self.width, scratch, buffer) }; + // STEP 3 & 4: Apply twiddle factors and transpose + unsafe { + array_utils::transpose_small_twiddle( + self.height, + self.width, + scratch, + buffer, + &self.twiddles, + ) + }; // STEP 5: perform FFTs of size `width` self.width_size_fft @@ -353,13 +356,16 @@ impl MixedRadixSmall { // STEP 2: perform FFTs of size `height` self.height_size_fft.process_with_scratch(output, scratch); - // STEP 3: Apply twiddle factors - for (element, twiddle) in output.iter_mut().zip(self.twiddles.iter()) { - *element = *element * twiddle; - } - - // STEP 4: transpose again - unsafe { array_utils::transpose_small(self.height, self.width, output, scratch) }; + // STEP 3 & 4: Apply twiddle factors and transpose + unsafe { + array_utils::transpose_small_twiddle( + self.height, + self.width, + output, + scratch, + &self.twiddles, + ) + }; // STEP 5: perform FFTs of size `width` self.width_size_fft.process_with_scratch(scratch, output); @@ -381,13 +387,16 @@ impl MixedRadixSmall { // STEP 2: perform FFTs of size `height` self.height_size_fft.process_with_scratch(output, input); - // STEP 3: Apply twiddle factors - for (element, twiddle) in output.iter_mut().zip(self.twiddles.iter()) { - *element = *element * twiddle; - } - - // STEP 4: transpose again - unsafe { array_utils::transpose_small(self.height, self.width, output, input) }; + // STEP 3 & 4: Apply twiddle factors and transpose + unsafe { + array_utils::transpose_small_twiddle( + self.height, + self.width, + output, + input, + &self.twiddles, + ) + }; // STEP 5: perform FFTs of size `width` self.width_size_fft.process_with_scratch(input, output); diff --git a/src/array_utils.rs b/src/array_utils.rs index bfae7df0..54ecf4b7 100644 --- a/src/array_utils.rs +++ b/src/array_utils.rs @@ -17,6 +17,36 @@ pub unsafe fn transpose_small(width: usize, height: usize, input: &[T], } } +pub unsafe fn transpose_small_twiddle( + width: usize, + height: usize, + input: &[Complex], + output: &mut [Complex], + twiddles: &[Complex], +) { + assert!(input.len() >= width * height); + assert!(output.len() >= width * height); + assert!(twiddles.len() >= width * height); + + #[cfg(all(target_arch = "aarch64", feature = "neon"))] + { + if crate::neon::neon_utils::transpose_small_twiddle(width, height, input, output, twiddles) + { + return; + } + } + + for y in 0..height { + for x in 0..width { + let in_idx = x + y * width; + let out_idx = y + x * height; + let val = *input.get_unchecked(in_idx); + let tw = *twiddles.get_unchecked(in_idx); + *output.get_unchecked_mut(out_idx) = val * tw; + } + } +} + #[allow(unused)] pub unsafe fn workaround_transmute(slice: &[T]) -> &[U] { let ptr = slice.as_ptr() as *const U; @@ -115,7 +145,7 @@ mod unit_tests { use num_traits::Zero; #[test] - fn test_transpose() { + fn test_transpose_small() { let sizes: Vec = (1..16).collect(); for &width in &sizes { @@ -141,6 +171,75 @@ mod unit_tests { } } } + + #[test] + fn test_transpose_small_twiddle() { + let sizes: Vec = (1..16).collect(); + + for &width in &sizes { + for &height in &sizes { + let len = width * height; + + // Test f32 + let input_f32: Vec> = random_signal(len); + let twiddles_f32: Vec> = random_signal(len); + let mut output_f32 = vec![Zero::zero(); len]; + unsafe { + transpose_small_twiddle( + width, + height, + &input_f32, + &mut output_f32, + &twiddles_f32, + ); + } + for x in 0..width { + for y in 0..height { + let expected = input_f32[x + y * width] * twiddles_f32[x + y * width]; + let actual = output_f32[y + x * height]; + assert!( + (actual.re - expected.re).abs() < 1e-5 + && (actual.im - expected.im).abs() < 1e-5, + "f32 mismatch at ({}, {}) for {}x{}", + x, + y, + width, + height + ); + } + } + + // Test f64 + let input_f64: Vec> = random_signal(len); + let twiddles_f64: Vec> = random_signal(len); + let mut output_f64 = vec![Zero::zero(); len]; + unsafe { + transpose_small_twiddle( + width, + height, + &input_f64, + &mut output_f64, + &twiddles_f64, + ); + } + for x in 0..width { + for y in 0..height { + let expected = input_f64[x + y * width] * twiddles_f64[x + y * width]; + let actual = output_f64[y + x * height]; + assert!( + (actual.re - expected.re).abs() < 1e-10 + && (actual.im - expected.im).abs() < 1e-10, + "f64 mismatch at ({}, {}) for {}x{}", + x, + y, + width, + height + ); + } + } + } + } + } } // A utility that validates the following conditions, then calls chunk_fn() on each chunk of buffer. Passes the entire scratch buffer with each call. diff --git a/src/lib.rs b/src/lib.rs index 2bcb30dd..89ee732a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -440,7 +440,7 @@ pub use self::sse::sse_planner::FftPlannerSse; // Algorithms implemented to use Neon instructions. Only compiled on AArch64, and only compiled if the "neon" feature flag is set. #[cfg(all(target_arch = "aarch64", feature = "neon"))] -mod neon; +pub(crate) mod neon; // If we're not on AArch64, or if the "neon" feature was disabled, keep a stub implementation around that has the same API, but does nothing // That way, users can write code using the Neon planner and compile it on any platform diff --git a/src/neon/mod.rs b/src/neon/mod.rs index 20b936cb..e7c7d2e6 100644 --- a/src/neon/mod.rs +++ b/src/neon/mod.rs @@ -8,7 +8,7 @@ pub mod neon_butterflies; pub mod neon_prime_butterflies; pub mod neon_radix4; -mod neon_utils; +pub(crate) mod neon_utils; pub mod neon_planner; diff --git a/src/neon/neon_utils.rs b/src/neon/neon_utils.rs index 0bc8db2e..a309355b 100644 --- a/src/neon/neon_utils.rs +++ b/src/neon/neon_utils.rs @@ -1,4 +1,7 @@ +use crate::neon::neon_vector::{NeonArray, NeonArrayMut, NeonVector}; +use crate::FftNum; use core::arch::aarch64::*; +use num_complex::Complex; // __ __ _ _ _________ _ _ _ // | \/ | __ _| |_| |__ |___ /___ \| |__ (_) |_ @@ -183,8 +186,12 @@ pub unsafe fn duplicate_hi_f32(values: float32x4_t) -> float32x4_t { )) } -// transpose a 2x2 complex matrix given as [x0, x1], [x2, x3] -// result is [x0, x2], [x1, x3] +// Transpose a 2x2 complex f32 matrix in NEON registers: +// left: [a0.re, a0.im, a1.re, a1.im] (row 0: a0 = col 0, a1 = col 1) +// right: [b0.re, b0.im, b1.re, b1.im] (row 1: b0 = col 0, b1 = col 1) +// Returns: +// [0]: [a0.re, a0.im, b0.re, b0.im] (col 0: a0 = row 0, b0 = row 1) +// [1]: [a1.re, a1.im, b1.re, b1.im] (col 1: a1 = row 0, b1 = row 1) #[inline(always)] pub unsafe fn transpose_complex_2x2_f32(left: float32x4_t, right: float32x4_t) -> [float32x4_t; 2] { let temp02 = extract_lo_lo_f32(left, right); @@ -246,6 +253,190 @@ impl Rotate90F64 { } } +/// Fused complex twiddle multiplication and matrix transpose for small f64 sizes. +/// +/// Computes `output[y + x * height] = input[x + y * width] * twiddles[x + y * width]`. +/// Processes 2x2 complex blocks with NEON SIMD vector instructions, +/// followed by remainder column and row handling. +#[inline(always)] +pub unsafe fn transpose_small_twiddle_f64( + input: impl NeonArray, + mut output: impl NeonArrayMut, + twiddles: impl NeonArray, + width: usize, + height: usize, +) { + // Loop 1: process 2 rows at a time + let mut y = 0; + while y + 2 <= height { + let in_r0 = y * width; + let in_r1 = (y + 1) * width; + + // Loop 2: process 2 columns at a time (2x2 complex block) + let mut x = 0; + while x + 2 <= width { + let out_c0 = y + x * height; + let out_c1 = y + (x + 1) * height; + + let a0 = input.load_complex(in_r0 + x); + let tw_a0 = twiddles.load_complex(in_r0 + x); + let a1 = input.load_complex(in_r0 + x + 1); + let tw_a1 = twiddles.load_complex(in_r0 + x + 1); + + let b0 = input.load_complex(in_r1 + x); + let tw_b0 = twiddles.load_complex(in_r1 + x); + let b1 = input.load_complex(in_r1 + x + 1); + let tw_b1 = twiddles.load_complex(in_r1 + x + 1); + + let res_a0 = NeonVector::mul_complex(a0, tw_a0); + let res_a1 = NeonVector::mul_complex(a1, tw_a1); + let res_b0 = NeonVector::mul_complex(b0, tw_b0); + let res_b1 = NeonVector::mul_complex(b1, tw_b1); + + output.store_complex(res_a0, out_c0); + output.store_complex(res_b0, out_c0 + 1); + output.store_complex(res_a1, out_c1); + output.store_complex(res_b1, out_c1 + 1); + + x += 2; + } + + // Remainder column when width is odd (2 rows x 1 column) + if x < width { + let out_c0 = y + x * height; + + let a0 = input.load_complex(in_r0 + x); + let tw_a0 = twiddles.load_complex(in_r0 + x); + let b0 = input.load_complex(in_r1 + x); + let tw_b0 = twiddles.load_complex(in_r1 + x); + + let res_a0 = NeonVector::mul_complex(a0, tw_a0); + let res_b0 = NeonVector::mul_complex(b0, tw_b0); + + output.store_complex(res_a0, out_c0); + output.store_complex(res_b0, out_c0 + 1); + } + + y += 2; + } + + // Remainder row when height is odd (1 row x width columns) + if y < height { + let in_r = y * width; + for x in 0..width { + let in_idx = in_r + x; + let out_idx = y + x * height; + let a = input.load_complex(in_idx); + let tw = twiddles.load_complex(in_idx); + let res = NeonVector::mul_complex(a, tw); + output.store_complex(res, out_idx); + } + } +} + +/// Fused complex twiddle multiplication and matrix transpose for small f32 sizes. +/// +/// Computes `output[y + x * height] = input[x + y * width] * twiddles[x + y * width]`. +/// Processes 2x2 complex blocks with NEON SIMD vector instructions, +/// followed by remainder column and row handling. +#[inline(always)] +pub unsafe fn transpose_small_twiddle_f32( + input: impl NeonArray, + mut output: impl NeonArrayMut, + twiddles: impl NeonArray, + width: usize, + height: usize, +) { + // Loop 1: process 2 rows at a time + let mut y = 0; + while y + 2 <= height { + let in_r0 = y * width; + let in_r1 = (y + 1) * width; + + // Loop 2: process 2 columns at a time (2x2 complex block) + let mut x = 0; + while x + 2 <= width { + let out_c0 = y + x * height; + let out_c1 = y + (x + 1) * height; + + let row0 = input.load_complex(in_r0 + x); + let tw_row0 = twiddles.load_complex(in_r0 + x); + let row1 = input.load_complex(in_r1 + x); + let tw_row1 = twiddles.load_complex(in_r1 + x); + + let res0 = NeonVector::mul_complex(row0, tw_row0); + let res1 = NeonVector::mul_complex(row1, tw_row1); + + let [col0, col1] = transpose_complex_2x2_f32(res0, res1); + + output.store_complex(col0, out_c0); + output.store_complex(col1, out_c1); + + x += 2; + } + + // Remainder column when width is odd (2 rows x 1 column) + if x < width { + let out_c0 = y + x * height; + + let a0 = vget_low_f32(input.load_partial_lo_complex(in_r0 + x)); + let b0 = vget_low_f32(input.load_partial_lo_complex(in_r1 + x)); + let val = vcombine_f32(a0, b0); + + let tw_a0 = vget_low_f32(twiddles.load_partial_lo_complex(in_r0 + x)); + let tw_b0 = vget_low_f32(twiddles.load_partial_lo_complex(in_r1 + x)); + let tw = vcombine_f32(tw_a0, tw_b0); + + let res = NeonVector::mul_complex(val, tw); + + output.store_complex(res, out_c0); + } + + y += 2; + } + + // Remainder row when height is odd (1 row x width columns) + if y < height { + let in_r = y * width; + for x in 0..width { + let in_idx = in_r + x; + let out_idx = y + x * height; + let a = input.load_partial_lo_complex(in_idx); + let tw = twiddles.load_partial_lo_complex(in_idx); + let res = NeonVector::mul_complex(a, tw); + output.store_partial_lo_complex(res, out_idx); + } + } +} + +pub unsafe fn transpose_small_twiddle( + width: usize, + height: usize, + input: &[Complex], + output: &mut [Complex], + twiddles: &[Complex], +) -> bool { + debug_assert!(input.len() >= width * height); + debug_assert!(output.len() >= width * height); + debug_assert!(twiddles.len() >= width * height); + + use std::any::TypeId; + if TypeId::of::() == TypeId::of::() { + let input: &[Complex] = crate::array_utils::workaround_transmute(input); + let output: &mut [Complex] = crate::array_utils::workaround_transmute_mut(output); + let twiddles: &[Complex] = crate::array_utils::workaround_transmute(twiddles); + transpose_small_twiddle_f64(input, output, twiddles, width, height); + return true; + } else if TypeId::of::() == TypeId::of::() { + let input: &[Complex] = crate::array_utils::workaround_transmute(input); + let output: &mut [Complex] = crate::array_utils::workaround_transmute_mut(output); + let twiddles: &[Complex] = crate::array_utils::workaround_transmute(twiddles); + transpose_small_twiddle_f32(input, output, twiddles, width, height); + return true; + } + false +} + #[cfg(test)] mod unit_tests { use super::*; @@ -298,4 +489,72 @@ mod unit_tests { assert_eq!(second, second_expected); } } + + #[test] + fn test_transpose_small_twiddle_neon() { + use num_traits::Zero; + for width in [1, 2, 3, 4, 5, 7, 8, 16, 31, 32, 33] { + for height in [1, 2, 3, 4, 5, 7, 8, 16, 31, 32, 33] { + let len = width * height; + + // f32 + let input_f32: Vec> = (0..len) + .map(|i| Complex::new((i % 7) as f32, (i % 5) as f32)) + .collect(); + let twiddles_f32: Vec> = (0..len) + .map(|i| Complex::new((i % 3) as f32, (i % 4) as f32)) + .collect(); + let mut out_f32 = vec![Complex::zero(); len]; + unsafe { + transpose_small_twiddle(width, height, &input_f32, &mut out_f32, &twiddles_f32); + } + for y in 0..height { + for x in 0..width { + let expected = input_f32[x + y * width] * twiddles_f32[x + y * width]; + let actual = out_f32[y + x * height]; + assert!( + (actual.re - expected.re).abs() < 1e-5 + && (actual.im - expected.im).abs() < 1e-5, + "f32 twiddle mismatch at ({}, {}) for {}x{}: expected {:?}, got {:?}", + x, + y, + width, + height, + expected, + actual + ); + } + } + + // f64 + let input_f64: Vec> = (0..len) + .map(|i| Complex::new((i % 7) as f64, (i % 5) as f64)) + .collect(); + let twiddles_f64: Vec> = (0..len) + .map(|i| Complex::new((i % 3) as f64, (i % 4) as f64)) + .collect(); + let mut out_f64 = vec![Complex::zero(); len]; + unsafe { + transpose_small_twiddle(width, height, &input_f64, &mut out_f64, &twiddles_f64); + } + for y in 0..height { + for x in 0..width { + let expected = input_f64[x + y * width] * twiddles_f64[x + y * width]; + let actual = out_f64[y + x * height]; + assert!( + (actual.re - expected.re).abs() < 1e-10 + && (actual.im - expected.im).abs() < 1e-10, + "f64 twiddle mismatch at ({}, {}) for {}x{}: expected {:?}, got {:?}", + x, + y, + width, + height, + expected, + actual + ); + } + } + } + } + } }