Whole-workspace rustfmt pass picked up while iterating on Mamba GPU backward work. Verified formatting-only via diff sampling; no logic changed. Co-Authored-By: Claude Sonnet 5 <[email protected]>
560 lines
17 KiB
Rust
560 lines
17 KiB
Rust
//! Mel spectrogram utilities for audio processing
|
|
//!
|
|
//! This module provides utilities for converting between linear and mel-scale
|
|
//! spectrograms, which are commonly used in speech processing and TTS systems.
|
|
|
|
use crate::error::{Result, TtsError};
|
|
use rtx_tensor::{Device, Result as TensorResult, Tensor, TensorError};
|
|
use serde::{Deserialize, Serialize};
|
|
|
|
/// Configuration for mel spectrogram computation
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct MelConfig {
|
|
/// Sample rate in Hz
|
|
pub sample_rate: usize,
|
|
/// FFT size
|
|
pub n_fft: usize,
|
|
/// Hop length between frames
|
|
pub hop_length: usize,
|
|
/// Window length
|
|
pub win_length: usize,
|
|
/// Number of mel frequency bins
|
|
pub n_mels: usize,
|
|
/// Minimum frequency in Hz
|
|
pub f_min: f32,
|
|
/// Maximum frequency in Hz (None means sample_rate / 2)
|
|
pub f_max: Option<f32>,
|
|
}
|
|
|
|
impl Default for MelConfig {
|
|
fn default() -> Self {
|
|
Self {
|
|
sample_rate: 22050,
|
|
n_fft: 1024,
|
|
hop_length: 256,
|
|
win_length: 1024,
|
|
n_mels: 80,
|
|
f_min: 0.0,
|
|
f_max: None, // Will use sample_rate / 2
|
|
}
|
|
}
|
|
}
|
|
|
|
impl MelConfig {
|
|
/// Get the maximum frequency, defaulting to Nyquist frequency
|
|
pub fn get_f_max(&self) -> f32 {
|
|
self.f_max.unwrap_or(self.sample_rate as f32 / 2.0)
|
|
}
|
|
|
|
/// Validate configuration parameters
|
|
pub fn validate(&self) -> Result<()> {
|
|
if self.sample_rate == 0 {
|
|
return Err(TtsError::InvalidConfig(
|
|
"Sample rate must be > 0".to_string(),
|
|
));
|
|
}
|
|
if self.n_fft == 0 {
|
|
return Err(TtsError::InvalidConfig("FFT size must be > 0".to_string()));
|
|
}
|
|
if self.hop_length == 0 {
|
|
return Err(TtsError::InvalidConfig(
|
|
"Hop length must be > 0".to_string(),
|
|
));
|
|
}
|
|
if self.win_length == 0 || self.win_length > self.n_fft {
|
|
return Err(TtsError::InvalidConfig(
|
|
"Window length must be > 0 and <= n_fft".to_string(),
|
|
));
|
|
}
|
|
if self.n_mels == 0 {
|
|
return Err(TtsError::InvalidConfig(
|
|
"Number of mels must be > 0".to_string(),
|
|
));
|
|
}
|
|
if self.f_min < 0.0 {
|
|
return Err(TtsError::InvalidConfig("f_min must be >= 0".to_string()));
|
|
}
|
|
let f_max = self.get_f_max();
|
|
if f_max <= self.f_min {
|
|
return Err(TtsError::InvalidConfig("f_max must be > f_min".to_string()));
|
|
}
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
/// Mel spectrogram processor
|
|
pub struct MelSpectrogram {
|
|
config: MelConfig,
|
|
mel_filterbank: Tensor,
|
|
device: Device,
|
|
}
|
|
|
|
impl MelSpectrogram {
|
|
/// Create a new mel spectrogram processor
|
|
pub fn new(config: MelConfig, device: &Device) -> Result<Self> {
|
|
config.validate()?;
|
|
|
|
let mel_filterbank = Self::create_mel_filterbank(&config, device)
|
|
.map_err(|e| TtsError::TensorError(e.to_string()))?;
|
|
|
|
Ok(Self {
|
|
config,
|
|
mel_filterbank,
|
|
device: device.clone(),
|
|
})
|
|
}
|
|
|
|
/// Create mel filterbank matrix
|
|
fn create_mel_filterbank(config: &MelConfig, device: &Device) -> TensorResult<Tensor> {
|
|
let n_freqs = config.n_fft / 2 + 1;
|
|
let f_max = config.get_f_max();
|
|
|
|
// Convert Hz to mel scale
|
|
let mel_min = Self::hz_to_mel(config.f_min);
|
|
let mel_max = Self::hz_to_mel(f_max);
|
|
|
|
// Create equally spaced mel points
|
|
let mut mel_points = Vec::with_capacity(config.n_mels + 2);
|
|
for i in 0..(config.n_mels + 2) {
|
|
let mel = mel_min + (mel_max - mel_min) * i as f32 / (config.n_mels + 1) as f32;
|
|
mel_points.push(mel);
|
|
}
|
|
|
|
// Convert mel points to Hz
|
|
let freq_points: Vec<f32> = mel_points.iter().map(|&m| Self::mel_to_hz(m)).collect();
|
|
|
|
// Convert Hz to FFT bin indices
|
|
let bin_points: Vec<f32> = freq_points
|
|
.iter()
|
|
.map(|&f| f * config.n_fft as f32 / config.sample_rate as f32)
|
|
.collect();
|
|
|
|
// Create triangular filters
|
|
let mut filterbank_data = vec![0.0f32; config.n_mels * n_freqs];
|
|
|
|
for mel_idx in 0..config.n_mels {
|
|
let left = bin_points[mel_idx];
|
|
let center = bin_points[mel_idx + 1];
|
|
let right = bin_points[mel_idx + 2];
|
|
|
|
for freq_idx in 0..n_freqs {
|
|
let bin = freq_idx as f32;
|
|
let value = if bin >= left && bin < center {
|
|
if center > left {
|
|
(bin - left) / (center - left)
|
|
} else {
|
|
0.0
|
|
}
|
|
} else if bin >= center && bin < right {
|
|
if right > center {
|
|
(right - bin) / (right - center)
|
|
} else {
|
|
0.0
|
|
}
|
|
} else {
|
|
0.0
|
|
};
|
|
|
|
filterbank_data[mel_idx * n_freqs + freq_idx] = value;
|
|
}
|
|
}
|
|
|
|
Tensor::from_data(filterbank_data, vec![config.n_mels, n_freqs], device)
|
|
}
|
|
|
|
/// Convert linear spectrogram to mel spectrogram
|
|
///
|
|
/// # Arguments
|
|
/// * `linear_spec` - Linear spectrogram [n_freqs, time_frames]
|
|
///
|
|
/// # Returns
|
|
/// Mel spectrogram [n_mels, time_frames]
|
|
pub fn spectrogram_to_mel(&self, linear_spec: &Tensor) -> Result<Tensor> {
|
|
// mel_spec = mel_fb @ linear_spec
|
|
// [n_mels, n_freqs] @ [n_freqs, time_frames] -> [n_mels, time_frames]
|
|
self.mel_filterbank
|
|
.matmul(linear_spec)
|
|
.map_err(|e| TtsError::TensorError(e.to_string()))
|
|
}
|
|
|
|
/// Convert mel spectrogram to linear spectrogram (approximate inverse)
|
|
///
|
|
/// Uses pseudo-inverse of mel filterbank for reconstruction
|
|
///
|
|
/// # Arguments
|
|
/// * `mel_spec` - Mel spectrogram [n_mels, time_frames]
|
|
///
|
|
/// # Returns
|
|
/// Linear spectrogram [n_freqs, time_frames]
|
|
pub fn mel_to_spectrogram(&self, mel_spec: &Tensor) -> Result<Tensor> {
|
|
// Compute pseudo-inverse: (A^T A)^-1 A^T
|
|
// For efficiency, we use a simple transpose approximation
|
|
// This is not perfect but works well for vocoding
|
|
let mel_fb_t = self
|
|
.mel_filterbank
|
|
.transpose(0, 1)
|
|
.map_err(|e| TtsError::TensorError(e.to_string()))?;
|
|
|
|
// linear_spec ≈ mel_fb^T @ mel_spec
|
|
// [n_freqs, n_mels] @ [n_mels, time_frames] -> [n_freqs, time_frames]
|
|
mel_fb_t
|
|
.matmul(mel_spec)
|
|
.map_err(|e| TtsError::TensorError(e.to_string()))
|
|
}
|
|
|
|
/// Convert audio to mel spectrogram
|
|
///
|
|
/// # Arguments
|
|
/// * `audio` - Audio waveform [samples]
|
|
///
|
|
/// # Returns
|
|
/// Mel spectrogram [n_mels, time_frames]
|
|
pub fn audio_to_mel(&self, audio: &Tensor) -> Result<Tensor> {
|
|
use rtx_tensor::signal::{MelOptions, WindowType, mel_spectrogram};
|
|
|
|
let options = MelOptions {
|
|
sample_rate: self.config.sample_rate as f32,
|
|
n_mels: self.config.n_mels,
|
|
f_min: self.config.f_min,
|
|
f_max: self.config.get_f_max(),
|
|
power: 2.0,
|
|
normalize: true,
|
|
};
|
|
|
|
mel_spectrogram(
|
|
audio,
|
|
self.config.n_fft,
|
|
self.config.hop_length,
|
|
Some(WindowType::Hann),
|
|
true,
|
|
&options,
|
|
)
|
|
.map_err(|e| TtsError::TensorError(e.to_string()))
|
|
}
|
|
|
|
/// Get the configuration
|
|
pub fn config(&self) -> &MelConfig {
|
|
&self.config
|
|
}
|
|
|
|
/// Get the mel filterbank
|
|
pub fn filterbank(&self) -> &Tensor {
|
|
&self.mel_filterbank
|
|
}
|
|
|
|
/// Convert frequency in Hz to mel scale
|
|
fn hz_to_mel(hz: f32) -> f32 {
|
|
2595.0 * (1.0 + hz / 700.0).log10()
|
|
}
|
|
|
|
/// Convert mel scale to frequency in Hz
|
|
fn mel_to_hz(mel: f32) -> f32 {
|
|
700.0 * (10.0f32.powf(mel / 2595.0) - 1.0)
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use approx::assert_relative_eq;
|
|
|
|
#[test]
|
|
fn test_mel_config_default() {
|
|
let config = MelConfig::default();
|
|
assert_eq!(config.sample_rate, 22050);
|
|
assert_eq!(config.n_fft, 1024);
|
|
assert_eq!(config.n_mels, 80);
|
|
assert_eq!(config.f_min, 0.0);
|
|
assert_eq!(config.get_f_max(), 11025.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_valid() {
|
|
let config = MelConfig::default();
|
|
assert!(config.validate().is_ok());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_sample_rate() {
|
|
let mut config = MelConfig::default();
|
|
config.sample_rate = 0;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_n_fft() {
|
|
let mut config = MelConfig::default();
|
|
config.n_fft = 0;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_hop_length() {
|
|
let mut config = MelConfig::default();
|
|
config.hop_length = 0;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_win_length() {
|
|
let mut config = MelConfig::default();
|
|
config.win_length = 0;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_win_length_too_large() {
|
|
let mut config = MelConfig::default();
|
|
config.win_length = config.n_fft + 1;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_n_mels() {
|
|
let mut config = MelConfig::default();
|
|
config.n_mels = 0;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_f_min() {
|
|
let mut config = MelConfig::default();
|
|
config.f_min = -1.0;
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_validate_invalid_f_max() {
|
|
let mut config = MelConfig::default();
|
|
config.f_max = Some(0.0);
|
|
assert!(config.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_config_serialize() {
|
|
let config = MelConfig::default();
|
|
let json = serde_json::to_string(&config).unwrap();
|
|
let deserialized: MelConfig = serde_json::from_str(&json).unwrap();
|
|
assert_eq!(config.sample_rate, deserialized.sample_rate);
|
|
assert_eq!(config.n_mels, deserialized.n_mels);
|
|
}
|
|
|
|
#[test]
|
|
fn test_hz_to_mel_conversions() {
|
|
// Test known conversions
|
|
let mel_0 = MelSpectrogram::hz_to_mel(0.0);
|
|
assert_relative_eq!(mel_0, 0.0, epsilon = 1e-4);
|
|
|
|
let mel_1000 = MelSpectrogram::hz_to_mel(1000.0);
|
|
// HTK mel formula: 2595 * log10(1 + hz/700) ≈ 401.9 mels at 1000 Hz
|
|
assert!(mel_1000 > 0.0);
|
|
assert!(mel_1000 < 1000.0); // 1000 Hz ≈ 401.9 mels with HTK formula
|
|
|
|
// Test round-trip conversion
|
|
let hz = 440.0; // A4 note
|
|
let mel = MelSpectrogram::hz_to_mel(hz);
|
|
let hz_back = MelSpectrogram::mel_to_hz(mel);
|
|
assert_relative_eq!(hz, hz_back, epsilon = 1e-3);
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_to_hz_conversions() {
|
|
let hz_0 = MelSpectrogram::mel_to_hz(0.0);
|
|
assert_relative_eq!(hz_0, 0.0, epsilon = 1e-4);
|
|
|
|
let hz_1000 = MelSpectrogram::mel_to_hz(1000.0);
|
|
assert!(hz_1000 > 0.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_spectrogram_creation() {
|
|
let device = Device::cpu();
|
|
let config = MelConfig::default();
|
|
let mel_spec = MelSpectrogram::new(config.clone(), &device);
|
|
assert!(mel_spec.is_ok());
|
|
|
|
let mel_spec = mel_spec.unwrap();
|
|
assert_eq!(mel_spec.config().n_mels, config.n_mels);
|
|
|
|
// Check filterbank shape
|
|
let fb_dims = mel_spec.filterbank().dims();
|
|
assert_eq!(fb_dims[0], config.n_mels);
|
|
assert_eq!(fb_dims[1], config.n_fft / 2 + 1);
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_spectrogram_creation_invalid_config() {
|
|
let device = Device::cpu();
|
|
let mut config = MelConfig::default();
|
|
config.n_fft = 0;
|
|
let result = MelSpectrogram::new(config, &device);
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_spectrogram_to_mel() {
|
|
let device = Device::cpu();
|
|
let config = MelConfig {
|
|
sample_rate: 16000,
|
|
n_fft: 512,
|
|
hop_length: 128,
|
|
win_length: 512,
|
|
n_mels: 40,
|
|
f_min: 0.0,
|
|
f_max: None,
|
|
};
|
|
let mel_proc = MelSpectrogram::new(config.clone(), &device).unwrap();
|
|
|
|
// Create a dummy linear spectrogram
|
|
let n_freqs = config.n_fft / 2 + 1;
|
|
let n_frames = 50;
|
|
let linear_spec = Tensor::randn(&[n_freqs, n_frames], &device).unwrap();
|
|
|
|
// Convert to mel
|
|
let mel_spec = mel_proc.spectrogram_to_mel(&linear_spec).unwrap();
|
|
let mel_dims = mel_spec.dims();
|
|
assert_eq!(mel_dims[0], config.n_mels);
|
|
assert_eq!(mel_dims[1], n_frames);
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_to_spectrogram() {
|
|
let device = Device::cpu();
|
|
let config = MelConfig {
|
|
sample_rate: 16000,
|
|
n_fft: 512,
|
|
hop_length: 128,
|
|
win_length: 512,
|
|
n_mels: 40,
|
|
f_min: 0.0,
|
|
f_max: None,
|
|
};
|
|
let mel_proc = MelSpectrogram::new(config.clone(), &device).unwrap();
|
|
|
|
// Create a dummy mel spectrogram
|
|
let n_frames = 50;
|
|
let mel_spec = Tensor::randn(&[config.n_mels, n_frames], &device).unwrap();
|
|
|
|
// Convert to linear
|
|
let linear_spec = mel_proc.mel_to_spectrogram(&mel_spec).unwrap();
|
|
let linear_dims = linear_spec.dims();
|
|
assert_eq!(linear_dims[0], config.n_fft / 2 + 1);
|
|
assert_eq!(linear_dims[1], n_frames);
|
|
}
|
|
|
|
#[test]
|
|
fn test_mel_round_trip() {
|
|
let device = Device::cpu();
|
|
let config = MelConfig {
|
|
sample_rate: 16000,
|
|
n_fft: 512,
|
|
hop_length: 128,
|
|
win_length: 512,
|
|
n_mels: 40,
|
|
f_min: 0.0,
|
|
f_max: None,
|
|
};
|
|
let mel_proc = MelSpectrogram::new(config.clone(), &device).unwrap();
|
|
|
|
let n_freqs = config.n_fft / 2 + 1;
|
|
let n_frames = 50;
|
|
let original_linear = Tensor::randn(&[n_freqs, n_frames], &device).unwrap();
|
|
|
|
// Forward: linear -> mel
|
|
let mel_spec = mel_proc.spectrogram_to_mel(&original_linear).unwrap();
|
|
|
|
// Backward: mel -> linear (approximate)
|
|
let reconstructed_linear = mel_proc.mel_to_spectrogram(&mel_spec).unwrap();
|
|
|
|
// Check shapes match
|
|
assert_eq!(original_linear.dims(), reconstructed_linear.dims());
|
|
}
|
|
|
|
#[test]
|
|
fn test_audio_to_mel() {
|
|
let device = Device::cpu();
|
|
let config = MelConfig::default();
|
|
let mel_proc = MelSpectrogram::new(config.clone(), &device).unwrap();
|
|
|
|
// Create a dummy audio signal (1 second)
|
|
let audio_len = config.sample_rate;
|
|
let audio = Tensor::randn(&[audio_len], &device).unwrap();
|
|
|
|
// Convert to mel
|
|
let mel_spec = mel_proc.audio_to_mel(&audio).unwrap();
|
|
let mel_dims = mel_spec.dims();
|
|
|
|
assert_eq!(mel_dims[0], config.n_mels);
|
|
assert!(mel_dims[1] > 0); // Should have time frames
|
|
}
|
|
|
|
#[test]
|
|
fn test_filterbank_properties() {
|
|
let device = Device::cpu();
|
|
let config = MelConfig::default();
|
|
let mel_proc = MelSpectrogram::new(config.clone(), &device).unwrap();
|
|
|
|
let fb = mel_proc.filterbank();
|
|
let fb_data = fb.to_cpu().unwrap();
|
|
|
|
// Check all values are non-negative
|
|
for &val in &fb_data {
|
|
assert!(val >= 0.0, "Filterbank values must be non-negative");
|
|
}
|
|
|
|
// Check each mel filter sums to approximately 1 (due to normalization in creation)
|
|
let n_freqs = config.n_fft / 2 + 1;
|
|
for mel_idx in 0..config.n_mels {
|
|
let mut sum = 0.0;
|
|
for freq_idx in 0..n_freqs {
|
|
sum += fb_data[mel_idx * n_freqs + freq_idx];
|
|
}
|
|
// Each filter should have non-zero energy
|
|
assert!(sum > 0.0, "Mel filter {} has zero energy", mel_idx);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_different_configs() {
|
|
let device = Device::cpu();
|
|
|
|
// Test with different configurations
|
|
let configs = vec![
|
|
MelConfig {
|
|
sample_rate: 16000,
|
|
n_fft: 512,
|
|
hop_length: 160,
|
|
win_length: 512,
|
|
n_mels: 40,
|
|
f_min: 0.0,
|
|
f_max: Some(8000.0),
|
|
},
|
|
MelConfig {
|
|
sample_rate: 22050,
|
|
n_fft: 1024,
|
|
hop_length: 256,
|
|
win_length: 1024,
|
|
n_mels: 80,
|
|
f_min: 80.0,
|
|
f_max: Some(7600.0),
|
|
},
|
|
MelConfig {
|
|
sample_rate: 44100,
|
|
n_fft: 2048,
|
|
hop_length: 512,
|
|
win_length: 2048,
|
|
n_mels: 128,
|
|
f_min: 20.0,
|
|
f_max: Some(20000.0),
|
|
},
|
|
];
|
|
|
|
for config in configs {
|
|
let mel_proc = MelSpectrogram::new(config.clone(), &device).unwrap();
|
|
let fb_dims = mel_proc.filterbank().dims();
|
|
assert_eq!(fb_dims[0], config.n_mels);
|
|
assert_eq!(fb_dims[1], config.n_fft / 2 + 1);
|
|
}
|
|
}
|
|
}
|