171 lines
6.9 KiB
Rust
171 lines
6.9 KiB
Rust
//! Standalone test for meta-learning implementation
|
||
//! This validates the meta-learning code without dependency issues
|
||
|
||
use rtx_transformers::meta::*;
|
||
use rtx_tensor::{Device, Tensor};
|
||
|
||
fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||
println!("Testing RTX Meta-Learning Implementation...");
|
||
|
||
let device = Device::cuda(0).unwrap_or(Device::default());
|
||
|
||
// Test 1: FewShotDataset creation
|
||
println!("✓ Testing FewShotDataset creation...");
|
||
let dataset = FewShotDataset::synthetic(5, 10, 32, &device)?;
|
||
assert_eq!(dataset.num_classes, 5);
|
||
assert_eq!(dataset.feature_dim, 32);
|
||
println!(" Created dataset with {} classes and {} feature dimensions",
|
||
dataset.num_classes, dataset.feature_dim);
|
||
|
||
// Test 2: Episode sampling
|
||
println!("✓ Testing episode sampling...");
|
||
let episode = dataset.sample_episode(3, 2, 3, &device)?;
|
||
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
|
||
println!(" Generated {}-way {}-shot episode with {} support and {} query samples",
|
||
episode.n_way, episode.k_shot, episode.support_set.len(), episode.query_set.len());
|
||
|
||
// Test 3: MAML configuration and creation
|
||
println!("✓ Testing MAML implementation...");
|
||
let maml_config = MAMLConfig {
|
||
inner_lr: 0.01,
|
||
outer_lr: 0.001,
|
||
inner_steps: 1,
|
||
first_order: false,
|
||
meta_batch_size: 4,
|
||
};
|
||
|
||
let mut maml = MAML::new(32, 16, 3, maml_config, &device)?;
|
||
println!(" Created MAML with input_dim={}, hidden_dim={}, output_dim={}",
|
||
32, 16, 3);
|
||
|
||
// Test inner loop adaptation
|
||
let meta_params = maml.meta_network.get_parameters();
|
||
let (adapted_params, inner_loss) = maml.inner_loop_adaptation(
|
||
meta_params, &episode, &device
|
||
)?;
|
||
assert!(inner_loss >= 0.0);
|
||
println!(" Inner loop adaptation completed with loss: {:.4}", inner_loss);
|
||
|
||
// Test query evaluation
|
||
let (outer_loss, accuracy) = maml.evaluate_query_set(
|
||
adapted_params, &episode, &device
|
||
)?;
|
||
assert!(outer_loss >= 0.0);
|
||
assert!(accuracy >= 0.0 && accuracy <= 1.0);
|
||
println!(" Query evaluation: loss={:.4}, accuracy={:.4}", outer_loss, accuracy);
|
||
|
||
// Test 4: Prototypical Networks
|
||
println!("✓ Testing Prototypical Networks...");
|
||
let proto_config = PrototypicalConfig {
|
||
learning_rate: 0.001,
|
||
feature_dim: 16,
|
||
distance_metric: DistanceMetric::Euclidean,
|
||
temperature: 1.0,
|
||
batch_size: 4,
|
||
};
|
||
|
||
let proto_net = PrototypicalNetworks::new(32, proto_config, &device)?;
|
||
println!(" Created Prototypical Networks with feature_dim={}", 16);
|
||
|
||
// Test prototype computation
|
||
let prototypes = proto_net.compute_prototypes(&episode, &device)?;
|
||
assert_eq!(prototypes.num_classes, episode.n_way);
|
||
assert_eq!(prototypes.feature_dim, 16);
|
||
println!(" Computed {} prototypes with {} dimensions",
|
||
prototypes.num_classes, prototypes.feature_dim);
|
||
|
||
// Test classification
|
||
let (logits, probabilities, proto_accuracy) = proto_net.classify(&episode, &device)?;
|
||
assert_eq!(logits.shape()[0], episode.query_set.len());
|
||
assert_eq!(logits.shape()[1], episode.n_way);
|
||
assert!(proto_accuracy >= 0.0 && proto_accuracy <= 1.0);
|
||
println!(" Classification accuracy: {:.4}", proto_accuracy);
|
||
|
||
// Test 5: Episode Sampler
|
||
println!("✓ Testing Episode Sampler...");
|
||
let sampler_config = EpisodeSamplerConfig {
|
||
n_way: 3,
|
||
k_shot: 1,
|
||
query_per_class: 2,
|
||
num_episodes: 5,
|
||
seed: Some(42),
|
||
};
|
||
|
||
let mut sampler = EpisodeSampler::new(dataset, sampler_config);
|
||
let episodes = sampler.sample_episodes(3, &device)?;
|
||
assert_eq!(episodes.len(), 3);
|
||
println!(" Generated {} episodes for meta-training", episodes.len());
|
||
|
||
// Test 6: Meta-learning Pipeline
|
||
println!("✓ Testing Meta-learning Pipeline...");
|
||
let dataset = FewShotDataset::synthetic(4, 8, 24, &device)?;
|
||
let pipeline_config = MetaLearningConfig {
|
||
input_dim: 24,
|
||
hidden_dim: 12,
|
||
output_dim: 4,
|
||
algorithm: MetaLearningAlgorithm::Prototypical,
|
||
sampler_config: EpisodeSamplerConfig {
|
||
n_way: 2,
|
||
k_shot: 1,
|
||
query_per_class: 2,
|
||
num_episodes: 5,
|
||
seed: Some(123),
|
||
},
|
||
..Default::default()
|
||
};
|
||
|
||
let mut pipeline = MetaLearningPipeline::new(dataset, pipeline_config, &device)?;
|
||
let results = pipeline.train(3, &device)?;
|
||
assert!(results.proto_stats.is_some());
|
||
println!(" Pipeline training completed with {} episodes processed",
|
||
results.proto_stats.unwrap().episodes_processed);
|
||
|
||
// Test 7: Distance metrics
|
||
println!("✓ Testing distance metrics...");
|
||
for metric in [DistanceMetric::Euclidean, DistanceMetric::Cosine, DistanceMetric::Manhattan] {
|
||
let config = PrototypicalConfig {
|
||
distance_metric: metric.clone(),
|
||
feature_dim: 8,
|
||
..Default::default()
|
||
};
|
||
let proto_net = PrototypicalNetworks::new(16, config, &device)?;
|
||
|
||
let queries = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0], &[1, 8], device.clone())?;
|
||
let prototypes = Tensor::from_vec(vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0], &[1, 8], device.clone())?;
|
||
|
||
let distances = match metric {
|
||
DistanceMetric::Euclidean => proto_net.compute_euclidean_distances(&queries, &prototypes)?,
|
||
DistanceMetric::Cosine => proto_net.compute_cosine_distances(&queries, &prototypes)?,
|
||
DistanceMetric::Manhattan => proto_net.compute_manhattan_distances(&queries, &prototypes)?,
|
||
};
|
||
|
||
assert_eq!(distances.shape(), &[1, 1]);
|
||
println!(" {:?} distance computed successfully", metric);
|
||
}
|
||
|
||
// Test 8: Evaluation metrics
|
||
println!("✓ Testing evaluation metrics...");
|
||
let mut metrics = FewShotMetrics::new();
|
||
metrics.add_episode_result(0.8, 0.3, &[0, 1]);
|
||
metrics.add_episode_result(0.9, 0.2, &[0, 1, 2]);
|
||
metrics.add_episode_result(0.7, 0.4, &[1, 2]);
|
||
|
||
assert_eq!(metrics.num_episodes, 3);
|
||
let expected_mean_acc = (0.8 + 0.9 + 0.7) / 3.0;
|
||
assert!((metrics.mean_accuracy - expected_mean_acc).abs() < 1e-6);
|
||
println!(" Evaluation metrics: mean_accuracy={:.4}, num_episodes={}",
|
||
metrics.mean_accuracy, metrics.num_episodes);
|
||
|
||
println!("\n🎉 All meta-learning tests passed successfully!");
|
||
println!("✅ MAML (Model-Agnostic Meta-Learning) implemented");
|
||
println!("✅ Prototypical Networks implemented");
|
||
println!("✅ Episode sampling utilities implemented");
|
||
println!("✅ Multiple distance metrics supported");
|
||
println!("✅ Evaluation framework implemented");
|
||
println!("✅ Complete meta-learning pipeline implemented");
|
||
|
||
Ok(())
|
||
} |