Add .gitignore, AGENTS.md, scripts directory, and src directory to initialize the repository.
37 lines
No EOL
1.3 KiB
Python
37 lines
No EOL
1.3 KiB
Python
"""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() |