use core::fmt;
use std::sync::Arc;
use rustfft::{Fft, FftPlanner, num_complex::Complex};
use crate::audio::lid::error::Result;
pub(crate) const N_MELS: usize = 60;
pub(crate) const N_FFT: usize = 400;
pub(crate) const HOP: usize = 160;
const N_FREQ: usize = N_FFT / 2 + 1;
const F_MAX: f64 = 8_000.0;
const AMIN: f64 = 1e-10;
const TOP_DB: f64 = 80.0;
pub(crate) struct MelExtractor {
window: [f32; N_FFT],
filterbank: Vec<f32>,
fft: Arc<dyn Fft<f64>>,
}
impl fmt::Debug for MelExtractor {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MelExtractor").finish_non_exhaustive()
}
}
impl MelExtractor {
fn periodic_hamming() -> [f32; N_FFT] {
let mut window = [0.0f32; N_FFT];
for (k, slot) in window.iter_mut().enumerate() {
let phase = core::f32::consts::TAU * (k as f32) / (N_FFT as f32);
let cosine = f64::from(phase).cos() as f32;
*slot = 0.54f32 - 0.46f32 * cosine;
}
window
}
fn hz_to_mel(hz: f64) -> f64 {
2595.0 * (1.0 + hz / 700.0).log10()
}
fn mel_to_hz(mel: f32) -> f32 {
let exponent = mel / 2595.0f32;
let power = 10.0f64.powf(f64::from(exponent)) as f32;
700.0f32 * (power - 1.0f32)
}
fn build_filterbank() -> Vec<f32> {
let mel_min = Self::hz_to_mel(0.0);
let mel_max = Self::hz_to_mel(F_MAX);
let step = (mel_max - mel_min) / (N_MELS + 1) as f64;
let mut edges_hz = [0.0f32; N_MELS + 2];
for (m, slot) in edges_hz.iter_mut().enumerate() {
let mel = ((m as f64) * step + mel_min) as f32;
*slot = Self::mel_to_hz(mel);
}
edges_hz[N_MELS + 1] = Self::mel_to_hz(mel_max as f32);
let mut filterbank = vec![0.0f32; N_FREQ * N_MELS];
for k in 0..N_FREQ {
let freq = (k as f64 * (F_MAX / (N_FREQ - 1) as f64)) as f32;
for m in 0..N_MELS {
let center = edges_hz[m + 1];
let band = edges_hz[m + 1] - edges_hz[m];
let slope = (freq - center) / band;
filterbank[k * N_MELS + m] = f32::max(0.0, f32::min(slope + 1.0, -slope + 1.0));
}
}
filterbank
}
pub(crate) fn new() -> Self {
let mut planner = FftPlanner::<f64>::new();
let fft = planner.plan_fft_forward(N_FFT);
Self {
window: Self::periodic_hamming(),
filterbank: Self::build_filterbank(),
fft,
}
}
pub(crate) fn extract_into(&self, samples: &[f32], out: &mut [f32]) -> Result<()> {
let frames = super::frame_count(samples.len());
debug_assert_eq!(out.len(), frames * N_MELS);
super::check_finite_samples(samples)?;
let half = N_FFT / 2;
let mut centered = vec![0.0f64; samples.len() + N_FFT];
for (dst, &src) in centered[half..half + samples.len()]
.iter_mut()
.zip(samples.iter())
{
*dst = f64::from(src);
}
let mut fft_buffer = vec![Complex::new(0.0f64, 0.0); N_FFT];
let mut fft_scratch = vec![Complex::new(0.0f64, 0.0); self.fft.get_inplace_scratch_len()];
let mut power = vec![0.0f64; N_FREQ];
let mut db = vec![0.0f64; frames * N_MELS];
let mut db_max = f64::NEG_INFINITY;
for t in 0..frames {
let start = t * HOP;
for ((slot, &sample), &weight) in fft_buffer
.iter_mut()
.zip(centered[start..start + N_FFT].iter())
.zip(self.window.iter())
{
*slot = Complex::new(sample * f64::from(weight), 0.0);
}
self
.fft
.process_with_scratch(&mut fft_buffer, &mut fft_scratch);
for (dst, bin) in power.iter_mut().zip(fft_buffer.iter().take(N_FREQ)) {
*dst = bin.re * bin.re + bin.im * bin.im;
}
let row = &mut db[t * N_MELS..(t + 1) * N_MELS];
for (m, slot) in row.iter_mut().enumerate() {
let mut energy = 0.0f64;
for (k, &bin_power) in power.iter().enumerate() {
energy += f64::from(self.filterbank[k * N_MELS + m]) * bin_power;
}
let decibels = 10.0 * energy.max(AMIN).log10();
*slot = decibels;
db_max = db_max.max(decibels);
}
}
let floor = db_max - TOP_DB;
for (dst, &value) in out.iter_mut().zip(db.iter()) {
*dst = value.max(floor) as f32;
}
Ok(())
}
}
#[cfg(test)]
mod tests;