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

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");
}
}