714 lines
23 KiB
Rust
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(())
|
|
}
|
|
}
|