use std::f32::consts::PI;
#[derive(Debug, Clone)]
pub struct MelOpts {
pub sample_rate: u32,
pub n_fft: usize,
pub win_length: usize,
pub hop_length: usize,
pub n_mels: usize,
pub preemph: f32,
pub log_zero_guard: f32,
pub low_freq: f32,
pub high_freq: f32,
}
impl MelOpts {
pub fn canary_default() -> Self {
Self {
sample_rate: 16_000,
n_fft: 512,
win_length: 400,
hop_length: 160,
n_mels: 128,
preemph: 0.97,
log_zero_guard: 2.0_f32.powi(-24),
low_freq: 0.0,
high_freq: 0.0,
}
}
pub fn effective_high_freq(&self) -> f32 {
if self.high_freq > 0.0 {
self.high_freq
} else {
(self.sample_rate as f32) / 2.0
}
}
}
pub struct MelFrontend {
opts: MelOpts,
window: Vec<f32>,
mel_filters: Vec<MelFilter>,
fft: std::sync::Arc<dyn rustfft::Fft<f32> + Send + Sync>,
}
#[derive(Debug, Clone)]
struct MelFilter {
start_bin: usize,
weights: Vec<f32>,
}
impl MelFrontend {
pub fn new(opts: MelOpts) -> Self {
let window = centred_hann_window(opts.win_length, opts.n_fft);
let mel_filters = build_slaney_mel_filters(
opts.sample_rate as f32,
opts.n_fft,
opts.n_mels,
opts.low_freq,
opts.effective_high_freq(),
);
let mut planner = rustfft::FftPlanner::<f32>::new();
let fft = planner.plan_fft_forward(opts.n_fft);
Self {
opts,
window,
mel_filters,
fft,
}
}
pub fn n_mels(&self) -> usize {
self.opts.n_mels
}
pub fn compute(&self, samples: &[f32]) -> (Vec<f32>, usize) {
let n_fft = self.opts.n_fft;
let hop = self.opts.hop_length;
let n_mels = self.opts.n_mels;
if samples.is_empty() {
return (Vec::new(), 0);
}
let coeff = self.opts.preemph;
let pre: Vec<f32> = if coeff == 0.0 {
samples.to_vec()
} else {
let mut out = Vec::with_capacity(samples.len());
out.push(samples[0]); for t in 1..samples.len() {
out.push(samples[t] - coeff * samples[t - 1]);
}
out
};
let pad = n_fft / 2;
let total = pre.len() + 2 * pad;
let mut padded = vec![0.0_f32; total];
padded[pad..pad + pre.len()].copy_from_slice(&pre);
if total < n_fft {
return (Vec::new(), 0);
}
let valid_starts = total - n_fft + 1;
let num_frames = valid_starts.div_ceil(hop);
let mut features = vec![0.0_f32; num_frames * n_mels];
let mut buf = vec![rustfft::num_complex::Complex32::new(0.0, 0.0); n_fft];
let n_bins = n_fft / 2 + 1;
for f in 0..num_frames {
let start = f * hop;
for i in 0..n_fft {
let s = padded[start + i];
buf[i] = rustfft::num_complex::Complex32::new(s * self.window[i], 0.0);
}
self.fft.process(&mut buf);
for (m, filt) in self.mel_filters.iter().enumerate() {
let mut energy = 0.0_f32;
for (k, w) in filt.weights.iter().enumerate() {
let bin = filt.start_bin + k;
if bin < n_bins {
let c = buf[bin];
energy += (c.re * c.re + c.im * c.im) * *w;
}
}
features[f * n_mels + m] = (energy + self.opts.log_zero_guard).ln();
}
}
let valid_frames = samples.len() / hop;
let valid_frames = valid_frames.min(num_frames);
features.truncate(valid_frames * n_mels);
let num_frames = valid_frames;
if num_frames < 2 {
return (features, num_frames);
}
for m in 0..n_mels {
let mut sum = 0.0_f32;
for f in 0..num_frames {
sum += features[f * n_mels + m];
}
let mean = sum / num_frames as f32;
let mut sq = 0.0_f32;
for f in 0..num_frames {
let d = features[f * n_mels + m] - mean;
sq += d * d;
}
let std = (sq / (num_frames - 1) as f32).sqrt();
let denom = std + 1e-5;
for f in 0..num_frames {
features[f * n_mels + m] = (features[f * n_mels + m] - mean) / denom;
}
}
(features, num_frames)
}
}
fn centred_hann_window(win_length: usize, n_fft: usize) -> Vec<f32> {
assert!(win_length <= n_fft);
let pad_total = n_fft - win_length;
let left = pad_total / 2;
let core = hann_window(win_length);
let mut out = vec![0.0_f32; n_fft];
out[left..left + win_length].copy_from_slice(&core);
out
}
fn hann_window(n: usize) -> Vec<f32> {
if n <= 1 {
return vec![1.0; n];
}
let denom = (n - 1) as f32;
(0..n)
.map(|i| 0.5 * (1.0 - (2.0 * PI * i as f32 / denom).cos()))
.collect()
}
fn hz_to_mel_slaney(hz: f32) -> f32 {
let f_sp = 200.0_f32 / 3.0;
let min_log_hz = 1000.0_f32;
let min_log_mel = (min_log_hz - 0.0) / f_sp;
let logstep = (6.4_f32).ln() / 27.0;
if hz < min_log_hz {
hz / f_sp
} else {
min_log_mel + (hz / min_log_hz).ln() / logstep
}
}
fn mel_to_hz_slaney(mel: f32) -> f32 {
let f_sp = 200.0_f32 / 3.0;
let min_log_hz = 1000.0_f32;
let min_log_mel = (min_log_hz - 0.0) / f_sp;
let logstep = (6.4_f32).ln() / 27.0;
if mel < min_log_mel {
mel * f_sp
} else {
min_log_hz * (logstep * (mel - min_log_mel)).exp()
}
}
fn build_slaney_mel_filters(
sample_rate: f32,
n_fft: usize,
n_mels: usize,
low_freq: f32,
high_freq: f32,
) -> Vec<MelFilter> {
let n_bins = n_fft / 2 + 1;
let bin_hz = sample_rate / n_fft as f32;
let mel_lo = hz_to_mel_slaney(low_freq);
let mel_hi = hz_to_mel_slaney(high_freq);
let mel_step = (mel_hi - mel_lo) / (n_mels + 1) as f32;
let mel_points: Vec<f32> = (0..=n_mels + 1)
.map(|i| mel_lo + i as f32 * mel_step)
.collect();
let hz_points: Vec<f32> = mel_points.iter().map(|m| mel_to_hz_slaney(*m)).collect();
let mut filters = Vec::with_capacity(n_mels);
for m in 0..n_mels {
let left_hz = hz_points[m];
let centre_hz = hz_points[m + 1];
let right_hz = hz_points[m + 2];
let enorm = 2.0 / (right_hz - left_hz).max(1e-10);
let mut start_bin: Option<usize> = None;
let mut weights_full: Vec<f32> = Vec::new();
for bin in 0..n_bins {
let hz = bin as f32 * bin_hz;
if hz <= left_hz || hz >= right_hz {
continue;
}
let w_unnorm = if hz <= centre_hz {
(hz - left_hz) / (centre_hz - left_hz).max(1e-10)
} else {
(right_hz - hz) / (right_hz - centre_hz).max(1e-10)
};
if !w_unnorm.is_finite() || w_unnorm <= 0.0 {
continue;
}
let w = w_unnorm * enorm;
if start_bin.is_none() {
start_bin = Some(bin);
}
let s = start_bin.unwrap();
while weights_full.len() < bin - s {
weights_full.push(0.0);
}
weights_full.push(w);
}
filters.push(MelFilter {
start_bin: start_bin.unwrap_or(0),
weights: weights_full,
});
}
filters
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f32, b: f32, tol: f32) -> bool {
(a - b).abs() <= tol
}
#[test]
fn hann_window_endpoints_are_zero() {
let w = hann_window(400);
assert!(approx_eq(w[0], 0.0, 1e-6), "{}", w[0]);
assert!(approx_eq(w[399], 0.0, 1e-6), "{}", w[399]);
}
#[test]
fn hann_window_peak_at_centre() {
let w = hann_window(400);
let mid = 199; assert!(w[mid] > 0.99 && w[mid] <= 1.0, "{}", w[mid]);
}
#[test]
fn hann_window_short_sizes() {
assert_eq!(hann_window(0), Vec::<f32>::new());
assert_eq!(hann_window(1), vec![1.0]);
}
#[test]
fn slaney_mel_known_anchors() {
assert!(approx_eq(hz_to_mel_slaney(0.0), 0.0, 1e-6));
assert!(approx_eq(hz_to_mel_slaney(1000.0), 15.0, 1e-4));
assert!(approx_eq(hz_to_mel_slaney(200.0), 3.0, 1e-4));
for hz in [50.0_f32, 200.0, 500.0, 1000.0, 4000.0, 8000.0] {
let rt = mel_to_hz_slaney(hz_to_mel_slaney(hz));
assert!(approx_eq(rt, hz, hz * 1e-4), "hz={hz} rt={rt}");
}
}
#[test]
fn slaney_mel_filterbank_has_correct_count_and_nonempty() {
let filters = build_slaney_mel_filters(16_000.0, 512, 80, 0.0, 8000.0);
assert_eq!(filters.len(), 80);
for (i, f) in filters.iter().enumerate() {
let any_nonzero = f.weights.iter().any(|&w| w > 0.0);
assert!(any_nonzero, "filter {i} is empty: {f:?}");
}
}
#[test]
fn slaney_mel_normalisation_keeps_filters_finite() {
let filters = build_slaney_mel_filters(16_000.0, 512, 80, 0.0, 8000.0);
for (i, f) in filters.iter().enumerate() {
for (k, &w) in f.weights.iter().enumerate() {
assert!(w.is_finite() && w >= 0.0, "filter {i} bin {k} = {w}");
}
}
}
#[test]
fn frontend_empty_input_returns_empty() {
let fe = MelFrontend::new(MelOpts::canary_default());
let (feats, n_frames) = fe.compute(&[]);
assert!(feats.is_empty());
assert_eq!(n_frames, 0);
}
#[test]
fn frontend_output_shape_matches_n_mels_x_frames() {
let fe = MelFrontend::new(MelOpts::canary_default());
let samples = vec![0.0_f32; 16_000];
let (feats, n_frames) = fe.compute(&samples);
assert_eq!(n_frames, 100);
assert_eq!(feats.len(), n_frames * 128);
}
#[test]
fn frontend_output_is_finite_for_silence() {
let fe = MelFrontend::new(MelOpts::canary_default());
let samples = vec![0.0_f32; 16_000];
let (feats, _n) = fe.compute(&samples);
for (i, x) in feats.iter().enumerate() {
assert!(x.is_finite(), "feature {i} = {x}");
}
}
#[test]
fn frontend_cmvn_zero_means_per_feature() {
let fe = MelFrontend::new(MelOpts::canary_default());
let mut state: u32 = 0xDEADBEEF;
let n = 16_000_usize;
let samples: Vec<f32> = (0..n)
.map(|_| {
state = state.wrapping_mul(1664525).wrapping_add(1013904223);
(state as f32 / u32::MAX as f32) * 2.0 - 1.0
})
.collect();
let (feats, n_frames) = fe.compute(&samples);
let n_mels = fe.n_mels();
for m in 0..n_mels {
let sum: f32 = (0..n_frames).map(|f| feats[f * n_mels + m]).sum();
let mean = sum / n_frames as f32;
assert!(mean.abs() < 1e-3, "mel {m} mean after CMVN = {mean}");
}
}
#[test]
fn frontend_default_opts_match_canary_card() {
let o = MelOpts::canary_default();
assert_eq!(o.sample_rate, 16_000);
assert_eq!(o.n_fft, 512);
assert_eq!(o.win_length, 400);
assert_eq!(o.hop_length, 160);
assert_eq!(o.n_mels, 128);
assert!(approx_eq(o.preemph, 0.97, 1e-9));
assert!(approx_eq(o.log_zero_guard, 2.0_f32.powi(-24), 1e-30));
}
}