Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,679 @@
//! 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");
}
}