//! Standalone test for Reptile meta-learning implementation use crate::meta::{Reptile, ReptileConfig, InterpolationStrategy, ReptileTaskSampler}; use crate::meta::meta_learning_tests::FewShotDataset; use rtx_tensor::Device; fn main() -> Result<(), Box> { println!("Testing Reptile Meta-Learning Implementation"); let device = Device::cuda(0).unwrap_or(Device::default()); // Test 1: Basic configuration and creation println!("1. Testing Reptile configuration and creation..."); let config = ReptileConfig::default(); assert_eq!(config.inner_lr, 0.01); assert_eq!(config.meta_step_size, 0.1); assert_eq!(config.num_inner_steps, 5); assert_eq!(config.num_ways, 5); assert_eq!(config.num_shots, 1); assert!(config.validate().is_ok()); let reptile = Reptile::new(64, 32, 5, config.clone(), &device)?; println!("✓ Reptile created successfully"); // Test 2: Dataset and task sampling println!("2. Testing task sampling..."); let dataset = FewShotDataset::synthetic(10, 20, 64, &device)?; let sampler = ReptileTaskSampler::new(dataset, config)?; let task = sampler.sample_task(&device)?; assert_eq!(task.n_way, 5); assert_eq!(task.k_shot, 1); println!("✓ Task sampling works correctly"); // Test 3: SGD task adaptation println!("3. Testing SGD task adaptation..."); let (adapted_params, adaptation_loss) = reptile.sgd_task_adaptation(&task, &device)?; assert!(adaptation_loss >= 0.0); assert!(!adapted_params.is_empty()); println!("✓ SGD task adaptation completed with loss: {:.4}", adaptation_loss); // Test 4: Parameter interpolation strategies println!("4. Testing interpolation strategies..."); // Standard Reptile let mut reptile_standard = Reptile::new(32, 16, 3, ReptileConfig { interpolation_strategy: InterpolationStrategy::Standard, num_ways: 3, ..Default::default() }, &device)?; let dataset_small = FewShotDataset::synthetic(5, 10, 32, &device)?; let task_small = dataset_small.sample_episode(3, 1, 3, &device)?; let stats = reptile_standard.meta_update_single_task(&task_small, &device)?; assert_eq!(stats.meta_updates, 1); assert_eq!(stats.tasks_processed, 1); println!("✓ Standard Reptile interpolation works"); // Parallel Reptile let mut reptile_parallel = Reptile::new(32, 16, 3, ReptileConfig { interpolation_strategy: InterpolationStrategy::Parallel, num_tasks: 2, num_ways: 3, ..Default::default() }, &device)?; let tasks = vec![ dataset_small.sample_episode(3, 1, 3, &device)?, dataset_small.sample_episode(3, 1, 3, &device)?, ]; let parallel_stats = reptile_parallel.meta_update(tasks, &device)?; assert_eq!(parallel_stats.meta_updates, 1); assert_eq!(parallel_stats.tasks_processed, 2); println!("✓ Parallel Reptile interpolation works"); // Tail averaging let mut reptile_tail = Reptile::new(32, 16, 3, ReptileConfig { interpolation_strategy: InterpolationStrategy::TailAveraging { window_size: 3 }, num_ways: 3, ..Default::default() }, &device)?; let tail_stats = reptile_tail.meta_update_single_task(&task_small, &device)?; assert_eq!(tail_stats.meta_updates, 1); println!("✓ Tail averaging interpolation works"); // Test 5: Few-shot evaluation println!("5. Testing few-shot evaluation..."); let (accuracy, loss) = reptile_standard.evaluate_few_shot(&task_small, &device)?; assert!(accuracy >= 0.0 && accuracy <= 1.0); assert!(loss >= 0.0); println!("✓ Few-shot evaluation: accuracy={:.4}, loss={:.4}", accuracy, loss); // Test 6: Hyperparameter validation println!("6. Testing hyperparameter validation..."); let invalid_config = ReptileConfig { inner_lr: -0.1, ..Default::default() }; assert!(invalid_config.validate().is_err()); let invalid_config = ReptileConfig { meta_step_size: 0.0, ..Default::default() }; assert!(invalid_config.validate().is_err()); let invalid_config = ReptileConfig { num_inner_steps: 0, ..Default::default() }; assert!(invalid_config.validate().is_err()); println!("✓ Hyperparameter validation works correctly"); // Test 7: Different N-way K-shot scenarios println!("7. Testing N-way K-shot scenarios..."); // 3-way 1-shot let config_3w1s = ReptileConfig { num_ways: 3, num_shots: 1, ..Default::default() }; assert!(config_3w1s.validate().is_ok()); let reptile_3w1s = Reptile::new(24, 12, 3, config_3w1s, &device)?; println!("✓ 3-way 1-shot configuration works"); // 5-way 5-shot let config_5w5s = ReptileConfig { num_ways: 5, num_shots: 5, ..Default::default() }; assert!(config_5w5s.validate().is_ok()); let reptile_5w5s = Reptile::new(24, 12, 5, config_5w5s, &device)?; println!("✓ 5-way 5-shot configuration works"); // Test 8: Statistics tracking println!("8. Testing statistics tracking..."); let mut reptile_stats = Reptile::new(16, 8, 2, ReptileConfig::default(), &device)?; let initial_stats = reptile_stats.get_stats(); assert_eq!(initial_stats.meta_updates, 0); assert_eq!(initial_stats.tasks_processed, 0); assert_eq!(initial_stats.avg_adaptation_loss, 0.0); reptile_stats.reset_stats(); let reset_stats = reptile_stats.get_stats(); assert_eq!(reset_stats.meta_updates, 0); println!("✓ Statistics tracking works correctly"); // Test 9: Parameter change magnitude println!("9. Testing parameter change magnitude..."); let original_params = reptile.meta_network.get_parameters(); let adapted_params = original_params.clone(); let magnitude = reptile.compute_parameter_change_magnitude(&original_params, &adapted_params, &device)?; assert!(magnitude >= 0.0); println!("✓ Parameter change magnitude: {:.6}", magnitude); println!("\n🎉 All Reptile tests passed successfully!"); println!("✅ Implementation supports:"); println!(" - Standard Reptile (Serial)"); println!(" - Parallel Reptile (Batch)"); println!(" - Tail Averaging Reptile"); println!(" - N-way K-shot learning"); println!(" - SGD task adaptation"); println!(" - Parameter interpolation"); println!(" - Few-shot evaluation"); println!(" - Comprehensive validation"); Ok(()) } #[cfg(test)] mod tests { use super::*; #[test] fn test_reptile_standalone() { main().expect("Reptile standalone test failed"); } }