longitudinal/scripts/07_train.py
Furen Xiao b6fa62a763 feat: initial project structure
Add .gitignore, AGENTS.md, scripts directory, and src directory to initialize the repository.
2026-09-25 16:00:37 +08:00

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()