176 lines
6.2 KiB
Rust
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(())
|
|
}
|