Download tinychess/parallel.py from cazyundee/training: direct link, hf CLI and curl.
- Browser
- Download file 6.93 kB
-
https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/parallel.py
- Command line
-
hf download hf://spaces/cazyundee/training/tinychess/parallel.py
-
curl -L -o parallel.py https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/parallel.py
6.93 kB
| """Multi-process self-play. | |
| `play_games` batches NN calls across games but runs in ONE process, so on a | |
| 192-core machine it uses a single core's worth of BLAS plus whatever intra-op | |
| threads torch grabs. The thread benchmark measured 1 proc x 64 threads = 123 | |
| moves/s versus 32 procs x 2 threads = 2258 moves/s (18x): this model is tiny, so | |
| the win comes from process-level parallelism, not thread-level. | |
| This module shards `n_games` across worker processes, each running the existing | |
| single-process `play_games` with its own seed, then concatenates the GameRecords. | |
| GameRecord holds only plain Python lists/strings, so it pickles cheaply. | |
| Nothing about the RL algorithm changes: the same games are played from the same | |
| distribution, just concurrently. Determinism is preserved per shard via | |
| `seed + shard_index`. | |
| """ | |
| from __future__ import annotations | |
| import os | |
| from dataclasses import replace | |
| import torch | |
| _WORKER = {} | |
| def _init_worker(payload_path, threads): | |
| """Rebuild the model once per worker, not once per task. | |
| The payload travels via a temp FILE rather than the initargs pickle: with | |
| the 'spawn' start method initargs are streamed down a pipe, and a | |
| multi-MB state_dict plus config raced the pipe buffer (BrokenPipeError). | |
| """ | |
| import io | |
| n = str(max(1, threads)) | |
| for v in ("OMP_NUM_THREADS", "MKL_NUM_THREADS", "OPENBLAS_NUM_THREADS", | |
| "NUMEXPR_NUM_THREADS", "VECLIB_MAXIMUM_THREADS"): | |
| os.environ[v] = n | |
| torch.set_num_threads(max(1, threads)) | |
| try: | |
| torch.set_num_interop_threads(1) | |
| except RuntimeError: | |
| pass | |
| from .model import TinyChess | |
| payload = torch.load(payload_path, map_location="cpu", weights_only=False) | |
| model = TinyChess(payload["cfg"]) | |
| model.load_state_dict(payload["state"]) | |
| model.eval() | |
| opp = None | |
| if payload.get("opp_state") is not None: | |
| opp = TinyChess(payload["opp_cfg"]) | |
| opp.load_state_dict(payload["opp_state"]) | |
| opp.eval() | |
| _WORKER["model"] = model | |
| _WORKER["opp"] = opp | |
| _WORKER["cfg"] = payload["sp_cfg"] | |
| def _run_shard(arg): | |
| n_games, seed = arg | |
| from .selfplay import play_games | |
| cfg = replace(_WORKER["cfg"], n_games=n_games, seed=seed) | |
| recs, stats = play_games(_WORKER["model"], _WORKER["opp"], cfg) | |
| return recs, stats | |
| def _pack(model, opp, sp_cfg, path): | |
| payload = {"cfg": model.cfg, "state": model.state_dict(), | |
| "opp_cfg": getattr(opp, "cfg", None), | |
| "opp_state": opp.state_dict() if opp is not None else None, | |
| "sp_cfg": sp_cfg} | |
| torch.save(payload, path) | |
| return path | |
| def play_games_parallel(model, opp, cfg, n_workers=None, threads_per_worker=2, | |
| pool=None): | |
| """Drop-in parallel replacement for play_games(model, opp, cfg). | |
| Returns (recs, stats) with the same shape as play_games. Falls back to the | |
| single-process path when n_workers <= 1 or there are too few games. | |
| """ | |
| from .selfplay import play_games, summarise | |
| n_games = cfg.n_games | |
| if n_workers is None: | |
| n_workers = max(1, (os.cpu_count() or 2) // max(1, threads_per_worker)) | |
| n_workers = max(1, min(n_workers, n_games)) | |
| if n_workers == 1: | |
| return play_games(model, opp, cfg) | |
| base, extra = divmod(n_games, n_workers) | |
| shards = [(base + (1 if i < extra else 0), cfg.seed * 1000 + i) | |
| for i in range(n_workers)] | |
| shards = [s for s in shards if s[0] > 0] | |
| # NOTE: use "spawn", not "fork". The parent process has already initialised | |
| # an OpenMP thread pool (torch does this on the first op). libgomp is not | |
| # fork-safe: forked children inherit a broken pool and deadlock. Observed | |
| # directly -- a 1x1 bench row completed, then the first multi-worker row | |
| # hung for 17+ minutes on the 192-core box. "spawn" starts clean | |
| # interpreters, and the pool is created once so the import cost is amortised. | |
| import tempfile | |
| import multiprocessing as mp | |
| ctx = mp.get_context("spawn") | |
| created = pool is None | |
| tmp = None | |
| if created: | |
| fd, tmp = tempfile.mkstemp(suffix=".pt", prefix="tc_sp_") | |
| os.close(fd) | |
| _pack(model, opp, cfg, tmp) | |
| pool = ctx.Pool(processes=len(shards), initializer=_init_worker, | |
| initargs=(tmp, threads_per_worker)) | |
| try: | |
| out = pool.map(_run_shard, shards) | |
| finally: | |
| if created: | |
| pool.close() | |
| pool.join() | |
| if tmp and os.path.exists(tmp): | |
| os.unlink(tmp) | |
| recs = [r for sub, _ in out for r in sub] | |
| stats = summarise(recs) | |
| stats["illegal_attempts"] = sum(s.get("illegal_attempts", 0) for _, s in out) | |
| stats["n_workers"] = len(shards) | |
| return recs, stats | |
| def _set_weights(path): | |
| payload = torch.load(path, map_location="cpu", weights_only=False) | |
| m = _WORKER.get("model") | |
| if m is not None: | |
| m.load_state_dict(payload["state"]) | |
| m.eval() | |
| if payload.get("opp_state") is not None: | |
| from .model import TinyChess | |
| o = TinyChess(payload["opp_cfg"]) | |
| o.load_state_dict(payload["opp_state"]) | |
| o.eval() | |
| _WORKER["opp"] = o | |
| else: | |
| _WORKER["opp"] = None | |
| _WORKER["cfg"] = payload["sp_cfg"] | |
| return True | |
| class SelfPlayPool: | |
| """Long-lived spawn pool. Spawning 96 interpreters costs seconds; doing it | |
| every iteration would dominate. Weights are re-broadcast per iteration.""" | |
| def __init__(self, model, opp, cfg, n_workers, threads_per_worker=2): | |
| import multiprocessing as mp | |
| import tempfile | |
| self.n_workers = max(1, n_workers) | |
| self.threads = threads_per_worker | |
| fd, self.tmp = tempfile.mkstemp(suffix=".pt", prefix="tc_pool_") | |
| os.close(fd) | |
| _pack(model, opp, cfg, self.tmp) | |
| ctx = mp.get_context("spawn") | |
| self.pool = ctx.Pool(processes=self.n_workers, initializer=_init_worker, | |
| initargs=(self.tmp, threads_per_worker)) | |
| def play(self, model, opp, cfg): | |
| from .selfplay import summarise | |
| _pack(model, opp, cfg, self.tmp) | |
| self.pool.map(_set_weights, [self.tmp] * self.n_workers, chunksize=1) | |
| n = cfg.n_games | |
| w = min(self.n_workers, n) | |
| base, extra = divmod(n, w) | |
| shards = [(base + (1 if i < extra else 0), cfg.seed * 1000 + i) | |
| for i in range(w)] | |
| shards = [x for x in shards if x[0] > 0] | |
| out = self.pool.map(_run_shard, shards) | |
| recs = [r for sub, _ in out for r in sub] | |
| stats = summarise(recs) | |
| stats["illegal_attempts"] = sum(s.get("illegal_attempts", 0) for _, s in out) | |
| stats["n_workers"] = len(shards) | |
| return recs, stats | |
| def close(self): | |
| try: | |
| self.pool.close() | |
| self.pool.join() | |
| finally: | |
| if os.path.exists(self.tmp): | |
| os.unlink(self.tmp) | |