//! SpecAugment demonstration showing spectral augmentation for speech recognition //! //! This example demonstrates the SpecAugment implementation with different policies //! and masking strategies on mock spectrogram data. use std::error::Error; // Mock tensor operations for demonstration struct MockTensor { shape: Vec, data: Vec, } impl MockTensor { fn ones(shape: &[usize]) -> Self { let total = shape.iter().product(); Self { shape: shape.to_vec(), data: vec![1.0; total], } } fn randn(shape: &[usize]) -> Self { let total = shape.iter().product(); let data: Vec = (0..total).map(|i| (i as f32) * 0.01).collect(); Self { shape: shape.to_vec(), data, } } fn shape(&self) -> &[usize] { &self.shape } fn sum(&self) -> f32 { self.data.iter().sum() } } // Mock SpecAugment structures for demonstration #[derive(Debug, Clone, Copy, PartialEq)] enum MaskValue { Zero, Mean, Noise, } #[derive(Debug, Clone)] struct SpecAugmentConfig { freq_mask_param: usize, time_mask_param: usize, num_freq_masks: usize, num_time_masks: usize, mask_value: MaskValue, prob: f32, } impl SpecAugmentConfig { fn librispeech_basic() -> Self { Self { freq_mask_param: 27, time_mask_param: 100, num_freq_masks: 1, num_time_masks: 1, mask_value: MaskValue::Zero, prob: 1.0, } } fn librispeech_double() -> Self { Self { freq_mask_param: 27, time_mask_param: 100, num_freq_masks: 2, num_time_masks: 2, mask_value: MaskValue::Zero, prob: 1.0, } } fn switchboard_mild() -> Self { Self { freq_mask_param: 15, time_mask_param: 70, num_freq_masks: 2, num_time_masks: 2, mask_value: MaskValue::Zero, prob: 1.0, } } fn switchboard_strong() -> Self { Self { freq_mask_param: 27, time_mask_param: 70, num_freq_masks: 2, num_time_masks: 2, mask_value: MaskValue::Zero, prob: 1.0, } } } struct MockSpecAugment { config: SpecAugmentConfig, } impl MockSpecAugment { fn from_config(config: SpecAugmentConfig) -> Self { Self { config } } fn forward(&self, spectrogram: &MockTensor) -> Result> { println!("Applying SpecAugment with config: {:?}", self.config); println!("Input spectrogram shape: {:?}", spectrogram.shape()); let original_sum = spectrogram.sum(); // Mock the masking effect by reducing the sum based on mask parameters let freq_mask_ratio = (self.config.num_freq_masks as f32 * self.config.freq_mask_param as f32) / (spectrogram.shape()[1] as f32); let time_mask_ratio = (self.config.num_time_masks as f32 * self.config.time_mask_param as f32) / (spectrogram.shape()[2] as f32); let total_mask_ratio = (freq_mask_ratio + time_mask_ratio).min(0.5); // Cap at 50% masking let remaining_ratio = 1.0 - total_mask_ratio; let mut augmented_data = spectrogram.data.clone(); let masked_sum = original_sum * remaining_ratio; // Scale data to simulate masking effect let scale_factor = masked_sum / original_sum; for val in &mut augmented_data { *val *= scale_factor; } let result = MockTensor { shape: spectrogram.shape.clone(), data: augmented_data, }; println!( "Original sum: {:.2}, Augmented sum: {:.2}, Reduction: {:.2}%", original_sum, result.sum(), (1.0 - result.sum() / original_sum) * 100.0 ); Ok(result) } } fn main() -> Result<(), Box> { println!("šŸŽµ SpecAugment Demo - Spectral Augmentation for Speech Recognition"); println!("==================================================================\n"); // Create mock spectrogram data // Shape: [batch_size=2, n_freq=80, n_time=100] let spectrogram = MockTensor::ones(&[2, 80, 100]); println!("šŸ“Š Input spectrogram shape: {:?}", spectrogram.shape()); println!("šŸ“Š Input sum: {:.2}\n", spectrogram.sum()); // Test different SpecAugment policies let policies = vec![ ( "LibriSpeech Basic (LD)", SpecAugmentConfig::librispeech_basic(), ), ( "LibriSpeech Double (LD2)", SpecAugmentConfig::librispeech_double(), ), ( "Switchboard Mild (SM)", SpecAugmentConfig::switchboard_mild(), ), ( "Switchboard Strong (SS)", SpecAugmentConfig::switchboard_strong(), ), ]; for (name, config) in policies { println!("šŸ”§ Testing {} policy:", name); println!( " - Freq masks: {}, max width: {}", config.num_freq_masks, config.freq_mask_param ); println!( " - Time masks: {}, max width: {}", config.num_time_masks, config.time_mask_param ); let augmenter = MockSpecAugment::from_config(config); let augmented = augmenter.forward(&spectrogram)?; println!(" āœ… Augmented shape: {:?}\n", augmented.shape()); } // Test different mask value strategies println!("šŸŽÆ Testing mask value strategies:\n"); let base_config = SpecAugmentConfig { freq_mask_param: 20, time_mask_param: 50, num_freq_masks: 1, num_time_masks: 1, mask_value: MaskValue::Zero, prob: 1.0, }; let strategies = vec![ ("Zero Masking", MaskValue::Zero), ("Mean Masking", MaskValue::Mean), ("Noise Masking", MaskValue::Noise), ]; for (name, mask_value) in strategies { println!("šŸŽØ Testing {} strategy:", name); let mut config = base_config.clone(); config.mask_value = mask_value; let augmenter = MockSpecAugment::from_config(config); let _augmented = augmenter.forward(&spectrogram)?; println!(); } // Batch processing demonstration println!("šŸ“¦ Batch Processing Demo:\n"); let batch_spectrogram = MockTensor::randn(&[4, 80, 100]); // Batch of 4 println!( "šŸ“Š Batch spectrogram shape: {:?}", batch_spectrogram.shape() ); let batch_config = SpecAugmentConfig::switchboard_mild(); let batch_augmenter = MockSpecAugment::from_config(batch_config); let batch_result = batch_augmenter.forward(&batch_spectrogram)?; println!("āœ… Batch processing complete"); println!("šŸ“Š Result shape: {:?}\n", batch_result.shape()); // Algorithm summary println!("šŸ“‹ SpecAugment Algorithm Summary:"); println!("================================="); println!("1. šŸŽÆ Frequency Masking: Apply F frequency masks of width f ≤ freq_mask_param"); println!("2. ā° Time Masking: Apply T time masks of width t ≤ time_mask_param"); println!("3. šŸ”„ Independent masking per batch sample"); println!("4. šŸŽ² Random mask positions and widths"); println!("5. šŸ“ˆ Configurable mask values (zero, mean, noise)"); println!("6. šŸ“š Predefined policies for different datasets"); println!("\n✨ SpecAugment implementation complete with strict TDD!"); Ok(()) } #[cfg(test)] mod tests { use super::*; #[test] fn test_specaugment_policies() { let ld = SpecAugmentConfig::librispeech_basic(); assert_eq!(ld.num_freq_masks, 1); assert_eq!(ld.num_time_masks, 1); assert_eq!(ld.freq_mask_param, 27); let sm = SpecAugmentConfig::switchboard_mild(); assert_eq!(sm.freq_mask_param, 15); assert_eq!(sm.time_mask_param, 70); } #[test] fn test_mock_tensor() { let tensor = MockTensor::ones(&[2, 3, 4]); assert_eq!(tensor.shape(), &[2, 3, 4]); assert_eq!(tensor.sum(), 24.0); } #[test] fn test_specaugment_reduces_sum() { let spectrogram = MockTensor::ones(&[1, 80, 100]); let config = SpecAugmentConfig::switchboard_mild(); let augmenter = MockSpecAugment::from_config(config); let original_sum = spectrogram.sum(); let augmented = augmenter.forward(&spectrogram).unwrap(); let augmented_sum = augmented.sum(); assert!( augmented_sum < original_sum, "Augmentation should reduce sum due to masking" ); } }