use core::fmt;
use std::sync::Arc;
use rustfft::{Fft, FftPlanner, num_complex::Complex};
use crate::audio::identity::{
N_FRAMES, N_MELS, SAMPLE_RATE_HZ, WINDOW_SAMPLES,
error::{Error, Result, WindowLength},
};
const N_FFT: usize = 512;
const WIN_LENGTH: usize = 400;
pub(super) const HOP: usize = 240;
const WINDOW_OFFSET: usize = (N_FFT - WIN_LENGTH) / 2;
const F_MIN: f64 = 20.0;
const F_MAX: f64 = 7600.0;
const PRE_EMPHASIS: f64 = 0.97;
const LOG_EPSILON: f64 = 1e-6;
const N_FREQ: usize = N_FFT / 2 + 1;
const CENTER_PAD: usize = N_FFT / 2;
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_hamming(n: usize) -> Vec<f64> {
let denom = n as f64;
(0..n)
.map(|k| 0.54 - 0.46 * (2.0 * std::f64::consts::PI * (k as f64) / denom).cos())
.collect()
}
fn padded_window() -> Vec<f64> {
let taps = Self::periodic_hamming(WIN_LENGTH);
let mut window = vec![0.0f64; N_FFT];
window[WINDOW_OFFSET..WINDOW_OFFSET + WIN_LENGTH].copy_from_slice(&taps);
window
}
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 all_freqs: Vec<f64> = (0..n_freq)
.map(|k| (k as f64) * (f64::from(sr) / 2.0) / ((n_freq - 1) as f64))
.collect();
let mel_min = Self::hz_to_htk_mel(fmin);
let mel_max = Self::hz_to_htk_mel(fmax);
let f_pts: Vec<f64> = (0..n_mels + 2)
.map(|i| {
let mel = mel_min + (mel_max - mel_min) * (i as f64) / ((n_mels + 1) as f64);
Self::htk_mel_to_hz(mel)
})
.collect();
let mut fb = vec![0.0f64; n_mels * n_freq];
for m in 0..n_mels {
let left_diff = f_pts[m + 1] - f_pts[m];
let right_diff = f_pts[m + 2] - f_pts[m + 1];
for (k, &f) in all_freqs.iter().enumerate() {
let down = (f - f_pts[m]) / left_diff;
let up = (f_pts[m + 2] - f) / right_diff;
fb[m * n_freq + k] = down.min(up).max(0.0);
}
}
fb
}
pub(crate) fn new() -> Self {
let window = Self::padded_window();
let filterbank = Self::build_htk_filterbank(SAMPLE_RATE_HZ, N_FFT, N_MELS, F_MIN, F_MAX);
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);
}
fn pre_emphasize(samples: &[f32]) -> Vec<f64> {
debug_assert!(samples.len() >= 2, "the fixed window is far longer than 2");
let mut out = Vec::with_capacity(samples.len());
out.push(f64::from(samples[0]) - PRE_EMPHASIS * f64::from(samples[1]));
for pair in samples.windows(2) {
out.push(f64::from(pair[1]) - PRE_EMPHASIS * f64::from(pair[0]));
}
out
}
fn center_pad(signal: &[f64]) -> Vec<f64> {
let len = signal.len();
debug_assert!(len > CENTER_PAD, "reflect padding needs len > pad");
let mut padded = Vec::with_capacity(len + 2 * CENTER_PAD);
for i in 0..CENTER_PAD {
padded.push(signal[CENTER_PAD - i]);
}
padded.extend_from_slice(signal);
for i in 0..CENTER_PAD {
padded.push(signal[len - 2 - i]);
}
padded
}
pub(crate) fn extract_into(&self, samples: &[f32], out: &mut [f32]) -> Result<()> {
debug_assert_eq!(out.len(), N_MELS * N_FRAMES);
if samples.len() != WINDOW_SAMPLES {
return Err(Error::WindowLength(WindowLength::new(
samples.len(),
WINDOW_SAMPLES,
)));
}
let emphasized = Self::pre_emphasize(samples);
let padded = Self::center_pad(&emphasized);
debug_assert_eq!(padded.len(), WINDOW_SAMPLES + 2 * CENTER_PAD);
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 log_mel = vec![0.0f64; N_MELS * N_FRAMES];
for t in 0..N_FRAMES {
let start = t * HOP;
frame.copy_from_slice(&padded[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);
log_mel[mel_bin * N_FRAMES + t] = (acc + LOG_EPSILON).ln();
}
}
for mel_bin in 0..N_MELS {
let row = &log_mel[mel_bin * N_FRAMES..(mel_bin + 1) * N_FRAMES];
let mean = row.iter().sum::<f64>() / (N_FRAMES as f64);
let dst = &mut out[mel_bin * N_FRAMES..(mel_bin + 1) * N_FRAMES];
for (o, &v) in dst.iter_mut().zip(row.iter()) {
*o = (v - mean) as f32;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests;