"""DDP training entrypoint (run via torchrun). torchrun --standalone --nproc_per_node=3 scripts/07_train.py \ --rows data/manifests/split_train.jsonl --val data/manifests/split_val.jsonl \ --epochs 40 --lr 3e-4 --ckpt-dir runs/round0 """ import os import sys sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import argparse from src.common import load_jsonl from src import training def main(): ap = argparse.ArgumentParser() ap.add_argument("--rows", required=True, help="comma-separated jsonl (labeled + pseudo rows)") ap.add_argument("--val", required=True) ap.add_argument("--epochs", type=int, default=40) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--batch", type=int, default=3) ap.add_argument("--patch", type=int, default=96) ap.add_argument("--workers", type=int, default=4) ap.add_argument("--ckpt-dir", required=True) ap.add_argument("--resume", default=None, help="checkpoint to warm-restart weights from") ap.add_argument("--val-every", type=int, default=2) ap.add_argument("--val-limit", type=int, default=60) args = ap.parse_args() rows = [] for f in args.rows.split(","): rows.extend(load_jsonl(f)) val_rows = load_jsonl(args.val) training.train_ddp(rows, val_rows, args) if __name__ == "__main__": main()