Initial commit
This commit is contained in:
@@ -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");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user