diff --git a/Cargo.lock b/Cargo.lock index fa335c3..09a96ae 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1398,6 +1398,7 @@ name = "sonora-fft" version = "0.2.0" dependencies = [ "proptest", + "sonora-simd", "test-strategy", ] diff --git a/crates/sonora-bench/Cargo.lock b/crates/sonora-bench/Cargo.lock index 02df083..246eef0 100644 --- a/crates/sonora-bench/Cargo.lock +++ b/crates/sonora-bench/Cargo.lock @@ -909,6 +909,9 @@ dependencies = [ [[package]] name = "sonora-fft" version = "0.2.0" +dependencies = [ + "sonora-simd", +] [[package]] name = "sonora-ns" diff --git a/crates/sonora-bench/tests/cpp_comparison.rs b/crates/sonora-bench/tests/cpp_comparison.rs index 5013e78..beb5210 100644 --- a/crates/sonora-bench/tests/cpp_comparison.rs +++ b/crates/sonora-bench/tests/cpp_comparison.rs @@ -3,10 +3,16 @@ //! Per-component tests verify close matching at each DSP stage. //! Full-pipeline tests verify end-to-end equivalence. //! -//! On ARM (NEON) both Rust and C++ use fused multiply-add producing bit-identical -//! results. On x86 the Rust `mul_add` intrinsic and C++ scalar arithmetic may -//! diverge by a small amount due to different FMA contraction behaviour between -//! LLVM and GCC. The tolerances below accommodate this. +//! On AArch64 both Rust and C++ use fused multiply-add, but the results are not +//! always bit-identical. For example, on aarch64-apple-darwin the optimized C++ +//! build computes each fft4g twiddle `cos`/`sin` pair with one +//! `__sincosf_stret` call, while Rust calls `cosf` and `sinf` separately (and +//! merges only some pairs in release builds); some table entries differ by +//! 1 ulp. On x86 without the `fma` target feature, Rust uses unfused +//! arithmetic in the C++ source order. The C++ reference is built with +//! `-march=native`, so on a CPU with FMA the C++ compiler may contract those +//! expressions into fused multiply-adds, and the results may diverge by a +//! small amount. The tolerances below accommodate this. use sonora::config::{EchoCanceller, GainController2, NoiseSuppression, TransparentModeType}; use sonora::high_pass_filter::HighPassFilter; diff --git a/crates/sonora-common-audio/src/cascaded_biquad_filter.rs b/crates/sonora-common-audio/src/cascaded_biquad_filter.rs index d506d51..a4e2e0a 100644 --- a/crates/sonora-common-audio/src/cascaded_biquad_filter.rs +++ b/crates/sonora-common-audio/src/cascaded_biquad_filter.rs @@ -2,6 +2,8 @@ //! //! Ported from `modules/audio_processing/utility/cascaded_biquad_filter.h/cc`. +use sonora_simd::NATIVE_FMA; + /// Coefficients for a single second-order (biquad) IIR section. /// /// Transfer function: `H(z) = (b[0] + b[1]*z^-1 + b[2]*z^-2) / (1 + a[0]*z^-1 + a[1]*z^-2)` @@ -52,11 +54,20 @@ impl CascadedBiQuadFilter { /// Filters `x` into `y` (separate input/output). pub fn process(&mut self, x: &[f32], y: &mut [f32]) { + self.process_with::(x, y); + } + + /// [`Self::process`], with the recursion fused by [`f32::mul_add`] if + /// `FMA`, else in the C++ expression and its operation order. The public + /// methods pass [`NATIVE_FMA`]; the tests run both forms. Always inlined, + /// so that LLVM optimizes this body as part of `process`. + #[inline(always)] + fn process_with(&mut self, x: &[f32], y: &mut [f32]) { if self.biquads.is_empty() { y.copy_from_slice(x); return; } - Self::apply_biquad(x, y, &mut self.biquads[0]); + Self::apply_biquad::(x, y, &mut self.biquads[0]); for k in 1..self.biquads.len() { // Split borrow: process y in-place through remaining stages. let (_, rest) = self.biquads.split_at_mut(k); @@ -73,14 +84,17 @@ impl CascadedBiQuadFilter { let mut m_y_1 = bq.y[1]; for v in y.iter_mut() { let tmp = *v; - // Use mul_add chains to emit fmadd/fmsub instructions. - *v = c_b_0.mul_add( - tmp, - c_b_1.mul_add( - m_x_0, - c_b_2.mul_add(m_x_1, (-c_a_0).mul_add(m_y_0, -c_a_1 * m_y_1)), - ), - ); + *v = if FMA { + c_b_0.mul_add( + tmp, + c_b_1.mul_add( + m_x_0, + c_b_2.mul_add(m_x_1, (-c_a_0).mul_add(m_y_0, -c_a_1 * m_y_1)), + ), + ) + } else { + c_b_0 * tmp + c_b_1 * m_x_0 + c_b_2 * m_x_1 - c_a_0 * m_y_0 - c_a_1 * m_y_1 + }; m_x_1 = m_x_0; m_x_0 = tmp; m_y_1 = m_y_0; @@ -92,7 +106,19 @@ impl CascadedBiQuadFilter { } /// Filters `y` in-place through all stages. + // rustc makes small functions that call nothing but intrinsics available + // for inlining in other crates; this wrapper calls + // `process_in_place_with`, so it needs `#[inline]` for that. + #[inline] pub fn process_in_place(&mut self, y: &mut [f32]) { + self.process_in_place_with::(y); + } + + /// [`Self::process_in_place`], in the form `FMA` selects; see + /// [`Self::process_with`]. Not forced inline, unlike `process_with`: other + /// crates inline `process_in_place`, and forcing this body into it would + /// change how they inline the filter. + fn process_in_place_with(&mut self, y: &mut [f32]) { for bq in &mut self.biquads { let c_b_0 = bq.coefficients.b[0]; let c_b_1 = bq.coefficients.b[1]; @@ -105,13 +131,17 @@ impl CascadedBiQuadFilter { let mut m_y_1 = bq.y[1]; for v in y.iter_mut() { let tmp = *v; - *v = c_b_0.mul_add( - tmp, - c_b_1.mul_add( - m_x_0, - c_b_2.mul_add(m_x_1, (-c_a_0).mul_add(m_y_0, -c_a_1 * m_y_1)), - ), - ); + *v = if FMA { + c_b_0.mul_add( + tmp, + c_b_1.mul_add( + m_x_0, + c_b_2.mul_add(m_x_1, (-c_a_0).mul_add(m_y_0, -c_a_1 * m_y_1)), + ), + ) + } else { + c_b_0 * tmp + c_b_1 * m_x_0 + c_b_2 * m_x_1 - c_a_0 * m_y_0 - c_a_1 * m_y_1 + }; m_x_1 = m_x_0; m_x_0 = tmp; m_y_1 = m_y_0; @@ -129,7 +159,7 @@ impl CascadedBiQuadFilter { } } - fn apply_biquad(x: &[f32], y: &mut [f32], bq: &mut BiQuad) { + fn apply_biquad(x: &[f32], y: &mut [f32], bq: &mut BiQuad) { debug_assert_eq!(x.len(), y.len()); let c_b_0 = bq.coefficients.b[0]; let c_b_1 = bq.coefficients.b[1]; @@ -142,13 +172,17 @@ impl CascadedBiQuadFilter { let mut m_y_1 = bq.y[1]; for (xi, yi) in x.iter().zip(y.iter_mut()) { let tmp = *xi; - *yi = c_b_0.mul_add( - tmp, - c_b_1.mul_add( - m_x_0, - c_b_2.mul_add(m_x_1, (-c_a_0).mul_add(m_y_0, -c_a_1 * m_y_1)), - ), - ); + *yi = if FMA { + c_b_0.mul_add( + tmp, + c_b_1.mul_add( + m_x_0, + c_b_2.mul_add(m_x_1, (-c_a_0).mul_add(m_y_0, -c_a_1 * m_y_1)), + ), + ) + } else { + c_b_0 * tmp + c_b_1 * m_x_0 + c_b_2 * m_x_1 - c_a_0 * m_y_0 - c_a_1 * m_y_1 + }; m_x_1 = m_x_0; m_x_0 = tmp; m_y_1 = m_y_0; @@ -230,6 +264,101 @@ mod tests { } } + /// Reference cascade: the fused `mul_add` chain, or the C++ expression + /// `c_b_0 * tmp + c_b_1 * m_x_0 + c_b_2 * m_x_1 - c_a_0 * m_y_0 - c_a_1 * m_y_1`. + fn reference_cascade(coeffs: &[BiQuadCoefficients], x: &[f32], fused: bool) -> Vec { + let mut y = x.to_vec(); + for c in coeffs { + let (mut x0, mut x1, mut y0, mut y1) = (0.0_f32, 0.0_f32, 0.0_f32, 0.0_f32); + for v in &mut y { + let tmp = *v; + *v = if fused { + c.b[0].mul_add( + tmp, + c.b[1].mul_add(x0, c.b[2].mul_add(x1, (-c.a[0]).mul_add(y0, -c.a[1] * y1))), + ) + } else { + c.b[0] * tmp + c.b[1] * x0 + c.b[2] * x1 - c.a[0] * y0 - c.a[1] * y1 + }; + x1 = x0; + x0 = tmp; + y1 = y0; + y0 = *v; + } + } + y + } + + /// A high-pass section, then the low-pass section above. Two stages + /// cover all three loops: `apply_biquad` and the in-place stage loop + /// inside `process_with`, and the loop in `process_in_place_with`. + fn two_stages() -> [BiQuadCoefficients; 2] { + [ + BiQuadCoefficients { + b: [0.972_613, -1.945_226, 0.972_613], + a: [-1.944_48, 0.945_976], + }, + lowpass_coefficients(), + ] + } + + /// Input on which the fused and the C++ forms of [`two_stages`] differ. + fn test_signal() -> Vec { + (0..32) + .map(|i| (i as f32 * 0.618_034).fract() - 0.5) + .collect() + } + + fn bits(v: &[f32]) -> Vec { + v.iter().map(|x| x.to_bits()).collect() + } + + /// Output bits of `process_with::` and `process_in_place_with::`, + /// each on a new filter. + fn run_forms(coeffs: &[BiQuadCoefficients], x: &[f32]) -> [Vec; 2] { + let mut y = vec![0.0_f32; x.len()]; + CascadedBiQuadFilter::new(coeffs).process_with::(x, &mut y); + let mut in_place = x.to_vec(); + CascadedBiQuadFilter::new(coeffs).process_in_place_with::(&mut in_place); + [bits(&y), bits(&in_place)] + } + + /// Both forms of all three loops, on every target: the fused form must + /// match the `mul_add` chain, and the plain form the C++ expression, bit + /// for bit. + #[test] + fn recursion_matches_fused_and_plain_references() { + let coeffs = two_stages(); + let input = test_signal(); + let fused = bits(&reference_cascade(&coeffs, &input, true)); + let plain = bits(&reference_cascade(&coeffs, &input, false)); + assert_ne!(fused, plain); + + assert_eq!(run_forms::(&coeffs, &input), [fused.clone(), fused]); + assert_eq!(run_forms::(&coeffs, &input), [plain.clone(), plain]); + } + + /// `process` and `process_in_place` run the [`NATIVE_FMA`] form. + #[test] + fn public_methods_use_native_fma() { + let coeffs = two_stages(); + let input = test_signal(); + // The forms differ on this input, so the comparison below can fail. + assert_ne!( + run_forms::(&coeffs, &input), + run_forms::(&coeffs, &input) + ); + + let mut y = vec![0.0_f32; input.len()]; + CascadedBiQuadFilter::new(&coeffs).process(&input, &mut y); + let mut in_place = input.clone(); + CascadedBiQuadFilter::new(&coeffs).process_in_place(&mut in_place); + assert_eq!( + [bits(&y), bits(&in_place)], + run_forms::(&coeffs, &input) + ); + } + #[test] fn multi_stage_filter() { let coeffs = [lowpass_coefficients(), lowpass_coefficients()]; diff --git a/crates/sonora-fft/Cargo.toml b/crates/sonora-fft/Cargo.toml index d37ca47..0bdf436 100644 --- a/crates/sonora-fft/Cargo.toml +++ b/crates/sonora-fft/Cargo.toml @@ -13,6 +13,9 @@ repository.workspace = true [lints] workspace = true +[dependencies] +sonora-simd = { workspace = true } + [dev-dependencies] proptest = { workspace = true } test-strategy = { workspace = true } diff --git a/crates/sonora-fft/src/fft4g.rs b/crates/sonora-fft/src/fft4g.rs index 69e0347..df13ed1 100644 --- a/crates/sonora-fft/src/fft4g.rs +++ b/crates/sonora-fft/src/fft4g.rs @@ -33,6 +33,8 @@ use std::f32::consts::FRAC_PI_4; use std::ptr; +use sonora_simd::NATIVE_FMA; + /// Variable-size real FFT using Ooura's fft4g algorithm. /// /// Supports power-of-2 sizes (`n >= 2`). Twiddle tables and bit-reversal @@ -109,15 +111,25 @@ impl Fft4g { /// /// Panics if `a.len() != n`. pub fn rdft(&self, a: &mut [f32]) { + self.rdft_with::(a); + } + + /// [`Self::rdft`], with the twiddle multiplications fused by + /// [`f32::mul_add`] if `FMA`, else in the plain expressions of the C + /// reference and their operation order. The public methods pass + /// [`NATIVE_FMA`]; the tests run both forms. Always inlined, so that + /// LLVM optimizes this body as part of `rdft`. + #[inline(always)] + fn rdft_with(&self, a: &mut [f32]) { assert_eq!(a.len(), self.n, "input length must be {}", self.n); let n = self.n; if n > 4 { apply_bitrv2(&self.bitrv_ip, self.bitrv_m, self.bitrv_long, a); - cftfsub(n, a, &self.w); - rftfsub(n, a, self.nc, &self.w[self.nw..]); + cftfsub::(n, a, &self.w); + rftfsub::(n, a, self.nc, &self.w[self.nw..]); } else if n == 4 { - cftfsub(n, a, &self.w); + cftfsub::(n, a, &self.w); } let xi = a[0] - a[1]; a[0] += a[1]; @@ -133,17 +145,23 @@ impl Fft4g { /// /// Panics if `a.len() != n`. pub fn irdft(&self, a: &mut [f32]) { + self.irdft_with::(a); + } + + /// [`Self::irdft`], in the form `FMA` selects; see [`Self::rdft_with`]. + #[inline(always)] + fn irdft_with(&self, a: &mut [f32]) { assert_eq!(a.len(), self.n, "input length must be {}", self.n); let n = self.n; a[1] = 0.5 * (a[0] - a[1]); a[0] -= a[1]; if n > 4 { - rftbsub(n, a, self.nc, &self.w[self.nw..]); + rftbsub::(n, a, self.nc, &self.w[self.nw..]); apply_bitrv2(&self.bitrv_ip, self.bitrv_m, self.bitrv_long, a); - cftbsub(n, a, &self.w); + cftbsub::(n, a, &self.w); } else if n == 4 { - cftfsub(n, a, &self.w); + cftfsub::(n, a, &self.w); } } } @@ -352,13 +370,13 @@ fn bitrv2(n: usize, ip: &mut [usize], a: &mut [f32]) { } /// Forward complex sub-transform (radix-4 decomposition). -fn cftfsub(n: usize, a: &mut [f32], w: &[f32]) { +fn cftfsub(n: usize, a: &mut [f32], w: &[f32]) { let mut l = 2; if n > 8 { - cft1st(n, a, w); + cft1st::(n, a, w); l = 8; while (l << 2) < n { - cftmdl(n, l, a, w); + cftmdl::(n, l, a, w); l <<= 2; } } @@ -401,13 +419,13 @@ fn cftfsub(n: usize, a: &mut [f32], w: &[f32]) { } /// Backward complex sub-transform (radix-4 decomposition). -fn cftbsub(n: usize, a: &mut [f32], w: &[f32]) { +fn cftbsub(n: usize, a: &mut [f32], w: &[f32]) { let mut l = 2; if n > 8 { - cft1st(n, a, w); + cft1st::(n, a, w); l = 8; while (l << 2) < n { - cftmdl(n, l, a, w); + cftmdl::(n, l, a, w); l <<= 2; } } @@ -453,7 +471,7 @@ fn cftbsub(n: usize, a: &mut [f32], w: &[f32]) { /// /// # Safety contract /// `a.len() >= n >= 16` and `w.len() >= n/4`. All indices stay within bounds. -fn cft1st(n: usize, a: &mut [f32], w: &[f32]) { +fn cft1st(n: usize, a: &mut [f32], w: &[f32]) { // SAFETY: All indices into `a` are in 0..n and all indices into `w` are // in 0..n/4. This is guaranteed by the Ooura algorithm structure and // validated by the test suite. @@ -506,8 +524,14 @@ fn cft1st(n: usize, a: &mut [f32], w: &[f32]) { let wk2i = get(w, k1 + 1); let wk1r = get(w, k2); let wk1i = get(w, k2 + 1); - let wk3r = (-2.0 * wk2i).mul_add(wk1i, wk1r); - let wk3i = (2.0 * wk2i).mul_add(wk1r, -wk1i); + let (wk3r, wk3i) = if FMA { + ( + (-2.0 * wk2i).mul_add(wk1i, wk1r), + (2.0 * wk2i).mul_add(wk1r, -wk1i), + ) + } else { + (wk1r - 2.0 * wk2i * wk1i, 2.0 * wk2i * wk1r - wk1i) + }; let x0r = get(a, j) + get(a, j + 2); let x0i = get(a, j + 1) + get(a, j + 3); @@ -521,21 +545,42 @@ fn cft1st(n: usize, a: &mut [f32], w: &[f32]) { set(a, j + 1, x0i + x2i); let x0r = x0r - x2r; let x0i = x0i - x2i; - set(a, j + 4, wk2r.mul_add(x0r, -wk2i * x0i)); - set(a, j + 5, wk2r.mul_add(x0i, wk2i * x0r)); + if FMA { + set(a, j + 4, wk2r.mul_add(x0r, -wk2i * x0i)); + set(a, j + 5, wk2r.mul_add(x0i, wk2i * x0r)); + } else { + set(a, j + 4, wk2r * x0r - wk2i * x0i); + set(a, j + 5, wk2r * x0i + wk2i * x0r); + } let x0r = x1r - x3i; let x0i = x1i + x3r; - set(a, j + 2, wk1r.mul_add(x0r, -wk1i * x0i)); - set(a, j + 3, wk1r.mul_add(x0i, wk1i * x0r)); + if FMA { + set(a, j + 2, wk1r.mul_add(x0r, -wk1i * x0i)); + set(a, j + 3, wk1r.mul_add(x0i, wk1i * x0r)); + } else { + set(a, j + 2, wk1r * x0r - wk1i * x0i); + set(a, j + 3, wk1r * x0i + wk1i * x0r); + } let x0r = x1r + x3i; let x0i = x1i - x3r; - set(a, j + 6, wk3r.mul_add(x0r, -wk3i * x0i)); - set(a, j + 7, wk3r.mul_add(x0i, wk3i * x0r)); + if FMA { + set(a, j + 6, wk3r.mul_add(x0r, -wk3i * x0i)); + set(a, j + 7, wk3r.mul_add(x0i, wk3i * x0r)); + } else { + set(a, j + 6, wk3r * x0r - wk3i * x0i); + set(a, j + 7, wk3r * x0i + wk3i * x0r); + } let wk1r = get(w, k2 + 2); let wk1i = get(w, k2 + 3); - let wk3r = (-2.0 * wk2r).mul_add(wk1i, wk1r); - let wk3i = (2.0 * wk2r).mul_add(wk1r, -wk1i); + let (wk3r, wk3i) = if FMA { + ( + (-2.0 * wk2r).mul_add(wk1i, wk1r), + (2.0 * wk2r).mul_add(wk1r, -wk1i), + ) + } else { + (wk1r - 2.0 * wk2r * wk1i, 2.0 * wk2r * wk1r - wk1i) + }; let x0r = get(a, j + 8) + get(a, j + 10); let x0i = get(a, j + 9) + get(a, j + 11); @@ -549,16 +594,31 @@ fn cft1st(n: usize, a: &mut [f32], w: &[f32]) { set(a, j + 9, x0i + x2i); let x0r = x0r - x2r; let x0i = x0i - x2i; - set(a, j + 12, (-wk2i).mul_add(x0r, -wk2r * x0i)); - set(a, j + 13, (-wk2i).mul_add(x0i, wk2r * x0r)); + if FMA { + set(a, j + 12, (-wk2i).mul_add(x0r, -wk2r * x0i)); + set(a, j + 13, (-wk2i).mul_add(x0i, wk2r * x0r)); + } else { + set(a, j + 12, -wk2i * x0r - wk2r * x0i); + set(a, j + 13, -wk2i * x0i + wk2r * x0r); + } let x0r = x1r - x3i; let x0i = x1i + x3r; - set(a, j + 10, wk1r.mul_add(x0r, -wk1i * x0i)); - set(a, j + 11, wk1r.mul_add(x0i, wk1i * x0r)); + if FMA { + set(a, j + 10, wk1r.mul_add(x0r, -wk1i * x0i)); + set(a, j + 11, wk1r.mul_add(x0i, wk1i * x0r)); + } else { + set(a, j + 10, wk1r * x0r - wk1i * x0i); + set(a, j + 11, wk1r * x0i + wk1i * x0r); + } let x0r = x1r + x3i; let x0i = x1i - x3r; - set(a, j + 14, wk3r.mul_add(x0r, -wk3i * x0i)); - set(a, j + 15, wk3r.mul_add(x0i, wk3i * x0r)); + if FMA { + set(a, j + 14, wk3r.mul_add(x0r, -wk3i * x0i)); + set(a, j + 15, wk3r.mul_add(x0i, wk3i * x0r)); + } else { + set(a, j + 14, wk3r * x0r - wk3i * x0i); + set(a, j + 15, wk3r * x0i + wk3i * x0r); + } j += 16; } @@ -569,7 +629,7 @@ fn cft1st(n: usize, a: &mut [f32], w: &[f32]) { /// /// # Safety contract /// All indices into `a` are in `0..n` and into `w` are in `0..n/4`. -fn cftmdl(n: usize, l: usize, a: &mut [f32], w: &[f32]) { +fn cftmdl(n: usize, l: usize, a: &mut [f32], w: &[f32]) { let m = l << 2; // SAFETY: All indices are bounded by `n` (see module-level doc). @@ -633,8 +693,14 @@ fn cftmdl(n: usize, l: usize, a: &mut [f32], w: &[f32]) { let wk2i = get(w, k1 + 1); let wk1r = get(w, k2); let wk1i = get(w, k2 + 1); - let wk3r = (-2.0 * wk2i).mul_add(wk1i, wk1r); - let wk3i = (2.0 * wk2i).mul_add(wk1r, -wk1i); + let (wk3r, wk3i) = if FMA { + ( + (-2.0 * wk2i).mul_add(wk1i, wk1r), + (2.0 * wk2i).mul_add(wk1r, -wk1i), + ) + } else { + (wk1r - 2.0 * wk2i * wk1i, 2.0 * wk2i * wk1r - wk1i) + }; for j in (k..l + k).step_by(2) { let j1 = j + l; @@ -652,22 +718,43 @@ fn cftmdl(n: usize, l: usize, a: &mut [f32], w: &[f32]) { set(a, j + 1, x0i + x2i); let x0r = x0r - x2r; let x0i = x0i - x2i; - set(a, j2, wk2r.mul_add(x0r, -wk2i * x0i)); - set(a, j2 + 1, wk2r.mul_add(x0i, wk2i * x0r)); + if FMA { + set(a, j2, wk2r.mul_add(x0r, -wk2i * x0i)); + set(a, j2 + 1, wk2r.mul_add(x0i, wk2i * x0r)); + } else { + set(a, j2, wk2r * x0r - wk2i * x0i); + set(a, j2 + 1, wk2r * x0i + wk2i * x0r); + } let x0r = x1r - x3i; let x0i = x1i + x3r; - set(a, j1, wk1r.mul_add(x0r, -wk1i * x0i)); - set(a, j1 + 1, wk1r.mul_add(x0i, wk1i * x0r)); + if FMA { + set(a, j1, wk1r.mul_add(x0r, -wk1i * x0i)); + set(a, j1 + 1, wk1r.mul_add(x0i, wk1i * x0r)); + } else { + set(a, j1, wk1r * x0r - wk1i * x0i); + set(a, j1 + 1, wk1r * x0i + wk1i * x0r); + } let x0r = x1r + x3i; let x0i = x1i - x3r; - set(a, j3, wk3r.mul_add(x0r, -wk3i * x0i)); - set(a, j3 + 1, wk3r.mul_add(x0i, wk3i * x0r)); + if FMA { + set(a, j3, wk3r.mul_add(x0r, -wk3i * x0i)); + set(a, j3 + 1, wk3r.mul_add(x0i, wk3i * x0r)); + } else { + set(a, j3, wk3r * x0r - wk3i * x0i); + set(a, j3 + 1, wk3r * x0i + wk3i * x0r); + } } let wk1r = get(w, k2 + 2); let wk1i = get(w, k2 + 3); - let wk3r = (-2.0 * wk2r).mul_add(wk1i, wk1r); - let wk3i = (2.0 * wk2r).mul_add(wk1r, -wk1i); + let (wk3r, wk3i) = if FMA { + ( + (-2.0 * wk2r).mul_add(wk1i, wk1r), + (2.0 * wk2r).mul_add(wk1r, -wk1i), + ) + } else { + (wk1r - 2.0 * wk2r * wk1i, 2.0 * wk2r * wk1r - wk1i) + }; for j in (k + m..l + (k + m)).step_by(2) { let j1 = j + l; @@ -685,16 +772,31 @@ fn cftmdl(n: usize, l: usize, a: &mut [f32], w: &[f32]) { set(a, j + 1, x0i + x2i); let x0r = x0r - x2r; let x0i = x0i - x2i; - set(a, j2, (-wk2i).mul_add(x0r, -wk2r * x0i)); - set(a, j2 + 1, (-wk2i).mul_add(x0i, wk2r * x0r)); + if FMA { + set(a, j2, (-wk2i).mul_add(x0r, -wk2r * x0i)); + set(a, j2 + 1, (-wk2i).mul_add(x0i, wk2r * x0r)); + } else { + set(a, j2, -wk2i * x0r - wk2r * x0i); + set(a, j2 + 1, -wk2i * x0i + wk2r * x0r); + } let x0r = x1r - x3i; let x0i = x1i + x3r; - set(a, j1, wk1r.mul_add(x0r, -wk1i * x0i)); - set(a, j1 + 1, wk1r.mul_add(x0i, wk1i * x0r)); + if FMA { + set(a, j1, wk1r.mul_add(x0r, -wk1i * x0i)); + set(a, j1 + 1, wk1r.mul_add(x0i, wk1i * x0r)); + } else { + set(a, j1, wk1r * x0r - wk1i * x0i); + set(a, j1 + 1, wk1r * x0i + wk1i * x0r); + } let x0r = x1r + x3i; let x0i = x1i - x3r; - set(a, j3, wk3r.mul_add(x0r, -wk3i * x0i)); - set(a, j3 + 1, wk3r.mul_add(x0i, wk3i * x0r)); + if FMA { + set(a, j3, wk3r.mul_add(x0r, -wk3i * x0i)); + set(a, j3 + 1, wk3r.mul_add(x0i, wk3i * x0r)); + } else { + set(a, j3, wk3r * x0r - wk3i * x0i); + set(a, j3 + 1, wk3r * x0i + wk3i * x0r); + } } k += m2; @@ -703,7 +805,7 @@ fn cftmdl(n: usize, l: usize, a: &mut [f32], w: &[f32]) { } /// Real FFT forward post-processing (split-radix real/imaginary separation). -fn rftfsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { +fn rftfsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { let m = n >> 1; let ks = 2 * nc / m; let mut kk = 0; @@ -717,8 +819,11 @@ fn rftfsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { let wki = get(c, kk); let xr = get(a, j) - get(a, k); let xi = get(a, j + 1) + get(a, k + 1); - let yr = wkr.mul_add(xr, -wki * xi); - let yi = wkr.mul_add(xi, wki * xr); + let (yr, yi) = if FMA { + (wkr.mul_add(xr, -wki * xi), wkr.mul_add(xi, wki * xr)) + } else { + (wkr * xr - wki * xi, wkr * xi + wki * xr) + }; set(a, j, get(a, j) - yr); set(a, j + 1, get(a, j + 1) - yi); set(a, k, get(a, k) + yr); @@ -729,7 +834,7 @@ fn rftfsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { } /// Real FFT backward pre-processing (split-radix real/imaginary recombination). -fn rftbsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { +fn rftbsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { let m = n >> 1; let ks = 2 * nc / m; let mut kk = 0; @@ -744,8 +849,11 @@ fn rftbsub(n: usize, a: &mut [f32], nc: usize, c: &[f32]) { let wki = get(c, kk); let xr = get(a, j) - get(a, k); let xi = get(a, j + 1) + get(a, k + 1); - let yr = wkr.mul_add(xr, wki * xi); - let yi = wkr.mul_add(xi, -wki * xr); + let (yr, yi) = if FMA { + (wkr.mul_add(xr, wki * xi), wkr.mul_add(xi, -wki * xr)) + } else { + (wkr * xr + wki * xi, wkr * xi - wki * xr) + }; set(a, j, get(a, j) - yr); set(a, j + 1, yi - get(a, j + 1)); set(a, k, get(a, k) + yr); @@ -858,6 +966,364 @@ mod tests { } } + /// One twiddled radix-4 butterfly of `cft1st` (`l = 2`) and `cftmdl`, on + /// the points `j`, `j + l`, `j + 2l` and `j + 3l`, with the twiddles the + /// C code loads at `k1` for the first or the `second` group. + fn reference_butterfly( + a: &mut [f32], + j: usize, + l: usize, + w: &[f32], + k1: usize, + second: bool, + fused: bool, + ) { + let cmul = |(wr, wi): (f32, f32), (xr, xi): (f32, f32)| { + if fused { + (wr.mul_add(xr, -wi * xi), wr.mul_add(xi, wi * xr)) + } else { + (wr * xr - wi * xi, wr * xi + wi * xr) + } + }; + let (k2, wk2r, wk2i) = (2 * k1, w[k1], w[k1 + 1]); + let (w2, s, (wk1r, wk1i)) = if second { + ((-wk2i, wk2r), wk2r, (w[k2 + 2], w[k2 + 3])) + } else { + ((wk2r, wk2i), wk2i, (w[k2], w[k2 + 1])) + }; + let w3 = if fused { + ( + (-2.0 * s).mul_add(wk1i, wk1r), + (2.0 * s).mul_add(wk1r, -wk1i), + ) + } else { + (wk1r - 2.0 * s * wk1i, 2.0 * s * wk1r - wk1i) + }; + let (j1, j2, j3) = (j + l, j + 2 * l, j + 3 * l); + let x0 = (a[j] + a[j1], a[j + 1] + a[j1 + 1]); + let x1 = (a[j] - a[j1], a[j + 1] - a[j1 + 1]); + let x2 = (a[j2] + a[j3], a[j2 + 1] + a[j3 + 1]); + let x3 = (a[j2] - a[j3], a[j2 + 1] - a[j3 + 1]); + let out = [ + (j, (x0.0 + x2.0, x0.1 + x2.1)), + (j2, cmul(w2, (x0.0 - x2.0, x0.1 - x2.1))), + (j1, cmul((wk1r, wk1i), (x1.0 - x3.1, x1.1 + x3.0))), + (j3, cmul(w3, (x1.0 + x3.1, x1.1 - x3.0))), + ]; + for (i, (re, im)) in out { + a[i] = re; + a[i + 1] = im; + } + } + + fn bits(v: &[f32]) -> Vec { + v.iter().map(|x| x.to_bits()).collect() + } + + /// Asserts that `fused` and `plain`, the outputs of a pass of butterflies + /// with span `l` (`cft1st` is `l = 2`), differ in every twiddled output: + /// at `j + l`, `j + 2l` and `j + 3l`, real and imaginary, in both groups, + /// for at least one butterfly. A statement that computes one of these in + /// the wrong form then fails the bit-exact comparison on its own. + fn assert_each_twiddled_output_differs(fused: &[f32], plain: &[f32], l: usize) { + let m = 4 * l; + // differs[group][quarter][part]: quarter q holds a[j + q * l], and + // quarter 0, a[j] = x0 + x2, has no twiddle. + let mut differs = [[[false; 2]; 4]; 2]; + // The first block of 2m values has no `mul_add`. + for p in 2 * m..fused.len() { + if fused[p].to_bits() != plain[p].to_bits() { + let (group, q) = (p % (2 * m) / m, p % m); + differs[group][q / l][q % 2] = true; + } + } + for (group, quarters) in differs.iter().enumerate() { + for (quarter, parts) in quarters.iter().enumerate().skip(1) { + for (part, &differs) in ["real", "imaginary"].iter().zip(parts) { + assert!( + differs, + "l {l}, group {group}: {part} a[j + {quarter}l] same in both forms" + ); + } + } + } + } + + /// `n` pseudo-random values in `-0.5..0.5`, the same on every target. + fn noise(n: usize) -> Vec { + let mut x = 1_u32; + (0..n) + .map(|_| { + x = x.wrapping_mul(1_664_525).wrapping_add(1_013_904_223); + (x >> 8) as f32 / (1 << 24) as f32 - 0.5 + }) + .collect() + } + + /// Both forms of the twiddled butterflies of `cft1st` and `cftmdl`, on + /// every target: the fused form must match the `mul_add` chains, and the + /// plain form the C expressions, bit for bit. The input is pseudo-random: + /// in the golden-ratio sequence of the other tests, values `l` apart + /// differ by almost the same amount, and some outputs then matched + /// between the forms in every butterfly. With this input and these + /// sizes, each twiddled output differs in dozens of butterflies. + #[test] + fn butterflies_match_fused_and_plain_references() { + // cft1st, n = 4096: iteration t runs at j = 16t with k1 = 2t, as two + // groups at l = 2. The untwiddled block a[0..16] is not checked. + let fft = Fft4g::new(4096); + let w = &fft.w; + let input = noise(4096); + let [fused, plain] = [true, false].map(|fused| { + let mut e = input.clone(); + for t in 1..256 { + reference_butterfly(&mut e, 16 * t, 2, w, 2 * t, false, fused); + reference_butterfly(&mut e, 16 * t + 8, 2, w, 2 * t, true, fused); + } + e + }); + assert_each_twiddled_output_differs(&fused, &plain, 2); + let mut a = input.clone(); + cft1st::(4096, &mut a, w); + assert_eq!(bits(&a[16..]), bits(&fused[16..])); + let mut a = input; + cft1st::(4096, &mut a, w); + assert_eq!(bits(&a[16..]), bits(&plain[16..])); + + // cftmdl, n = 8192, l = 8: iteration t runs at k = 64t with k1 = 2t. + // The untwiddled block a[0..64] is not checked. + let fft = Fft4g::new(8192); + let w = &fft.w; + let input = noise(8192); + let [fused, plain] = [true, false].map(|fused| { + let mut e = input.clone(); + for t in 1..128 { + for j in (64 * t..64 * t + 8).step_by(2) { + reference_butterfly(&mut e, j, 8, w, 2 * t, false, fused); + reference_butterfly(&mut e, j + 32, 8, w, 2 * t, true, fused); + } + } + e + }); + assert_each_twiddled_output_differs(&fused, &plain, 8); + let mut a = input.clone(); + cftmdl::(8192, 8, &mut a, w); + assert_eq!(bits(&a[64..]), bits(&fused[64..])); + let mut a = input; + cftmdl::(8192, 8, &mut a, w); + assert_eq!(bits(&a[64..]), bits(&plain[64..])); + } + + /// Both forms of `rftfsub` and `rftbsub`, on every target: the fused + /// form must match the `mul_add` result, and the plain form the C + /// expressions, bit for bit. + #[test] + fn rft_sub_matches_fused_and_plain_references() { + // With n = 8 each routine runs one iteration: j = 2, k = 6, + // wkr = 0.5 - c[1], wki = c[1]. a[2] = a[3] = 0 makes a[2] and a[3] + // carry yr and yi exactly. + let c = [0.0_f32, 0.3]; + let input = [0.4_f32, -0.8, 0.0, 0.0, 0.6, 0.9, 0.09, 0.65]; + let (wkr, wki) = (0.5 - c[1], c[1]); + let (xr, xi) = (input[2] - input[6], input[3] + input[7]); + + // rftfsub: C `yr = wkr * xr - wki * xi; yi = wkr * xi + wki * xr;` + let fused = (wkr.mul_add(xr, -wki * xi), wkr.mul_add(xi, wki * xr)); + let plain = (wkr * xr - wki * xi, wkr * xi + wki * xr); + assert_ne!(fused.0.to_bits(), plain.0.to_bits()); + assert_ne!(fused.1.to_bits(), plain.1.to_bits()); + let expected = |(yr, yi): (f32, f32)| { + let mut e = input; + e[2] -= yr; + e[3] -= yi; + e[6] += yr; + e[7] -= yi; + e.map(f32::to_bits) + }; + let mut a = input; + rftfsub::(8, &mut a, 2, &c); + assert_eq!(a.map(f32::to_bits), expected(fused)); + let mut a = input; + rftfsub::(8, &mut a, 2, &c); + assert_eq!(a.map(f32::to_bits), expected(plain)); + + // rftbsub: C `yr = wkr * xr + wki * xi; yi = wkr * xi - wki * xr;` + let fused = (wkr.mul_add(xr, wki * xi), wkr.mul_add(xi, -wki * xr)); + let plain = (wkr * xr + wki * xi, wkr * xi - wki * xr); + assert_ne!(fused.0.to_bits(), plain.0.to_bits()); + assert_ne!(fused.1.to_bits(), plain.1.to_bits()); + let expected = |(yr, yi): (f32, f32)| { + let mut e = input; + e[1] = -e[1]; + e[2] -= yr; + e[3] = yi - e[3]; + e[6] += yr; + e[7] = yi - e[7]; + e[5] = -e[5]; + e.map(f32::to_bits) + }; + let mut a = input; + rftbsub::(8, &mut a, 2, &c); + assert_eq!(a.map(f32::to_bits), expected(fused)); + let mut a = input; + rftbsub::(8, &mut a, 2, &c); + assert_eq!(a.map(f32::to_bits), expected(plain)); + } + + /// Output bits of `rdft_with::` and `irdft_with::` on `x`. + fn run_forms(fft: &Fft4g, x: &[f32]) -> [Vec; 2] { + let mut forward = x.to_vec(); + fft.rdft_with::(&mut forward); + let mut inverse = x.to_vec(); + fft.irdft_with::(&mut inverse); + [bits(&forward), bits(&inverse)] + } + + /// `rdft` and `irdft` run the [`NATIVE_FMA`] form. + #[test] + fn rdft_and_irdft_use_native_fma() { + // n = 512 runs cft1st, cftmdl, and rftfsub or rftbsub. + let fft = Fft4g::new(512); + let input: Vec = (0..512) + .map(|i| (i as f32 * 0.618_034).fract() - 0.5) + .collect(); + // The forms differ on this input, so the comparison below can fail. + let [fused, plain] = [ + run_forms::(&fft, &input), + run_forms::(&fft, &input), + ]; + assert!(fused[0] != plain[0] && fused[1] != plain[1]); + + let mut forward = input.clone(); + fft.rdft(&mut forward); + let mut inverse = input.clone(); + fft.irdft(&mut inverse); + assert_eq!( + [bits(&forward), bits(&inverse)], + run_forms::(&fft, &input) + ); + } + + /// The last pass of `cftfsub`, or of `cftbsub` if `backward`, for + /// `n = 4 * l`, copied from them. It has no `mul_add`, so both forms + /// share it. + fn last_radix4_pass(a: &mut [f32], l: usize, backward: bool) { + for j in (0..l).step_by(2) { + let (j1, j2, j3) = (j + l, j + 2 * l, j + 3 * l); + let x0r = a[j] + a[j1]; + let x1r = a[j] - a[j1]; + let x2r = a[j2] + a[j3]; + let x2i = a[j2 + 1] + a[j3 + 1]; + let x3r = a[j2] - a[j3]; + let x3i = a[j2 + 1] - a[j3 + 1]; + if backward { + let x0i = -a[j + 1] - a[j1 + 1]; + let x1i = -a[j + 1] + a[j1 + 1]; + a[j] = x0r + x2r; + a[j + 1] = x0i - x2i; + a[j2] = x0r - x2r; + a[j2 + 1] = x0i + x2i; + a[j1] = x1r - x3i; + a[j1 + 1] = x1i - x3r; + a[j3] = x1r + x3i; + a[j3 + 1] = x1i + x3r; + } else { + let x0i = a[j + 1] + a[j1 + 1]; + let x1i = a[j + 1] - a[j1 + 1]; + a[j] = x0r + x2r; + a[j + 1] = x0i + x2i; + a[j2] = x0r - x2r; + a[j2 + 1] = x0i - x2i; + a[j1] = x1r - x3i; + a[j1 + 1] = x1i + x3r; + a[j3] = x1r + x3i; + a[j3 + 1] = x1i - x3r; + } + } + } + + /// Output bits of `rdft_with` on `x` at n = 512, or of `irdft_with` if + /// `inverse`, built from the kernels without the call chain in between: + /// `cft1st`, both `cftmdl` passes, and `rftfsub` or `rftbsub`, each in + /// the form `forms` gives for it. + fn composed(fft: &Fft4g, x: &[f32], inverse: bool, forms: [bool; 3]) -> Vec { + assert_eq!(fft.n, 512); + let (n, nc, w, c) = (fft.n, fft.nc, &fft.w[..], &fft.w[fft.nw..]); + let rft_sub = |a: &mut [f32]| match (inverse, forms[2]) { + (false, true) => rftfsub::(n, a, nc, c), + (false, false) => rftfsub::(n, a, nc, c), + (true, true) => rftbsub::(n, a, nc, c), + (true, false) => rftbsub::(n, a, nc, c), + }; + let mut a = x.to_vec(); + if inverse { + a[1] = 0.5 * (a[0] - a[1]); + a[0] -= a[1]; + rft_sub(&mut a); + } + apply_bitrv2(&fft.bitrv_ip, fft.bitrv_m, fft.bitrv_long, &mut a); + if forms[0] { + cft1st::(n, &mut a, w); + } else { + cft1st::(n, &mut a, w); + } + for l in [8, 32] { + if forms[1] { + cftmdl::(n, l, &mut a, w); + } else { + cftmdl::(n, l, &mut a, w); + } + } + last_radix4_pass(&mut a, 128, inverse); + if !inverse { + rft_sub(&mut a); + let xi = a[0] - a[1]; + a[0] += a[1]; + a[1] = xi; + } + bits(&a) + } + + /// `rdft_with` and `irdft_with` pass their form to every kernel: on every + /// target, each form matches the kernels composed in that form, bit for + /// bit. + #[test] + fn transforms_pass_the_form_to_every_kernel() { + let fft = Fft4g::new(512); + let input: Vec = (0..512) + .map(|i| (i as f32 * 0.618_034).fract() - 0.5) + .collect(); + for fma in [true, false] { + let expected = [false, true].map(|inverse| composed(&fft, &input, inverse, [fma; 3])); + // Each `::` in the call chain sets the form of one of these: + // `cft1st`; both `cftmdl` passes; all three (the call of `cftfsub` + // or `cftbsub`); `rftfsub` or `rftbsub`. Each of them, alone in + // the other form, changes both transforms, so a wrong `::` + // fails the comparison below. + for flipped in [ + [true, false, false], + [false, true, false], + [true, true, false], + [false, false, true], + ] { + let forms = flipped.map(|flip| flip != fma); + for (inverse, expected) in [false, true].into_iter().zip(&expected) { + assert_ne!( + &composed(&fft, &input, inverse, forms), + expected, + "fma {fma}, inverse {inverse}, forms {forms:?}" + ); + } + } + let actual = if fma { + run_forms::(&fft, &input) + } else { + run_forms::(&fft, &input) + }; + assert_eq!(actual, expected, "fma {fma}"); + } + } + #[test] #[should_panic(expected = "power of 2")] fn rejects_non_power_of_two() { diff --git a/crates/sonora-simd/src/lib.rs b/crates/sonora-simd/src/lib.rs index 74f111e..6781856 100644 --- a/crates/sonora-simd/src/lib.rs +++ b/crates/sonora-simd/src/lib.rs @@ -385,6 +385,24 @@ pub fn detect_backend() -> SimdBackend { SimdBackend::Scalar } +/// Whether the scalar kernels of the Sonora crates use [`f32::mul_add`]. +/// +/// True on AArch64 (`aarch64` and `arm64ec`), and on x86 built with the `fma` +/// target feature: there `mul_add` is one instruction. Every other target +/// uses the plain C++ expressions, in their operation order. Targets without +/// native FMA take this path because `mul_add` would be an `fmaf` library +/// call per operation. Other targets with native FMA, such as riscv64gc, +/// powerpc64 and armv7 with VFPv4, take it by policy, so that they evaluate +/// the expressions as C++ built without floating-point contraction does. +/// +/// This is a property of the compilation target, fixed at build time. Unlike +/// [`detect_backend`], it does not depend on the CPU the code runs on. +pub const NATIVE_FMA: bool = cfg!(any( + target_arch = "aarch64", + target_arch = "arm64ec", + target_feature = "fma" +)); + #[cfg(test)] mod tests { use super::*; diff --git a/crates/sonora/src/three_band_filter_bank.rs b/crates/sonora/src/three_band_filter_bank.rs index 2d86dd4..6171508 100644 --- a/crates/sonora/src/three_band_filter_bank.rs +++ b/crates/sonora/src/three_band_filter_bank.rs @@ -5,6 +5,8 @@ //! //! Ported from `modules/audio_processing/three_band_filter_bank.h/cc`. +use sonora_simd::NATIVE_FMA; + const SQRT_3: f32 = 1.732_050_8; const SPARSITY: usize = 4; @@ -57,8 +59,20 @@ const DCT_MODULATION: [[f32; NUM_BANDS]; NUM_NON_ZERO_FILTERS] = [ /// Polyphase filter core: filters `input` through `filter` with shift `in_shift`, /// using and updating `state`. /// +/// If `FMA`, the unrolled 4-tap sums of Parts 1 and 3 use [`f32::mul_add`]; +/// otherwise they use plain arithmetic, in the C++ operation order. C++ starts +/// that sum from `0.0`; leaving it out changes only the sign of an all-zero +/// sum, which `analysis` and `synthesis` lose when they add the result into +/// zeroed buffers. The public methods pass [`NATIVE_FMA`]; the tests run both +/// forms. +/// /// Direct port of C++ `FilterCore` in `three_band_filter_bank.cc`. -fn filter_core( +// LLVM's inline cost for this body sits at its threshold, so without +// `inline(always)` whether `analysis` and `synthesis` inline it depends on +// the codegen-unit partition (with Cargo's default 16 units it did on +// aarch64 and did not on x86_64). +#[inline(always)] +fn filter_core( filter: &[f32; FILTER_SIZE], input: &[f32; SPLIT_BAND_SIZE], in_shift: usize, @@ -79,13 +93,20 @@ fn filter_core( #[allow(clippy::needless_range_loop, reason = "index used in arithmetic")] for k in 0..in_shift { let j = MEMORY_SIZE + k - in_shift; - output[k] = f0.mul_add( - state[j], - f1.mul_add( - state[j - STRIDE], - f2.mul_add(state[j - 2 * STRIDE], f3 * state[j - 3 * STRIDE]), - ), - ); + output[k] = if FMA { + f0.mul_add( + state[j], + f1.mul_add( + state[j - STRIDE], + f2.mul_add(state[j - 2 * STRIDE], f3 * state[j - 3 * STRIDE]), + ), + ) + } else { + f0 * state[j] + + f1 * state[j - STRIDE] + + f2 * state[j - 2 * STRIDE] + + f3 * state[j - 3 * STRIDE] + }; } // Part 2: transition samples (partially from input, partially from state). @@ -112,13 +133,20 @@ fn filter_core( #[allow(clippy::needless_range_loop, reason = "index used in arithmetic")] for k in (FILTER_SIZE * STRIDE)..SPLIT_BAND_SIZE { let base = k - in_shift; - output[k] = f0.mul_add( - input[base], - f1.mul_add( - input[base - STRIDE], - f2.mul_add(input[base - 2 * STRIDE], f3 * input[base - 3 * STRIDE]), - ), - ); + output[k] = if FMA { + f0.mul_add( + input[base], + f1.mul_add( + input[base - STRIDE], + f2.mul_add(input[base - 2 * STRIDE], f3 * input[base - 3 * STRIDE]), + ), + ) + } else { + f0 * input[base] + + f1 * input[base - STRIDE] + + f2 * input[base - 2 * STRIDE] + + f3 * input[base - 3 * STRIDE] + }; } // Update state from end of input. @@ -151,6 +179,18 @@ impl ThreeBandFilterBank { &mut self, input: &[f32; FULL_BAND_SIZE], output: &mut [[f32; SPLIT_BAND_SIZE]; NUM_BANDS], + ) { + self.analysis_with::(input, output); + } + + /// [`Self::analysis`], with [`filter_core`] in the form `FMA` selects. + /// Always inlined, so that LLVM optimizes this body as part of + /// `analysis`. + #[inline(always)] + fn analysis_with( + &mut self, + input: &[f32; FULL_BAND_SIZE], + output: &mut [[f32; SPLIT_BAND_SIZE]; NUM_BANDS], ) { // Initialize output to zero. for band in output.iter_mut() { @@ -184,7 +224,7 @@ impl ThreeBandFilterBank { // Filter. let mut out_subsampled = [0.0f32; SPLIT_BAND_SIZE]; - filter_core( + filter_core::( filter, &in_subsampled, in_shift, @@ -208,6 +248,18 @@ impl ThreeBandFilterBank { &mut self, input: &[[f32; SPLIT_BAND_SIZE]; NUM_BANDS], output: &mut [f32; FULL_BAND_SIZE], + ) { + self.synthesis_with::(input, output); + } + + /// [`Self::synthesis`], with [`filter_core`] in the form `FMA` selects. + /// Always inlined, so that LLVM optimizes this body as part of + /// `synthesis`. + #[inline(always)] + fn synthesis_with( + &mut self, + input: &[[f32; SPLIT_BAND_SIZE]; NUM_BANDS], + output: &mut [f32; FULL_BAND_SIZE], ) { output.fill(0.0); @@ -240,7 +292,7 @@ impl ThreeBandFilterBank { // Filter. let mut out_subsampled = [0.0f32; SPLIT_BAND_SIZE]; - filter_core( + filter_core::( filter, &in_subsampled, in_shift, @@ -334,6 +386,112 @@ mod tests { ); } + /// Both forms of the unrolled Parts 1 and 3 of `filter_core`, on every + /// target: the fused form must match the `mul_add` chain, and the plain + /// form the C++ tap order, bit for bit. + #[test] + fn filter_core_matches_fused_and_plain_references() { + use std::array::from_fn; + + let filter = &FILTER_COEFFS[1]; + let state: [f32; MEMORY_SIZE] = from_fn(|i| (i as f32 * 0.618_034).fract() - 0.5); + let input: [f32; SPLIT_BAND_SIZE] = from_fn(|i| (i as f32 * 0.414_214).fract() - 0.5); + // Sample n of the stream is history[MEMORY_SIZE + n]; n < 0 is state. + let history: Vec = state.iter().chain(&input).copied().collect(); + + let (mut part1_differs, mut part3_differs) = (false, false); + // in_shift >= 1 so that Part 1 (state only) runs. + for in_shift in 1..STRIDE { + let mut fused_output = [0.0_f32; SPLIT_BAND_SIZE]; + filter_core::( + filter, + &input, + in_shift, + &mut fused_output, + &mut state.clone(), + ); + let mut plain_output = [0.0_f32; SPLIT_BAND_SIZE]; + filter_core::( + filter, + &input, + in_shift, + &mut plain_output, + &mut state.clone(), + ); + + // Part 2 (taps split between state and input) is unchanged. + let parts_1_and_3 = (0..in_shift).chain(FILTER_SIZE * STRIDE..SPLIT_BAND_SIZE); + for k in parts_1_and_3 { + let t: [f32; FILTER_SIZE] = + from_fn(|i| history[MEMORY_SIZE + k - in_shift - i * STRIDE]); + let f = filter; + let fused = f[0].mul_add(t[0], f[1].mul_add(t[1], f[2].mul_add(t[2], f[3] * t[3]))); + // C++ accumulates `out[k] += in[j] * filter[i]` for i = 0..3. + let plain = f[0] * t[0] + f[1] * t[1] + f[2] * t[2] + f[3] * t[3]; + if fused.to_bits() != plain.to_bits() { + if k < in_shift { + part1_differs = true; + } else { + part3_differs = true; + } + } + assert_eq!( + fused_output[k].to_bits(), + fused.to_bits(), + "fused, in_shift {in_shift}, k {k}" + ); + assert_eq!( + plain_output[k].to_bits(), + plain.to_bits(), + "plain, in_shift {in_shift}, k {k}" + ); + } + } + assert!(part1_differs && part3_differs); + } + + fn bits(v: &[f32]) -> Vec { + v.iter().map(|x| x.to_bits()).collect() + } + + /// Output bits of `analysis_with::` on `input` and of + /// `synthesis_with::` on `bands`, each on a new filter bank. + fn run_forms( + input: &[f32; FULL_BAND_SIZE], + bands: &[[f32; SPLIT_BAND_SIZE]; NUM_BANDS], + ) -> [Vec; 2] { + let mut analysis = [[0.0_f32; SPLIT_BAND_SIZE]; NUM_BANDS]; + ThreeBandFilterBank::new().analysis_with::(input, &mut analysis); + let mut synthesis = [0.0_f32; FULL_BAND_SIZE]; + ThreeBandFilterBank::new().synthesis_with::(bands, &mut synthesis); + [bits(analysis.as_flattened()), bits(&synthesis)] + } + + /// `analysis` and `synthesis` run the [`NATIVE_FMA`] form. + #[test] + fn analysis_and_synthesis_use_native_fma() { + use std::array::from_fn; + + let input: [f32; FULL_BAND_SIZE] = from_fn(|i| (i as f32 * 0.618_034).fract() - 0.5); + let bands: [[f32; SPLIT_BAND_SIZE]; NUM_BANDS] = + from_fn(|b| from_fn(|i| ((b * SPLIT_BAND_SIZE + i) as f32 * 0.414_214).fract() - 0.5)); + // The forms differ on these inputs, so the comparison below can fail. + let [fused, plain] = [ + run_forms::(&input, &bands), + run_forms::(&input, &bands), + ]; + assert!(fused[0] != plain[0] && fused[1] != plain[1]); + + let mut analysis = [[0.0_f32; SPLIT_BAND_SIZE]; NUM_BANDS]; + ThreeBandFilterBank::new().analysis(&input, &mut analysis); + let mut synthesis = [0.0_f32; FULL_BAND_SIZE]; + ThreeBandFilterBank::new().synthesis(&bands, &mut synthesis); + assert_eq!( + [bits(analysis.as_flattened()), bits(&synthesis)], + run_forms::(&input, &bands) + ); + } + #[test] fn zero_input_produces_zero_output() { let mut fb = ThreeBandFilterBank::new(); diff --git a/fuzz/Cargo.lock b/fuzz/Cargo.lock index 0eee20c..d587bef 100644 --- a/fuzz/Cargo.lock +++ b/fuzz/Cargo.lock @@ -496,6 +496,9 @@ dependencies = [ [[package]] name = "sonora-fft" version = "0.2.0" +dependencies = [ + "sonora-simd", +] [[package]] name = "sonora-fuzz"