//! VICReg (Variance-Invariance-Covariance Regularization) Implementation //! //! Self-supervised learning method that explicitly avoids collapse through three loss terms: //! - Invariance: Similar representations for augmented views (MSE loss) //! - Variance: Maintains variance ≥ γ in each dimension (hinge loss) //! - Covariance: Decorrelates different dimensions (Frobenius norm of off-diagonal covariance) //! //! Based on "VICReg: Variance-Invariance-Covariance Regularization for Self-Supervised Learning" //! (Bardes et al., 2021). //! //! Key features: //! - No momentum encoders or large batches required //! - Explicit variance and covariance regularization prevents collapse //! - Expander network (projector) transforms representations //! - Support for different backbone architectures //! - Symmetric loss computation on both augmented views use crate::prelude::*; use super::byol::Backbone; use std::sync::Arc; use parking_lot::RwLock; /// VICReg configuration parameters #[derive(Debug, Clone)] pub struct VICRegConfig { /// Dimension of backbone output features pub backbone_dim: usize, /// Expander network layer dimensions pub expander_dims: Vec, /// λ - Coefficient for invariance loss (similarity between views) pub sim_coeff: f32, /// μ - Coefficient for variance loss (prevent collapse) pub std_coeff: f32, /// ν - Coefficient for covariance loss (decorrelate features) pub cov_coeff: f32, /// γ - Target standard deviation (variance regularization) pub variance_target: f32, /// ε - Epsilon for numerical stability pub epsilon: f32, } impl Default for VICRegConfig { fn default() -> Self { Self { backbone_dim: 2048, expander_dims: vec![8192, 8192, 8192], sim_coeff: 25.0, std_coeff: 25.0, cov_coeff: 1.0, variance_target: 1.0, epsilon: 1e-4, } } } impl VICRegConfig { /// Create new VICReg configuration pub fn new(backbone_dim: usize, expander_dims: Vec) -> Self { Self { backbone_dim, expander_dims, ..Default::default() } } /// Set backbone dimension pub fn with_backbone_dim(mut self, backbone_dim: usize) -> Self { self.backbone_dim = backbone_dim; self } /// Set invariance loss coefficient (λ) pub fn with_sim_coeff(mut self, sim_coeff: f32) -> Self { self.sim_coeff = sim_coeff; self } /// Set variance loss coefficient (μ) pub fn with_std_coeff(mut self, std_coeff: f32) -> Self { self.std_coeff = std_coeff; self } /// Set covariance loss coefficient (ν) pub fn with_cov_coeff(mut self, cov_coeff: f32) -> Self { self.cov_coeff = cov_coeff; self } /// Set variance target (γ) pub fn with_variance_target(mut self, variance_target: f32) -> Self { self.variance_target = variance_target; self } /// Set epsilon for numerical stability pub fn with_epsilon(mut self, epsilon: f32) -> Self { self.epsilon = epsilon; self } } /// Expander network that maps backbone features to embedding space /// Similar to projector but terminology from VICReg paper #[derive(Debug)] pub struct ExpanderNetwork { layers: Vec, input_dim: usize, layer_dims: Vec, device: Device, } #[derive(Debug)] struct LinearLayer { weight: Arc>, bias: Arc>, } impl ExpanderNetwork { /// Create new expander network pub fn new( input_dim: usize, layer_dims: Vec, device: &Device, ) -> Result { let mut layers = Vec::new(); let mut current_dim = input_dim; for &output_dim in layer_dims.iter() { // Initialize weights with Xavier/Glorot normal initialization let scale = (2.0 / (current_dim + output_dim) as f32).sqrt(); let weight = Tensor::randn(vec![current_dim, output_dim], DType::F32, device)? .mul(&Tensor::full(vec![current_dim, output_dim], scale, DType::F32, device)?)?; let bias = Tensor::zeros(vec![output_dim], device)?; layers.push(LinearLayer { weight: Arc::new(RwLock::new(weight)), bias: Arc::new(RwLock::new(bias)), }); current_dim = output_dim; } Ok(Self { layers, input_dim, layer_dims, device: device.clone(), }) } /// Forward pass through expander network pub fn forward(&self, input: &Tensor) -> Result { let mut x = input.clone(); for (i, layer) in self.layers.iter().enumerate() { // Linear transformation let weight = layer.weight.read(); let bias = layer.bias.read(); x = x.matmul(&*weight)?.add(&*bias)?; // ReLU activation (except for the last layer) if i < self.layers.len() - 1 { x = x.relu()?; } } Ok(x) } pub fn input_dim(&self) -> usize { self.input_dim } pub fn output_dim(&self) -> usize { *self.layer_dims.last().unwrap_or(&self.input_dim) } pub fn layer_dims(&self) -> &[usize] { &self.layer_dims } } /// Result of VICReg loss computation #[derive(Debug, Clone)] pub struct VICRegLossResult { pub total_loss: Tensor, pub invariance_loss: Tensor, pub variance_loss_y1: Tensor, pub variance_loss_y2: Tensor, pub covariance_loss_y1: Tensor, pub covariance_loss_y2: Tensor, } /// Training result from VICReg training step #[derive(Debug, Clone)] pub struct VICRegTrainingResult { pub total_loss: f32, pub invariance_loss: f32, pub variance_loss: f32, pub covariance_loss: f32, } /// VICReg trainer combining backbone and expander networks pub struct VICRegTrainer { backbone: B, expander: ExpanderNetwork, config: VICRegConfig, device: Device, } impl VICRegTrainer { /// Create new VICReg trainer pub fn new( backbone: B, config: VICRegConfig, device: &Device, ) -> Result { let backbone_dim = backbone.output_dim(); let mut config = config; config.backbone_dim = backbone_dim; let expander = ExpanderNetwork::new( backbone_dim, config.expander_dims.clone(), device, )?; Ok(Self { backbone, expander, config, device: device.clone(), }) } /// Forward pass through both networks pub fn forward(&self, x1: &Tensor, x2: &Tensor) -> Result { // Get representations from backbone let z1 = self.backbone.forward(x1)?; let z2 = self.backbone.forward(x2)?; // Transform through expander network let y1 = self.expander.forward(&z1)?; let y2 = self.expander.forward(&z2)?; // Compute VICReg loss compute_vicreg_loss(&y1, &y2, &self.config) } /// Training step with augmentation (simplified for test) pub fn train_step(&mut self, images: &Tensor) -> Result { // In practice, you would apply different augmentations here // For testing, we'll use the same image with small noise as "augmentation" let noise = Tensor::randn(images.shape().to_vec(), &self.device)? .mul(&Tensor::full(images.shape().to_vec(), 0.01, DType::F32, &self.device)?)?; let x2 = images.add(&noise)?; let loss_result = self.forward(images, &x2)?; // In real implementation, you would: // 1. Compute gradients // 2. Update parameters // 3. Zero gradients Ok(VICRegTrainingResult { total_loss: loss_result.total_loss.to_scalar::()?, invariance_loss: loss_result.invariance_loss.to_scalar::()?, variance_loss: (loss_result.variance_loss_y1.to_scalar::()? + loss_result.variance_loss_y2.to_scalar::()?) * 0.5, covariance_loss: (loss_result.covariance_loss_y1.to_scalar::()? + loss_result.covariance_loss_y2.to_scalar::()?) * 0.5, }) } pub fn config(&self) -> &VICRegConfig { &self.config } } /// Compute invariance loss: MSE between representations pub fn compute_invariance_loss(y1: &Tensor, y2: &Tensor) -> Result { let diff = y1.sub(y2)?; let squared_diff = diff.mul(&diff)?; // Mean over all dimensions let total_elements = squared_diff.shape().iter().product::() as f32; let sum = squared_diff.sum(None)?; sum.div(&Tensor::full(vec![], total_elements, DType::F32, squared_diff.device())?) } /// Compute variance loss: Hinge loss to maintain std ≥ γ pub fn compute_variance_loss(y: &Tensor, gamma: f32, epsilon: f32) -> Result { let batch_size = y.shape()[0] as f32; let feature_dim = y.shape()[1]; // Compute mean along batch dimension let mean = y.sum_axis(0)?.div(&Tensor::full(vec![feature_dim], batch_size, DType::F32, y.device())?)?; // Center the features: y - mean let centered = y.sub(&mean.unsqueeze(0)?)?; // Compute variance: E[(y - mean)²] let squared_diff = centered.mul(¢ered)?; let variance = squared_diff.sum_axis(0)?.div(&Tensor::full(vec![feature_dim], batch_size, DType::F32, y.device())?)?; // Compute standard deviation with epsilon for stability let eps_tensor = Tensor::full(vec![feature_dim], epsilon, DType::F32, y.device())?; let variance_eps = variance.add(&eps_tensor)?; let std = variance_eps.sqrt()?; // Hinge loss: ReLU(γ - std) let gamma_tensor = Tensor::full(vec![feature_dim], gamma, DType::F32, y.device())?; let hinge = gamma_tensor.sub(&std)?; let hinge_loss = hinge.relu()?; // Mean over all features hinge_loss.mean_all() } /// Compute covariance loss: Frobenius norm of off-diagonal covariance matrix pub fn compute_covariance_loss(y: &Tensor, epsilon: f32) -> Result { let batch_size = y.shape()[0] as f32; let feature_dim = y.shape()[1]; // Center the features let mean = y.sum_axis(0)?.div(&Tensor::full(vec![feature_dim], batch_size, DType::F32, y.device())?)?; let centered = y.sub(&mean.unsqueeze(0)?)?; // Compute covariance matrix: (1/N) * X^T * X let covariance = centered.transpose(0, 1)?.matmul(¢ered)? .div(&Tensor::full(vec![feature_dim, feature_dim], batch_size, DType::F32, y.device())?)?; // Add epsilon to diagonal for numerical stability let eps_tensor = Tensor::full(vec![feature_dim], epsilon, DType::F32, y.device())?; let eye = create_identity_matrix(feature_dim, y.device())?; let eps_eye = eye.mul(&eps_tensor.unsqueeze(1)?)?; let stable_cov = covariance.add(&eps_eye)?; // Zero out diagonal elements (we only want off-diagonal) let off_diagonal_mask = create_off_diagonal_mask(feature_dim, y.device())?; let off_diagonal_cov = stable_cov.mul(&off_diagonal_mask)?; // Compute Frobenius norm of off-diagonal elements let squared = off_diagonal_cov.mul(&off_diagonal_cov)?; squared.sum(None) } /// Compute complete VICReg loss with all three terms pub fn compute_vicreg_loss( y1: &Tensor, y2: &Tensor, config: &VICRegConfig ) -> Result { // Invariance loss: similarity between representations let invariance_loss = compute_invariance_loss(y1, y2)?; // Variance losses: prevent dimensional collapse let variance_loss_y1 = compute_variance_loss(y1, config.variance_target, config.epsilon)?; let variance_loss_y2 = compute_variance_loss(y2, config.variance_target, config.epsilon)?; // Covariance losses: decorrelate features let covariance_loss_y1 = compute_covariance_loss(y1, config.epsilon)?; let covariance_loss_y2 = compute_covariance_loss(y2, config.epsilon)?; // Weighted total loss let sim_term = invariance_loss.mul(&Tensor::full(vec![], config.sim_coeff, DType::F32, y1.device())?)?; let std_term1 = variance_loss_y1.mul(&Tensor::full(vec![], config.std_coeff, DType::F32, y1.device())?)?; let std_term2 = variance_loss_y2.mul(&Tensor::full(vec![], config.std_coeff, DType::F32, y1.device())?)?; let cov_term1 = covariance_loss_y1.mul(&Tensor::full(vec![], config.cov_coeff, DType::F32, y1.device())?)?; let cov_term2 = covariance_loss_y2.mul(&Tensor::full(vec![], config.cov_coeff, DType::F32, y1.device())?)?; let total_loss = sim_term .add(&std_term1)? .add(&std_term2)? .add(&cov_term1)? .add(&cov_term2)?; Ok(VICRegLossResult { total_loss, invariance_loss, variance_loss_y1, variance_loss_y2, covariance_loss_y1, covariance_loss_y2, }) } /// Helper function to create identity matrix fn create_identity_matrix(size: usize, device: &Device) -> Result { let mut data = vec![0.0f32; size * size]; for i in 0..size { data[i * size + i] = 1.0; } Tensor::from_slice(&data, vec![size, size], DType::F32, device) } /// Helper function to create off-diagonal mask (1s off-diagonal, 0s on diagonal) fn create_off_diagonal_mask(size: usize, device: &Device) -> Result { let mut data = vec![1.0f32; size * size]; for i in 0..size { data[i * size + i] = 0.0; } Tensor::from_slice(&data, vec![size, size], DType::F32, device) }