"""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)