use crate::model::audio_encoder::{
HOP_LEN, LOG_MEL_EPS, N_FFT, NORM_VAR_EPS, PREEMPH, SAMPLE_RATE, WINDOW_LEN,
};
use rustfft::FftPlanner;
use rustfft::num_complex::Complex32;
pub const N_FFT_BINS: usize = N_FFT / 2 + 1;
const _: () = assert!(
WINDOW_LEN <= N_FFT,
"WINDOW_LEN must be <= N_FFT (audio_encoder constants)"
);
pub fn build_mel_filterbank(n_mel: usize, n_fft: usize, sample_rate: usize) -> Vec<f32> {
assert!(n_mel > 0, "n_mel must be > 0");
assert!(n_fft > 1, "n_fft must be > 1");
assert!(sample_rate > 0, "sample_rate must be > 0");
let n_fft_bins = n_fft / 2 + 1;
let bin_hz_step = sample_rate as f64 / n_fft as f64;
let fmin = 0.0_f64;
let fmax = 0.5_f64 * sample_rate as f64;
let min_log_hz = 1000.0_f64;
let lin_slope = 3.0 / 200.0;
let min_log_mel = min_log_hz * lin_slope;
let log_step = 6.4_f64.ln() / 27.0;
let hz_to_mel = |f_hz: f64| -> f64 {
if f_hz < min_log_hz {
f_hz * lin_slope
} else {
min_log_mel + (f_hz / min_log_hz).ln() / log_step
}
};
let mel_to_hz = |m: f64| -> f64 {
if m < min_log_mel {
m / lin_slope
} else {
min_log_hz * ((m - min_log_mel) * log_step).exp()
}
};
let m_lo = hz_to_mel(fmin);
let m_hi = hz_to_mel(fmax);
let mut hz_pts = Vec::with_capacity(n_mel + 2);
for i in 0..(n_mel + 2) {
let m = m_lo + (m_hi - m_lo) * (i as f64 / (n_mel + 1) as f64);
hz_pts.push(mel_to_hz(m));
}
let mut filters = vec![0.0f32; n_mel * n_fft_bins];
for m in 0..n_mel {
let f_left = hz_pts[m];
let f_center = hz_pts[m + 1];
let f_right = hz_pts[m + 2];
let denom_l = (f_center - f_left).max(1e-30);
let denom_r = (f_right - f_center).max(1e-30);
let enorm = 2.0 / (f_right - f_left).max(1e-30);
let row = &mut filters[m * n_fft_bins..(m + 1) * n_fft_bins];
for (k, slot) in row.iter_mut().enumerate() {
let f = k as f64 * bin_hz_step;
let w = if f >= f_left && f <= f_center {
(f - f_left) / denom_l
} else if f > f_center && f <= f_right {
(f_right - f) / denom_r
} else {
0.0
};
*slot = (w * enorm) as f32;
}
}
filters
}
pub fn build_hann_window(length: usize) -> Vec<f32> {
let mut w = Vec::with_capacity(length);
let denom = length as f64;
for i in 0..length {
let v = 0.5 * (1.0 - (2.0 * std::f64::consts::PI * i as f64 / denom).cos());
w.push(v as f32);
}
w
}
pub fn build_padded_hann_window() -> Vec<f32> {
let raw = build_hann_window(WINDOW_LEN);
let lo = (N_FFT - WINDOW_LEN) / 2;
let mut padded = vec![0.0f32; N_FFT];
padded[lo..lo + WINDOW_LEN].copy_from_slice(&raw);
padded
}
pub(crate) fn effective_n_len(n_samples: usize, n_frames: usize) -> usize {
(n_samples / HOP_LEN).min(n_frames)
}
pub fn n_frames_for(n_samples: usize) -> usize {
if n_samples == 0 {
return 0;
}
let n_samples_padded = match n_samples.checked_add(2 * (N_FFT / 2)) {
Some(v) => v,
None => return 0,
};
if n_samples_padded < N_FFT {
0
} else {
(n_samples_padded - N_FFT) / HOP_LEN + 1
}
}
pub fn log_mel_spectrogram(pcm: &[f32], n_mel_bins: usize) -> (Vec<f32>, usize) {
if pcm.is_empty() || n_mel_bins == 0 {
return (Vec::new(), 0);
}
let n_samples_in = pcm.len();
let pad_amount = N_FFT / 2;
let n_samples_padded = match n_samples_in.checked_add(2 * pad_amount) {
Some(val) => val,
None => return (Vec::new(), 0),
};
let mut samples = vec![0.0f32; n_samples_padded];
samples[pad_amount..pad_amount + n_samples_in].copy_from_slice(pcm);
let inner_end = n_samples_padded - pad_amount;
let mut prev = samples[pad_amount];
for s in samples[pad_amount + 1..inner_end].iter_mut() {
let cur = *s;
*s = cur - PREEMPH * prev;
prev = cur;
}
let hann = build_padded_hann_window();
let filters = build_mel_filterbank(n_mel_bins, N_FFT, SAMPLE_RATE as usize);
let mut planner = FftPlanner::<f32>::new();
let fft = planner.plan_fft_forward(N_FFT);
let n_frames = n_frames_for(n_samples_in);
if n_frames == 0 {
return (Vec::new(), 0);
}
let mut mel = vec![0.0f32; n_mel_bins * n_frames];
let mut fft_buf: Vec<Complex32> = vec![Complex32::new(0.0, 0.0); N_FFT];
let mut power_spec = vec![0.0f64; N_FFT_BINS];
for ti in 0..n_frames {
let offset = ti * HOP_LEN;
let frame_samples = &samples[offset..offset + N_FFT];
for (fb, (&h, &s)) in fft_buf.iter_mut().zip(hann.iter().zip(frame_samples)) {
*fb = Complex32::new(h * s, 0.0);
}
fft.process(&mut fft_buf);
for (p, c) in power_spec.iter_mut().zip(fft_buf.iter().take(N_FFT_BINS)) {
*p = c.re as f64 * c.re as f64 + c.im as f64 * c.im as f64;
}
for mi in 0..n_mel_bins {
let frow = &filters[mi * N_FFT_BINS..(mi + 1) * N_FFT_BINS];
let sum: f64 = power_spec
.iter()
.zip(frow)
.map(|(&p, &f)| p * f as f64)
.sum();
mel[mi * n_frames + ti] = (sum + LOG_MEL_EPS as f64).ln() as f32;
}
}
let effective_n_len = effective_n_len(n_samples_in, n_frames);
for mi in 0..n_mel_bins {
let row = &mut mel[mi * n_frames..(mi + 1) * n_frames];
if effective_n_len > 1 {
let mut mean_sum = 0.0f64;
for &v in &row[..effective_n_len] {
mean_sum += v as f64;
}
let mean = mean_sum / effective_n_len as f64;
let mut var_sum = 0.0f64;
for &v in &row[..effective_n_len] {
let d = v as f64 - mean;
var_sum += d * d;
}
let var = var_sum / (effective_n_len - 1) as f64; let inv_std = 1.0 / (var + NORM_VAR_EPS).sqrt();
for v in row[..effective_n_len].iter_mut() {
*v = ((*v as f64 - mean) * inv_std) as f32;
}
for v in row[effective_n_len..].iter_mut() {
*v = 0.0;
}
} else {
for v in row.iter_mut() {
*v = 0.0;
}
}
}
let mut mel_time_major = vec![0.0f32; n_frames * n_mel_bins];
for mi in 0..n_mel_bins {
for ti in 0..n_frames {
mel_time_major[ti * n_mel_bins + mi] = mel[mi * n_frames + ti];
}
}
(mel_time_major, n_frames)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hann_window_periodic_endpoints() {
let w = build_hann_window(8);
assert!((w[0] - 0.0).abs() < 1e-6, "w[0] = {}", w[0]);
assert!((w[4] - 1.0).abs() < 1e-6, "w[4] = {}", w[4]);
assert!((w[1] - w[7]).abs() < 1e-6);
assert!((w[2] - w[6]).abs() < 1e-6);
assert!((w[3] - w[5]).abs() < 1e-6);
}
#[test]
fn hann_window_lfm2a_dims() {
let w = build_hann_window(WINDOW_LEN);
assert_eq!(w.len(), 400);
assert!((w[200] - 1.0).abs() < 1e-6);
}
#[test]
fn mel_filterbank_shape_and_positive() {
let n_mel = 32;
let f = build_mel_filterbank(n_mel, N_FFT, SAMPLE_RATE as usize);
assert_eq!(f.len(), n_mel * N_FFT_BINS);
for mi in 0..n_mel {
let row = &f[mi * N_FFT_BINS..(mi + 1) * N_FFT_BINS];
assert!(row.iter().any(|&v| v > 0.0), "mel row {mi} is all zero");
}
}
#[test]
fn mel_filterbank_peak_indices_monotonic() {
let n_mel = 32;
let f = build_mel_filterbank(n_mel, N_FFT, SAMPLE_RATE as usize);
let mut prev_peak = 0;
for mi in 0..n_mel {
let row = &f[mi * N_FFT_BINS..(mi + 1) * N_FFT_BINS];
let (peak_k, _) = row
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap();
assert!(
peak_k >= prev_peak,
"filter {mi} peak at bin {peak_k} < prev {prev_peak}"
);
prev_peak = peak_k;
}
}
#[test]
fn log_mel_spectrogram_sine_wave_smoke() {
let n_mel = 80;
let dur_sec = 1.0;
let n_samples = (SAMPLE_RATE as f32 * dur_sec) as usize;
let freq_hz = 1000.0_f32;
let pcm: Vec<f32> = (0..n_samples)
.map(|i| (2.0 * std::f32::consts::PI * freq_hz * i as f32 / SAMPLE_RATE as f32).sin())
.collect();
let (mel, n_frames) = log_mel_spectrogram(&pcm, n_mel);
assert_eq!(mel.len(), n_frames * n_mel);
assert!(n_frames > 0);
for (i, &v) in mel.iter().enumerate() {
assert!(v.is_finite(), "mel[{i}] = {v} (not finite)");
}
let mut min = f32::INFINITY;
let mut max = f32::NEG_INFINITY;
for &v in &mel {
min = min.min(v);
max = max.max(v);
}
assert!(
max - min > 0.01,
"mel output has no variation (min={min}, max={max})"
);
}
#[test]
fn n_frames_for_matches_log_mel_spectrogram() {
for n in [0usize, 1, 159, 160, 161, 320, 1599, 1600, 16000] {
let pcm = vec![0.1f32; n];
let (_, actual) = log_mel_spectrogram(&pcm, 8);
assert_eq!(
n_frames_for(n),
actual,
"n_frames_for({n}) disagrees with log_mel_spectrogram"
);
}
}
#[test]
fn padded_hann_window_is_centered() {
let w = build_padded_hann_window();
assert_eq!(w.len(), N_FFT);
let lo = (N_FFT - WINDOW_LEN) / 2;
assert!(
(w[N_FFT / 2] - 1.0).abs() < 1e-6,
"peak should land at N_FFT/2, got {}",
w[N_FFT / 2]
);
assert!(w[..lo].iter().all(|&v| v == 0.0), "low flank is not zero");
assert!(
w[lo + WINDOW_LEN..].iter().all(|&v| v == 0.0),
"high flank is not zero"
);
assert!(w[lo + 1..lo + WINDOW_LEN].iter().all(|&v| v > 0.0));
}
#[test]
fn effective_n_len_bounds_the_nonzero_frames() {
let n_mel = 16;
let n = 5000;
let pcm: Vec<f32> = (0..n)
.map(|i| (i as f32 * 0.03).sin() + (i as f32 * 0.11).cos())
.collect();
let (mel, n_frames) = log_mel_spectrogram(&pcm, n_mel);
let eff = 31;
assert_eq!(effective_n_len(n, n_frames), eff);
assert!(eff > 1 && eff < n_frames, "eff {eff} of {n_frames} frames");
assert!(
mel[eff * n_mel..].iter().all(|&v| v == 0.0),
"frames at or past effective_n_len {eff} are not all zero"
);
assert!(
mel[(eff - 1) * n_mel..eff * n_mel]
.iter()
.any(|&v| v != 0.0),
"the last live frame ({}) is entirely zero",
eff - 1
);
}
#[test]
fn log_mel_spectrogram_empty_input_is_empty_output() {
let (mel, n_frames) = log_mel_spectrogram(&[], 80);
assert_eq!(n_frames, 0);
assert!(mel.is_empty());
}
}