Files
rustytorch/crates/models/rtx-tts/src/vocoder/mel.rs
T
osobhandClaude Sonnet 5 4aaa36a57a style: cargo fmt --workspace (whitespace/wrapping only, no semantic change)
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]>
2026-08-10 07:09:36 -07:00

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);
}
}
}