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