//! 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> { 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::() / test_accuracies.len() as f32; let mean_loss = test_losses.iter().sum::() / 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(()) }