166 lines
6.5 KiB
Rust
166 lines
6.5 KiB
Rust
//! 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<dyn std::error::Error>> {
|
|
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");
|
|
}
|
|
} |