//! Griffin-Lim Algorithm for phase reconstruction //! //! This module implements the Griffin-Lim algorithm, an iterative method //! for reconstructing audio waveforms from magnitude spectrograms. //! //! The algorithm is particularly useful as a fast, non-neural baseline vocoder. use super::{MelConfig, MelSpectrogram, Vocoder}; use crate::error::{Result, TtsError}; use rtx_tensor::signal::{WindowType, istft, stft}; use rtx_tensor::{ComplexTensor, Device, Tensor}; use serde::{Deserialize, Serialize}; /// Configuration for Griffin-Lim algorithm #[derive(Debug, Clone, Serialize, Deserialize)] pub struct GriffinLimConfig { /// FFT size pub n_fft: usize, /// Hop length between frames pub hop_length: usize, /// Window length pub win_length: usize, /// Number of iterations pub n_iter: usize, /// Sample rate in Hz pub sample_rate: usize, /// Mel configuration for converting mel to linear spectrogram pub mel_config: MelConfig, } impl Default for GriffinLimConfig { fn default() -> Self { Self { n_fft: 1024, hop_length: 256, win_length: 1024, n_iter: 60, sample_rate: 22050, mel_config: MelConfig::default(), } } } impl GriffinLimConfig { /// Validate configuration pub fn validate(&self) -> Result<()> { 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_iter == 0 { return Err(TtsError::InvalidConfig( "Number of iterations must be > 0".to_string(), )); } if self.sample_rate == 0 { return Err(TtsError::InvalidConfig( "Sample rate must be > 0".to_string(), )); } self.mel_config.validate()?; Ok(()) } } /// Griffin-Lim vocoder implementation pub struct GriffinLim { config: GriffinLimConfig, mel_processor: MelSpectrogram, device: Device, } impl GriffinLim { /// Create a new Griffin-Lim vocoder pub fn new(config: GriffinLimConfig, device: &Device) -> Result { config.validate()?; let mel_processor = MelSpectrogram::new(config.mel_config.clone(), device)?; Ok(Self { config, mel_processor, device: device.clone(), }) } /// Convert mel spectrogram to linear spectrogram fn mel_to_linear(&self, mel_spec: &Tensor) -> Result { self.mel_processor.mel_to_spectrogram(mel_spec) } /// Reconstruct phase using Griffin-Lim algorithm /// /// # Arguments /// * `magnitude` - Magnitude spectrogram [n_freqs, time_frames] /// /// # Returns /// Reconstructed audio waveform fn reconstruct_phase(&self, magnitude: &Tensor) -> Result { let dims = magnitude.dims(); if dims.len() != 2 { return Err(TtsError::InvalidInput( "Magnitude must be 2D [n_freqs, time_frames]".to_string(), )); } let n_freqs = dims[0]; let n_frames = dims[1]; let expected_n_freqs = self.config.n_fft / 2 + 1; if n_freqs != expected_n_freqs { return Err(TtsError::InvalidInput(format!( "Expected {} frequency bins, got {}", expected_n_freqs, n_freqs ))); } // Initialize with random phase let phase = Tensor::randn(&[n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; // Convert to radians (0 to 2π) let phase_data = phase .to_cpu() .map_err(|e| TtsError::TensorError(e.to_string()))?; let phase_rad: Vec = phase_data .iter() .map(|&p| p * 2.0 * std::f32::consts::PI) .collect(); let mut phase_tensor = Tensor::from_data(phase_rad, vec![n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; // Get magnitude data let mag_data = magnitude .to_cpu() .map_err(|e| TtsError::TensorError(e.to_string()))?; // Griffin-Lim iterations for _iter in 0..self.config.n_iter { // Construct complex spectrum from magnitude and phase let phase_data = phase_tensor .to_cpu() .map_err(|e| TtsError::TensorError(e.to_string()))?; let mut real_data = Vec::with_capacity(n_freqs * n_frames); let mut imag_data = Vec::with_capacity(n_freqs * n_frames); for i in 0..(n_freqs * n_frames) { let mag = mag_data[i]; let phase = phase_data[i]; real_data.push(mag * phase.cos()); imag_data.push(mag * phase.sin()); } let real = Tensor::from_data(real_data, vec![n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; let imag = Tensor::from_data(imag_data, vec![n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; let complex_spec = ComplexTensor::from_real_imag(real, imag) .map_err(|e| TtsError::TensorError(e.to_string()))?; // ISTFT to get waveform let waveform = istft( &complex_spec, self.config.win_length, self.config.hop_length, Some(WindowType::Hann), None, ) .map_err(|e| TtsError::VocodingError(format!("ISTFT failed: {}", e)))?; // STFT to get new spectrum let stft_result = stft( &waveform, self.config.win_length, self.config.hop_length, Some(WindowType::Hann), true, ) .map_err(|e| TtsError::VocodingError(format!("STFT failed: {}", e)))?; // Extract phase from new spectrum let new_real = stft_result .stft .real() .to_cpu() .map_err(|e| TtsError::TensorError(e.to_string()))?; let new_imag = stft_result .stft .imag() .to_cpu() .map_err(|e| TtsError::TensorError(e.to_string()))?; let new_phase_data: Vec = new_real .iter() .zip(new_imag.iter()) .map(|(&r, &i)| i.atan2(r)) .collect(); phase_tensor = Tensor::from_data(new_phase_data, vec![n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; } // Final reconstruction let phase_data = phase_tensor .to_cpu() .map_err(|e| TtsError::TensorError(e.to_string()))?; let mut real_data = Vec::with_capacity(n_freqs * n_frames); let mut imag_data = Vec::with_capacity(n_freqs * n_frames); for i in 0..(n_freqs * n_frames) { let mag = mag_data[i]; let phase = phase_data[i]; real_data.push(mag * phase.cos()); imag_data.push(mag * phase.sin()); } let real = Tensor::from_data(real_data, vec![n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; let imag = Tensor::from_data(imag_data, vec![n_freqs, n_frames], &self.device) .map_err(|e| TtsError::TensorError(e.to_string()))?; let final_spec = ComplexTensor::from_real_imag(real, imag) .map_err(|e| TtsError::TensorError(e.to_string()))?; istft( &final_spec, self.config.win_length, self.config.hop_length, Some(WindowType::Hann), None, ) .map_err(|e| TtsError::VocodingError(format!("Final ISTFT failed: {}", e))) } /// Get configuration pub fn config(&self) -> &GriffinLimConfig { &self.config } } impl Vocoder for GriffinLim { fn synthesize(&self, mel: &Tensor) -> Result { // Convert mel to linear spectrogram let linear_spec = self.mel_to_linear(mel)?; // Reconstruct phase and synthesize audio self.reconstruct_phase(&linear_spec) } fn get_sample_rate(&self) -> usize { self.config.sample_rate } } #[cfg(test)] mod tests { use super::*; use approx::assert_relative_eq; #[test] fn test_griffin_lim_config_default() { let config = GriffinLimConfig::default(); assert_eq!(config.n_fft, 1024); assert_eq!(config.hop_length, 256); assert_eq!(config.win_length, 1024); assert_eq!(config.n_iter, 60); assert_eq!(config.sample_rate, 22050); } #[test] fn test_griffin_lim_config_validate_valid() { let config = GriffinLimConfig::default(); assert!(config.validate().is_ok()); } #[test] fn test_griffin_lim_config_validate_invalid_n_fft() { let mut config = GriffinLimConfig::default(); config.n_fft = 0; assert!(config.validate().is_err()); } #[test] fn test_griffin_lim_config_validate_invalid_hop_length() { let mut config = GriffinLimConfig::default(); config.hop_length = 0; assert!(config.validate().is_err()); } #[test] fn test_griffin_lim_config_validate_invalid_win_length() { let mut config = GriffinLimConfig::default(); config.win_length = 0; assert!(config.validate().is_err()); } #[test] fn test_griffin_lim_config_validate_win_length_too_large() { let mut config = GriffinLimConfig::default(); config.win_length = config.n_fft + 1; assert!(config.validate().is_err()); } #[test] fn test_griffin_lim_config_validate_invalid_n_iter() { let mut config = GriffinLimConfig::default(); config.n_iter = 0; assert!(config.validate().is_err()); } #[test] fn test_griffin_lim_config_validate_invalid_sample_rate() { let mut config = GriffinLimConfig::default(); config.sample_rate = 0; assert!(config.validate().is_err()); } #[test] fn test_griffin_lim_config_serialize() { let config = GriffinLimConfig::default(); let json = serde_json::to_string(&config).unwrap(); let deserialized: GriffinLimConfig = serde_json::from_str(&json).unwrap(); assert_eq!(config.n_fft, deserialized.n_fft); assert_eq!(config.n_iter, deserialized.n_iter); } #[test] fn test_griffin_lim_creation() { let device = Device::cpu(); let config = GriffinLimConfig::default(); let gl = GriffinLim::new(config.clone(), &device); assert!(gl.is_ok()); let gl = gl.unwrap(); assert_eq!(gl.get_sample_rate(), config.sample_rate); assert_eq!(gl.config().n_fft, config.n_fft); } #[test] fn test_griffin_lim_creation_invalid_config() { let device = Device::cpu(); let mut config = GriffinLimConfig::default(); config.n_fft = 0; let result = GriffinLim::new(config, &device); assert!(result.is_err()); } #[test] fn test_mel_to_linear() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 10, sample_rate: 16000, mel_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 gl = GriffinLim::new(config.clone(), &device).unwrap(); // Create a dummy mel spectrogram let n_frames = 50; let mel_spec = Tensor::randn(&[config.mel_config.n_mels, n_frames], &device).unwrap(); // Convert to linear let linear_spec = gl.mel_to_linear(&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_reconstruct_phase_shape() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 5, // Use fewer iterations for faster test sample_rate: 16000, mel_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 gl = GriffinLim::new(config.clone(), &device).unwrap(); // Create a magnitude spectrogram let n_freqs = config.n_fft / 2 + 1; let n_frames = 50; let magnitude = Tensor::randn(&[n_freqs, n_frames], &device).unwrap(); // Reconstruct phase let audio = gl.reconstruct_phase(&magnitude).unwrap(); let audio_dims = audio.dims(); assert_eq!(audio_dims.len(), 1); assert!(audio_dims[0] > 0); } #[test] fn test_reconstruct_phase_invalid_shape() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 5, sample_rate: 16000, mel_config: MelConfig::default(), }; let gl = GriffinLim::new(config, &device).unwrap(); // Create 1D tensor (invalid) let magnitude = Tensor::randn(&[100], &device).unwrap(); let result = gl.reconstruct_phase(&magnitude); assert!(result.is_err()); } #[test] fn test_reconstruct_phase_wrong_n_freqs() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 5, sample_rate: 16000, mel_config: MelConfig::default(), }; let gl = GriffinLim::new(config, &device).unwrap(); // Create magnitude with wrong number of frequency bins let magnitude = Tensor::randn(&[100, 50], &device).unwrap(); let result = gl.reconstruct_phase(&magnitude); assert!(result.is_err()); } #[test] fn test_synthesize() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 5, // Use fewer iterations for faster test sample_rate: 16000, mel_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 gl = GriffinLim::new(config.clone(), &device).unwrap(); // Create a dummy mel spectrogram let n_frames = 50; let mel_spec = Tensor::randn(&[config.mel_config.n_mels, n_frames], &device).unwrap(); // Synthesize audio let audio = gl.synthesize(&mel_spec).unwrap(); let audio_dims = audio.dims(); assert_eq!(audio_dims.len(), 1); assert!(audio_dims[0] > 0); } #[test] fn test_synthesize_invalid_mel_shape() { let device = Device::cpu(); let config = GriffinLimConfig::default(); let gl = GriffinLim::new(config, &device).unwrap(); // Create invalid mel spectrogram (wrong number of mels) let mel_spec = Tensor::randn(&[40, 50], &device).unwrap(); let result = gl.synthesize(&mel_spec); assert!(result.is_err()); } #[test] fn test_vocoder_trait_implementation() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 5, sample_rate: 16000, mel_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 gl = GriffinLim::new(config.clone(), &device).unwrap(); // Test Vocoder trait methods assert_eq!(gl.get_sample_rate(), 16000); let mel_spec = Tensor::randn(&[40, 50], &device).unwrap(); let audio = gl.synthesize(&mel_spec).unwrap(); assert!(audio.dims()[0] > 0); } #[test] fn test_different_iterations() { let device = Device::cpu(); let iterations = vec![1, 5, 10, 20]; let base_config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 1, // Will be overridden sample_rate: 16000, mel_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_spec = Tensor::randn(&[40, 30], &device).unwrap(); for n_iter in iterations { let mut config = base_config.clone(); config.n_iter = n_iter; let gl = GriffinLim::new(config, &device).unwrap(); let audio = gl.synthesize(&mel_spec).unwrap(); assert!(audio.dims()[0] > 0); } } #[test] fn test_different_sample_rates() { let device = Device::cpu(); let sample_rates = vec![8000, 16000, 22050, 44100]; for sample_rate in sample_rates { let config = GriffinLimConfig { n_fft: 512, hop_length: 128, win_length: 512, n_iter: 5, sample_rate, mel_config: MelConfig { sample_rate, n_fft: 512, hop_length: 128, win_length: 512, n_mels: 40, f_min: 0.0, f_max: None, }, }; let gl = GriffinLim::new(config.clone(), &device).unwrap(); assert_eq!(gl.get_sample_rate(), sample_rate); } } #[test] fn test_deterministic_reconstruction() { let device = Device::cpu(); let config = GriffinLimConfig { n_fft: 256, hop_length: 64, win_length: 256, n_iter: 3, sample_rate: 16000, mel_config: MelConfig { sample_rate: 16000, n_fft: 256, hop_length: 64, win_length: 256, n_mels: 20, f_min: 0.0, f_max: None, }, }; let gl = GriffinLim::new(config, &device).unwrap(); // Create a fixed mel spectrogram let mel_data = vec![0.5f32; 20 * 30]; let mel_spec = Tensor::from_data(mel_data, vec![20, 30], &device).unwrap(); // Synthesize twice - note: results may differ due to random phase initialization let audio1 = gl.synthesize(&mel_spec).unwrap(); let audio2 = gl.synthesize(&mel_spec).unwrap(); // Both should produce valid audio assert_eq!(audio1.dims().len(), 1); assert_eq!(audio2.dims().len(), 1); } }