//! 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 Evaluation mode (default: both)"); eprintln!(" --dim Feature dimension (default: 192)"); eprintln!(" --classes Number of classes (default: 10)"); eprintln!(" --train Number of training samples (default: 500)"); eprintln!(" --test Number of test samples (default: 100)"); eprintln!(" --epochs Linear probe epochs (default: 20)"); eprintln!(" --lr Linear probe learning rate (default: 1e-3)"); eprintln!(" --k k for k-NN (default: 5)"); eprintln!(" --seed LCG seed (default: 42)"); eprintln!(" --features Load train features from file"); eprintln!(" --labels Load train labels from file"); eprintln!(" --test-features Load test features from file"); eprintln!(" --test-labels Load test labels from file"); eprintln!(" --output Save results as CSV (default: stdout)"); eprintln!(" --help Show this message"); } fn parse_args() -> Option<( JepaEvalConfig, Option, Option, Option, Option, Option, )> { let args: Vec = 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 = None; let mut labels_path: Option = None; let mut test_features_path: Option = None; let mut test_labels_path: Option = None; let mut output_path: Option = 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, train_dim: usize, train_labels: Vec, test_feats: Vec, _test_dim: usize, test_labels: Vec, ) -> 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); } } } }