//! VICReg Self-Supervised Learning Demo //! //! Demonstrates VICReg (Variance-Invariance-Covariance Regularization) training //! with complete TDD implementation following the existing patterns. use rtx_transformers::prelude::*; use rtx_transformers::ssl::*; /// Simple backbone for demo purposes struct SimpleBackbone { output_dim: usize, device: Device, } impl Backbone for SimpleBackbone { fn forward(&self, x: &Tensor) -> Result { let batch_size = x.shape()[0]; // Simulate backbone output (normally would be actual convolution/transformer) Tensor::randn(vec![batch_size, self.output_dim], DType::F32, &self.device).map_err(|e| { crate::error::TransformerError::Generic(format!("Tensor creation failed: {}", e)) }) } fn output_dim(&self) -> usize { self.output_dim } } fn main() -> Result<()> { println!("๐Ÿ”ฅ VICReg Self-Supervised Learning Demo"); println!("=========================================="); let device = Device::cuda(0).unwrap_or(Device::default()); // Configure VICReg with paper's hyperparameters let config = VICRegConfig::new(2048, vec![8192, 8192, 8192]) .with_sim_coeff(25.0) // ฮป - invariance loss weight .with_std_coeff(25.0) // ฮผ - variance loss weight .with_cov_coeff(1.0) // ฮฝ - covariance loss weight .with_variance_target(1.0) // ฮณ - target standard deviation .with_epsilon(1e-4); // ฮต - numerical stability println!("โœ… VICReg Configuration:"); println!(" - Backbone dim: {}", config.backbone_dim); println!(" - Expander dims: {:?}", config.expander_dims); println!(" - Sim coeff (ฮป): {}", config.sim_coeff); println!(" - Std coeff (ฮผ): {}", config.std_coeff); println!(" - Cov coeff (ฮฝ): {}", config.cov_coeff); println!(" - Variance target (ฮณ): {}", config.variance_target); // Create backbone let backbone = SimpleBackbone { output_dim: 2048, device: device.clone(), }; // Create VICReg trainer let mut trainer = VICRegTrainer::new(backbone, config, &device)?; println!("\n๐Ÿ“Š Training VICReg Model:"); println!("========================="); // Training loop for epoch in 0..5 { // Create batch of "images" (normally these would be real images) let batch_size = 32; let images = Tensor::randn(vec![batch_size, 3, 224, 224], DType::F32, &device)?; // Training step let result = trainer.train_step(&images)?; println!( "Epoch {}: Total Loss: {:.6}, Inv: {:.6}, Var: {:.6}, Cov: {:.6}", epoch + 1, result.total_loss, result.invariance_loss, result.variance_loss, result.covariance_loss ); } println!("\n๐Ÿงช Testing Individual Loss Components:"); println!("====================================="); // Test individual loss components let batch_size = 64; let feature_dim = 8192; // Create random representations let y1 = Tensor::randn(vec![batch_size, feature_dim], DType::F32, &device)?; let y2 = Tensor::randn(vec![batch_size, feature_dim], DType::F32, &device)?; // Test invariance loss let inv_loss = compute_invariance_loss(&y1, &y2)?; println!( "๐Ÿ“ˆ Invariance Loss (MSE): {:.6}", inv_loss.to_scalar::()? ); // Test variance loss let var_loss1 = compute_variance_loss(&y1, 1.0, 1e-4)?; let var_loss2 = compute_variance_loss(&y2, 1.0, 1e-4)?; println!("๐Ÿ“Š Variance Loss Y1: {:.6}", var_loss1.to_scalar::()?); println!("๐Ÿ“Š Variance Loss Y2: {:.6}", var_loss2.to_scalar::()?); // Test covariance loss let cov_loss1 = compute_covariance_loss(&y1, 1e-4)?; let cov_loss2 = compute_covariance_loss(&y2, 1e-4)?; println!( "๐Ÿ”— Covariance Loss Y1: {:.6}", cov_loss1.to_scalar::()? ); println!( "๐Ÿ”— Covariance Loss Y2: {:.6}", cov_loss2.to_scalar::()? ); // Test complete VICReg loss let vicreg_config = VICRegConfig::default(); let loss_result = compute_vicreg_loss(&y1, &y2, &vicreg_config)?; println!("\n๐ŸŽฏ Complete VICReg Loss Breakdown:"); println!( " - Total: {:.6}", loss_result.total_loss.to_scalar::()? ); println!( " - Invariance: {:.6}", loss_result.invariance_loss.to_scalar::()? ); println!( " - Variance Y1: {:.6}", loss_result.variance_loss_y1.to_scalar::()? ); println!( " - Variance Y2: {:.6}", loss_result.variance_loss_y2.to_scalar::()? ); println!( " - Covariance Y1: {:.6}", loss_result.covariance_loss_y1.to_scalar::()? ); println!( " - Covariance Y2: {:.6}", loss_result.covariance_loss_y2.to_scalar::()? ); // Test expander network println!("\n๐Ÿ”ง Testing Expander Network:"); println!("============================"); let input_dim = 2048; let expander_dims = vec![8192, 8192, 8192]; let expander = ExpanderNetwork::new(input_dim, expander_dims.clone(), &device)?; let input = Tensor::randn(vec![32, input_dim], DType::F32, &device)?; let output = expander.forward(&input)?; println!("โœ… Expander Network:"); println!(" - Input dim: {}", expander.input_dim()); println!(" - Output dim: {}", expander.output_dim()); println!(" - Layer dims: {:?}", expander.layer_dims()); println!(" - Input shape: {:?}", input.shape()); println!(" - Output shape: {:?}", output.shape()); // Test edge cases println!("\n๐Ÿงฎ Testing Edge Cases:"); println!("====================="); // Test identical inputs (should give zero invariance loss) let identical_loss = compute_invariance_loss(&y1, &y1)?; println!( "๐ŸŽฏ Identical inputs invariance loss: {:.10}", identical_loss.to_scalar::()? ); // Test different batch sizes for batch_size in [1, 8, 64, 128] { let y_test = Tensor::randn(vec![batch_size, 256], DType::F32, &device)?; let var_loss = compute_variance_loss(&y_test, 1.0, 1e-4)?; let cov_loss = compute_covariance_loss(&y_test, 1e-4)?; println!( "๐Ÿ“ Batch size {}: Var={:.6}, Cov={:.6}", batch_size, var_loss.to_scalar::()?, cov_loss.to_scalar::()? ); } println!("\n๐ŸŽ‰ VICReg Demo completed successfully!"); println!("====================================="); println!("๐Ÿ”ฌ TDD Implementation verified with:"); println!(" โœ… Configuration management"); println!(" โœ… Expander network architecture"); println!(" โœ… Three loss components (invariance, variance, covariance)"); println!(" โœ… Complete training pipeline"); println!(" โœ… Edge case handling"); println!(" โœ… Backbone integration"); Ok(()) }