Files
rustytorch/crates/training/rtx-compress/tests/knowledge_distillation_tests.rs
T
2026-03-04 00:08:42 +00:00

714 lines
23 KiB
Rust

//! 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<Tensor> {
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<Vec<Tensor>> {
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<Tensor> {
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(())
}
}