use core::fmt;
use std::sync::Arc;
use rustfft::{Fft, FftPlanner, num_complex::Complex};
use crate::audio::ced::{
WINDOW_SAMPLES,
error::{AudioTooLong, Error, Result},
};
pub(crate) const N_MELS: usize = 64;
pub(crate) const N_FRAMES: usize = 1 + WINDOW_SAMPLES / HOP;
const N_FFT: usize = 512; const HOP: usize = 160; const SR: u32 = 16_000;
const FMIN: f64 = 0.0;
const FMAX: f64 = 8_000.0; const AMIN: f64 = 1e-10; const TOP_DB: f64 = 120.0; const N_FREQ: usize = N_FFT / 2 + 1;
pub(crate) struct MelExtractor {
window: Vec<f64>, filterbank: Vec<f64>, 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_hann(n: usize) -> Vec<f64> {
let denom = n as f64;
(0..n)
.map(|k| 0.5 - 0.5 * (2.0 * std::f64::consts::PI * (k as f64) / denom).cos())
.collect()
}
fn hz_to_htk_mel(hz: f64) -> f64 {
2595.0 * (1.0 + hz / 700.0).log10()
}
fn htk_mel_to_hz(mel: f64) -> f64 {
700.0 * (10f64.powf(mel / 2595.0) - 1.0)
}
fn build_htk_filterbank(sr: u32, n_fft: usize, n_mels: usize, fmin: f64, fmax: f64) -> Vec<f64> {
let n_freq = n_fft / 2 + 1;
let mel_min = Self::hz_to_htk_mel(fmin);
let mel_max = Self::hz_to_htk_mel(fmax);
let mel_points: Vec<f64> = (0..n_mels + 2)
.map(|i| mel_min + (mel_max - mel_min) * (i as f64) / (n_mels + 1) as f64)
.collect();
let hz_points: Vec<f64> = mel_points.iter().map(|&m| Self::htk_mel_to_hz(m)).collect();
let bin_hz: Vec<f64> = (0..n_freq)
.map(|k| (k as f64) * (sr as f64) / (n_fft as f64))
.collect();
let mut fb = vec![0.0f64; n_mels * n_freq];
for m in 0..n_mels {
let left = hz_points[m];
let center = hz_points[m + 1];
let right = hz_points[m + 2];
let inv_left_diff = 1.0 / (center - left);
let inv_right_diff = 1.0 / (right - center);
for (k, &f) in bin_hz.iter().enumerate() {
let weight = if f >= left && f <= center {
(f - left) * inv_left_diff
} else if f >= center && f <= right {
(right - f) * inv_right_diff
} else {
0.0
};
fb[m * n_freq + k] = weight;
}
}
fb
}
pub(crate) fn new() -> Self {
let window = Self::periodic_hann(N_FFT);
let filterbank = Self::build_htk_filterbank(SR, N_FFT, N_MELS, FMIN, FMAX);
let mut planner = FftPlanner::<f64>::new();
let fft = planner.plan_fft_forward(N_FFT);
Self {
window,
filterbank,
fft,
}
}
fn power_spectrum(fft_input: &[Complex<f64>], power: &mut [f64]) {
for (dst, c) in power.iter_mut().zip(fft_input.iter().take(N_FREQ)) {
*dst = c.re * c.re + c.im * c.im;
}
}
fn mel_filterbank_dot(weights: &[f64], power: &[f64]) -> f64 {
weights.iter().zip(power.iter()).map(|(w, p)| w * p).sum()
}
fn stft_one_frame_power(
&self,
frame: &[f64],
fft_input: &mut [Complex<f64>],
fft_scratch: &mut [Complex<f64>],
power: &mut [f64],
) {
for ((dst, &s), &w) in fft_input
.iter_mut()
.zip(frame.iter())
.zip(self.window.iter())
{
*dst = Complex::new(s * w, 0.0);
}
self.fft.process_with_scratch(fft_input, fft_scratch);
Self::power_spectrum(fft_input, power);
}
pub(crate) fn extract_into(&self, samples: &[f32], out: &mut [f32]) -> Result<()> {
debug_assert_eq!(out.len(), N_MELS * N_FRAMES);
if samples.is_empty() {
return Err(Error::EmptyAudio);
}
if samples.len() > WINDOW_SAMPLES {
return Err(Error::AudioTooLong(AudioTooLong::new(
samples.len(),
WINDOW_SAMPLES,
)));
}
let mut padded: Vec<f64> = Vec::with_capacity(WINDOW_SAMPLES);
padded.extend(samples.iter().map(|&s| s as f64));
padded.resize(WINDOW_SAMPLES, 0.0);
let half_fft = N_FFT / 2;
let mut centered: Vec<f64> = Vec::with_capacity(WINDOW_SAMPLES + 2 * half_fft);
for i in 0..half_fft {
centered.push(padded[half_fft - i]);
}
centered.extend_from_slice(&padded);
for i in 0..half_fft {
centered.push(padded[WINDOW_SAMPLES - 2 - i]);
}
debug_assert_eq!(centered.len(), WINDOW_SAMPLES + 2 * half_fft);
let mut frame = vec![0.0f64; N_FFT];
let mut power = vec![0.0f64; N_FREQ];
let mut fft_input = 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 db = vec![0.0f64; N_MELS * N_FRAMES];
let mut db_max = f64::MIN;
for t in 0..N_FRAMES {
let start = t * HOP;
frame.copy_from_slice(¢ered[start..start + N_FFT]);
self.stft_one_frame_power(&frame, &mut fft_input, &mut fft_scratch, &mut power);
for mel_bin in 0..N_MELS {
let row = &self.filterbank[mel_bin * N_FREQ..(mel_bin + 1) * N_FREQ];
let acc = Self::mel_filterbank_dot(row, &power);
let v = 10.0 * acc.max(AMIN).log10();
db[mel_bin * N_FRAMES + t] = v;
db_max = db_max.max(v);
}
}
let floor = db_max - TOP_DB;
for (dst, &v) in out.iter_mut().zip(db.iter()) {
*dst = v.max(floor) as f32;
}
Ok(())
}
}
#[cfg(test)]
mod tests;