use alloc::vec;
use alloc::vec::Vec;
use libm::{expf, log2f, logf, powf};
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
pub enum MelScale {
Htk,
Slaney,
}
const SLANEY_F_SP: f32 = 200.0 / 3.0;
const SLANEY_MIN_LOG_HZ: f32 = 1000.0;
const SLANEY_LOGSTEP: f32 = 0.068_751_97_f32;
const SLANEY_MIN_LOG_MEL: f32 = SLANEY_MIN_LOG_HZ / SLANEY_F_SP;
impl MelScale {
#[inline]
fn hz_to_mel(self, hz: f32) -> f32 {
match self {
MelScale::Htk => 2595.0 * core::f32::consts::LOG10_2 * log2f(1.0 + hz / 700.0),
MelScale::Slaney => {
if hz < SLANEY_MIN_LOG_HZ {
hz / SLANEY_F_SP
} else {
SLANEY_MIN_LOG_MEL + logf(hz / SLANEY_MIN_LOG_HZ) / SLANEY_LOGSTEP
}
}
}
}
#[inline]
fn mel_to_hz(self, mel: f32) -> f32 {
match self {
MelScale::Htk => 700.0 * (powf(10.0, mel / 2595.0) - 1.0),
MelScale::Slaney => {
if mel < SLANEY_MIN_LOG_MEL {
SLANEY_F_SP * mel
} else {
SLANEY_MIN_LOG_HZ * expf(SLANEY_LOGSTEP * (mel - SLANEY_MIN_LOG_MEL))
}
}
}
}
}
#[derive(Clone, Debug)]
struct MelBand {
start_bin: usize,
weights: Vec<f32>,
}
#[derive(Clone, Debug)]
pub struct MelFilterBank {
pub n_mels: usize,
pub n_fft: usize,
pub sr: u32,
pub fmin: f32,
pub fmax: f32,
pub scale: MelScale,
sparse: Vec<MelBand>,
}
impl MelFilterBank {
#[must_use]
pub fn new(
n_mels: usize,
n_fft: usize,
sr: u32,
fmin: f32,
fmax: f32,
scale: MelScale,
) -> Self {
Self::try_new(n_mels, n_fft, sr, fmin, fmax, scale).expect("invalid MelFilterBank config")
}
pub fn try_new(
n_mels: usize,
n_fft: usize,
sr: u32,
fmin: f32,
fmax: f32,
scale: MelScale,
) -> crate::Result<Self> {
if n_mels == 0 {
return Err(crate::AfpError::Config("n_mels must be > 0".into()));
}
if n_fft < 2 || !n_fft.is_multiple_of(2) {
return Err(crate::AfpError::Config(
"n_fft must be even and >= 2".into(),
));
}
if fmin < 0.0 || fmin.is_nan() {
return Err(crate::AfpError::Config("fmin must be >= 0".into()));
}
if fmin >= fmax || fmax.is_nan() {
return Err(crate::AfpError::Config(
"fmin must be strictly less than fmax".into(),
));
}
let n_bins = n_fft / 2 + 1;
let mel_min = scale.hz_to_mel(fmin);
let mel_max = scale.hz_to_mel(fmax);
let n_points = n_mels + 2;
let mut hz_points = Vec::with_capacity(n_points);
for k in 0..n_points {
let mel = mel_min + (mel_max - mel_min) * k as f32 / (n_points - 1) as f32;
hz_points.push(scale.mel_to_hz(mel));
}
let bin_hz = sr as f32 / n_fft as f32;
let mut sparse = Vec::with_capacity(n_mels);
for k in 0..n_mels {
let left = hz_points[k];
let centre = hz_points[k + 1];
let right = hz_points[k + 2];
let norm = 2.0 / (right - left).max(1e-10);
let first_bin = ((left / bin_hz).floor() as usize + 1).min(n_bins);
let last_bin_raw = (right / bin_hz).ceil() as usize;
let last_bin = if last_bin_raw == 0 {
0
} else {
(last_bin_raw - 1).min(n_bins - 1)
};
if first_bin <= last_bin && first_bin < n_bins {
let mut weights = Vec::with_capacity(last_bin - first_bin + 1);
for b in first_bin..=last_bin {
let f = b as f32 * bin_hz;
let w = if f <= left || f >= right {
0.0
} else if f <= centre {
norm * (f - left) / (centre - left).max(1e-10)
} else {
norm * (right - f) / (right - centre).max(1e-10)
};
weights.push(w);
}
let first_nz = weights.iter().position(|&w| w != 0.0);
let last_nz = weights.iter().rposition(|&w| w != 0.0);
match (first_nz, last_nz) {
(Some(f), Some(l)) => {
weights.copy_within(f..=l, 0);
weights.truncate(l - f + 1);
sparse.push(MelBand {
start_bin: first_bin + f,
weights,
});
}
_ => {
sparse.push(MelBand {
start_bin: 0,
weights: Vec::new(),
});
}
}
} else {
sparse.push(MelBand {
start_bin: 0,
weights: Vec::new(),
});
}
}
Ok(Self {
n_mels,
n_fft,
sr,
fmin,
fmax,
scale,
sparse,
})
}
#[must_use]
pub const fn n_bins(&self) -> usize {
self.n_fft / 2 + 1
}
#[must_use]
pub fn matrix(&self) -> Vec<f32> {
let n_bins = self.n_bins();
let mut mat = vec![0.0_f32; self.n_mels * n_bins];
for (k, band) in self.sparse.iter().enumerate() {
for (j, &w) in band.weights.iter().enumerate() {
mat[k * n_bins + band.start_bin + j] = w;
}
}
mat
}
pub fn log_mel(&self, magnitude: &[f32], out: &mut [f32]) {
assert_eq!(
magnitude.len(),
self.n_bins(),
"magnitude length must equal n_bins"
);
assert_eq!(out.len(), self.n_mels, "out length must equal n_mels");
for (k, slot) in out.iter_mut().enumerate() {
let band = &self.sparse[k];
let acc = super::dot_sq_wide(
&band.weights,
&magnitude[band.start_bin..band.start_bin + band.weights.len()],
);
*slot = core::f32::consts::LOG10_2 * log2f(acc + 1e-10);
}
}
pub fn log_mel_from_power(&self, power: &[f32], out: &mut [f32]) {
assert_eq!(power.len(), self.n_bins(), "power length must equal n_bins");
assert_eq!(out.len(), self.n_mels, "out length must equal n_mels");
for (k, slot) in out.iter_mut().enumerate() {
let band = &self.sparse[k];
let acc = super::dot_wide(
&band.weights,
&power[band.start_bin..band.start_bin + band.weights.len()],
);
*slot = core::f32::consts::LOG10_2 * log2f(acc + 1e-10);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn htk_round_trip() {
for &hz in &[0.0_f32, 100.0, 440.0, 1_000.0, 5_000.0, 11_025.0] {
let m = MelScale::Htk.hz_to_mel(hz);
assert_relative_eq!(MelScale::Htk.mel_to_hz(m), hz, max_relative = 1e-5);
}
}
#[test]
fn slaney_round_trip() {
for &hz in &[
0.0_f32, 100.0, 440.0, 999.0, 1_000.0, 1_001.0, 5_000.0, 11_025.0,
] {
let m = MelScale::Slaney.hz_to_mel(hz);
assert_relative_eq!(MelScale::Slaney.mel_to_hz(m), hz, max_relative = 1e-4);
}
}
#[test]
fn matrix_dimensions() {
let fb = MelFilterBank::new(64, 1024, 16_000, 0.0, 8_000.0, MelScale::Htk);
assert_eq!(fb.n_bins(), 513);
assert_eq!(fb.matrix().len(), 64 * 513);
}
#[test]
fn each_filter_has_a_peak_in_band() {
let fb = MelFilterBank::new(40, 2048, 22_050, 0.0, 11_025.0, MelScale::Slaney);
let n_bins = fb.n_bins();
let mat = fb.matrix();
for k in 0..fb.n_mels {
let row = &mat[k * n_bins..(k + 1) * n_bins];
let max = row.iter().cloned().fold(0.0_f32, f32::max);
assert!(max > 0.0, "filter {k} is all-zero");
}
}
#[test]
fn log_mel_floor_at_silence() {
let fb = MelFilterBank::new(16, 512, 16_000, 0.0, 8_000.0, MelScale::Htk);
let zeros = vec![0.0_f32; fb.n_bins()];
let mut out = vec![0.0_f32; fb.n_mels];
fb.log_mel(&zeros, &mut out);
for v in out {
assert_relative_eq!(v, -10.0, max_relative = 1e-5);
}
}
#[test]
fn htk_and_slaney_diverge_above_1khz() {
let lo = 500.0_f32;
let hi = 4_000.0_f32;
let m_htk_lo = MelScale::Htk.hz_to_mel(lo);
let m_sla_lo = MelScale::Slaney.hz_to_mel(lo);
let m_htk_hi = MelScale::Htk.hz_to_mel(hi);
let m_sla_hi = MelScale::Slaney.hz_to_mel(hi);
let diff_lo = (m_htk_lo - m_sla_lo).abs();
let diff_hi = (m_htk_hi - m_sla_hi).abs();
assert!(
diff_hi > diff_lo,
"expected divergence to grow above 1 kHz: lo={diff_lo} hi={diff_hi}",
);
}
#[test]
fn matrix_rows_are_non_negative() {
let fb = MelFilterBank::new(64, 2048, 22_050, 0.0, 11_025.0, MelScale::Slaney);
for w in fb.matrix() {
assert!(w >= 0.0, "negative weight in mel matrix: {w}");
}
}
#[test]
fn log_mel_from_power_matches_log_mel_on_squared_input() {
let fb = MelFilterBank::new(32, 1024, 16_000, 0.0, 8_000.0, MelScale::Slaney);
let n_bins = fb.n_bins();
let mag: Vec<f32> = (0..n_bins)
.map(|b| ((b as f32 * 0.073).sin().abs() + 0.001) * (1 + b % 7) as f32)
.collect();
let pow: Vec<f32> = mag.iter().map(|m| m * m).collect();
let mut out_mag = vec![0.0_f32; fb.n_mels];
let mut out_pow = vec![0.0_f32; fb.n_mels];
fb.log_mel(&mag, &mut out_mag);
fb.log_mel_from_power(&pow, &mut out_pow);
for (a, b) in out_mag.iter().zip(out_pow.iter()) {
assert_relative_eq!(*a, *b, max_relative = 1e-6);
}
}
#[test]
fn log_mel_picks_up_dirac_in_band() {
let fb = MelFilterBank::new(40, 2048, 22_050, 0.0, 11_025.0, MelScale::Slaney);
let mut mag = vec![0.0_f32; fb.n_bins()];
mag[200] = 1.0;
let mut out = vec![0.0_f32; fb.n_mels];
fb.log_mel(&mag, &mut out);
let max = out.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
assert!(max > -9.0, "no band responded: max={max}");
}
#[test]
#[should_panic(expected = "n_mels must be > 0")]
fn mel_filter_bank_panics_on_zero_n_mels() {
let _ = MelFilterBank::new(0, 1024, 16_000, 0.0, 8_000.0, MelScale::Slaney);
}
#[test]
#[should_panic(expected = "n_fft must be even and >= 2")]
fn mel_filter_bank_panics_on_odd_n_fft() {
let _ = MelFilterBank::new(64, 1023, 16_000, 0.0, 8_000.0, MelScale::Slaney);
}
#[test]
#[should_panic(expected = "n_fft must be even and >= 2")]
fn mel_filter_bank_panics_on_n_fft_below_two() {
let _ = MelFilterBank::new(64, 1, 16_000, 0.0, 8_000.0, MelScale::Slaney);
}
#[test]
#[should_panic(expected = "fmin must be strictly less than fmax")]
fn mel_filter_bank_panics_when_fmin_equals_fmax() {
let _ = MelFilterBank::new(64, 1024, 16_000, 1_000.0, 1_000.0, MelScale::Slaney);
}
#[test]
#[should_panic(expected = "fmin must be strictly less than fmax")]
fn mel_filter_bank_panics_when_fmin_above_fmax() {
let _ = MelFilterBank::new(64, 1024, 16_000, 4_000.0, 1_000.0, MelScale::Slaney);
}
#[test]
#[should_panic(expected = "fmin must be >= 0")]
fn mel_filter_bank_panics_on_negative_fmin() {
let _ = MelFilterBank::new(64, 1024, 16_000, -10.0, 8_000.0, MelScale::Slaney);
}
}