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