//! VICReg SSL Framework Integration Example //! //! Demonstrates how VICReg integrates with the unified SSL training framework, //! similar to other SSL methods like Barlow Twins, BYOL, and MoCo v3. use rtx_transformers::prelude::*; use rtx_transformers::ssl::*; fn main() -> Result<()> { println!("šŸ”„ VICReg SSL Framework Integration Demo"); println!("========================================"); let device = Device::cuda(0).unwrap_or(Device::default()); // Demo 1: VICReg through SSL Framework demo_ssl_framework_integration(&device)?; // Demo 2: Comparison with other SSL methods demo_ssl_method_comparison(&device)?; // Demo 3: VICReg-specific metrics demo_vicreg_metrics(&device)?; println!("āœ… All VICReg SSL integration demos completed!"); Ok(()) } /// Demo 1: VICReg integrated with the SSL training framework fn demo_ssl_framework_integration(device: &Device) -> Result<()> { println!("\nšŸ“š Demo 1: SSL Framework Integration"); println!("-----------------------------------"); // Create VICReg configuration let vicreg_config = VICRegConfig::new(2048, vec![8192, 8192, 8192]) .with_sim_coeff(25.0) .with_std_coeff(25.0) .with_cov_coeff(1.0) .with_variance_target(1.0); // Create SSL training configuration with VICReg let ssl_config = SSLTrainingConfig { method: SSLMethod::VICReg(vicreg_config), learning_rate: 1e-3, batch_size: 32, epochs: 100, ..Default::default() }; println!("āœ… SSL Configuration created with VICReg method"); println!(" - Learning rate: {}", ssl_config.learning_rate); println!(" - Batch size: {}", ssl_config.batch_size); println!(" - Epochs: {}", ssl_config.epochs); // In a real implementation, you would create the SSL trainer: // let backbone = VisionBackbone::resnet50(device); // let trainer = SSLTrainer::new(backbone, None, ssl_config, device)?; println!(" - VICReg integrated as SSL method āœ…"); Ok(()) } /// Demo 2: Compare VICReg configuration with other SSL methods fn demo_ssl_method_comparison(device: &Device) -> Result<()> { println!("\nšŸ“Š Demo 2: SSL Method Comparison"); println!("--------------------------------"); // VICReg configuration let vicreg_config = VICRegConfig::default(); let vicreg_method = SSLMethod::VICReg(vicreg_config); // Barlow Twins configuration (for comparison) let barlow_config = BarlowTwinsConfig::default(); let barlow_method = SSLMethod::BarlowTwins(barlow_config); println!("šŸ”¬ Method Characteristics:"); println!(" VICReg:"); println!(" - No momentum encoders needed"); println!(" - Explicit variance & covariance regularization"); println!(" - Three loss terms (invariance, variance, covariance)"); println!(" - Hyperparameters: Ī»=25, μ=25, ν=1"); println!(" Barlow Twins:"); println!(" - Cross-correlation matrix approach"); println!(" - Redundancy reduction principle"); println!(" - Two loss terms (invariance, redundancy)"); println!(" - Hyperparameter: Ī»=0.005"); match vicreg_method { SSLMethod::VICReg(config) => { println!( "āœ… VICReg method configured with {} expander layers", config.expander_dims.len() ); } _ => unreachable!(), } Ok(()) } /// Demo 3: VICReg-specific metrics tracking fn demo_vicreg_metrics(device: &Device) -> Result<()> { println!("\nšŸ“ˆ Demo 3: VICReg Metrics"); println!("------------------------"); // Create sample representations let batch_size = 64; let feature_dim = 8192; let y1 = Tensor::randn(vec![batch_size, feature_dim], DType::F32, device)?; let y2 = Tensor::randn(vec![batch_size, feature_dim], DType::F32, device)?; // Compute VICReg loss components let config = VICRegConfig::default(); let loss_result = compute_vicreg_loss(&y1, &y2, &config)?; // Create VICReg-specific metrics (as would be done in SSL trainer) let vicreg_metrics = MethodMetrics::VICReg { invariance_loss: loss_result.invariance_loss.to_scalar::()?, variance_loss: (loss_result.variance_loss_y1.to_scalar::()? + loss_result.variance_loss_y2.to_scalar::()?) / 2.0, covariance_loss: (loss_result.covariance_loss_y1.to_scalar::()? + loss_result.covariance_loss_y2.to_scalar::()?) / 2.0, }; // Display metrics as they would appear in training logs match vicreg_metrics { MethodMetrics::VICReg { invariance_loss, variance_loss, covariance_loss, } => { println!("šŸ“Š VICReg Training Metrics:"); println!(" - Invariance Loss: {:.6}", invariance_loss); println!(" - Variance Loss: {:.6}", variance_loss); println!(" - Covariance Loss: {:.6}", covariance_loss); println!( " - Total Loss: {:.6}", loss_result.total_loss.to_scalar::()? ); // Interpretation println!("\nšŸ” Metric Interpretation:"); println!(" - Invariance: Lower = more similar augmented views"); println!(" - Variance: Lower = better dimensional variance maintenance"); println!(" - Covariance: Lower = better feature decorrelation"); } _ => unreachable!(), } // Demonstrate loss component weighting println!("\nāš–ļø Loss Component Weighting:"); println!( " - Invariance Ɨ Ī»({}) = {:.6}", config.sim_coeff, loss_result.invariance_loss.to_scalar::()? * config.sim_coeff ); println!( " - Variance Ɨ μ({}) = {:.6}", config.std_coeff, (loss_result.variance_loss_y1.to_scalar::()? + loss_result.variance_loss_y2.to_scalar::()?) * config.std_coeff ); println!( " - Covariance Ɨ ν({}) = {:.6}", config.cov_coeff, (loss_result.covariance_loss_y1.to_scalar::()? + loss_result.covariance_loss_y2.to_scalar::()?) * config.cov_coeff ); Ok(()) } /// Example of how VICReg would be used in a complete training loop #[allow(dead_code)] fn example_training_loop() -> Result<()> { let device = Device::cuda(0).unwrap_or(Device::default()); // VICReg configuration let vicreg_config = VICRegConfig::new(2048, vec![8192, 8192, 8192]) .with_sim_coeff(25.0) .with_std_coeff(25.0) .with_cov_coeff(1.0); // SSL training configuration let ssl_config = SSLTrainingConfig { method: SSLMethod::VICReg(vicreg_config), learning_rate: 1e-3, batch_size: 256, // VICReg works well with moderate batch sizes epochs: 300, ..Default::default() }; // In practice: // let backbone = VisionBackbone::resnet50(&device); // let mut trainer = SSLTrainer::new(backbone, None, ssl_config, &device)?; // Training loop: // for epoch in 0..ssl_config.epochs { // // Load batch of images // let images = load_batch(ssl_config.batch_size)?; // // // Training step // let metrics = trainer.train_step(&images, Some(epoch))?; // // // Log VICReg-specific metrics // if let MethodMetrics::VICReg { invariance_loss, variance_loss, covariance_loss } = metrics.method_metrics { // log::info!("Epoch {}: Inv={:.4}, Var={:.4}, Cov={:.4}", // epoch, invariance_loss, variance_loss, covariance_loss); // } // } Ok(()) }