679 lines
23 KiB
Rust
679 lines
23 KiB
Rust
//! 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");
|
||
}
|
||
} |