use core::fmt;
use std::sync::Arc;
use rustfft::{Fft, FftPlanner, num_complex::Complex};
use crate::embeddings::clap::error::{Error, Result};
pub const T_FRAMES: usize = 1001;
pub const N_MELS: usize = 64;
pub const SAMPLE_RATE_HZ: u32 = 48_000;
pub const TARGET_SAMPLES: usize = 480_000;
const N_FFT: usize = 1024;
const HOP: usize = 480;
const FMIN: f64 = 50.0;
const FMAX: f64 = 14_000.0;
const POWER_TO_DB_AMIN: f64 = 1e-10;
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_slaney_mel(hz: f64) -> f64 {
const F_MIN: f64 = 0.0;
const F_SP: f64 = 200.0 / 3.0;
const MIN_LOG_HZ: f64 = 1000.0;
const MIN_LOG_MEL: f64 = (MIN_LOG_HZ - F_MIN) / F_SP;
let logstep = (6.4_f64).ln() / 27.0;
if hz < MIN_LOG_HZ {
(hz - F_MIN) / F_SP
} else {
MIN_LOG_MEL + (hz / MIN_LOG_HZ).ln() / logstep
}
}
fn slaney_mel_to_hz(mel: f64) -> f64 {
const F_MIN: f64 = 0.0;
const F_SP: f64 = 200.0 / 3.0;
const MIN_LOG_HZ: f64 = 1000.0;
const MIN_LOG_MEL: f64 = (MIN_LOG_HZ - F_MIN) / F_SP;
let logstep = (6.4_f64).ln() / 27.0;
if mel < MIN_LOG_MEL {
F_MIN + F_SP * mel
} else {
MIN_LOG_HZ * (logstep * (mel - MIN_LOG_MEL)).exp()
}
}
fn build_mel_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_slaney_mel(fmin);
let mel_max = Self::hz_to_slaney_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::slaney_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);
let slaney_norm = 2.0 / (right - left);
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 * slaney_norm;
}
}
fb
}
pub(crate) fn new() -> Self {
let window = Self::periodic_hann(N_FFT);
let filterbank = Self::build_mel_filterbank(SAMPLE_RATE_HZ, 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 * T_FRAMES);
if samples.is_empty() {
return Err(Error::EmptyAudio);
}
let mut padded: Vec<f64> = Vec::with_capacity(TARGET_SAMPLES);
if samples.len() >= TARGET_SAMPLES {
padded.extend(samples[..TARGET_SAMPLES].iter().map(|&s| s as f64));
} else {
let n_repeat = TARGET_SAMPLES / samples.len();
for _ in 0..n_repeat {
padded.extend(samples.iter().map(|&s| s as f64));
}
padded.resize(TARGET_SAMPLES, 0.0);
}
let half_fft = N_FFT / 2;
let mut centered: Vec<f64> = Vec::with_capacity(TARGET_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[TARGET_SAMPLES - 2 - i]);
}
debug_assert_eq!(centered.len(), TARGET_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()];
for t in 0..T_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 db = 10.0 * acc.max(POWER_TO_DB_AMIN).log10();
out[t * N_MELS + mel_bin] = db as f32;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests;