use rustfft::num_complex::Complex;
use super::*;
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_sample_wav() -> Vec<f32> {
let mut reader =
hound::WavReader::open("tests/clap/fixtures/mel/sample.wav").expect("open sample.wav");
let spec = reader.spec();
assert_eq!(spec.sample_rate, 48_000, "sample.wav must be 48 kHz");
assert_eq!(spec.channels, 1, "sample.wav must be mono");
match spec.sample_format {
hound::SampleFormat::Int => {
let scale = 1.0 / (1_i64 << (spec.bits_per_sample - 1)) as f32;
reader
.samples::<i32>()
.map(|s| s.unwrap() as f32 * scale)
.collect()
}
hound::SampleFormat::Float => reader.samples::<f32>().map(|s| s.unwrap()).collect(),
}
}
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")
}
fn nan_prop_min(xs: impl IntoIterator<Item = f32>) -> f32 {
xs.into_iter()
.reduce(|a, b| {
if a.is_nan() || b.is_nan() {
f32::NAN
} else {
a.min(b)
}
})
.expect("nan_prop_min over an empty iterator")
}
#[test]
fn target_samples_is_ten_seconds_at_sample_rate() {
assert_eq!(TARGET_SAMPLES, 10 * SAMPLE_RATE_HZ as usize);
}
#[test]
fn hann_window_periodic_length_1024() {
let win = MelExtractor::periodic_hann(1024);
assert_eq!(win.len(), 1024);
assert_eq!(win[0], 0.0);
assert!(
win[1023] > 0.0 && win[1023] < 1e-3,
"periodic Hann last sample should be positive but small; got {}",
win[1023]
);
for &v in &win {
assert!((0.0..=1.0 + 1e-7).contains(&v));
}
let max_idx = win
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap()
.0;
assert_eq!(max_idx, 512, "peak must be exactly at index N/2 = 512");
assert_eq!(win[512], 1.0);
assert!(win[513] < 1.0 && win[513] > 0.999);
}
#[test]
fn stft_peaks_at_expected_bin() {
let mel = MelExtractor::new();
let sr = 48_000_f64;
let freq = 1000.0_f64;
let frame: Vec<f64> = (0..N_FFT)
.map(|k| (2.0 * std::f64::consts::PI * freq * (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!(
peak_bin == 21 || peak_bin == 22,
"expected peak at bin 21 or 22, got {peak_bin}"
);
}
#[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 power_to_db_applied_once() {
let mel = MelExtractor::new();
let sr = 48_000_f32;
let samples: Vec<f32> = (0..TARGET_SAMPLES)
.map(|k| (2.0 * std::f32::consts::PI * 1000.0 * (k as f32) / sr).sin())
.collect();
let mut out = vec![0.0f32; N_MELS * T_FRAMES];
mel.extract_into(&samples, &mut out).unwrap();
let max = nan_prop_max(out.iter().copied());
let min = nan_prop_min(out.iter().copied());
assert!(
max > 20.0 && max < 50.0,
"unit-sine mel should peak near 29.3 dB; got max = {max}"
);
assert!(
(-100.0 - 1e-3..-50.0).contains(&min),
"amin floor should clip silent bins to -100 dB; got min = {min}"
);
}
#[test]
fn extract_into_rejects_empty_input() {
let mel = MelExtractor::new();
let mut out = vec![0.0f32; N_MELS * T_FRAMES];
let err = mel.extract_into(&[], &mut out).unwrap_err();
assert!(matches!(err, Error::EmptyAudio), "got {err:?}");
}
#[test]
fn short_clip_is_repeat_padded() {
let mel = MelExtractor::new();
let samples = vec![0.25f32; 48_000];
let mut out = vec![0.0f32; N_MELS * T_FRAMES];
mel.extract_into(&samples, &mut out).unwrap();
let mid_a = &out[500 * N_MELS..501 * N_MELS];
let mid_b = &out[600 * N_MELS..601 * N_MELS];
let max_diff = nan_prop_max(mid_a.iter().zip(mid_b.iter()).map(|(a, b)| (a - b).abs()));
assert!(
max_diff < 1e-3,
"repeat-padded constant clip should give stable interior rows; diff = {max_diff}"
);
}
#[test]
fn filterbank_rows_match_librosa() {
let fb = MelExtractor::build_mel_filterbank(48_000, 1024, 64, 50.0, 14_000.0);
for &row_idx in &[0usize, 10, 32] {
let expected = read_npy_f32_shaped(
&format!("tests/clap/fixtures/mel/filterbank_row_{row_idx}.npy"),
&[N_FREQ as u64],
);
let actual = &fb[row_idx * N_FREQ..(row_idx + 1) * N_FREQ];
let max_diff = nan_prop_max(
actual
.iter()
.zip(expected.iter())
.map(|(a, b)| (*a as f32 - b).abs()),
);
assert!(
max_diff < 1e-6,
"filterbank row {row_idx} max_abs_diff = {max_diff:.3e}"
);
}
}
#[test]
fn extract_into_matches_textclap_golden_mel() {
const PARITY_MAX_ABS_DIFF: f32 = 1e-5;
let golden = read_npy_f32_shaped(
"tests/clap/fixtures/mel/golden_mel.npy",
&[T_FRAMES as u64, N_MELS as u64],
);
let samples = read_sample_wav();
let mel = MelExtractor::new();
let mut out = vec![0.0f32; N_MELS * T_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] clapkit-vs-textclap-golden max_abs_diff = {max_diff:.6e}");
assert!(
max_diff <= PARITY_MAX_ABS_DIFF,
"mel parity vs textclap golden regressed: max_abs_diff = {max_diff:.3e} > {PARITY_MAX_ABS_DIFF:.3e}"
);
}