Files
rustytorch/crates/training/rtx-transformers/examples/expert_dropout_demo.rs
T
2026-03-04 00:08:42 +00:00

304 lines
9.1 KiB
Rust

//! Expert Dropout Demo
//!
//! This example demonstrates the usage of expert dropout with MoE layers,
//! showcasing different dropout strategies and their effects on expert utilization.
use anyhow::Result;
use rtx_tensor::{DType, Device, Tensor};
use rtx_transformers::layers::{
DropoutScheduler, DropoutStatistics, DropoutStrategy, ExpertDropoutConfig, ExpertDropoutLayer,
ExpertImportanceScorer, ExpertOutputs, MoEConfig, Router,
};
fn main() -> Result<()> {
println!("🚀 Expert Dropout Demo");
println!("======================");
let device = Device::cuda(0).unwrap_or(Device::default());
let moe_config = MoEConfig::new(8, 2, 768, 3072);
let batch_size = 4;
let seq_len = 16;
// Demo 1: Random Dropout Strategy
println!("\n📊 Demo 1: Random Dropout Strategy");
demo_random_dropout(&moe_config, &device, batch_size, seq_len)?;
// Demo 2: Block Dropout Strategy
println!("\n🧱 Demo 2: Block Dropout Strategy");
demo_block_dropout(&moe_config, &device, batch_size, seq_len)?;
// Demo 3: Load-Aware Dropout Strategy
println!("\n⚖️ Demo 3: Load-Aware Dropout Strategy");
demo_load_aware_dropout(&moe_config, &device, batch_size, seq_len)?;
// Demo 4: Progressive Dropout Strategy
println!("\n📈 Demo 4: Progressive Dropout Strategy");
demo_progressive_dropout(&moe_config, &device, batch_size, seq_len)?;
// Demo 5: Dropout Scheduling
println!("\n⏰ Demo 5: Dropout Scheduling");
demo_dropout_scheduling(&moe_config, &device, batch_size, seq_len)?;
// Demo 6: Expert Importance Scoring
println!("\n⭐ Demo 6: Expert Importance Scoring");
demo_importance_scoring()?;
// Demo 7: MoE Integration
println!("\n🔗 Demo 7: MoE Integration");
demo_moe_integration(&moe_config, &device, batch_size, seq_len)?;
println!("\n✅ Expert Dropout Demo Complete!");
Ok(())
}
fn demo_random_dropout(
moe_config: &MoEConfig,
device: &Device,
batch_size: usize,
seq_len: usize,
) -> Result<()> {
let dropout_config = ExpertDropoutConfig::new(0.25, DropoutStrategy::Random);
let mut dropout_layer = ExpertDropoutLayer::new(dropout_config, moe_config.clone(), device)?;
dropout_layer.set_training(true);
let expert_outputs = create_mock_expert_outputs(
batch_size,
seq_len,
moe_config.num_experts,
moe_config.hidden_dim,
device,
);
// Run multiple iterations to show randomness
for i in 1..=3 {
let result = dropout_layer.forward(&expert_outputs)?;
let active_count = result.active_experts.iter().filter(|&&x| x).count();
println!(
" Iteration {}: {}/{} experts active",
i, active_count, moe_config.num_experts
);
println!(" Dropout rate: {:.2}", result.dropout_stats.dropout_rate);
}
Ok(())
}
fn demo_block_dropout(
moe_config: &MoEConfig,
device: &Device,
batch_size: usize,
seq_len: usize,
) -> Result<()> {
let mut dropout_config = ExpertDropoutConfig::new(0.25, DropoutStrategy::Block);
dropout_config.block_size = Some(2);
let mut dropout_layer = ExpertDropoutLayer::new(dropout_config, moe_config.clone(), device)?;
dropout_layer.set_training(true);
let expert_outputs = create_mock_expert_outputs(
batch_size,
seq_len,
moe_config.num_experts,
moe_config.hidden_dim,
device,
);
let result = dropout_layer.forward(&expert_outputs)?;
let active_count = result.active_experts.iter().filter(|&&x| x).count();
println!(
" Block dropout result: {}/{} experts active",
active_count, moe_config.num_experts
);
println!(" Active experts: {:?}", result.active_experts);
println!(" Dropout rate: {:.2}", result.dropout_stats.dropout_rate);
Ok(())
}
fn demo_load_aware_dropout(
moe_config: &MoEConfig,
device: &Device,
batch_size: usize,
seq_len: usize,
) -> Result<()> {
let dropout_config = ExpertDropoutConfig::new(0.3, DropoutStrategy::LoadAware);
let mut dropout_layer = ExpertDropoutLayer::new(dropout_config, moe_config.clone(), device)?;
dropout_layer.set_training(true);
// Simulate expert loads (some experts are underutilized)
let expert_loads = vec![0.9, 0.8, 0.1, 0.2, 0.7, 0.05, 0.15, 0.6];
dropout_layer.update_expert_loads(&expert_loads);
let expert_outputs = create_mock_expert_outputs(
batch_size,
seq_len,
moe_config.num_experts,
moe_config.hidden_dim,
device,
);
let result = dropout_layer.forward(&expert_outputs)?;
println!(" Expert loads: {:?}", expert_loads);
println!(" Active experts: {:?}", result.active_experts);
println!(" Experts 2, 5, 6 (low load) should be more likely to be dropped");
Ok(())
}
fn demo_progressive_dropout(
moe_config: &MoEConfig,
device: &Device,
batch_size: usize,
seq_len: usize,
) -> Result<()> {
let dropout_config = ExpertDropoutConfig::new(0.25, DropoutStrategy::Progressive);
let mut dropout_layer = ExpertDropoutLayer::new(dropout_config, moe_config.clone(), device)?;
dropout_layer.set_training(true);
let expert_outputs = create_mock_expert_outputs(
batch_size,
seq_len,
moe_config.num_experts,
moe_config.hidden_dim,
device,
);
// Show how dropout pattern changes with training step
for step in [100, 200, 300] {
dropout_layer.set_training_step(step);
let result = dropout_layer.forward(&expert_outputs)?;
let active_count = result.active_experts.iter().filter(|&&x| x).count();
println!(
" Step {}: {}/{} experts active, pattern: {:?}",
step, active_count, moe_config.num_experts, result.active_experts
);
}
Ok(())
}
fn demo_dropout_scheduling(
moe_config: &MoEConfig,
device: &Device,
batch_size: usize,
seq_len: usize,
) -> Result<()> {
let dropout_config = ExpertDropoutConfig::new(0.5, DropoutStrategy::Random);
let mut dropout_layer = ExpertDropoutLayer::new(dropout_config, moe_config.clone(), device)?;
dropout_layer.set_training(true);
// Create scheduler: start at 50% dropout, end at 10%
let scheduler = DropoutScheduler::new(0.5, 0.1, 1000);
dropout_layer.set_scheduler(scheduler.clone());
println!(" Scheduler: 50% → 10% dropout over 1000 steps");
for step in [0, 250, 500, 750, 1000] {
let rate = scheduler.get_dropout_rate(step);
println!(" Step {}: scheduled dropout rate = {:.2}", step, rate);
}
Ok(())
}
fn demo_importance_scoring() -> Result<()> {
let mut scorer = ExpertImportanceScorer::new(8);
println!(
" Initial importance scores: {:?}",
scorer.get_importance_scores()
);
// Simulate routing weights where expert 4 is most important
let routing_weights = vec![0.05, 0.1, 0.08, 0.12, 0.4, 0.03, 0.07, 0.15];
scorer.update_scores(&routing_weights);
println!(
" Updated importance scores: {:?}",
scorer.get_importance_scores()
);
let max_idx = scorer
.get_importance_scores()
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.total_cmp(b))
.unwrap()
.0;
println!(" Most important expert: #{} (expected: #4)", max_idx);
Ok(())
}
fn demo_moe_integration(
moe_config: &MoEConfig,
device: &Device,
batch_size: usize,
seq_len: usize,
) -> Result<()> {
let router = Router::new(moe_config.clone(), device)?;
let dropout_config = ExpertDropoutConfig::new(0.2, DropoutStrategy::Random);
let mut dropout_layer = ExpertDropoutLayer::new(dropout_config, moe_config.clone(), device)?;
dropout_layer.set_training(true);
// Create input tensor
let input = create_mock_input(batch_size, seq_len, moe_config.hidden_dim, device);
// Get routing information
let routing_info = router.route(&input)?;
println!(
" Original routing - {} experts available",
moe_config.num_experts
);
// Apply expert dropout to routing
let modified_routing = dropout_layer.apply_to_routing(&routing_info)?;
let active_count = modified_routing
.active_expert_mask
.iter()
.filter(|&&x| x)
.count();
println!(
" After dropout - {}/{} experts active",
active_count, moe_config.num_experts
);
println!(
" Active expert mask: {:?}",
modified_routing.active_expert_mask
);
Ok(())
}
// Helper functions
fn create_mock_expert_outputs(
batch_size: usize,
seq_len: usize,
num_experts: usize,
hidden_dim: usize,
device: &Device,
) -> ExpertOutputs {
let mut outputs = Vec::new();
for _ in 0..num_experts {
let expert_output = Tensor::randn(&[batch_size, seq_len, hidden_dim], DType::F32, device)
.expect("Failed to create tensor");
outputs.push(expert_output);
}
ExpertOutputs { outputs }
}
fn create_mock_input(
batch_size: usize,
seq_len: usize,
hidden_dim: usize,
device: &Device,
) -> Tensor {
Tensor::randn(&[batch_size, seq_len, hidden_dim], DType::F32, device)
.expect("Failed to create input tensor")
}