training / scripts /bench_learn.py
cazyundee's picture
add learn-phase thread benchmark
3f4a221 verified
Raw History Blame Contribute Delete
2.62 kB
#!/usr/bin/env python3
"""Does the learn phase actually scale with torch threads?
main_v5 spends 390 s/iter in learn vs 45 s in self-play with --threads 48.
If the update step doesn't scale past a few threads, those cores are wasted and
the run should instead use more parallel ARMS (or a smaller batch).
"""
import argparse, json, os, sys, time
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
from tinychess.checkpoint import load_model_only
from tinychess.config import REGISTRY
from tinychess.model import TinyChess
from tinychess.selfplay import SelfPlayConfig, play_games
from tinychess.replay import ReplayBuffer
import train_selfplay as T
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", default="")
ap.add_argument("--batch", type=int, default=1024)
ap.add_argument("--games", type=int, default=24)
ap.add_argument("--reps", type=int, default=6)
ap.add_argument("--threads", default="1,2,4,8,16,32,48")
ap.add_argument("--out", default="runs/bench_learn.json")
a = ap.parse_args()
if a.ckpt and os.path.exists(a.ckpt):
m, _ = load_model_only(a.ckpt)
else:
m = TinyChess(REGISTRY["baseline"]())
cfg = SelfPlayConfig(n_games=a.games, max_plies=160, steps=4,
adjudicate="material", seed=0)
recs, _ = play_games(m, None, cfg)
buf = ReplayBuffer()
for r in recs:
buf.add_game(r, gen=0)
print(f"buffer {len(buf)} samples", flush=True)
class A: pass
a2 = A(); a2.steps = 4; a2.norm_adv = 1; a2.ppo_clip = 0.2
a2.c_entropy = 0.01; a2.c_value = 1.0; a2.clip = 1.0; a2.c_compute = 0.0
a2.terminal_weight = 0.0; a2.terminal_tau = 10.0
opt = torch.optim.Adam(m.parameters(), lr=1e-6)
res = {"cpu_count": os.cpu_count(), "batch": a.batch, "runs": []}
for nt in [int(x) for x in a.threads.split(",")]:
torch.set_num_threads(nt)
T.train_step(m, opt, buf.sample(a.batch), a2) # warmup
t = time.time()
for _ in range(a.reps):
T.train_step(m, opt, buf.sample(a.batch), a2)
ms = (time.time() - t) / a.reps * 1000
res["runs"].append({"threads": nt, "ms_per_update": round(ms, 1)})
print(f"threads={nt:3d}: {ms:8.1f} ms/update", flush=True)
json.dump(res, open(a.out, "w"), indent=2)
best = min(res["runs"], key=lambda r: r["ms_per_update"])
res["best"] = best
json.dump(res, open(a.out, "w"), indent=2)
print(f"\nBEST: {best['threads']} threads = {best['ms_per_update']} ms/update")
if __name__ == "__main__":
main()