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]>
353 lines
11 KiB
Rust
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);
|
|
}
|
|
}
|
|
}
|
|
}
|