use rustfft::num_complex::Complex;
use super::*;
macro_rules! fixture {
($name:literal) => {
concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/identity/fixtures/mel/",
$name
)
};
}
fn read_npy_f32_shaped(path: &str, expected_shape: &[u64]) -> Vec<f32> {
let bytes = std::fs::read(path).unwrap_or_else(|e| panic!("read {path}: {e}"));
let npy = npyz::NpyFile::new(&bytes[..]).unwrap_or_else(|e| panic!("parse npy {path}: {e}"));
assert_eq!(
npy.shape(),
expected_shape,
"{path}: declared NPY shape {:?} != expected {expected_shape:?}",
npy.shape()
);
let data = npy
.into_vec::<f32>()
.unwrap_or_else(|e| panic!("decode npy {path}: {e}"));
let expected_len = expected_shape.iter().product::<u64>() as usize;
assert_eq!(
data.len(),
expected_len,
"{path}: decoded {} elements, declared shape implies {expected_len}",
data.len()
);
if let Some(i) = data.iter().position(|v| !v.is_finite()) {
panic!("{path}: non-finite element at flat index {i}: {}", data[i]);
}
data
}
fn read_golden_wav(path: &str) -> Vec<f32> {
let mut reader = hound::WavReader::open(path).unwrap_or_else(|e| panic!("open {path}: {e}"));
let spec = reader.spec();
assert_eq!(spec.sample_rate, SAMPLE_RATE_HZ, "{path}: sample rate");
assert_eq!(spec.channels, 1, "{path}: channel count");
assert_eq!(spec.bits_per_sample, 16, "{path}: bit depth");
assert_eq!(
spec.sample_format,
hound::SampleFormat::Int,
"{path}: sample format"
);
let samples: Vec<f32> = reader
.samples::<i16>()
.map(|s| f32::from(s.expect("decode sample")) / 32_768.0)
.collect();
assert_eq!(samples.len(), WINDOW_SAMPLES, "{path}: sample count");
samples
}
fn nan_prop_max(xs: impl IntoIterator<Item = f32>) -> f32 {
xs.into_iter()
.reduce(|a, b| {
if a.is_nan() || b.is_nan() {
f32::NAN
} else {
a.max(b)
}
})
.expect("nan_prop_max over an empty iterator")
}
const GOLDEN_CLIPS: [(&str, &str); 3] = [
(fixture!("tone_220.wav"), fixture!("tone_220_mel.npy")),
(fixture!("clipped.wav"), fixture!("clipped_mel.npy")),
(fixture!("formant.wav"), fixture!("formant_mel.npy")),
];
#[test]
fn frame_grid_covers_the_padded_window_exactly() {
assert_eq!(N_FRAMES, 401);
assert_eq!(N_FRAMES, 1 + WINDOW_SAMPLES / HOP);
assert_eq!(CENTER_PAD, N_FFT / 2);
assert_eq!(
(N_FRAMES - 1) * HOP + N_FFT,
WINDOW_SAMPLES + 2 * CENTER_PAD
);
assert_eq!(N_FREQ, N_FFT / 2 + 1);
}
#[test]
fn window_taps_are_periodic_hamming() {
const WINDOW_MAX_ABS_DIFF: f32 = 1e-6;
let golden = read_npy_f32_shaped(fixture!("window.npy"), &[WIN_LENGTH as u64]);
let taps = MelExtractor::periodic_hamming(WIN_LENGTH);
assert_eq!(taps.len(), WIN_LENGTH);
let max_diff = nan_prop_max(
taps
.iter()
.zip(golden.iter())
.map(|(a, b)| (*a as f32 - b).abs()),
);
eprintln!("[mel] window vs checkpoint buffer max_abs_diff = {max_diff:.3e}");
assert!(
max_diff <= WINDOW_MAX_ABS_DIFF,
"analysis window diverged from the checkpoint's own: {max_diff:.3e} > {WINDOW_MAX_ABS_DIFF:.3e}"
);
assert!(
(taps[0] - 0.08).abs() < 1e-12,
"hamming's first tap is 0.08; hann's is 0 — got {}",
taps[0]
);
assert!(
(taps[WIN_LENGTH / 2] - 1.0).abs() < 1e-12,
"the periodic form peaks at EXACTLY 1.0 at k = n/2; the symmetric one \
reaches only ~0.9999964 — got {}",
taps[WIN_LENGTH / 2]
);
}
#[test]
fn window_is_zero_padded_and_centred_in_the_fft_frame() {
assert_eq!(WINDOW_OFFSET, 56);
let window = MelExtractor::padded_window();
assert_eq!(window.len(), N_FFT);
assert!(
window[..WINDOW_OFFSET].iter().all(|&v| v == 0.0),
"the leading {WINDOW_OFFSET} samples must be exactly zero"
);
assert!(
window[WINDOW_OFFSET + WIN_LENGTH..]
.iter()
.all(|&v| v == 0.0),
"the trailing samples must be exactly zero"
);
assert!(
(window[WINDOW_OFFSET] - 0.08).abs() < 1e-12,
"the window's first tap belongs at index {WINDOW_OFFSET}"
);
assert!(
(window[WINDOW_OFFSET + WIN_LENGTH / 2] - 1.0).abs() < 1e-12,
"the window's peak belongs at index {}",
WINDOW_OFFSET + WIN_LENGTH / 2
);
}
#[test]
fn filterbank_matches_the_checkpoints_own_mel_scale() {
const FILTERBANK_MAX_ABS_DIFF: f32 = 1e-5;
let golden = read_npy_f32_shaped(fixture!("filterbank.npy"), &[N_MELS as u64, N_FREQ as u64]);
let fb = MelExtractor::build_htk_filterbank(SAMPLE_RATE_HZ, N_FFT, N_MELS, F_MIN, F_MAX);
assert_eq!(fb.len(), N_MELS * N_FREQ);
let max_diff = nan_prop_max(
fb.iter()
.zip(golden.iter())
.map(|(a, b)| (*a as f32 - b).abs()),
);
eprintln!("[mel] filterbank vs checkpoint buffer max_abs_diff = {max_diff:.3e}");
assert!(
max_diff <= FILTERBANK_MAX_ABS_DIFF,
"mel filterbank diverged from the checkpoint's own: \
{max_diff:.3e} > {FILTERBANK_MAX_ABS_DIFF:.3e}"
);
let peak = nan_prop_max(fb.iter().map(|v| *v as f32));
assert!(
(peak - 1.0).abs() < 1e-3,
"unnormalized triangles peak at 1.0; got {peak}"
);
}
#[test]
fn pre_emphasis_reflects_its_first_neighbour_and_filters_the_rest() {
let x: [f32; 5] = [1.0, 2.0, 4.0, 8.0, 16.0];
let y = MelExtractor::pre_emphasize(&x);
assert_eq!(y.len(), x.len(), "pre-emphasis is length-preserving");
assert!((y[0] - (1.0 - 0.97 * 2.0)).abs() < 1e-12, "got {}", y[0]);
assert!(
(y[0] - 0.03).abs() > 0.9,
"y[0] must not be the replicate-pad value"
);
for n in 1..x.len() {
let want = f64::from(x[n]) - 0.97 * f64::from(x[n - 1]);
assert!((y[n] - want).abs() < 1e-12, "n = {n}: got {}", y[n]);
}
assert!(
(PRE_EMPHASIS - 0.97).abs() < 1e-12,
"the coefficient is 0.97"
);
}
#[test]
fn center_padding_reflects_without_repeating_the_edge() {
let signal: Vec<f64> = (0..1000).map(|i| i as f64).collect();
let padded = MelExtractor::center_pad(&signal);
assert_eq!(padded.len(), signal.len() + 2 * CENTER_PAD);
assert_eq!(padded[0], signal[CENTER_PAD]);
assert_eq!(padded[CENTER_PAD - 1], signal[1]);
assert_eq!(padded[CENTER_PAD], signal[0]);
let tail = CENTER_PAD + signal.len();
assert_eq!(padded[tail - 1], signal[signal.len() - 1]);
assert_eq!(padded[tail], signal[signal.len() - 2]);
assert_eq!(
padded[padded.len() - 1],
signal[signal.len() - 1 - CENTER_PAD]
);
}
#[test]
fn power_spectrum_is_exact_magnitude_squared() {
let mut input = vec![Complex::new(0.0f64, 0.0); N_FFT];
input[0] = Complex::new(3.0, 4.0);
let mut power = vec![0.0f64; N_FREQ];
MelExtractor::power_spectrum(&input, &mut power);
assert_eq!(power[0], 25.0);
}
#[test]
fn stft_peaks_at_the_expected_bin() {
let mel = MelExtractor::new();
let sr = f64::from(SAMPLE_RATE_HZ);
let frame: Vec<f64> = (0..N_FFT)
.map(|k| (std::f64::consts::TAU * 1000.0 * (k as f64) / sr).sin())
.collect();
let mut power = vec![0.0f64; N_FREQ];
let mut fft_input = vec![Complex::new(0.0f64, 0.0); N_FFT];
let mut fft_scratch = vec![Complex::new(0.0f64, 0.0); mel.fft.get_inplace_scratch_len()];
mel.stft_one_frame_power(&frame, &mut fft_input, &mut fft_scratch, &mut power);
let (peak_bin, _) = power
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap();
assert_eq!(peak_bin, 32, "1 kHz belongs in bin 32");
}
#[test]
fn a_silent_window_produces_a_zero_mel() {
const SILENCE_MAX_ABS: f32 = 1e-12;
let mel = MelExtractor::new();
let mut out = vec![f32::NAN; N_MELS * N_FRAMES];
mel
.extract_into(&vec![0.0f32; WINDOW_SAMPLES], &mut out)
.expect("extract silence");
let worst = nan_prop_max(out.iter().map(|v| v.abs()));
eprintln!("[mel] silence max|mel| = {worst:.3e}");
assert!(
worst <= SILENCE_MAX_ABS,
"a mean-normalized log-mel of silence is zero to summation rounding; \
worst was {worst:.3e}"
);
}
#[test]
fn mean_normalization_is_per_mel_bin_over_time() {
let mel = MelExtractor::new();
let sr = f64::from(SAMPLE_RATE_HZ);
let samples: Vec<f32> = (0..WINDOW_SAMPLES)
.map(|i| {
let t = i as f64 / sr;
let env = 0.5 + 0.4 * (std::f64::consts::TAU * 0.7 * t).sin();
(env * (std::f64::consts::TAU * 700.0 * t).sin()) as f32
})
.collect();
let mut out = vec![0.0f32; N_MELS * N_FRAMES];
mel.extract_into(&samples, &mut out).expect("extract tone");
let mut worst_bin_mean = 0.0f64;
for bin in 0..N_MELS {
let row = &out[bin * N_FRAMES..(bin + 1) * N_FRAMES];
let mean = row.iter().map(|v| f64::from(*v)).sum::<f64>() / (N_FRAMES as f64);
worst_bin_mean = worst_bin_mean.max(mean.abs());
}
assert!(
worst_bin_mean < 1e-4,
"every mel bin's mean over time must be ~0; worst was {worst_bin_mean:.3e}"
);
let span = nan_prop_max(out.iter().copied()) - (-nan_prop_max(out.iter().map(|v| -v)));
assert!(
span > 1.0,
"the test signal must exercise the mel; span {span}"
);
}
#[test]
fn mel_matches_the_committed_goldens() {
const PARITY_MAX_ABS_DIFF: f32 = 1.5e-4;
let mel = MelExtractor::new();
let mut worst = 0.0f32;
for (wav_path, mel_path) in GOLDEN_CLIPS {
let samples = read_golden_wav(wav_path);
let golden = read_npy_f32_shaped(mel_path, &[N_MELS as u64, N_FRAMES as u64]);
let mut out = vec![0.0f32; N_MELS * N_FRAMES];
mel.extract_into(&samples, &mut out).expect("extract_into");
let max_diff = nan_prop_max(out.iter().zip(golden.iter()).map(|(a, b)| (a - b).abs()));
eprintln!("[mel] {mel_path}: max_abs_diff = {max_diff:.6e}");
assert!(
max_diff <= PARITY_MAX_ABS_DIFF,
"front-end parity regressed on {mel_path}: \
max_abs_diff = {max_diff:.3e} > {PARITY_MAX_ABS_DIFF:.3e}"
);
worst = worst.max(max_diff);
}
eprintln!("[mel] worst over the committed corpus = {worst:.6e}");
}
#[test]
fn extract_into_refuses_anything_but_one_exact_window() {
let mel = MelExtractor::new();
let mut out = vec![0.0f32; N_MELS * N_FRAMES];
for len in [0usize, 1, WINDOW_SAMPLES - 1, WINDOW_SAMPLES + 1] {
let err = mel
.extract_into(&vec![0.0f32; len], &mut out)
.expect_err("must refuse a clip that is not one window");
assert!(
matches!(err, Error::WindowLength(w) if w.got() == len && w.expected() == WINDOW_SAMPLES),
"len {len}: got {err:?}"
);
}
assert!(
mel
.extract_into(&vec![0.0f32; WINDOW_SAMPLES], &mut out)
.is_ok()
);
}