//! Comprehensive tests for knowledge distillation framework //! //! Tests cover: //! - Response-based distillation (temperature scaling, soft targets) //! - Feature-based distillation (intermediate layer matching) //! - Attention transfer mechanisms (attention maps alignment) //! - Multi-teacher distillation #![cfg(feature = "disabled_tests")] //! - Progressive distillation (teacher → intermediate → student) //! - Task-specific distillation losses //! - Online distillation (peer teaching) use rtx_autograd::{backward, tensor_with_grad}; use rtx_compress::{ Result, distillation::{ AttentionTransferConfig, DistillationConfig, DistillationLoss, DistillationMethod, DistillationResult, FeatureMatchingConfig, KnowledgeDistiller, MultiTeacherConfig, ProgressiveDistillationConfig, StudentModel, TeacherModel, }, }; use rtx_tensor::{Device, Tensor}; use std::collections::HashMap; #[cfg(test)] mod knowledge_distillation_tests { use super::*; fn create_test_logits( batch_size: usize, num_classes: usize, temperature: f32, ) -> Result { let device = Device::try_default()?; let mut data = Vec::new(); for b in 0..batch_size { for c in 0..num_classes { // Create realistic logit distributions let base_logit = if c == (b % num_classes) { 5.0 / temperature // Correct class gets higher logit } else { ((c * 17 + b * 31) % 100) as f32 / 100.0 - 2.0 }; data.push(base_logit); } } Tensor::from_slice(&data, &[batch_size, num_classes], &device) } fn create_test_features(batch_size: usize, feature_dims: &[usize]) -> Result> { let device = Device::try_default()?; let mut features = Vec::new(); for (layer_idx, &dim) in feature_dims.iter().enumerate() { let mut data = Vec::new(); for b in 0..batch_size { for f in 0..dim { let val = ((b * 100 + f * 7 + layer_idx * 13) % 200) as f32 / 200.0 - 0.5; data.push(val); } } features.push(Tensor::from_slice(&data, &[batch_size, dim], &device)?); } Ok(features) } fn create_attention_maps( batch_size: usize, num_heads: usize, seq_len: usize, ) -> Result { let device = Device::try_default()?; let mut data = Vec::new(); for b in 0..batch_size { for h in 0..num_heads { for i in 0..seq_len { for j in 0..seq_len { // Create attention-like patterns (higher for diagonal and nearby positions) let distance = ((i as i32 - j as i32).abs() + 1) as f32; let attention = (1.0 / distance).exp() + 0.1; data.push(attention); } } } } Tensor::from_slice(&data, &[batch_size, num_heads, seq_len, seq_len], &device) } #[test] fn test_response_based_distillation() -> Result<()> { let config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 4.0, alpha: 0.7, // 70% distillation loss, 30% task loss }, DistillationLoss::KullbackLeibler, ); let distiller = KnowledgeDistiller::new(config)?; // Create teacher and student logits let teacher_logits = create_test_logits(32, 10, 4.0)?; let student_logits = create_test_logits(32, 10, 4.0)?; // Create ground truth labels let device = Device::try_default()?; let labels = Tensor::from_slice( &[ 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, ], &[32], &device, )?; let result = distiller.compute_distillation_loss( &student_logits, &teacher_logits, Some(&labels), None, // No features for response-based )?; // Verify loss components assert!(result.total_loss > 0.0); assert!(result.distillation_loss > 0.0); assert!(result.task_loss.unwrap_or(0.0) >= 0.0); // Distillation loss should dominate due to alpha=0.7 assert!(result.distillation_loss > result.task_loss.unwrap_or(0.0) * 0.5); Ok(()) } #[test] fn test_temperature_scaling_effect() -> Result<()> { let low_temp_config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 1.0, alpha: 1.0, // Pure distillation }, DistillationLoss::KullbackLeibler, ); let high_temp_config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 10.0, alpha: 1.0, }, DistillationLoss::KullbackLeibler, ); let low_temp_distiller = KnowledgeDistiller::new(low_temp_config)?; let high_temp_distiller = KnowledgeDistiller::new(high_temp_config)?; let teacher_logits = create_test_logits(16, 5, 1.0)?; let student_logits = create_test_logits(16, 5, 1.0)?; let low_temp_result = low_temp_distiller.compute_distillation_loss( &student_logits, &teacher_logits, None, None, )?; let high_temp_result = high_temp_distiller.compute_distillation_loss( &student_logits, &teacher_logits, None, None, )?; // Higher temperature should generally produce smoother gradients and different loss magnitudes assert_ne!( low_temp_result.distillation_loss, high_temp_result.distillation_loss ); Ok(()) } #[test] fn test_feature_based_distillation() -> Result<()> { let feature_config = FeatureMatchingConfig::new( vec![64, 128, 256], // Feature dimensions for each layer vec![0.3, 0.4, 0.3], // Weights for each layer "mse".to_string(), // MSE loss for feature matching ); let config = DistillationConfig::new( DistillationMethod::FeatureBased { feature_config: feature_config, intermediate_matching: true, }, DistillationLoss::MeanSquaredError, ); let distiller = KnowledgeDistiller::new(config)?; let teacher_features = create_test_features(16, &[64, 128, 256])?; let student_features = create_test_features(16, &[64, 128, 256])?; let result = distiller.compute_feature_distillation_loss(&student_features, &teacher_features)?; assert!(result.total_loss > 0.0); assert!(result.feature_losses.len() == 3); // Each layer should contribute to the loss for &loss in &result.feature_losses { assert!(loss > 0.0); } Ok(()) } #[test] fn test_attention_transfer() -> Result<()> { let attention_config = AttentionTransferConfig::new( 8, // Number of attention heads 64, // Sequence length true, // Use attention entropy regularization 0.5, // Attention loss weight ); let config = DistillationConfig::new( DistillationMethod::AttentionTransfer { attention_config: attention_config, transfer_heads_individually: true, }, DistillationLoss::AttentionTransfer, ); let distiller = KnowledgeDistiller::new(config)?; let teacher_attention = create_attention_maps(8, 8, 64)?; let student_attention = create_attention_maps(8, 8, 64)?; let result = distiller .compute_attention_distillation_loss(&student_attention, &teacher_attention)?; assert!(result.total_loss > 0.0); assert!(result.attention_alignment_loss > 0.0); if let Some(entropy_loss) = result.attention_entropy_loss { assert!(entropy_loss >= 0.0); } Ok(()) } #[test] fn test_multi_teacher_distillation() -> Result<()> { let teachers_config = MultiTeacherConfig::new( vec![0.4, 0.3, 0.3], // Teacher weights (should sum to 1.0) "weighted_average".to_string(), // Aggregation method true, // Use teacher agreement loss ); let config = DistillationConfig::new( DistillationMethod::MultiTeacher { teachers_config: teachers_config, }, DistillationLoss::KullbackLeibler, ); let distiller = KnowledgeDistiller::new(config)?; // Create multiple teacher outputs let teacher_logits_1 = create_test_logits(16, 10, 3.0)?; let teacher_logits_2 = create_test_logits(16, 10, 3.5)?; let teacher_logits_3 = create_test_logits(16, 10, 4.0)?; let teacher_outputs = vec![teacher_logits_1, teacher_logits_2, teacher_logits_3]; let student_logits = create_test_logits(16, 10, 3.0)?; let result = distiller.compute_multi_teacher_distillation_loss( &student_logits, &teacher_outputs, None, )?; assert!(result.total_loss > 0.0); assert!(result.distillation_loss > 0.0); if let Some(agreement_loss) = result.teacher_agreement_loss { assert!(agreement_loss >= 0.0); } Ok(()) } #[test] fn test_progressive_distillation() -> Result<()> { let progressive_config = ProgressiveDistillationConfig::new( vec![256, 128, 64], // Model sizes in progressive order 10, // Number of steps between stages 0.1, // Size penalty coefficient ); let config = DistillationConfig::new( DistillationMethod::Progressive { progressive_config: progressive_config, }, DistillationLoss::KullbackLeibler, ); let mut distiller = KnowledgeDistiller::new(config)?; // Simulate progressive distillation stages let large_teacher_logits = create_test_logits(8, 10, 4.0)?; let medium_student_logits = create_test_logits(8, 10, 4.0)?; let small_student_logits = create_test_logits(8, 10, 4.0)?; // Stage 1: Large → Medium distiller.set_progressive_stage(0)?; let stage1_result = distiller.compute_distillation_loss( &medium_student_logits, &large_teacher_logits, None, None, )?; // Stage 2: Medium → Small distiller.set_progressive_stage(1)?; let stage2_result = distiller.compute_distillation_loss( &small_student_logits, &medium_student_logits, // Medium becomes teacher None, None, )?; assert!(stage1_result.total_loss > 0.0); assert!(stage2_result.total_loss > 0.0); // Size penalties should be different for different stages assert_ne!( stage1_result.regularization_loss, stage2_result.regularization_loss ); Ok(()) } #[test] fn test_online_distillation() -> Result<()> { let config = DistillationConfig::new( DistillationMethod::Online { peer_learning_rate: 0.1, ensemble_temperature: 3.0, mutual_learning: true, }, DistillationLoss::KullbackLeibler, ); let distiller = KnowledgeDistiller::new(config)?; // Create multiple peer models let peer_outputs = vec![ create_test_logits(12, 5, 3.0)?, create_test_logits(12, 5, 3.0)?, create_test_logits(12, 5, 3.0)?, ]; let result = distiller.compute_online_distillation_loss(&peer_outputs)?; assert!(result.total_loss > 0.0); assert!(result.peer_agreement_losses.len() == 3); // Each peer should have agreement loss with ensemble for &loss in &result.peer_agreement_losses { assert!(loss >= 0.0); } Ok(()) } #[test] fn test_different_distillation_losses() -> Result<()> { let student_logits = create_test_logits(16, 8, 2.0)?; let teacher_logits = create_test_logits(16, 8, 2.0)?; // Test KL Divergence let kl_config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 4.0, alpha: 1.0, }, DistillationLoss::KullbackLeibler, ); let kl_distiller = KnowledgeDistiller::new(kl_config)?; let kl_result = kl_distiller.compute_distillation_loss(&student_logits, &teacher_logits, None, None)?; // Test MSE let mse_config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 4.0, alpha: 1.0, }, DistillationLoss::MeanSquaredError, ); let mse_distiller = KnowledgeDistiller::new(mse_config)?; let mse_result = mse_distiller.compute_distillation_loss( &student_logits, &teacher_logits, None, None, )?; // Test Cosine Similarity let cosine_config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 4.0, alpha: 1.0, }, DistillationLoss::CosineSimilarity, ); let cosine_distiller = KnowledgeDistiller::new(cosine_config)?; let cosine_result = cosine_distiller.compute_distillation_loss( &student_logits, &teacher_logits, None, None, )?; // All losses should be positive but different assert!(kl_result.distillation_loss > 0.0); assert!(mse_result.distillation_loss > 0.0); assert!(cosine_result.distillation_loss > 0.0); // Losses should be different assert_ne!(kl_result.distillation_loss, mse_result.distillation_loss); assert_ne!( mse_result.distillation_loss, cosine_result.distillation_loss ); Ok(()) } #[test] fn test_layer_wise_feature_alignment() -> Result<()> { let feature_config = FeatureMatchingConfig::new( vec![32, 64, 128, 256], // Different layer sizes vec![0.1, 0.2, 0.3, 0.4], // Increasing weights for deeper layers "cosine".to_string(), ); let config = DistillationConfig::new( DistillationMethod::FeatureBased { feature_config: feature_config, intermediate_matching: true, }, DistillationLoss::CosineSimilarity, ); let distiller = KnowledgeDistiller::new(config)?; let teacher_features = create_test_features(10, &[32, 64, 128, 256])?; let student_features = create_test_features(10, &[32, 64, 128, 256])?; let result = distiller.compute_feature_distillation_loss(&student_features, &teacher_features)?; assert_eq!(result.feature_losses.len(), 4); // Verify weighted combination (deeper layers should contribute more) let weighted_sum = result.feature_losses[0] * 0.1 + result.feature_losses[1] * 0.2 + result.feature_losses[2] * 0.3 + result.feature_losses[3] * 0.4; assert!((result.total_loss - weighted_sum).abs() < 0.01); Ok(()) } #[test] fn test_adaptive_temperature_scaling() -> Result<()> { let config = DistillationConfig::new_with_adaptive_temperature( DistillationMethod::ResponseBased { temperature: 5.0, // Initial temperature alpha: 0.8, }, DistillationLoss::KullbackLeibler, true, // Enable adaptive scaling 0.1, // Temperature adaptation rate ); let mut distiller = KnowledgeDistiller::new(config)?; let teacher_logits = create_test_logits(20, 6, 1.0)?; let student_logits = create_test_logits(20, 6, 1.0)?; // Compute loss multiple times to trigger adaptation let initial_result = distiller.compute_distillation_loss(&student_logits, &teacher_logits, None, None)?; // Simulate training progress distiller.update_temperature_adaptation(0.05)?; // Reduce temperature let adapted_result = distiller.compute_distillation_loss(&student_logits, &teacher_logits, None, None)?; // Loss should change with temperature adaptation assert_ne!( initial_result.distillation_loss, adapted_result.distillation_loss ); Ok(()) } #[test] fn test_distillation_with_regularization() -> Result<()> { let config = DistillationConfig::new_with_regularization( DistillationMethod::ResponseBased { temperature: 3.0, alpha: 0.7, }, DistillationLoss::KullbackLeibler, true, // Enable L2 regularization 0.01, // L2 coefficient true, // Enable sparsity regularization 0.001, // Sparsity coefficient ); let distiller = KnowledgeDistiller::new(config)?; let teacher_logits = create_test_logits(16, 12, 3.0)?; let student_logits = create_test_logits(16, 12, 3.0)?; // Create some student parameters for regularization let student_params = vec![ create_test_features(1, &[128])?[0].clone(), create_test_features(1, &[256])?[0].clone(), ]; let result = distiller.compute_distillation_loss_with_regularization( &student_logits, &teacher_logits, None, &student_params, )?; assert!(result.total_loss > 0.0); assert!(result.distillation_loss > 0.0); // Should have regularization components if let Some(reg_loss) = result.regularization_loss { assert!(reg_loss >= 0.0); } Ok(()) } #[test] fn test_cross_modal_distillation() -> Result<()> { // Test distillation between different modalities (e.g., vision → text) let config = DistillationConfig::new( DistillationMethod::CrossModal { projection_dim: 128, modality_alignment_loss: true, temperature: 2.0, }, DistillationLoss::CosineSimilarity, ); let distiller = KnowledgeDistiller::new(config)?; // Teacher: vision features (2D) let teacher_features = create_test_features(8, &[512])?; // Image features // Student: text features (different dimensionality) let student_features = create_test_features(8, &[256])?; // Text features let result = distiller .compute_cross_modal_distillation_loss(&student_features[0], &teacher_features[0])?; assert!(result.total_loss > 0.0); assert!(result.alignment_loss > 0.0); if let Some(projection_loss) = result.projection_loss { assert!(projection_loss >= 0.0); } Ok(()) } #[test] fn test_distillation_gradient_flow() -> Result<()> { let config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 4.0, alpha: 0.5, }, DistillationLoss::KullbackLeibler, ); let distiller = KnowledgeDistiller::new(config)?; // Create student logits with gradient tracking let student_logits = tensor_with_grad(create_test_logits(8, 4, 4.0)?); let teacher_logits = create_test_logits(8, 4, 4.0)?; let result = distiller.compute_distillation_loss(&student_logits, &teacher_logits, None, None)?; // Backward pass to check gradient flow backward(vec![result.total_loss_tensor], HashMap::new())?; // Verify that gradients were computed let student_grad = student_logits.grad(); assert!( student_grad.is_some(), "Student logits should have gradients" ); Ok(()) } #[test] fn test_batch_size_invariance() -> Result<()> { let config = DistillationConfig::new( DistillationMethod::ResponseBased { temperature: 3.0, alpha: 1.0, }, DistillationLoss::KullbackLeibler, ); let distiller = KnowledgeDistiller::new(config)?; // Test with different batch sizes let small_batch_teacher = create_test_logits(4, 5, 3.0)?; let small_batch_student = create_test_logits(4, 5, 3.0)?; let large_batch_teacher = create_test_logits(16, 5, 3.0)?; let large_batch_student = create_test_logits(16, 5, 3.0)?; let small_result = distiller.compute_distillation_loss( &small_batch_student, &small_batch_teacher, None, None, )?; let large_result = distiller.compute_distillation_loss( &large_batch_student, &large_batch_teacher, None, None, )?; // Loss should scale reasonably with batch size let small_per_sample = small_result.distillation_loss / 4.0; let large_per_sample = large_result.distillation_loss / 16.0; // Per-sample loss should be roughly similar assert!((small_per_sample - large_per_sample).abs() / small_per_sample < 0.5); Ok(()) } #[test] fn test_distillation_scheduling() -> Result<()> { let config = DistillationConfig::new_with_scheduling( DistillationMethod::ResponseBased { temperature: 5.0, alpha: 0.9, }, DistillationLoss::KullbackLeibler, "cosine".to_string(), // Cosine annealing schedule 1000, // Total training steps ); let mut distiller = KnowledgeDistiller::new(config)?; let teacher_logits = create_test_logits(12, 7, 5.0)?; let student_logits = create_test_logits(12, 7, 5.0)?; // Test at different training steps let steps = [0, 250, 500, 750, 999]; let mut alpha_values = Vec::new(); for &step in &steps { distiller.set_training_step(step)?; let result = distiller.compute_distillation_loss( &student_logits, &teacher_logits, None, None, )?; // Extract effective alpha (distillation weight) let effective_alpha = distiller.get_current_alpha(); alpha_values.push(effective_alpha); } // Cosine schedule should vary alpha over training let initial_alpha = alpha_values[0]; let final_alpha = alpha_values[alpha_values.len() - 1]; assert_ne!(initial_alpha, final_alpha); assert!(alpha_values.iter().all(|&a| a >= 0.0 && a <= 1.0)); Ok(()) } }