//! Episode Sampling Utilities for Meta-Learning //! //! Provides utilities for: //! - Task distribution sampling //! - Episode generation for N-way K-shot learning //! - Few-shot dataset handling //! - Evaluation metrics for meta-learning use super::meta_learning_tests::{FewShotDataset, Episode}; use rtx_tensor::{Tensor, Device}; use std::collections::HashMap; use crate::{TransformerError, Result}; /// Configuration for episode sampling #[derive(Debug, Clone)] pub struct EpisodeSamplerConfig { /// Number of classes per episode (N-way) pub n_way: usize, /// Number of support examples per class (K-shot) pub k_shot: usize, /// Number of query examples per class pub query_per_class: usize, /// Total number of episodes to sample pub num_episodes: usize, /// Random seed for reproducible sampling pub seed: Option, } impl Default for EpisodeSamplerConfig { fn default() -> Self { Self { n_way: 5, k_shot: 1, query_per_class: 5, num_episodes: 100, seed: Some(42), } } } /// Episode sampler for generating few-shot learning tasks #[derive(Debug)] pub struct EpisodeSampler { /// Dataset to sample from pub dataset: FewShotDataset, /// Sampling configuration pub config: EpisodeSamplerConfig, /// Current sampling state pub state: SamplingState, } /// Internal sampling state #[derive(Debug, Clone, Default)] pub struct SamplingState { /// Number of episodes generated so far pub episodes_generated: usize, /// Random number generator state (simplified) pub rng_state: u64, } impl EpisodeSampler { /// Create new episode sampler pub fn new(dataset: FewShotDataset, config: EpisodeSamplerConfig) -> Self { let initial_rng_state = config.seed.unwrap_or(12345); Self { dataset, config, state: SamplingState { episodes_generated: 0, rng_state: initial_rng_state, }, } } /// Sample a single episode pub fn sample_episode(&mut self, device: &Device) -> Result { if self.config.n_way > self.dataset.num_classes { return Err(TransformerError::InvalidInput( format!("N-way ({}) cannot exceed number of classes ({})", self.config.n_way, self.dataset.num_classes) )); } // Check if we have enough samples per class let min_samples_needed = self.config.k_shot + self.config.query_per_class; for samples in self.dataset.classes.values() { if samples.len() < min_samples_needed { return Err(TransformerError::InvalidInput( format!("Need at least {} samples per class, found {}", min_samples_needed, samples.len()) )); } } // Select N classes (deterministic based on RNG state for reproducibility) let selected_classes = self.select_classes()?; let mut support_set = Vec::new(); let mut support_labels = Vec::new(); let mut query_set = Vec::new(); let mut query_labels = Vec::new(); for (new_label, &original_class) in selected_classes.iter().enumerate() { let class_samples = &self.dataset.classes[&original_class]; // Sample indices for this class (deterministic) let indices = self.sample_class_indices(class_samples.len(), min_samples_needed)?; // Support set for i in 0..self.config.k_shot { let sample_idx = indices[i]; support_set.push(class_samples[sample_idx].clone()); support_labels.push(new_label); } // Query set for i in self.config.k_shot..min_samples_needed { let sample_idx = indices[i]; query_set.push(class_samples[sample_idx].clone()); query_labels.push(new_label); } } self.state.episodes_generated += 1; Ok(Episode { support_set, support_labels, query_set, query_labels, n_way: self.config.n_way, k_shot: self.config.k_shot, }) } /// Select N classes for this episode fn select_classes(&mut self) -> Result> { let mut selected = Vec::new(); let available_classes: Vec = self.dataset.classes.keys().cloned().collect(); // Simple deterministic selection based on RNG state for i in 0..self.config.n_way { let class_idx = (self.state.rng_state + i as u64) as usize % available_classes.len(); let mut class_id = available_classes[class_idx]; // Avoid duplicates while selected.contains(&class_id) { class_id = (class_id + 1) % self.dataset.num_classes; } selected.push(class_id); } // Update RNG state self.state.rng_state = self.state.rng_state.wrapping_mul(1103515245).wrapping_add(12345); Ok(selected) } /// Sample indices for a given class fn sample_class_indices(&mut self, class_size: usize, num_needed: usize) -> Result> { if num_needed > class_size { return Err(TransformerError::InvalidInput( format!("Cannot sample {} indices from class of size {}", num_needed, class_size) )); } // Simple deterministic sampling without replacement let mut indices = Vec::new(); let mut used = std::collections::HashSet::new(); for i in 0..num_needed { let mut idx = (self.state.rng_state + i as u64) as usize % class_size; // Avoid duplicates while used.contains(&idx) { idx = (idx + 1) % class_size; } indices.push(idx); used.insert(idx); } // Update RNG state self.state.rng_state = self.state.rng_state.wrapping_mul(1103515245).wrapping_add(12345); Ok(indices) } /// Sample a batch of episodes pub fn sample_episodes(&mut self, batch_size: usize, device: &Device) -> Result> { let mut episodes = Vec::new(); for _ in 0..batch_size { let episode = self.sample_episode(device)?; episodes.push(episode); } Ok(episodes) } /// Sample episodes for meta-training pub fn sample_meta_train_episodes(&mut self, device: &Device) -> Result> { self.sample_episodes(self.config.num_episodes, device) } /// Reset sampling state pub fn reset(&mut self) { self.state.episodes_generated = 0; self.state.rng_state = self.config.seed.unwrap_or(12345); } /// Get sampling statistics pub fn get_stats(&self) -> SamplingStats { SamplingStats { episodes_generated: self.state.episodes_generated, n_way: self.config.n_way, k_shot: self.config.k_shot, query_per_class: self.config.query_per_class, total_classes: self.dataset.num_classes, feature_dim: self.dataset.feature_dim, } } } /// Statistics for episode sampling #[derive(Debug, Clone)] pub struct SamplingStats { /// Number of episodes generated pub episodes_generated: usize, /// Number of classes per episode pub n_way: usize, /// Number of support examples per class pub k_shot: usize, /// Number of query examples per class pub query_per_class: usize, /// Total number of classes in dataset pub total_classes: usize, /// Feature dimension pub feature_dim: usize, } /// Task distribution for meta-learning #[derive(Debug)] pub struct TaskDistribution { /// Different episode configurations pub task_configs: Vec, /// Weights for sampling different task types pub task_weights: Vec, /// Current task type pub current_task_type: usize, } impl TaskDistribution { /// Create task distribution with multiple episode types pub fn multi_task(base_config: EpisodeSamplerConfig) -> Self { let mut task_configs = Vec::new(); // Different N-way configurations for n_way in [2, 3, 5] { let mut config = base_config.clone(); config.n_way = n_way; task_configs.push(config); } // Different K-shot configurations for k_shot in [1, 2, 5] { let mut config = base_config.clone(); config.k_shot = k_shot; task_configs.push(config); } let num_tasks = task_configs.len(); let task_weights = vec![1.0 / num_tasks as f32; num_tasks]; Self { task_configs, task_weights, current_task_type: 0, } } /// Sample next task configuration pub fn sample_task_config(&mut self) -> &EpisodeSamplerConfig { // Simple round-robin for now let config = &self.task_configs[self.current_task_type]; self.current_task_type = (self.current_task_type + 1) % self.task_configs.len(); config } } /// Evaluation metrics for few-shot learning #[derive(Debug, Clone, Default)] pub struct FewShotMetrics { /// Total number of episodes evaluated pub num_episodes: usize, /// Mean accuracy across episodes pub mean_accuracy: f32, /// Standard deviation of accuracy pub std_accuracy: f32, /// Confidence interval (95%) pub confidence_interval: (f32, f32), /// Per-class accuracies pub per_class_accuracy: HashMap, /// Mean loss across episodes pub mean_loss: f32, } impl FewShotMetrics { /// Create new metrics tracker pub fn new() -> Self { Self::default() } /// Add episode result pub fn add_episode_result(&mut self, accuracy: f32, loss: f32, class_labels: &[usize]) { // Update running statistics let old_count = self.num_episodes as f32; let new_count = old_count + 1.0; // Update mean accuracy let old_mean_acc = self.mean_accuracy; self.mean_accuracy = (old_mean_acc * old_count + accuracy) / new_count; // Update mean loss let old_mean_loss = self.mean_loss; self.mean_loss = (old_mean_loss * old_count + loss) / new_count; // Update per-class accuracies (simplified) for &class_id in class_labels { let current_acc = self.per_class_accuracy.get(&class_id).unwrap_or(&0.0); self.per_class_accuracy.insert(class_id, (*current_acc + accuracy) / 2.0); } self.num_episodes += 1; // Update standard deviation (simplified) if self.num_episodes > 1 { self.std_accuracy = 0.1; // Placeholder - would compute properly in real implementation // 95% confidence interval let margin = 1.96 * self.std_accuracy / (self.num_episodes as f32).sqrt(); self.confidence_interval = ( self.mean_accuracy - margin, self.mean_accuracy + margin, ); } } /// Compute final statistics pub fn finalize(&mut self) -> FewShotMetrics { self.clone() } } /// Utilities for meta-learning evaluation pub struct MetaLearningEvaluator; impl MetaLearningEvaluator { /// Evaluate few-shot performance pub fn evaluate_few_shot_performance( episodes: &[Episode], predictions: &[Vec], losses: &[f32], ) -> Result { if episodes.len() != predictions.len() || episodes.len() != losses.len() { return Err(TransformerError::InvalidInput( "Episodes, predictions, and losses must have same length".into() )); } let mut metrics = FewShotMetrics::new(); for (i, episode) in episodes.iter().enumerate() { let pred = &predictions[i]; let loss = losses[i]; // Compute accuracy for this episode let correct = episode.query_labels .iter() .zip(pred.iter()) .filter(|(&true_label, &pred_label)| true_label == pred_label) .count(); let accuracy = correct as f32 / episode.query_labels.len() as f32; // Get unique class labels for this episode let mut class_labels: Vec = episode.query_labels.clone(); class_labels.sort(); class_labels.dedup(); metrics.add_episode_result(accuracy, loss, &class_labels); } Ok(metrics.finalize()) } /// Compare two meta-learning approaches pub fn compare_approaches( metrics_a: &FewShotMetrics, metrics_b: &FewShotMetrics, ) -> MetaLearningComparison { MetaLearningComparison { accuracy_difference: metrics_a.mean_accuracy - metrics_b.mean_accuracy, loss_difference: metrics_a.mean_loss - metrics_b.mean_loss, significance: if (metrics_a.mean_accuracy - metrics_b.mean_accuracy).abs() > 0.05 { "significant".to_string() } else { "not significant".to_string() }, better_approach: if metrics_a.mean_accuracy > metrics_b.mean_accuracy { "A".to_string() } else { "B".to_string() }, } } } /// Comparison between two meta-learning approaches #[derive(Debug, Clone)] pub struct MetaLearningComparison { /// Difference in accuracy (A - B) pub accuracy_difference: f32, /// Difference in loss (A - B) pub loss_difference: f32, /// Statistical significance pub significance: String, /// Which approach is better pub better_approach: String, } #[cfg(all(test, feature = "disabled_tests"))] mod tests { use super::*; #[test] fn test_episode_sampler_creation() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(5, 10, 8, &device).unwrap(); let config = EpisodeSamplerConfig::default(); let sampler = EpisodeSampler::new(dataset, config); assert_eq!(sampler.config.n_way, 5); assert_eq!(sampler.config.k_shot, 1); assert_eq!(sampler.state.episodes_generated, 0); } #[test] fn test_single_episode_sampling() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(5, 10, 8, &device).unwrap(); let config = EpisodeSamplerConfig { n_way: 3, k_shot: 2, query_per_class: 3, ..Default::default() }; let mut sampler = EpisodeSampler::new(dataset, config); let episode = sampler.sample_episode(&device).unwrap(); 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 assert_eq!(sampler.state.episodes_generated, 1); } #[test] fn test_batch_episode_sampling() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(4, 8, 6, &device).unwrap(); let config = EpisodeSamplerConfig { n_way: 2, k_shot: 1, query_per_class: 2, ..Default::default() }; let mut sampler = EpisodeSampler::new(dataset, config); let episodes = sampler.sample_episodes(3, &device).unwrap(); assert_eq!(episodes.len(), 3); assert_eq!(sampler.state.episodes_generated, 3); for episode in &episodes { assert_eq!(episode.n_way, 2); assert_eq!(episode.k_shot, 1); assert_eq!(episode.support_set.len(), 2); // 2 classes × 1 shot assert_eq!(episode.query_set.len(), 4); // 2 classes × 2 queries } } #[test] fn test_insufficient_samples_error() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(3, 2, 4, &device).unwrap(); // Only 2 samples per class let config = EpisodeSamplerConfig { k_shot: 1, query_per_class: 2, // Need 3 total, but only have 2 ..Default::default() }; let mut sampler = EpisodeSampler::new(dataset, config); let result = sampler.sample_episode(&device); assert!(result.is_err()); } #[test] fn test_too_many_classes_error() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(3, 5, 4, &device).unwrap(); // Only 3 classes let config = EpisodeSamplerConfig { n_way: 5, // Want 5 classes but only have 3 ..Default::default() }; let mut sampler = EpisodeSampler::new(dataset, config); let result = sampler.sample_episode(&device); assert!(result.is_err()); } #[test] fn test_reproducible_sampling() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(4, 6, 5, &device).unwrap(); let config = EpisodeSamplerConfig { n_way: 2, k_shot: 1, query_per_class: 2, seed: Some(12345), ..Default::default() }; // Sample with same seed let mut sampler1 = EpisodeSampler::new(dataset.clone(), config.clone()); let mut sampler2 = EpisodeSampler::new(dataset, config); let episode1 = sampler1.sample_episode(&device).unwrap(); let episode2 = sampler2.sample_episode(&device).unwrap(); // Episodes should have same structure assert_eq!(episode1.support_labels, episode2.support_labels); assert_eq!(episode1.query_labels, episode2.query_labels); } #[test] fn test_sampling_stats() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(5, 8, 10, &device).unwrap(); let config = EpisodeSamplerConfig { n_way: 3, k_shot: 2, query_per_class: 1, ..Default::default() }; let mut sampler = EpisodeSampler::new(dataset, config); let stats_before = sampler.get_stats(); assert_eq!(stats_before.episodes_generated, 0); sampler.sample_episode(&device).unwrap(); sampler.sample_episode(&device).unwrap(); let stats_after = sampler.get_stats(); assert_eq!(stats_after.episodes_generated, 2); assert_eq!(stats_after.n_way, 3); assert_eq!(stats_after.k_shot, 2); assert_eq!(stats_after.total_classes, 5); assert_eq!(stats_after.feature_dim, 10); } #[test] fn test_sampler_reset() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(4, 6, 6, &device).unwrap(); let config = EpisodeSamplerConfig::default(); let mut sampler = EpisodeSampler::new(dataset, config); sampler.sample_episode(&device).unwrap(); assert_eq!(sampler.state.episodes_generated, 1); sampler.reset(); assert_eq!(sampler.state.episodes_generated, 0); } #[test] fn test_task_distribution() { let base_config = EpisodeSamplerConfig { n_way: 3, k_shot: 1, ..Default::default() }; let mut task_dist = TaskDistribution::multi_task(base_config); assert!(!task_dist.task_configs.is_empty()); assert_eq!(task_dist.task_configs.len(), task_dist.task_weights.len()); // Sample different task configurations let config1 = task_dist.sample_task_config(); let config2 = task_dist.sample_task_config(); // Should get different configurations (eventually) for _ in 0..10 { let config = task_dist.sample_task_config(); if config.n_way != config1.n_way || config.k_shot != config1.k_shot { return; // Found different config, test passes } } } #[test] fn test_few_shot_metrics() { let mut metrics = FewShotMetrics::new(); assert_eq!(metrics.num_episodes, 0); assert_eq!(metrics.mean_accuracy, 0.0); // Add some episode results metrics.add_episode_result(0.8, 0.5, &[0, 1]); metrics.add_episode_result(0.9, 0.3, &[0, 1, 2]); assert_eq!(metrics.num_episodes, 2); assert_eq!(metrics.mean_accuracy, 0.85); // (0.8 + 0.9) / 2 assert_eq!(metrics.mean_loss, 0.4); // (0.5 + 0.3) / 2 } #[test] fn test_meta_learning_evaluator() { let device = Device::cuda(0).unwrap_or(Device::default()); let dataset = FewShotDataset::synthetic(3, 4, 5, &device).unwrap(); // Create test episodes let episode1 = dataset.sample_episode(2, 1, 2, &device).unwrap(); let episode2 = dataset.sample_episode(2, 1, 2, &device).unwrap(); let episodes = vec![episode1, episode2]; // Create test predictions (perfect accuracy) let predictions = vec![ episodes[0].query_labels.clone(), episodes[1].query_labels.clone(), ]; let losses = vec![0.1, 0.2]; let metrics = MetaLearningEvaluator::evaluate_few_shot_performance( &episodes, &predictions, &losses, ).unwrap(); assert_eq!(metrics.num_episodes, 2); assert_eq!(metrics.mean_accuracy, 1.0); // Perfect predictions assert_eq!(metrics.mean_loss, 0.15); // (0.1 + 0.2) / 2 } #[test] fn test_approach_comparison() { let metrics_a = FewShotMetrics { mean_accuracy: 0.85, mean_loss: 0.3, ..Default::default() }; let metrics_b = FewShotMetrics { mean_accuracy: 0.75, mean_loss: 0.4, ..Default::default() }; let comparison = MetaLearningEvaluator::compare_approaches(&metrics_a, &metrics_b); assert_eq!(comparison.accuracy_difference, 0.1); assert_eq!(comparison.loss_difference, -0.1); assert_eq!(comparison.better_approach, "A"); assert_eq!(comparison.significance, "significant"); } }