training / scripts /bench_parallel.py
cazyundee's picture
bench: __main__ guard for spawn
07f4740 verified
Raw History Blame Contribute Delete
1.84 kB
#!/usr/bin/env python3
"""Find the best (workers x threads) for self-play on this machine."""
import argparse, json, os, sys, time
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
from tinychess.config import REGISTRY
from tinychess.model import TinyChess
from tinychess.selfplay import SelfPlayConfig
from tinychess.parallel import play_games_parallel
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--games", type=int, default=128)
ap.add_argument("--max-plies", type=int, default=80)
ap.add_argument("--grid", default="1x1,8x2,16x2,32x2,48x2,64x2,96x2,48x4")
ap.add_argument("--out", default="runs/bench_parallel.json")
a = ap.parse_args()
m = TinyChess(REGISTRY["baseline"]()); m.eval()
res = {"cpu_count": os.cpu_count(), "games": a.games, "runs": []}
for spec in a.grid.split(","):
w, t = (int(x) for x in spec.split("x"))
cfg = SelfPlayConfig(n_games=a.games, max_plies=a.max_plies, steps=4,
adjudicate="material", seed=0)
t0 = time.time()
recs, st = play_games_parallel(m, None, cfg, n_workers=w, threads_per_worker=t)
dt = time.time() - t0
moves = sum(len(r.moves) for r in recs)
row = {"workers": w, "threads": t, "sec": round(dt, 2),
"moves_per_s": round(moves / dt, 1), "games": len(recs)}
res["runs"].append(row)
print(f"{w:3d}w x {t}t {dt:7.2f}s {row['moves_per_s']:8.1f} moves/s", flush=True)
json.dump(res, open(a.out, "w"), indent=2)
best = max(res["runs"], key=lambda r: r["moves_per_s"])
res["best"] = best
json.dump(res, open(a.out, "w"), indent=2)
print(f"\nBEST: {best['workers']}w x {best['threads']}t = {best['moves_per_s']} moves/s")
if __name__ == "__main__":
main()