Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions crates/sonora-bench/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

14 changes: 10 additions & 4 deletions crates/sonora-bench/tests/cpp_comparison.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
177 changes: 153 additions & 24 deletions crates/sonora-common-audio/src/cascaded_biquad_filter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)`
Expand Down Expand Up @@ -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::<NATIVE_FMA>(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<const FMA: bool>(&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::<FMA>(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);
Expand All @@ -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;
Expand All @@ -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::<NATIVE_FMA>(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<const FMA: bool>(&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];
Expand All @@ -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;
Expand All @@ -129,7 +159,7 @@ impl CascadedBiQuadFilter {
}
}

fn apply_biquad(x: &[f32], y: &mut [f32], bq: &mut BiQuad) {
fn apply_biquad<const FMA: bool>(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];
Expand All @@ -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;
Expand Down Expand Up @@ -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<f32> {
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<f32> {
(0..32)
.map(|i| (i as f32 * 0.618_034).fract() - 0.5)
.collect()
}

fn bits(v: &[f32]) -> Vec<u32> {
v.iter().map(|x| x.to_bits()).collect()
}

/// Output bits of `process_with::<FMA>` and `process_in_place_with::<FMA>`,
/// each on a new filter.
fn run_forms<const FMA: bool>(coeffs: &[BiQuadCoefficients], x: &[f32]) -> [Vec<u32>; 2] {
let mut y = vec![0.0_f32; x.len()];
CascadedBiQuadFilter::new(coeffs).process_with::<FMA>(x, &mut y);
let mut in_place = x.to_vec();
CascadedBiQuadFilter::new(coeffs).process_in_place_with::<FMA>(&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::<true>(&coeffs, &input), [fused.clone(), fused]);
assert_eq!(run_forms::<false>(&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::<true>(&coeffs, &input),
run_forms::<false>(&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::<NATIVE_FMA>(&coeffs, &input)
);
}

#[test]
fn multi_stage_filter() {
let coeffs = [lowpass_coefficients(), lowpass_coefficients()];
Expand Down
3 changes: 3 additions & 0 deletions crates/sonora-fft/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,9 @@ repository.workspace = true
[lints]
workspace = true

[dependencies]
sonora-simd = { workspace = true }

[dev-dependencies]
proptest = { workspace = true }
test-strategy = { workspace = true }
Loading
Loading