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

171 lines
6.9 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.
//! 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(())
}