Download code/scripts/bench_size.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 3.81 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/bench_size.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/bench_size.py
-
curl -L -o bench_size.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/bench_size.py
3.81 kB
| """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") | |