//! 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, } 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 { 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 { 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 = mel_points.iter().map(|&m| Self::mel_to_hz(m)).collect(); // Convert Hz to FFT bin indices let bin_points: Vec = 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 { // 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 { // 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 { 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); } } }