use rustfft::num_complex::Complex;
use rustfft::{Fft, FftPlanner};
use std::f32::consts::PI;
use std::sync::Arc;
#[derive(Clone, Debug)]
struct SparseMelBand {
start: usize,
weights: Vec<f32>,
}
pub struct MelSpectrogram {
n_fft: usize,
hop_length: usize,
window: Vec<f32>,
mel_bands: Vec<SparseMelBand>,
fft: Arc<dyn Fft<f32>>,
}
impl Default for MelSpectrogram {
fn default() -> Self {
Self::new()
}
}
impl MelSpectrogram {
pub fn new() -> Self {
let n_fft = super::N_FFT;
let hop_length = super::HOP_LENGTH;
let n_mels = super::N_MELS;
let sample_rate = 16000.0_f32;
let fmin = 0.0_f32;
let fmax = sample_rate / 2.0;
let window: Vec<f32> = (0..n_fft)
.map(|n| 0.5 * (1.0 - (2.0 * PI * n as f32 / (n_fft - 1) as f32).cos()))
.collect();
let mel_filterbank = Self::create_mel_filterbank(n_fft, n_mels, sample_rate, fmin, fmax);
let mel_bands = Self::sparsify_mel_filterbank(&mel_filterbank, n_mels, n_fft / 2 + 1);
let mut planner = FftPlanner::<f32>::new();
let fft = planner.plan_fft_forward(n_fft);
Self {
n_fft,
hop_length,
window,
mel_bands,
fft,
}
}
fn hz_to_mel(hz: f32) -> f32 {
2595.0 * (1.0 + hz / 700.0).log10()
}
fn mel_to_hz(mel: f32) -> f32 {
700.0 * (10.0_f32.powf(mel / 2595.0) - 1.0)
}
fn create_mel_filterbank(
n_fft: usize,
n_mels: usize,
sample_rate: f32,
fmin: f32,
fmax: f32,
) -> Vec<f32> {
let n_freqs = n_fft / 2 + 1;
let mel_min = Self::hz_to_mel(fmin);
let mel_max = Self::hz_to_mel(fmax);
let mel_points: Vec<f32> = (0..=(n_mels + 1))
.map(|i| mel_min + (mel_max - mel_min) * i as f32 / (n_mels + 1) as f32)
.collect();
let hz_points: Vec<f32> = mel_points.iter().map(|&m| Self::mel_to_hz(m)).collect();
let bin_points: Vec<f32> = hz_points
.iter()
.map(|&hz| hz * n_fft as f32 / sample_rate)
.collect();
let mut filterbank = vec![0.0_f32; n_mels * n_freqs];
for m in 0..n_mels {
let f_left = bin_points[m];
let f_center = bin_points[m + 1];
let f_right = bin_points[m + 2];
let row_start = m * n_freqs;
for k in 0..n_freqs {
let freq = k as f32;
let val = if freq >= f_left && freq <= f_center && f_center > f_left {
(freq - f_left) / (f_center - f_left)
} else if freq > f_center && freq <= f_right && f_right > f_center {
(f_right - freq) / (f_right - f_center)
} else {
0.0
};
filterbank[row_start + k] = val;
}
}
filterbank
}
fn sparsify_mel_filterbank(
filterbank: &[f32],
n_mels: usize,
n_freqs: usize,
) -> Vec<SparseMelBand> {
debug_assert_eq!(filterbank.len(), n_mels * n_freqs);
let mut bands = Vec::with_capacity(n_mels);
for m in 0..n_mels {
let row = &filterbank[m * n_freqs..(m + 1) * n_freqs];
let first = row.iter().position(|&w| w != 0.0);
let last = row.iter().rposition(|&w| w != 0.0);
match (first, last) {
(Some(start), Some(end)) => bands.push(SparseMelBand {
start,
weights: row[start..=end].to_vec(),
}),
_ => bands.push(SparseMelBand {
start: 0,
weights: vec![0.0],
}),
}
}
bands
}
pub fn compute(&self, samples: &[f32]) -> (Vec<f32>, usize) {
let n_freqs = self.n_fft / 2 + 1;
let mut fft_input = vec![Complex::new(0.0_f32, 0.0); self.n_fft];
let mut power = vec![0.0_f32; n_freqs];
let mut output = Vec::new();
let num_frames =
self.compute_with_buffers(samples, &mut fft_input, &mut power, &mut output);
(output, num_frames)
}
pub fn compute_with_buffers(
&self,
samples: &[f32],
fft_input: &mut Vec<Complex<f32>>,
power: &mut Vec<f32>,
output: &mut Vec<f32>,
) -> usize {
let n_freqs = self.n_fft / 2 + 1;
let n_mels = self.mel_bands.len();
if samples.len() < self.n_fft {
output.resize(n_mels, 0.0);
return 1;
}
let num_frames = (samples.len() - self.n_fft) / self.hop_length + 1;
output.resize(n_mels * num_frames, 0.0_f32);
if fft_input.len() < self.n_fft {
fft_input.resize(self.n_fft, Complex::new(0.0_f32, 0.0));
}
if power.len() < n_freqs {
power.resize(n_freqs, 0.0_f32);
}
for frame_idx in 0..num_frames {
let start = frame_idx * self.hop_length;
for i in 0..self.n_fft {
let sample = if start + i < samples.len() {
samples[start + i]
} else {
0.0
};
fft_input[i] = Complex::new(sample * self.window[i], 0.0);
}
self.fft.process(&mut fft_input[..self.n_fft]);
for k in 0..n_freqs {
power[k] = fft_input[k].norm_sqr();
}
for (m, band) in self.mel_bands.iter().enumerate() {
let mut mel_energy: f32 = 0.0;
let end = band.start + band.weights.len();
debug_assert!(end <= n_freqs);
for (i, &w) in band.weights.iter().enumerate() {
mel_energy += w * power[band.start + i];
}
output[m * num_frames + frame_idx] = (mel_energy.max(1e-10)).ln();
}
}
num_frames
}
}
#[cfg(all(test, not(miri)))]
mod tests {
use super::*;
#[test]
fn test_default_delegates_to_new() {
let silence = vec![0.0_f32; 3200];
let (a, fa) = MelSpectrogram::default().compute(&silence);
let (b, fb) = MelSpectrogram::new().compute(&silence);
assert_eq!(fa, fb);
assert_eq!(a, b);
}
#[test]
fn test_silence() {
let mel = MelSpectrogram::new();
let silence = vec![0.0_f32; 16000]; let (features, num_frames) = mel.compute(&silence);
assert!(num_frames > 0);
assert_eq!(features.len(), 64 * num_frames);
let floor = (1e-10_f32).ln();
for &v in &features {
assert!((v - floor).abs() < 0.01, "Expected ~{floor}, got {v}");
}
}
#[test]
fn test_output_shape() {
let mel = MelSpectrogram::new();
let samples = vec![0.0_f32; 3200]; let (features, num_frames) = mel.compute(&samples);
assert_eq!(num_frames, 19);
assert_eq!(features.len(), 64 * 19);
}
#[test]
fn test_too_short() {
let mel = MelSpectrogram::new();
let samples = vec![0.0_f32; 100]; let (features, num_frames) = mel.compute(&samples);
assert_eq!(num_frames, 1);
assert_eq!(features.len(), 64);
}
#[test]
fn test_sine_wave() {
let mel = MelSpectrogram::new();
let samples: Vec<f32> = (0..16000)
.map(|i| (2.0 * std::f32::consts::PI * 440.0 * i as f32 / 16000.0).sin())
.collect();
let (features, num_frames) = mel.compute(&samples);
assert!(num_frames > 0);
let floor = (1e-10_f32).ln();
let non_floor = features
.iter()
.filter(|&&v| (v - floor).abs() > 1.0)
.count();
assert!(
non_floor > 0,
"Expected some non-floor values for sine wave"
);
}
#[test]
fn test_sparsify_mel_filterbank_matches_dense_dot() {
let n_fft = crate::inference::N_FFT;
let n_mels = crate::inference::N_MELS;
let n_freqs = n_fft / 2 + 1;
let dense = MelSpectrogram::create_mel_filterbank(n_fft, n_mels, 16000.0, 0.0, 8000.0);
let bands = MelSpectrogram::sparsify_mel_filterbank(&dense, n_mels, n_freqs);
let power: Vec<f32> = (0..n_freqs)
.map(|k| ((k as f32) * 0.01).sin().abs() + 0.1)
.collect();
for m in 0..n_mels {
let mut dense_e = 0.0_f32;
let row = &dense[m * n_freqs..(m + 1) * n_freqs];
for (k, &p) in power.iter().enumerate() {
dense_e += row[k] * p;
}
let band = &bands[m];
let mut sparse_e = 0.0_f32;
for (i, &w) in band.weights.iter().enumerate() {
sparse_e += w * power[band.start + i];
}
assert!(
(dense_e - sparse_e).abs() <= 1e-5 * dense_e.max(1.0),
"band {m}: dense={dense_e} sparse={sparse_e}"
);
}
let mean_nz = bands.iter().map(|b| b.weights.len()).sum::<usize>() as f32 / n_mels as f32;
assert!(
mean_nz < (n_freqs as f32) * 0.25,
"expected sparse bands, mean nz={mean_nz} of {n_freqs}"
);
}
#[test]
fn test_sparse_compute_matches_dense_reference() {
let n_fft = crate::inference::N_FFT;
let n_mels = crate::inference::N_MELS;
let hop = crate::inference::HOP_LENGTH;
let n_freqs = n_fft / 2 + 1;
let dense = MelSpectrogram::create_mel_filterbank(n_fft, n_mels, 16000.0, 0.0, 8000.0);
let mel = MelSpectrogram::new();
let samples: Vec<f32> = (0..3200)
.map(|i| (2.0 * std::f32::consts::PI * 440.0 * i as f32 / 16000.0).sin())
.collect();
let (sparse_feats, n_frames) = mel.compute(&samples);
let mut planner = FftPlanner::<f32>::new();
let fft = planner.plan_fft_forward(n_fft);
let window: Vec<f32> = (0..n_fft)
.map(|n| {
0.5 * (1.0 - (2.0 * std::f32::consts::PI * n as f32 / (n_fft - 1) as f32).cos())
})
.collect();
let mut dense_feats = vec![0.0_f32; n_mels * n_frames];
let mut fft_input = vec![Complex::new(0.0_f32, 0.0); n_fft];
let mut power = vec![0.0_f32; n_freqs];
for frame_idx in 0..n_frames {
let start = frame_idx * hop;
for i in 0..n_fft {
let sample = samples.get(start + i).copied().unwrap_or(0.0);
fft_input[i] = Complex::new(sample * window[i], 0.0);
}
fft.process(&mut fft_input[..n_fft]);
for k in 0..n_freqs {
power[k] = fft_input[k].norm_sqr();
}
for m in 0..n_mels {
let mut e = 0.0_f32;
let row = &dense[m * n_freqs..(m + 1) * n_freqs];
for (k, &p) in power.iter().enumerate() {
e += row[k] * p;
}
dense_feats[m * n_frames + frame_idx] = e.max(1e-10).ln();
}
}
assert_eq!(sparse_feats.len(), dense_feats.len());
for (i, (s, d)) in sparse_feats.iter().zip(dense_feats.iter()).enumerate() {
let tol = 1e-4_f32 * d.abs().max(1.0);
assert!(
(s - d).abs() <= tol,
"feat[{i}]: sparse={s} dense={d} tol={tol}"
);
}
}
}