File size: 6,926 Bytes
b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 b97f43c fb6e523 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | """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)
|