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")
|