training / tinychess /parallel.py
cazyundee's picture
spawn-based persistent self-play pool (fork/libgomp deadlock fix)
fb6e523 verified
Raw History Blame Contribute Delete
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)