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

176 lines
6.2 KiB
Rust

//! FOMAML (First-Order Model-Agnostic Meta-Learning) Demonstration
//!
//! This example demonstrates the complete FOMAML pipeline for few-shot learning,
//! including task sampling, inner loop adaptation, and meta-parameter updates.
use rtx_tensor::Device;
use rtx_transformers::meta::*;
fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("🚀 FOMAML Demo - First-Order Model-Agnostic Meta-Learning");
println!("=========================================================");
// Setup device
let device = Device::cuda(0).unwrap_or(Device::default());
println!("📱 Using device: {:?}", device);
// FOMAML Configuration
let fomaml_config = FOMAMLConfig {
inner_lr: 0.01, // Inner loop learning rate
outer_lr: 0.001, // Meta learning rate
num_inner_steps: 5, // Adaptation steps
num_tasks_per_batch: 4, // Tasks per meta-update
first_order: true, // First-order approximation (key feature of FOMAML)
num_ways: 5, // 5-way classification
num_shots: 1, // 1-shot learning
num_queries: 15, // Query set size
};
println!("⚙️ FOMAML Configuration:");
println!(" - Inner LR: {}", fomaml_config.inner_lr);
println!(" - Outer LR: {}", fomaml_config.outer_lr);
println!(" - Inner steps: {}", fomaml_config.num_inner_steps);
println!(
" - Tasks per batch: {}",
fomaml_config.num_tasks_per_batch
);
println!(" - First-order: {}", fomaml_config.first_order);
println!(
" - N-way: {}, K-shot: {}",
fomaml_config.num_ways, fomaml_config.num_shots
);
// Validate configuration
fomaml_config.validate()?;
println!("✅ Configuration validated successfully");
// Create synthetic few-shot dataset
let num_classes = 10;
let samples_per_class = 20;
let feature_dim = 64;
let dataset = FewShotDataset::synthetic(num_classes, samples_per_class, feature_dim, &device)?;
println!("📊 Created synthetic dataset:");
println!(" - Classes: {}", num_classes);
println!(" - Samples per class: {}", samples_per_class);
println!(" - Feature dimension: {}", feature_dim);
// Initialize FOMAML model
let hidden_dim = 32;
let output_dim = fomaml_config.num_ways;
let mut fomaml = FOMAML::new(
feature_dim,
hidden_dim,
output_dim,
fomaml_config.clone(),
&device,
)?;
println!("🧠 Initialized FOMAML model:");
println!(" - Input dim: {}", feature_dim);
println!(" - Hidden dim: {}", hidden_dim);
println!(" - Output dim: {}", output_dim);
// Create task sampler
let sampler_config = TaskSamplerConfig {
n_way: fomaml_config.num_ways,
k_shot: fomaml_config.num_shots,
query_per_class: fomaml_config.num_queries / fomaml_config.num_ways,
seed: Some(42),
};
let sampler = TaskSampler::new(dataset, sampler_config)?;
println!(
"🎯 Created task sampler for {}-way {}-shot learning",
fomaml_config.num_ways, fomaml_config.num_shots
);
// Training phase
println!("\n🏋️ Training Phase");
println!("================");
let num_meta_updates = 10;
for meta_update in 1..=num_meta_updates {
// Sample tasks for this meta-batch
let tasks = sampler.sample_fomaml_tasks(&fomaml_config, &device)?;
// Perform FOMAML meta-update
let stats = fomaml.meta_update(tasks, &device)?;
if meta_update % 2 == 0 {
println!(
"Meta-update {}/{}: Loss={:.4}, Accuracy={:.2}%",
meta_update,
num_meta_updates,
stats.avg_outer_loss,
stats.avg_query_accuracy * 100.0
);
}
}
let final_stats = fomaml.get_stats();
println!("📈 Final Training Statistics:");
println!(" - Meta-updates: {}", final_stats.meta_updates);
println!(" - Tasks processed: {}", final_stats.tasks_processed);
println!(" - Avg inner loss: {:.4}", final_stats.avg_inner_loss);
println!(" - Avg outer loss: {:.4}", final_stats.avg_outer_loss);
println!(
" - Avg accuracy: {:.2}%",
final_stats.avg_query_accuracy * 100.0
);
// Evaluation phase
println!("\n🎯 Evaluation Phase");
println!("==================");
let num_test_episodes = 5;
let mut test_accuracies = Vec::new();
let mut test_losses = Vec::new();
for episode in 1..=num_test_episodes {
// Sample a test task
let test_task = sampler.sample_task(&device)?;
// Fast adaptation and evaluation
let (accuracy, loss) = fomaml.evaluate_few_shot(&test_task, &device)?;
test_accuracies.push(accuracy);
test_losses.push(loss);
println!(
"Test episode {}: Accuracy={:.2}%, Loss={:.4}",
episode,
accuracy * 100.0,
loss
);
}
let mean_accuracy = test_accuracies.iter().sum::<f32>() / test_accuracies.len() as f32;
let mean_loss = test_losses.iter().sum::<f32>() / test_losses.len() as f32;
println!("\n📊 Test Results Summary:");
println!(" - Mean accuracy: {:.2}%", mean_accuracy * 100.0);
println!(" - Mean loss: {:.4}", mean_loss);
// Demonstrate key FOMAML features
println!("\n🔍 FOMAML Key Features Demonstrated:");
println!("====================================");
println!("✅ First-order approximation (computational efficiency)");
println!("✅ Task sampling for N-way K-shot learning");
println!("✅ Inner loop adaptation with gradient detachment");
println!("✅ Outer loop meta-parameter updates");
println!("✅ Few-shot learning evaluation");
println!("✅ Batch processing of multiple tasks");
// Performance comparison note
println!("\n⚡ Performance Benefits of FOMAML:");
println!("- Faster than full MAML (no second derivatives)");
println!("- More memory efficient");
println!("- Suitable for larger models and more adaptation steps");
println!("- Maintains competitive few-shot learning performance");
println!("\n🎉 FOMAML demo completed successfully!");
Ok(())
}