use ndarray::Array2;
use num::Complex;
use rustfft::{Fft, FftPlanner};
use std::f64::consts::PI;
use std::sync::Arc;
use crate::mel::SparseMelFilterbank;
#[derive(Clone, Debug)]
pub struct FbankConfig {
pub sample_rate: f64,
pub num_mel_bins: usize,
pub frame_length_ms: f64,
pub frame_shift_ms: f64,
pub dither: f64,
pub energy_floor: f64,
pub use_energy: bool,
pub use_log_fbank: bool,
pub use_power: bool,
pub preemphasis: f64,
pub apply_cmn: bool,
pub low_freq: f64,
pub high_freq: f64,
}
impl Default for FbankConfig {
fn default() -> Self {
Self {
sample_rate: 16000.0,
num_mel_bins: 80,
frame_length_ms: 25.0,
frame_shift_ms: 10.0,
dither: 0.0,
energy_floor: 0.0, use_energy: false,
use_log_fbank: true,
use_power: true,
preemphasis: 0.97,
apply_cmn: true,
low_freq: 20.0,
high_freq: 0.0, }
}
}
impl FbankConfig {
pub fn frame_length_samples(&self) -> usize {
((self.frame_length_ms / 1000.0) * self.sample_rate).round() as usize
}
pub fn frame_shift_samples(&self) -> usize {
((self.frame_shift_ms / 1000.0) * self.sample_rate).round() as usize
}
pub fn fft_size(&self) -> usize {
let frame_len = self.frame_length_samples();
frame_len.next_power_of_two()
}
}
pub struct Fbank {
config: FbankConfig,
mel_filters: Array2<f64>,
sparse_mel_filters: SparseMelFilterbank,
fft: Arc<dyn Fft<f64>>,
window: Vec<f64>,
}
impl Fbank {
pub fn new(config: FbankConfig) -> Self {
let fft_size = config.fft_size();
let frame_len = config.frame_length_samples();
let window: Vec<f64> = (0..frame_len)
.map(|i| {
let a = 2.0 * PI * i as f64 / (frame_len - 1) as f64;
(0.5 - 0.5 * a.cos()).powf(0.85)
})
.collect();
let high_freq = if config.high_freq == 0.0 {
config.sample_rate / 2.0
} else {
config.high_freq
};
let mel_filters = kaldi_mel_filterbank(
config.sample_rate,
fft_size,
config.num_mel_bins,
config.low_freq,
high_freq,
);
let sparse_mel_filters = SparseMelFilterbank::from_dense(&mel_filters);
let mut planner = FftPlanner::new();
let fft = planner.plan_fft_forward(fft_size);
Self {
config,
mel_filters,
sparse_mel_filters,
fft,
window,
}
}
pub fn compute(&self, samples: &[f32]) -> Array2<f32> {
let frame_len = self.config.frame_length_samples();
let frame_shift = self.config.frame_shift_samples();
let fft_size = self.config.fft_size();
let preemph = self.config.preemphasis;
if samples.len() < frame_len {
return Array2::zeros((0, self.config.num_mel_bins));
}
let num_frames = 1 + (samples.len() - frame_len) / frame_shift;
let mut features = Array2::zeros((num_frames, self.config.num_mel_bins));
let mut complex_buf = vec![Complex::new(0.0, 0.0); fft_size];
let mut scratch_buf = vec![Complex::new(0.0, 0.0); self.fft.get_inplace_scratch_len()];
let mut frame_buf = vec![0.0f64; frame_len];
let mut power_spectrum = vec![0.0f64; fft_size / 2 + 1];
let mut mel_energies = vec![0.0f64; self.config.num_mel_bins];
for frame_idx in 0..num_frames {
let start = frame_idx * frame_shift;
let end = start + frame_len;
let frame_slice = &samples[start..end];
let mean: f64 = frame_slice.iter().map(|&x| x as f64).sum::<f64>() / frame_len as f64;
for (i, &sample) in frame_slice.iter().enumerate() {
frame_buf[i] = sample as f64 - mean;
}
if preemph > 0.0 {
for i in (1..frame_len).rev() {
frame_buf[i] -= preemph * frame_buf[i - 1];
}
if start > 0 {
frame_buf[0] -= preemph * (samples[start - 1] as f64 - mean);
}
}
for (i, &sample) in frame_buf.iter().enumerate() {
complex_buf[i] = Complex::new(sample * self.window[i], 0.0);
}
for i in frame_len..fft_size {
complex_buf[i] = Complex::new(0.0, 0.0);
}
self.fft
.process_with_scratch(&mut complex_buf, &mut scratch_buf);
for (i, c) in complex_buf.iter().take(fft_size / 2 + 1).enumerate() {
power_spectrum[i] = if self.config.use_power {
c.norm_sqr()
} else {
c.norm()
};
}
self.sparse_mel_filters
.project_power_f64(&power_spectrum, &mut mel_energies);
for (mel_idx, mel_energy) in mel_energies.iter_mut().enumerate() {
let floor = if self.config.energy_floor > 0.0 {
self.config.energy_floor
} else {
f32::EPSILON as f64 };
*mel_energy = (*mel_energy).max(floor);
if self.config.use_log_fbank {
*mel_energy = mel_energy.ln();
}
features[[frame_idx, mel_idx]] = *mel_energy as f32;
}
}
if self.config.apply_cmn && num_frames > 0 {
for mel_idx in 0..self.config.num_mel_bins {
let mean: f32 = features.column(mel_idx).mean().unwrap_or(0.0);
for frame_idx in 0..num_frames {
features[[frame_idx, mel_idx]] -= mean;
}
}
}
features
}
pub fn config(&self) -> &FbankConfig {
&self.config
}
pub fn dense_filterbank(&self) -> &Array2<f64> {
&self.mel_filters
}
}
fn kaldi_mel_filterbank(
sample_rate: f64,
fft_size: usize,
num_mel_bins: usize,
low_freq: f64,
high_freq: f64,
) -> Array2<f64> {
let num_fft_bins = fft_size / 2 + 1;
let mel_low = hz_to_mel(low_freq);
let mel_high = hz_to_mel(high_freq);
let mel_points: Vec<f64> = (0..=num_mel_bins + 1)
.map(|i| mel_low + (mel_high - mel_low) * i as f64 / (num_mel_bins + 1) as f64)
.collect();
let hz_points: Vec<f64> = mel_points.iter().map(|&m| mel_to_hz(m)).collect();
let mut filters = Array2::zeros((num_mel_bins, num_fft_bins));
for mel_idx in 0..num_mel_bins {
let left_hz = hz_points[mel_idx];
let center_hz = hz_points[mel_idx + 1];
let right_hz = hz_points[mel_idx + 2];
if center_hz <= left_hz || right_hz <= center_hz {
continue;
}
for freq_idx in 0..num_fft_bins {
let freq_hz = freq_idx as f64 * sample_rate / fft_size as f64;
if freq_hz > left_hz && freq_hz <= center_hz {
filters[[mel_idx, freq_idx]] = (freq_hz - left_hz) / (center_hz - left_hz);
} else if freq_hz > center_hz && freq_hz < right_hz {
filters[[mel_idx, freq_idx]] = (right_hz - freq_hz) / (right_hz - center_hz);
}
}
}
filters
}
fn hz_to_mel(hz: f64) -> f64 {
1127.0 * (1.0 + hz / 700.0).ln()
}
fn mel_to_hz(mel: f64) -> f64 {
700.0 * ((mel / 1127.0).exp() - 1.0)
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray_npy::NpzReader;
use std::fs::File;
use std::io::Read;
fn find_wav_data_offset(wav_bytes: &[u8]) -> Option<usize> {
if wav_bytes.len() < 12 {
return None;
}
let mut pos = 12; while pos + 8 <= wav_bytes.len() {
let chunk_id = &wav_bytes[pos..pos + 4];
let chunk_size = u32::from_le_bytes([
wav_bytes[pos + 4],
wav_bytes[pos + 5],
wav_bytes[pos + 6],
wav_bytes[pos + 7],
]) as usize;
if chunk_id == b"data" {
return Some(pos + 8); }
pos += 8 + chunk_size;
if chunk_size % 2 != 0 {
pos += 1;
}
}
None
}
#[test]
fn test_fbank_config_defaults() {
let config = FbankConfig::default();
assert_eq!(config.sample_rate, 16000.0);
assert_eq!(config.num_mel_bins, 80);
assert_eq!(config.frame_length_samples(), 400);
assert_eq!(config.frame_shift_samples(), 160);
assert_eq!(config.fft_size(), 512);
}
#[test]
fn test_hz_to_mel() {
assert!((hz_to_mel(0.0) - 0.0).abs() < 1e-6);
assert!((hz_to_mel(1000.0) - 999.98).abs() < 1.0);
assert!((hz_to_mel(8000.0) - 2840.0).abs() < 1.0);
}
#[test]
fn test_mel_to_hz() {
for hz in [0.0, 500.0, 1000.0, 4000.0, 8000.0] {
let mel = hz_to_mel(hz);
let hz_back = mel_to_hz(mel);
assert!(
(hz - hz_back).abs() < 1e-6,
"Round-trip failed for Hz={}",
hz
);
}
}
#[test]
fn test_fbank_basic() {
let config = FbankConfig::default();
let fbank = Fbank::new(config);
let samples = vec![0.0f32; 16000];
let features = fbank.compute(&samples);
assert_eq!(features.shape()[1], 80); assert!(features.shape()[0] > 90 && features.shape()[0] < 100);
}
#[test]
fn test_sparse_projection_matches_dense_filterbank() {
let config = FbankConfig::default();
let fbank = Fbank::new(config);
let power_spectrum = (0..fbank.mel_filters.ncols())
.map(|idx| ((idx as f64 + 1.0) * 0.013).sin().abs())
.collect::<Vec<_>>();
let mut sparse = vec![0.0; fbank.config.num_mel_bins];
fbank
.sparse_mel_filters
.project_power_f64(&power_spectrum, &mut sparse);
for mel_idx in 0..fbank.config.num_mel_bins {
let dense = fbank
.mel_filters
.row(mel_idx)
.iter()
.zip(power_spectrum.iter())
.map(|(filter, power)| filter * power)
.sum::<f64>();
assert!(
(sparse[mel_idx] - dense).abs() <= 1e-12,
"mel {mel_idx}: sparse {}, dense {}",
sparse[mel_idx],
dense
);
}
assert!(
fbank.sparse_mel_filters.non_zero_weights()
< fbank.sparse_mel_filters.dense_weights() / 10
);
}
#[test]
fn test_fbank_vs_kaldi_golden() {
let npz_path = "./testdata/kaldi_native_fbank_jfk.npz";
if !std::path::Path::new(npz_path).exists() {
eprintln!("Skipping golden test: {} not found", npz_path);
return;
}
let f = File::open(npz_path).unwrap();
let mut npz = NpzReader::new(f).unwrap();
let golden: Array2<f32> = npz.by_name("features").unwrap();
let wav_path = "./testdata/jfk_f32le.wav";
let mut wav_file = File::open(wav_path).unwrap();
let mut wav_bytes = Vec::new();
wav_file.read_to_end(&mut wav_bytes).unwrap();
let data_offset =
find_wav_data_offset(&wav_bytes).expect("Could not find 'data' chunk in WAV file");
let samples: Vec<f32> = wav_bytes[data_offset..]
.chunks_exact(4)
.map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.collect();
let config = FbankConfig {
apply_cmn: true,
..FbankConfig::default()
};
let fbank = Fbank::new(config);
let computed = fbank.compute(&samples);
let golden_t = golden.t();
eprintln!("Computed shape: {:?}", computed.shape());
eprintln!("Golden shape: {:?}", golden_t.shape());
assert_eq!(
computed.shape()[0],
golden_t.shape()[0],
"Frame count mismatch: computed {} vs golden {}",
computed.shape()[0],
golden_t.shape()[0]
);
let num_check = computed.shape()[0].min(50);
let mut max_diff = 0.0f32;
let mut sum_diff = 0.0f32;
let mut count = 0;
for frame_idx in 0..num_check {
for mel_idx in 0..80 {
let diff = (computed[[frame_idx, mel_idx]] - golden_t[[frame_idx, mel_idx]]).abs();
max_diff = max_diff.max(diff);
sum_diff += diff;
count += 1;
}
}
let avg_diff = sum_diff / count as f32;
eprintln!("Max difference: {:.4}", max_diff);
eprintln!("Avg difference: {:.4}", avg_diff);
eprintln!("\nFirst frame comparison (computed vs golden):");
for mel_idx in 0..5 {
eprintln!(
" mel[{}]: {:.4} vs {:.4}",
mel_idx,
computed[[0, mel_idx]],
golden_t[[0, mel_idx]]
);
}
eprintln!("\nNote: This is an approximation of kaldi fbank.");
eprintln!("For exact kaldi compatibility, use TorchScript-traced fbank model.");
let all_finite = computed.iter().all(|&x| x.is_finite());
assert!(all_finite, "Computed features contain non-finite values");
let variance: f32 = computed.iter().map(|&x| x * x).sum::<f32>() / computed.len() as f32;
assert!(variance > 0.1, "Output variance too low: {}", variance);
}
#[test]
fn debug_filterbank() {
let config = FbankConfig::default();
let fbank = Fbank::new(config);
println!("\nFilterbank check:");
for mel_idx in 0..10 {
let row = fbank.mel_filters.row(mel_idx);
let sum: f64 = row.iter().sum();
let nonzero: usize = row.iter().filter(|&&x| x > 0.0).count();
println!(
" Filter {}: sum={:.4}, nonzero_bins={}",
mel_idx, sum, nonzero
);
}
println!("\nFirst 3 filters (non-zero weights):");
for mel_idx in 0..3 {
let row = fbank.mel_filters.row(mel_idx);
print!(" Filter {}: ", mel_idx);
for (i, &w) in row.iter().enumerate() {
if w > 0.0 {
print!("bin{}={:.3} ", i, w);
}
}
println!();
}
}
#[test]
fn debug_compute_steps() {
let mut wav_file = File::open("./testdata/jfk_f32le.wav").unwrap();
let mut wav_bytes = Vec::new();
wav_file.read_to_end(&mut wav_bytes).unwrap();
let data_offset =
find_wav_data_offset(&wav_bytes).expect("Could not find 'data' chunk in WAV file");
let samples: Vec<f32> = wav_bytes[data_offset..]
.chunks_exact(4)
.map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.collect();
println!("\nFirst 10 audio samples: {:?}", &samples[..10]);
let config = FbankConfig {
apply_cmn: false,
..FbankConfig::default()
};
let fbank = Fbank::new(config);
let features = fbank.compute(&samples);
println!("\nFrame 0 (silent), first 5 mel bins (no CMN):");
for i in 0..5 {
println!(" mel[{}]: {:.4}", i, features[[0, i]]);
}
println!("\nExpected (from kaldi): -15.94 for all");
println!(
"f32::EPSILON: {:e}, ln(EPSILON): {:.4}",
f32::EPSILON,
(f32::EPSILON as f64).ln()
);
}
}