Files
rustytorch/crates/training/rtx-transformers/examples/jepa_eval.rs
T
osobhandClaude Sonnet 5 4aaa36a57a style: cargo fmt --workspace (whitespace/wrapping only, no semantic change)
Whole-workspace rustfmt pass picked up while iterating on Mamba GPU
backward work. Verified formatting-only via diff sampling; no logic
changed.

Co-Authored-By: Claude Sonnet 5 <[email protected]>
2026-08-10 07:09:36 -07:00

353 lines
11 KiB
Rust

//! JEPA evaluation entry point.
//!
//! Usage:
//! cargo run --example jepa_eval -p rtx-transformers -- --mode linear-probe
//! cargo run --example jepa_eval -p rtx-transformers -- --mode knn --k 10
//! cargo run --example jepa_eval -p rtx-transformers -- --mode both --dim 768 --classes 100
//! cargo run --example jepa_eval -p rtx-transformers -- --features train.txt --labels train_labels.txt \
//! --test-features test.txt --test-labels test_labels.txt
//! cargo run --example jepa_eval -p rtx-transformers -- --output results.csv
use rtx_transformers::ssl::jepa::{EvalMethod, JepaEvaluator};
use rtx_transformers::ssl::jepa_eval::{
EvalMode, EvalSuiteResult, JepaEvalConfig, load_features_txt, load_labels_txt, run_eval_suite,
save_eval_csv,
};
fn print_usage() {
eprintln!("JEPA Evaluation Harness");
eprintln!();
eprintln!("USAGE:");
eprintln!(" jepa_eval [OPTIONS]");
eprintln!();
eprintln!("OPTIONS:");
eprintln!(" --mode <linear-probe|knn|both> Evaluation mode (default: both)");
eprintln!(" --dim <N> Feature dimension (default: 192)");
eprintln!(" --classes <N> Number of classes (default: 10)");
eprintln!(" --train <N> Number of training samples (default: 500)");
eprintln!(" --test <N> Number of test samples (default: 100)");
eprintln!(" --epochs <N> Linear probe epochs (default: 20)");
eprintln!(" --lr <F> Linear probe learning rate (default: 1e-3)");
eprintln!(" --k <N> k for k-NN (default: 5)");
eprintln!(" --seed <N> LCG seed (default: 42)");
eprintln!(" --features <PATH> Load train features from file");
eprintln!(" --labels <PATH> Load train labels from file");
eprintln!(" --test-features <PATH> Load test features from file");
eprintln!(" --test-labels <PATH> Load test labels from file");
eprintln!(" --output <PATH> Save results as CSV (default: stdout)");
eprintln!(" --help Show this message");
}
fn parse_args() -> Option<(
JepaEvalConfig,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
)> {
let args: Vec<String> = std::env::args().skip(1).collect();
if args.iter().any(|a| a == "--help" || a == "-h") {
print_usage();
return None;
}
let mut mode = EvalMode::Both;
let mut feature_dim: usize = 192;
let mut num_classes: usize = 10;
let mut num_train: usize = 500;
let mut num_test: usize = 100;
let mut linear_probe_epochs: usize = 20;
let mut linear_probe_lr: f64 = 1e-3;
let mut knn_k: usize = 5;
let mut seed: u64 = 42;
let mut features_path: Option<String> = None;
let mut labels_path: Option<String> = None;
let mut test_features_path: Option<String> = None;
let mut test_labels_path: Option<String> = None;
let mut output_path: Option<String> = None;
let mut i = 0;
while i < args.len() {
match args[i].as_str() {
"--mode" => {
i += 1;
if i < args.len() {
mode = match args[i].as_str() {
"linear-probe" | "linear_probe" => EvalMode::LinearProbe,
"knn" | "KNN" => EvalMode::KNN,
"both" => EvalMode::Both,
other => {
eprintln!("Unknown mode '{}', expected linear-probe|knn|both", other);
return None;
}
};
}
}
"--dim" => {
i += 1;
if i < args.len() {
feature_dim = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --dim value '{}'", args[i]);
192
});
}
}
"--classes" => {
i += 1;
if i < args.len() {
num_classes = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --classes value '{}'", args[i]);
10
});
}
}
"--train" => {
i += 1;
if i < args.len() {
num_train = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --train value '{}'", args[i]);
500
});
}
}
"--test" => {
i += 1;
if i < args.len() {
num_test = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --test value '{}'", args[i]);
100
});
}
}
"--epochs" => {
i += 1;
if i < args.len() {
linear_probe_epochs = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --epochs value '{}'", args[i]);
20
});
}
}
"--lr" => {
i += 1;
if i < args.len() {
linear_probe_lr = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --lr value '{}'", args[i]);
1e-3
});
}
}
"--k" => {
i += 1;
if i < args.len() {
knn_k = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --k value '{}'", args[i]);
5
});
}
}
"--seed" => {
i += 1;
if i < args.len() {
seed = args[i].parse().unwrap_or_else(|_| {
eprintln!("Invalid --seed value '{}'", args[i]);
42
});
}
}
"--features" => {
i += 1;
if i < args.len() {
features_path = Some(args[i].clone());
}
}
"--labels" => {
i += 1;
if i < args.len() {
labels_path = Some(args[i].clone());
}
}
"--test-features" => {
i += 1;
if i < args.len() {
test_features_path = Some(args[i].clone());
}
}
"--test-labels" => {
i += 1;
if i < args.len() {
test_labels_path = Some(args[i].clone());
}
}
"--output" => {
i += 1;
if i < args.len() {
output_path = Some(args[i].clone());
}
}
unknown => {
eprintln!("Unknown argument '{}'. Use --help for usage.", unknown);
return None;
}
}
i += 1;
}
let config = JepaEvalConfig {
feature_dim,
num_classes,
num_train,
num_test,
linear_probe_epochs,
linear_probe_lr,
knn_k,
seed,
mode,
};
Some((
config,
features_path,
labels_path,
test_features_path,
test_labels_path,
output_path,
))
}
/// Run evaluation on externally loaded features.
fn run_from_files(
config: &JepaEvalConfig,
train_feats: Vec<f32>,
train_dim: usize,
train_labels: Vec<usize>,
test_feats: Vec<f32>,
_test_dim: usize,
test_labels: Vec<usize>,
) -> EvalSuiteResult {
use std::time::Instant;
let evaluator = JepaEvaluator::new(train_dim);
let num_train = train_labels.len();
let num_test = test_labels.len();
let start = Instant::now();
let linear_probe = if config.mode == EvalMode::LinearProbe || config.mode == EvalMode::Both {
Some(evaluator.linear_probe(
&train_feats,
&train_labels,
&test_feats,
&test_labels,
config.num_classes,
config.linear_probe_epochs,
config.linear_probe_lr,
))
} else {
None
};
let knn = if config.mode == EvalMode::KNN || config.mode == EvalMode::Both {
Some(evaluator.knn_eval(
&train_feats,
&train_labels,
&test_feats,
&test_labels,
config.knn_k,
))
} else {
None
};
let wall_ms = start.elapsed().as_secs_f32() * 1000.0;
EvalSuiteResult {
linear_probe,
knn,
feature_dim: train_dim,
num_train,
num_test,
wall_ms,
}
}
fn main() {
let parsed = match parse_args() {
Some(p) => p,
None => std::process::exit(1),
};
let (config, features_path, labels_path, test_features_path, test_labels_path, output_path) =
parsed;
let result = if let Some(ref feat_path) = features_path {
// File-based evaluation
let lbl_path = labels_path.as_deref().unwrap_or_else(|| {
eprintln!("--labels required when --features is given");
std::process::exit(1);
});
let tfeat_path = test_features_path.as_deref().unwrap_or_else(|| {
eprintln!("--test-features required when --features is given");
std::process::exit(1);
});
let tlbl_path = test_labels_path.as_deref().unwrap_or_else(|| {
eprintln!("--test-labels required when --features is given");
std::process::exit(1);
});
let (train_feats, train_dim) = load_features_txt(feat_path).unwrap_or_else(|e| {
eprintln!("Error loading train features: {}", e);
std::process::exit(1);
});
let train_labels = load_labels_txt(lbl_path).unwrap_or_else(|e| {
eprintln!("Error loading train labels: {}", e);
std::process::exit(1);
});
let (test_feats, test_dim) = load_features_txt(tfeat_path).unwrap_or_else(|e| {
eprintln!("Error loading test features: {}", e);
std::process::exit(1);
});
let test_labels = load_labels_txt(tlbl_path).unwrap_or_else(|e| {
eprintln!("Error loading test labels: {}", e);
std::process::exit(1);
});
if train_dim != test_dim {
eprintln!(
"Feature dim mismatch: train={} test={}",
train_dim, test_dim
);
std::process::exit(1);
}
run_from_files(
&config,
train_feats,
train_dim,
train_labels,
test_feats,
test_dim,
test_labels,
)
} else {
// Synthetic evaluation
run_eval_suite(&config)
};
// Print summary to stdout
println!("{}", result.summary());
// Optional CSV output
if let Some(ref out_path) = output_path {
match save_eval_csv(&result, out_path) {
Ok(()) => eprintln!("Results saved to '{}'", out_path),
Err(e) => {
eprintln!("Failed to save CSV: {}", e);
std::process::exit(1);
}
}
}
}