//! Standalone test for meta-learning implementation //! This validates the meta-learning code without dependency issues use rtx_transformers::meta::*; use rtx_tensor::{Device, Tensor}; fn main() -> Result<(), Box> { println!("Testing RTX Meta-Learning Implementation..."); let device = Device::cuda(0).unwrap_or(Device::default()); // Test 1: FewShotDataset creation println!("✓ Testing FewShotDataset creation..."); let dataset = FewShotDataset::synthetic(5, 10, 32, &device)?; assert_eq!(dataset.num_classes, 5); assert_eq!(dataset.feature_dim, 32); println!(" Created dataset with {} classes and {} feature dimensions", dataset.num_classes, dataset.feature_dim); // Test 2: Episode sampling println!("✓ Testing episode sampling..."); let episode = dataset.sample_episode(3, 2, 3, &device)?; assert_eq!(episode.n_way, 3); assert_eq!(episode.k_shot, 2); assert_eq!(episode.support_set.len(), 6); // 3 classes × 2 shots assert_eq!(episode.query_set.len(), 9); // 3 classes × 3 queries println!(" Generated {}-way {}-shot episode with {} support and {} query samples", episode.n_way, episode.k_shot, episode.support_set.len(), episode.query_set.len()); // Test 3: MAML configuration and creation println!("✓ Testing MAML implementation..."); let maml_config = MAMLConfig { inner_lr: 0.01, outer_lr: 0.001, inner_steps: 1, first_order: false, meta_batch_size: 4, }; let mut maml = MAML::new(32, 16, 3, maml_config, &device)?; println!(" Created MAML with input_dim={}, hidden_dim={}, output_dim={}", 32, 16, 3); // Test inner loop adaptation let meta_params = maml.meta_network.get_parameters(); let (adapted_params, inner_loss) = maml.inner_loop_adaptation( meta_params, &episode, &device )?; assert!(inner_loss >= 0.0); println!(" Inner loop adaptation completed with loss: {:.4}", inner_loss); // Test query evaluation let (outer_loss, accuracy) = maml.evaluate_query_set( adapted_params, &episode, &device )?; assert!(outer_loss >= 0.0); assert!(accuracy >= 0.0 && accuracy <= 1.0); println!(" Query evaluation: loss={:.4}, accuracy={:.4}", outer_loss, accuracy); // Test 4: Prototypical Networks println!("✓ Testing Prototypical Networks..."); let proto_config = PrototypicalConfig { learning_rate: 0.001, feature_dim: 16, distance_metric: DistanceMetric::Euclidean, temperature: 1.0, batch_size: 4, }; let proto_net = PrototypicalNetworks::new(32, proto_config, &device)?; println!(" Created Prototypical Networks with feature_dim={}", 16); // Test prototype computation let prototypes = proto_net.compute_prototypes(&episode, &device)?; assert_eq!(prototypes.num_classes, episode.n_way); assert_eq!(prototypes.feature_dim, 16); println!(" Computed {} prototypes with {} dimensions", prototypes.num_classes, prototypes.feature_dim); // Test classification let (logits, probabilities, proto_accuracy) = proto_net.classify(&episode, &device)?; assert_eq!(logits.shape()[0], episode.query_set.len()); assert_eq!(logits.shape()[1], episode.n_way); assert!(proto_accuracy >= 0.0 && proto_accuracy <= 1.0); println!(" Classification accuracy: {:.4}", proto_accuracy); // Test 5: Episode Sampler println!("✓ Testing Episode Sampler..."); let sampler_config = EpisodeSamplerConfig { n_way: 3, k_shot: 1, query_per_class: 2, num_episodes: 5, seed: Some(42), }; let mut sampler = EpisodeSampler::new(dataset, sampler_config); let episodes = sampler.sample_episodes(3, &device)?; assert_eq!(episodes.len(), 3); println!(" Generated {} episodes for meta-training", episodes.len()); // Test 6: Meta-learning Pipeline println!("✓ Testing Meta-learning Pipeline..."); let dataset = FewShotDataset::synthetic(4, 8, 24, &device)?; let pipeline_config = MetaLearningConfig { input_dim: 24, hidden_dim: 12, output_dim: 4, algorithm: MetaLearningAlgorithm::Prototypical, sampler_config: EpisodeSamplerConfig { n_way: 2, k_shot: 1, query_per_class: 2, num_episodes: 5, seed: Some(123), }, ..Default::default() }; let mut pipeline = MetaLearningPipeline::new(dataset, pipeline_config, &device)?; let results = pipeline.train(3, &device)?; assert!(results.proto_stats.is_some()); println!(" Pipeline training completed with {} episodes processed", results.proto_stats.unwrap().episodes_processed); // Test 7: Distance metrics println!("✓ Testing distance metrics..."); for metric in [DistanceMetric::Euclidean, DistanceMetric::Cosine, DistanceMetric::Manhattan] { let config = PrototypicalConfig { distance_metric: metric.clone(), feature_dim: 8, ..Default::default() }; let proto_net = PrototypicalNetworks::new(16, config, &device)?; let queries = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0], &[1, 8], device.clone())?; let prototypes = Tensor::from_vec(vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0], &[1, 8], device.clone())?; let distances = match metric { DistanceMetric::Euclidean => proto_net.compute_euclidean_distances(&queries, &prototypes)?, DistanceMetric::Cosine => proto_net.compute_cosine_distances(&queries, &prototypes)?, DistanceMetric::Manhattan => proto_net.compute_manhattan_distances(&queries, &prototypes)?, }; assert_eq!(distances.shape(), &[1, 1]); println!(" {:?} distance computed successfully", metric); } // Test 8: Evaluation metrics println!("✓ Testing evaluation metrics..."); let mut metrics = FewShotMetrics::new(); metrics.add_episode_result(0.8, 0.3, &[0, 1]); metrics.add_episode_result(0.9, 0.2, &[0, 1, 2]); metrics.add_episode_result(0.7, 0.4, &[1, 2]); assert_eq!(metrics.num_episodes, 3); let expected_mean_acc = (0.8 + 0.9 + 0.7) / 3.0; assert!((metrics.mean_accuracy - expected_mean_acc).abs() < 1e-6); println!(" Evaluation metrics: mean_accuracy={:.4}, num_episodes={}", metrics.mean_accuracy, metrics.num_episodes); println!("\n🎉 All meta-learning tests passed successfully!"); println!("✅ MAML (Model-Agnostic Meta-Learning) implemented"); println!("✅ Prototypical Networks implemented"); println!("✅ Episode sampling utilities implemented"); println!("✅ Multiple distance metrics supported"); println!("✅ Evaluation framework implemented"); println!("✅ Complete meta-learning pipeline implemented"); Ok(()) }