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

679 lines
23 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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<u64>,
}
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<Episode> {
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<Vec<usize>> {
let mut selected = Vec::new();
let available_classes: Vec<usize> = 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<Vec<usize>> {
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<Vec<Episode>> {
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<Vec<Episode>> {
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<EpisodeSamplerConfig>,
/// Weights for sampling different task types
pub task_weights: Vec<f32>,
/// 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<usize, f32>,
/// 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<usize>],
losses: &[f32],
) -> Result<FewShotMetrics> {
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<usize> = 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");
}
}