#![allow(unsafe_code)]
#![allow(dead_code)]
#[inline]
fn reconcile_mean(simd_sum: f64, samples: &[f32]) -> f64 {
if simd_sum.is_finite() {
simd_sum
} else {
scalar::scalar_mean(samples)
}
}
#[inline]
fn reconcile_var_sum(simd_var_sum: f64, samples: &[f32], mean: f64) -> f64 {
if simd_var_sum.is_finite() {
simd_var_sum
} else {
scalar::scalar_var_sum(samples, mean)
}
}
#[inline]
fn simd_path_loses_precision(mean: f64, inv_std: f64) -> bool {
const SIMD_PRECISION_TOLERANCE: f64 = 5e-5;
const F32_MANTISSA_RELATIVE_ULP: f64 = 1.0_f64 / (1u64 << 23) as f64;
let max_per_sample_err = mean.abs() * F32_MANTISSA_RELATIVE_ULP * inv_std;
max_per_sample_err > SIMD_PRECISION_TOLERANCE
}
const SIMD_SAFE_MAX_ABS: f32 = 1.0e4;
#[inline]
fn samples_within_simd_safe_range(samples: &[f32]) -> bool {
const SIGN_MASK: u32 = 0x7FFF_FFFF;
let threshold_bits = SIMD_SAFE_MAX_ABS.to_bits();
let inf_bits = f32::INFINITY.to_bits(); let max_abs_bits = samples
.iter()
.copied()
.map(|s| s.to_bits() & SIGN_MASK)
.fold(0_u32, u32::max);
max_abs_bits < inf_bits && max_abs_bits <= threshold_bits
}
#[inline]
pub fn zero_mean_unit_var_normalize(samples: &[f32]) -> Vec<f32> {
if !samples_within_simd_safe_range(samples) {
return scalar::zero_mean_unit_var_normalize(samples);
}
cfg_select! {
target_arch = "aarch64" => {
unsafe { neon::zero_mean_unit_var_normalize(samples) }
}
target_arch = "x86_64" => x86_dispatch(samples),
_ => scalar::zero_mean_unit_var_normalize(samples),
}
}
#[cfg(target_arch = "x86_64")]
#[inline]
fn x86_dispatch(samples: &[f32]) -> Vec<f32> {
if std::is_x86_feature_detected!("avx512f") {
return unsafe { x86_avx512::zero_mean_unit_var_normalize(samples) };
}
if std::is_x86_feature_detected!("avx2") {
return unsafe { x86_avx2::zero_mean_unit_var_normalize(samples) };
}
if std::is_x86_feature_detected!("sse4.1") {
return unsafe { x86_sse41::zero_mean_unit_var_normalize(samples) };
}
scalar::zero_mean_unit_var_normalize(samples)
}
pub mod scalar {
pub fn zero_mean_unit_var_normalize(samples: &[f32]) -> Vec<f32> {
if samples.is_empty() {
return Vec::new();
}
let n = samples.len() as f64;
let mean = scalar_mean(samples) / n;
let var = scalar_var_sum(samples, mean) / n;
let inv_std = 1.0_f64 / (var + 1e-7_f64).sqrt();
let mut out = Vec::with_capacity(samples.len());
for &s in samples {
out.push(((s as f64 - mean) * inv_std) as f32);
}
out
}
pub(super) fn scalar_mean(samples: &[f32]) -> f64 {
let mut sum = 0.0_f64;
for &s in samples {
sum += s as f64;
}
sum
}
pub(super) fn scalar_var_sum(samples: &[f32], mean: f64) -> f64 {
let mut var_sum = 0.0_f64;
for &s in samples {
let d = s as f64 - mean;
var_sum += d * d;
}
var_sum
}
}
#[inline]
pub(crate) fn normalize_with_silence_mask(samples: &[f32], speech_mask: &[bool]) -> Vec<f32> {
debug_assert_eq!(samples.len(), speech_mask.len());
if samples.is_empty() {
return Vec::new();
}
let mut sum = 0.0_f64;
let mut count: usize = 0;
for (s, &is_speech) in samples.iter().zip(speech_mask.iter()) {
if is_speech {
sum += *s as f64;
count += 1;
}
}
if count == 0 {
return vec![0.0_f32; samples.len()];
}
let mean = sum / count as f64;
let mut var_sum = 0.0_f64;
for (s, &is_speech) in samples.iter().zip(speech_mask.iter()) {
if is_speech {
let d = *s as f64 - mean;
var_sum += d * d;
}
}
let var = var_sum / count as f64;
let inv_std = 1.0_f64 / (var + 1e-7_f64).sqrt();
let mut out = Vec::with_capacity(samples.len());
for (s, &is_speech) in samples.iter().zip(speech_mask.iter()) {
if is_speech {
out.push(((*s as f64 - mean) * inv_std) as f32);
} else {
out.push(0.0_f32);
}
}
out
}
#[cfg(target_arch = "aarch64")]
#[doc(hidden)]
pub mod neon {
use core::arch::aarch64::*;
#[inline]
#[target_feature(enable = "neon")]
pub unsafe fn zero_mean_unit_var_normalize(samples: &[f32]) -> Vec<f32> {
if samples.is_empty() {
return Vec::new();
}
let n = samples.len();
let nf = n as f64;
let mut sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 4 <= n {
let v = vld1q_f32(samples.as_ptr().add(i));
sum += vaddvq_f32(v) as f64;
i += 4;
}
}
while i < n {
sum += samples[i] as f64;
i += 1;
}
let mean = super::reconcile_mean(sum, samples) / nf;
let mean_f32 = mean as f32;
let mean_v = unsafe { vdupq_n_f32(mean_f32) };
let mut var_sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 4 <= n {
let v = vld1q_f32(samples.as_ptr().add(i));
let d = vsubq_f32(v, mean_v);
let sq = vmulq_f32(d, d);
var_sum += vaddvq_f32(sq) as f64;
i += 4;
}
}
while i < n {
let d = samples[i] as f64 - mean;
var_sum += d * d;
i += 1;
}
let var = super::reconcile_var_sum(var_sum, samples, mean) / nf;
let inv_std = 1.0_f64 / (var + 1e-7_f64).sqrt();
if super::simd_path_loses_precision(mean, inv_std) {
return super::scalar::zero_mean_unit_var_normalize(samples);
}
let inv_std_f32 = inv_std as f32;
let inv_v = unsafe { vdupq_n_f32(inv_std_f32) };
let mut out: Vec<f32> = Vec::with_capacity(n);
let out_ptr = out.as_mut_ptr();
let mut i = 0usize;
unsafe {
while i + 4 <= n {
let v = vld1q_f32(samples.as_ptr().add(i));
let normed = vmulq_f32(vsubq_f32(v, mean_v), inv_v);
vst1q_f32(out_ptr.add(i), normed);
i += 4;
}
while i < n {
let s = *samples.as_ptr().add(i);
*out_ptr.add(i) = (s - mean_f32) * inv_std_f32;
i += 1;
}
out.set_len(n);
}
out
}
}
#[cfg(target_arch = "x86_64")]
#[doc(hidden)]
pub mod x86_sse41 {
use core::arch::x86_64::*;
#[inline]
#[target_feature(enable = "sse4.1")]
unsafe fn hsum_ps(v: __m128) -> f32 {
unsafe {
let shuf = _mm_movehdup_ps(v); let sums = _mm_add_ps(v, shuf); let shuf2 = _mm_movehl_ps(sums, sums); let sums = _mm_add_ss(sums, shuf2); _mm_cvtss_f32(sums)
}
}
#[inline]
#[target_feature(enable = "sse4.1")]
pub unsafe fn zero_mean_unit_var_normalize(samples: &[f32]) -> Vec<f32> {
if samples.is_empty() {
return Vec::new();
}
let n = samples.len();
let nf = n as f64;
let mut sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 4 <= n {
let v = _mm_loadu_ps(samples.as_ptr().add(i));
sum += hsum_ps(v) as f64;
i += 4;
}
}
while i < n {
sum += samples[i] as f64;
i += 1;
}
let mean = super::reconcile_mean(sum, samples) / nf;
let mean_f32 = mean as f32;
let mean_v = unsafe { _mm_set1_ps(mean_f32) };
let mut var_sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 4 <= n {
let v = _mm_loadu_ps(samples.as_ptr().add(i));
let d = _mm_sub_ps(v, mean_v);
let sq = _mm_mul_ps(d, d);
var_sum += hsum_ps(sq) as f64;
i += 4;
}
}
while i < n {
let d = samples[i] as f64 - mean;
var_sum += d * d;
i += 1;
}
let var = super::reconcile_var_sum(var_sum, samples, mean) / nf;
let inv_std = 1.0_f64 / (var + 1e-7_f64).sqrt();
if super::simd_path_loses_precision(mean, inv_std) {
return super::scalar::zero_mean_unit_var_normalize(samples);
}
let inv_std_f32 = inv_std as f32;
let inv_v = unsafe { _mm_set1_ps(inv_std_f32) };
let mut out: Vec<f32> = Vec::with_capacity(n);
let out_ptr = out.as_mut_ptr();
let mut i = 0usize;
unsafe {
while i + 4 <= n {
let v = _mm_loadu_ps(samples.as_ptr().add(i));
let normed = _mm_mul_ps(_mm_sub_ps(v, mean_v), inv_v);
_mm_storeu_ps(out_ptr.add(i), normed);
i += 4;
}
while i < n {
let s = *samples.as_ptr().add(i);
*out_ptr.add(i) = (s - mean_f32) * inv_std_f32;
i += 1;
}
out.set_len(n);
}
out
}
}
#[cfg(target_arch = "x86_64")]
#[doc(hidden)]
pub mod x86_avx2 {
use core::arch::x86_64::*;
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn hsum_ps(v: __m256) -> f32 {
unsafe {
let lo = _mm256_castps256_ps128(v);
let hi = _mm256_extractf128_ps(v, 1);
let s = _mm_add_ps(lo, hi);
let shuf = _mm_movehdup_ps(s);
let sums = _mm_add_ps(s, shuf);
let shuf2 = _mm_movehl_ps(sums, sums);
let sums = _mm_add_ss(sums, shuf2);
_mm_cvtss_f32(sums)
}
}
#[inline]
#[target_feature(enable = "avx2")]
pub unsafe fn zero_mean_unit_var_normalize(samples: &[f32]) -> Vec<f32> {
if samples.is_empty() {
return Vec::new();
}
let n = samples.len();
let nf = n as f64;
let mut sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 8 <= n {
let v = _mm256_loadu_ps(samples.as_ptr().add(i));
sum += hsum_ps(v) as f64;
i += 8;
}
}
while i < n {
sum += samples[i] as f64;
i += 1;
}
let mean = super::reconcile_mean(sum, samples) / nf;
let mean_f32 = mean as f32;
let mean_v = unsafe { _mm256_set1_ps(mean_f32) };
let mut var_sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 8 <= n {
let v = _mm256_loadu_ps(samples.as_ptr().add(i));
let d = _mm256_sub_ps(v, mean_v);
let sq = _mm256_mul_ps(d, d);
var_sum += hsum_ps(sq) as f64;
i += 8;
}
}
while i < n {
let d = samples[i] as f64 - mean;
var_sum += d * d;
i += 1;
}
let var = super::reconcile_var_sum(var_sum, samples, mean) / nf;
let inv_std = 1.0_f64 / (var + 1e-7_f64).sqrt();
if super::simd_path_loses_precision(mean, inv_std) {
return super::scalar::zero_mean_unit_var_normalize(samples);
}
let inv_std_f32 = inv_std as f32;
let inv_v = unsafe { _mm256_set1_ps(inv_std_f32) };
let mut out: Vec<f32> = Vec::with_capacity(n);
let out_ptr = out.as_mut_ptr();
let mut i = 0usize;
unsafe {
while i + 8 <= n {
let v = _mm256_loadu_ps(samples.as_ptr().add(i));
let normed = _mm256_mul_ps(_mm256_sub_ps(v, mean_v), inv_v);
_mm256_storeu_ps(out_ptr.add(i), normed);
i += 8;
}
while i < n {
let s = *samples.as_ptr().add(i);
*out_ptr.add(i) = (s - mean_f32) * inv_std_f32;
i += 1;
}
out.set_len(n);
}
out
}
}
#[cfg(target_arch = "x86_64")]
#[doc(hidden)]
pub mod x86_avx512 {
use core::arch::x86_64::*;
#[inline]
#[target_feature(enable = "avx512f")]
pub unsafe fn zero_mean_unit_var_normalize(samples: &[f32]) -> Vec<f32> {
if samples.is_empty() {
return Vec::new();
}
let n = samples.len();
let nf = n as f64;
let mut sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 16 <= n {
let v = _mm512_loadu_ps(samples.as_ptr().add(i));
sum += _mm512_reduce_add_ps(v) as f64;
i += 16;
}
}
while i < n {
sum += samples[i] as f64;
i += 1;
}
let mean = super::reconcile_mean(sum, samples) / nf;
let mean_f32 = mean as f32;
let mean_v = unsafe { _mm512_set1_ps(mean_f32) };
let mut var_sum = 0.0_f64;
let mut i = 0usize;
unsafe {
while i + 16 <= n {
let v = _mm512_loadu_ps(samples.as_ptr().add(i));
let d = _mm512_sub_ps(v, mean_v);
let sq = _mm512_mul_ps(d, d);
var_sum += _mm512_reduce_add_ps(sq) as f64;
i += 16;
}
}
while i < n {
let d = samples[i] as f64 - mean;
var_sum += d * d;
i += 1;
}
let var = super::reconcile_var_sum(var_sum, samples, mean) / nf;
let inv_std = 1.0_f64 / (var + 1e-7_f64).sqrt();
if super::simd_path_loses_precision(mean, inv_std) {
return super::scalar::zero_mean_unit_var_normalize(samples);
}
let inv_std_f32 = inv_std as f32;
let inv_v = unsafe { _mm512_set1_ps(inv_std_f32) };
let mut out: Vec<f32> = Vec::with_capacity(n);
let out_ptr = out.as_mut_ptr();
let mut i = 0usize;
unsafe {
while i + 16 <= n {
let v = _mm512_loadu_ps(samples.as_ptr().add(i));
let normed = _mm512_mul_ps(_mm512_sub_ps(v, mean_v), inv_v);
_mm512_storeu_ps(out_ptr.add(i), normed);
i += 16;
}
while i < n {
let s = *samples.as_ptr().add(i);
*out_ptr.add(i) = (s - mean_f32) * inv_std_f32;
i += 1;
}
out.set_len(n);
}
out
}
}
#[cfg(test)]
mod tests {
use super::*;
fn synth_input(n: usize) -> Vec<f32> {
let mut samples: Vec<f32> = Vec::with_capacity(n);
for i in 0..n {
let t = i as f32 / 16_000.0;
samples.push(2.7 * (2.0 * core::f32::consts::PI * 220.0 * t).sin() + 0.13);
}
samples
}
fn assert_matches_scalar(simd: &[f32], scalar: &[f32]) {
assert_eq!(simd.len(), scalar.len(), "SIMD length mismatch");
let mut max_abs = 0.0_f32;
for (a, b) in simd.iter().zip(scalar.iter()) {
let d = (a - b).abs();
if d > max_abs {
max_abs = d;
}
}
assert!(
max_abs < 1e-4,
"SIMD deviates from scalar: max abs error = {max_abs}",
);
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_matches_scalar() {
let samples = synth_input(480_001);
let s = scalar::zero_mean_unit_var_normalize(&samples);
let v = unsafe { neon::zero_mean_unit_var_normalize(&samples) };
assert_matches_scalar(&v, &s);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn x86_sse41_matches_scalar() {
if !std::is_x86_feature_detected!("sse4.1") {
return; }
let samples = synth_input(480_001);
let s = scalar::zero_mean_unit_var_normalize(&samples);
let v = unsafe { x86_sse41::zero_mean_unit_var_normalize(&samples) };
assert_matches_scalar(&v, &s);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn x86_avx2_matches_scalar() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let samples = synth_input(480_001);
let s = scalar::zero_mean_unit_var_normalize(&samples);
let v = unsafe { x86_avx2::zero_mean_unit_var_normalize(&samples) };
assert_matches_scalar(&v, &s);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn x86_avx512_matches_scalar() {
if !std::is_x86_feature_detected!("avx512f") {
return;
}
let samples = synth_input(480_001);
let s = scalar::zero_mean_unit_var_normalize(&samples);
let v = unsafe { x86_avx512::zero_mean_unit_var_normalize(&samples) };
assert_matches_scalar(&v, &s);
}
#[test]
fn empty_returns_empty() {
assert!(zero_mean_unit_var_normalize(&[]).is_empty());
}
#[test]
fn constant_signal_normalises_to_zero() {
let xs = vec![3.7_f32; 100];
let out = zero_mean_unit_var_normalize(&xs);
assert!(out.iter().all(|&v| v.abs() < 1e-3));
}
fn high_dynamic_range_input() -> Vec<f32> {
let mut xs = Vec::with_capacity(33);
for _ in 0..16 {
xs.push(1e20_f32);
xs.push(-1e20_f32);
}
xs.push(1e19_f32); xs
}
#[test]
fn dispatched_high_dynamic_range_matches_scalar() {
let xs = high_dynamic_range_input();
let s = scalar::zero_mean_unit_var_normalize(&xs);
let d = zero_mean_unit_var_normalize(&xs);
assert_matches_scalar(&d, &s);
assert!(s.iter().all(|x| x.is_finite()));
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_high_dynamic_range_matches_scalar() {
let xs = high_dynamic_range_input();
let s = scalar::zero_mean_unit_var_normalize(&xs);
let v = unsafe { neon::zero_mean_unit_var_normalize(&xs) };
assert_matches_scalar(&v, &s);
assert!(v.iter().all(|x| x.is_finite()));
}
#[cfg(target_arch = "x86_64")]
#[test]
fn x86_sse41_high_dynamic_range_matches_scalar() {
if !std::is_x86_feature_detected!("sse4.1") {
return;
}
let xs = high_dynamic_range_input();
let s = scalar::zero_mean_unit_var_normalize(&xs);
let v = unsafe { x86_sse41::zero_mean_unit_var_normalize(&xs) };
assert_matches_scalar(&v, &s);
assert!(v.iter().all(|x| x.is_finite()));
}
#[cfg(target_arch = "x86_64")]
#[test]
fn x86_avx2_high_dynamic_range_matches_scalar() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let xs = high_dynamic_range_input();
let s = scalar::zero_mean_unit_var_normalize(&xs);
let v = unsafe { x86_avx2::zero_mean_unit_var_normalize(&xs) };
assert_matches_scalar(&v, &s);
assert!(v.iter().all(|x| x.is_finite()));
}
#[cfg(target_arch = "x86_64")]
#[test]
fn x86_avx512_high_dynamic_range_matches_scalar() {
if !std::is_x86_feature_detected!("avx512f") {
return;
}
let xs = high_dynamic_range_input();
let s = scalar::zero_mean_unit_var_normalize(&xs);
let v = unsafe { x86_avx512::zero_mean_unit_var_normalize(&xs) };
assert_matches_scalar(&v, &s);
assert!(v.iter().all(|x| x.is_finite()));
}
fn finite_high_magnitude_input() -> Vec<f32> {
let base = 1.0e8_f32;
let mut xs = Vec::with_capacity(1024);
for i in 0..1024 {
xs.push(if i % 2 == 0 { base } else { -base });
xs.push(if i % 3 == 0 { base + 0.5 } else { -base + 0.25 });
}
xs
}
#[test]
fn dispatched_finite_high_magnitude_matches_scalar() {
let xs = finite_high_magnitude_input();
let s = scalar::zero_mean_unit_var_normalize(&xs);
let d = zero_mean_unit_var_normalize(&xs);
assert_matches_scalar(&d, &s);
assert!(s.iter().all(|x| x.is_finite()));
}
#[test]
fn precision_guard_recognises_safe_and_unsafe_inputs() {
assert!(samples_within_simd_safe_range(&[
0.5, -0.7, 1e3, -1e3, 9999.0
]));
assert!(!samples_within_simd_safe_range(&[1.0, 1e5]));
assert!(!samples_within_simd_safe_range(&[1.0, f32::NAN]));
assert!(!samples_within_simd_safe_range(&[1.0, f32::INFINITY]));
assert!(samples_within_simd_safe_range(&[]));
}
#[test]
fn silence_mask_normalize_keeps_masked_positions_at_zero() {
let mut samples = vec![0.0_f32; 16];
samples[0] = 0.5;
samples[1] = 0.6;
samples[2] = 0.4;
samples[3] = 0.5;
samples[12] = 0.5;
samples[13] = 0.6;
samples[14] = 0.4;
samples[15] = 0.5;
let mut speech_mask = vec![false; 16];
for slot in &mut speech_mask[0..4] {
*slot = true;
}
for slot in &mut speech_mask[12..16] {
*slot = true;
}
let normed = normalize_with_silence_mask(&samples, &speech_mask);
assert_eq!(normed.len(), samples.len());
for i in 4..12 {
assert_eq!(
normed[i], 0.0_f32,
"masked silence at index {i} must stay exactly 0; got {}",
normed[i],
);
}
let any_nonzero_speech =
normed[..4].iter().any(|&v| v != 0.0) || normed[12..].iter().any(|&v| v != 0.0);
assert!(any_nonzero_speech, "speech samples must not all be zero");
}
#[test]
fn silence_mask_normalize_all_silence_yields_zeros() {
let samples = vec![0.5_f32, 0.6, 0.4, 0.5];
let speech_mask = vec![false; 4];
let normed = normalize_with_silence_mask(&samples, &speech_mask);
assert_eq!(normed, vec![0.0_f32; 4]);
}
#[test]
fn silence_mask_normalize_all_speech_matches_regular_normalize() {
let samples = vec![0.5_f32, 0.6, 0.4, 0.5, -0.3, -0.1, 0.2, 0.8];
let speech_mask = vec![true; samples.len()];
let masked = normalize_with_silence_mask(&samples, &speech_mask);
let regular = scalar::zero_mean_unit_var_normalize(&samples);
assert_matches_scalar(&masked, ®ular);
}
#[test]
fn dispatched_low_variance_near_one_matches_scalar() {
let next_up_one = f32::from_bits(1.0_f32.to_bits() + 1);
let mut xs = Vec::with_capacity(64);
for _ in 0..32 {
xs.push(1.0_f32);
xs.push(next_up_one);
}
let s = scalar::zero_mean_unit_var_normalize(&xs);
let d = zero_mean_unit_var_normalize(&xs);
assert_matches_scalar(&d, &s);
}
#[test]
fn dispatched_low_variance_at_high_magnitude_matches_scalar() {
let mag = 100.0_f32;
let next_up = f32::from_bits(mag.to_bits() + 1);
let mut xs = Vec::with_capacity(64);
for _ in 0..32 {
xs.push(mag);
xs.push(next_up);
}
let s = scalar::zero_mean_unit_var_normalize(&xs);
let d = zero_mean_unit_var_normalize(&xs);
assert_matches_scalar(&d, &s);
}
#[test]
fn simd_path_loses_precision_fires_only_on_low_variance() {
assert!(simd_path_loses_precision(1.0, 3_162.0));
assert!(!simd_path_loses_precision(0.01, 5.0));
assert!(!simd_path_loses_precision(0.0, 1e9));
assert!(simd_path_loses_precision(1e3, 5.0));
}
}