File size: 2,616 Bytes
3f4a221
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Does the learn phase actually scale with torch threads?

main_v5 spends 390 s/iter in learn vs 45 s in self-play with --threads 48.
If the update step doesn't scale past a few threads, those cores are wasted and
the run should instead use more parallel ARMS (or a smaller batch).
"""
import argparse, json, os, sys, time
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
from tinychess.checkpoint import load_model_only
from tinychess.config import REGISTRY
from tinychess.model import TinyChess
from tinychess.selfplay import SelfPlayConfig, play_games
from tinychess.replay import ReplayBuffer
import train_selfplay as T


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--ckpt", default="")
    ap.add_argument("--batch", type=int, default=1024)
    ap.add_argument("--games", type=int, default=24)
    ap.add_argument("--reps", type=int, default=6)
    ap.add_argument("--threads", default="1,2,4,8,16,32,48")
    ap.add_argument("--out", default="runs/bench_learn.json")
    a = ap.parse_args()

    if a.ckpt and os.path.exists(a.ckpt):
        m, _ = load_model_only(a.ckpt)
    else:
        m = TinyChess(REGISTRY["baseline"]())
    cfg = SelfPlayConfig(n_games=a.games, max_plies=160, steps=4,
                         adjudicate="material", seed=0)
    recs, _ = play_games(m, None, cfg)
    buf = ReplayBuffer()
    for r in recs:
        buf.add_game(r, gen=0)
    print(f"buffer {len(buf)} samples", flush=True)

    class A: pass
    a2 = A(); a2.steps = 4; a2.norm_adv = 1; a2.ppo_clip = 0.2
    a2.c_entropy = 0.01; a2.c_value = 1.0; a2.clip = 1.0; a2.c_compute = 0.0
    a2.terminal_weight = 0.0; a2.terminal_tau = 10.0
    opt = torch.optim.Adam(m.parameters(), lr=1e-6)

    res = {"cpu_count": os.cpu_count(), "batch": a.batch, "runs": []}
    for nt in [int(x) for x in a.threads.split(",")]:
        torch.set_num_threads(nt)
        T.train_step(m, opt, buf.sample(a.batch), a2)      # warmup
        t = time.time()
        for _ in range(a.reps):
            T.train_step(m, opt, buf.sample(a.batch), a2)
        ms = (time.time() - t) / a.reps * 1000
        res["runs"].append({"threads": nt, "ms_per_update": round(ms, 1)})
        print(f"threads={nt:3d}: {ms:8.1f} ms/update", flush=True)
        json.dump(res, open(a.out, "w"), indent=2)
    best = min(res["runs"], key=lambda r: r["ms_per_update"])
    res["best"] = best
    json.dump(res, open(a.out, "w"), indent=2)
    print(f"\nBEST: {best['threads']} threads = {best['ms_per_update']} ms/update")


if __name__ == "__main__":
    main()