#!/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()