"""Compare val losses of runs at their last common eval step (matched tokens) and report tok/s. Usage: $TA_PY scripts/compare_runs.py runA runB ...""" import json, sys runs = {} for d in sys.argv[1:]: recs = [json.loads(l) for l in open(f"{d}/log.jsonl")] runs[d] = recs common = set.intersection(*[{r["step"] for r in recs if "val" in r} for recs in runs.values()]) if not common: sys.exit("no common eval step yet") s = max(common) for d, recs in runs.items(): v = next(r for r in recs if r["step"] == s and "val" in r) toks = [r["tok_s"] for r in recs[-20:]] mean = sum(v["val"].values()) / len(v["val"]) print(f"{d}: step {s} tokens {v['tokens']/1e6:.0f}M mean_val {mean:.4f} tok/s {sum(toks)/len(toks):.0f} " + " ".join(f"{k}={x:.3f}" for k, x in sorted(v["val"].items())))