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]>
97 lines
3.6 KiB
Rust
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);
|
|
}
|