File size: 3,811 Bytes
4397e12
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Measure sustained training throughput of candidate sizes on the B70 and turn it into a size pick.

Random tokens with random document cuts; full step = compiled fwd+bwd + optimizer step.
Usage: source env.sh && $TA_PY scripts/bench_size.py --seconds 120
"""
import argparse, json, math, os, time

import torch

from tiny_agent.model import ModelConfig, TinyAgentLM, make_block_mask
from tiny_agent.optim import build_optimizers

CONFIGS = {
    "S": dict(d_model=512, n_layers=12, n_heads=8),
    "M": dict(d_model=640, n_layers=16, n_heads=10),
    "L": dict(d_model=768, n_layers=18, n_heads=12),
    "XL": dict(d_model=1024, n_layers=20, n_heads=16),
}


def random_docs(B, T, mean_len, dev):
    starts = (torch.rand(B, T, device=dev) < 1.0 / mean_len)
    starts[:, 0] = True
    return starts.long().cumsum(1)


def run(name, kw, engram, seconds, T, micro_B, accum):
    dev = "xpu"
    cfg = ModelConfig(**kw, engram_layers=(1,) if engram else ())
    torch.manual_seed(0)
    model = TinyAgentLM(cfg).to(dev)
    opts = build_optimizers(model, lr=3e-3)
    cmodel = torch.compile(model)
    counts = model.param_counts()
    torch.xpu.reset_peak_memory_stats()

    def step():
        for _ in range(accum):
            idx = torch.randint(0, cfg.vocab_size, (micro_B, T + 1), device=dev)
            doc = random_docs(micro_B, T, 600, dev)
            bm = make_block_mask(doc, cfg.swa_window)
            with torch.autocast("xpu", dtype=torch.bfloat16):
                loss = cmodel(idx[:, :-1], doc, bm, idx[:, 1:])
            (loss / accum).backward()
        for o in opts:
            o.step()
            o.zero_grad(set_to_none=True)
        return loss

    t0 = time.time()
    for _ in range(3):
        step()
    torch.xpu.synchronize()
    warm = time.time() - t0
    n, t0 = 0, time.time()
    while time.time() - t0 < seconds:
        loss = step()
        n += 1
    loss.item()
    torch.xpu.synchronize()
    dt = time.time() - t0
    toks = n * accum * micro_B * T / dt
    N = counts["backbone"]
    # 6*N*D (+ attention ~ 12*L*T*d_attn per token, causal-halved) for the FLOP rate
    attn = 6 * cfg.n_layers * T * cfg.n_heads * cfg.head_dim  # fwd+bwd, causal half included
    flops = (6 * (N + cfg.d_model * cfg.vocab_size) + attn) * toks
    r = dict(name=name, engram=engram, **counts, tok_s=round(toks), tflops=round(flops / 1e12, 1),
             peak_gib=round(torch.xpu.max_memory_allocated() / 2**30, 1), warmup_s=round(warm),
             steps=n, micro_B=micro_B, accum=accum, T=T,
             kv_bytes_tok_8k=model.kv_bytes_per_token(8192))
    day = toks * 86400
    r["tokens_24h_B"] = round(day / 1e9, 2)
    r["tok_per_param_24h"] = round(day / N, 1)
    print(json.dumps(r), flush=True)
    del model, cmodel, opts
    torch.xpu.empty_cache()
    return r


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--seconds", type=float, default=120)
    ap.add_argument("--T", type=int, default=2048)
    ap.add_argument("--micro_B", type=int, default=8)
    ap.add_argument("--accum", type=int, default=2)
    ap.add_argument("--configs", default="S,M,L,XL")
    ap.add_argument("--engram", default="0,1")
    ap.add_argument("--out", default=os.path.join(os.environ.get("TA_DATA", "."), "logs", "bench_size.jsonl"))
    a = ap.parse_args()
    with open(a.out, "a") as f:
        for name in a.configs.split(","):
            for e in a.engram.split(","):
                try:
                    r = run(name, CONFIGS[name], e == "1", a.seconds, a.T, a.micro_B, a.accum)
                except torch.OutOfMemoryError as ex:
                    r = dict(name=name, engram=e == "1", error="OOM")
                    print(r, flush=True)
                    torch.xpu.empty_cache()
                f.write(json.dumps(r) + "\n")