Files
rustytorch/crates/training/rtx-transformers/examples/jepa_train.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

97 lines
3.6 KiB
Rust

//! JEPA training entry point.
//!
//! Usage:
//! cargo run --example jepa_train -p rtx-transformers -- --config path/to/config.toml
//! cargo run --example jepa_train -p rtx-transformers -- --size small --steps 1000
//! cargo run --example jepa_train -p rtx-transformers -- --dry-run
//!
//! With no args, runs a 100-step dry run with ViT-Tiny and synthetic data.
use rtx_transformers::ssl::jepa_runner::{JepaRunConfig, run_jepa_benchmark, run_jepa_training};
fn main() {
let args: Vec<String> = std::env::args().collect();
let mut config = JepaRunConfig::default();
config.total_steps = 100; // default to short dry run
config.log_every = 10;
let mut i = 1;
while i < args.len() {
match args[i].as_str() {
"--config" => {
i += 1;
let content = std::fs::read_to_string(&args[i])
.unwrap_or_else(|e| panic!("Cannot read config {}: {}", args[i], e));
config = rtx_transformers::ssl::jepa_runner::parse_config_from_str(&content)
.unwrap_or_else(|e| panic!("Config parse error: {}", e));
}
"--size" => {
i += 1;
config.vit_size = match args[i].as_str() {
"tiny" => rtx_transformers::ssl::jepa_runner::ViTSizeStr::Tiny,
"small" => rtx_transformers::ssl::jepa_runner::ViTSizeStr::Small,
"base" => rtx_transformers::ssl::jepa_runner::ViTSizeStr::Base,
"large" => rtx_transformers::ssl::jepa_runner::ViTSizeStr::Large,
"huge" => rtx_transformers::ssl::jepa_runner::ViTSizeStr::Huge,
s => panic!("Unknown size: {}", s),
};
}
"--steps" => {
i += 1;
config.total_steps = args[i].parse().unwrap_or_else(|_| panic!("Invalid steps"));
}
"--dry-run" => {
config.total_steps = 10;
config.log_every = 1;
}
"--benchmark" => {
config.benchmark_mode = true;
}
"--benchmark-steps" => {
i += 1;
if i < args.len() {
config.benchmark_steps = args[i].parse().unwrap_or(50);
}
}
"--benchmark-patches" => {
i += 1;
if i < args.len() {
config.benchmark_patches = args[i].parse().unwrap_or(196);
}
}
"--gpu" => {
config.use_gpu = true;
}
"--device" => {
i += 1;
if let Some(v) = args.get(i) {
config.gpu_device_id = v.parse().unwrap_or(0);
}
}
_ => {}
}
i += 1;
}
if config.benchmark_mode {
let result = run_jepa_benchmark(&config);
result.print_summary();
return;
}
println!(
"Starting JEPA training: {:?} for {} steps",
config.vit_size, config.total_steps
);
let summary = run_jepa_training(config);
println!("\n=== Training Complete ===");
println!("Steps: {}", summary.total_steps);
println!("Final loss: {:.4}", summary.final_loss);
println!("Mean loss (last 100): {:.4}", summary.mean_loss);
println!("Throughput: {:.1} steps/s", summary.steps_per_second);
println!("Tokens/sec: {:.0}", summary.tokens_per_sec);
println!("Wall time: {:.1}s", summary.wall_time_seconds);
println!("Checkpoints saved: {}", summary.checkpoints_saved);
}