Download scripts/bench_parallel.py from cazyundee/training: direct link, hf CLI and curl.
- Browser
- Download file 1.84 kB
-
https://huggingface.co/spaces/cazyundee/training/resolve/main/scripts/bench_parallel.py
- Command line
-
hf download hf://spaces/cazyundee/training/scripts/bench_parallel.py
-
curl -L -o bench_parallel.py https://huggingface.co/spaces/cazyundee/training/resolve/main/scripts/bench_parallel.py
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() | |