File size: 3,182 Bytes
11f47cf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""
Measure self-play throughput for (parallel experiments x threads each).

The research machine has 16 cores. The question is whether one 16-thread job or
several narrow jobs give more total positions/second. Tiny models usually
parallelise badly across threads, so N x 1-thread often wins -- but measure it.

    python scripts/bench_threads.py --total-cores 16 --games 8 --plies 40
"""
from __future__ import annotations

import argparse
import itertools
import json
import multiprocessing as mp
import os
import sys
import time

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))


def worker(threads, games, plies, steps, seed, q):
    os.environ["OMP_NUM_THREADS"] = str(threads)
    os.environ["MKL_NUM_THREADS"] = str(threads)
    import torch
    torch.set_num_threads(threads)
    from tinychess.config import baseline_config
    from tinychess.model import TinyChess
    from tinychess.selfplay import SelfPlayConfig, play_games
    torch.manual_seed(seed)
    m = TinyChess(baseline_config())
    t = time.time()
    _, st = play_games(m, cfg=SelfPlayConfig(n_games=games, max_plies=plies,
                                             steps=steps, seed=seed,
                                             adjudicate="material"))
    q.put({"moves": st["total_moves"], "seconds": time.time() - t})


def run(nproc, threads, games, plies, steps):
    q = mp.Queue()
    ps = [mp.Process(target=worker, args=(threads, games, plies, steps, 100 + i, q))
          for i in range(nproc)]
    t0 = time.time()
    for p in ps:
        p.start()
    res = [q.get() for _ in ps]
    for p in ps:
        p.join()
    wall = time.time() - t0
    moves = sum(r["moves"] for r in res)
    return {"procs": nproc, "threads": threads, "wall": round(wall, 2),
            "moves": moves, "moves_per_s": round(moves / wall, 1)}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--total-cores", type=int, default=os.cpu_count())
    ap.add_argument("--games", type=int, default=8)
    ap.add_argument("--plies", type=int, default=40)
    ap.add_argument("--steps", type=int, default=4)
    ap.add_argument("--out", default="runs/bench_threads.json")
    a = ap.parse_args()

    C = a.total_cores
    combos = []
    n = 1
    while n <= C:
        combos.append((n, max(1, C // n)))
        n *= 2
    print(f"cores={C}  benchmarking {len(combos)} configurations")
    print(f"{'procs':>6} {'thr':>4} {'wall(s)':>9} {'moves':>8} {'moves/s':>9}")
    rows = []
    for nproc, thr in combos:
        r = run(nproc, thr, a.games, a.plies, a.steps)
        rows.append(r)
        print(f"{r['procs']:>6} {r['threads']:>4} {r['wall']:>9.2f} "
              f"{r['moves']:>8} {r['moves_per_s']:>9}")
    best = max(rows, key=lambda r: r["moves_per_s"])
    print(f"\nBEST: {best['procs']} process(es) x {best['threads']} thread(s) "
          f"= {best['moves_per_s']} moves/s")
    os.makedirs(os.path.dirname(os.path.abspath(a.out)), exist_ok=True)
    json.dump({"cores": C, "rows": rows, "best": best}, open(a.out, "w"), indent=2)


if __name__ == "__main__":
    mp.set_start_method("spawn", force=True)
    main()