//! 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 = 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); }