Download code/harness/bench.py from dreddnafious/mini-agi-replication: direct link, hf CLI and curl.
- Browser
- Download file 4.51 kB
-
https://huggingface.co/spaces/dreddnafious/mini-agi-replication/resolve/main/code/harness/bench.py
- Command line
-
hf download hf://spaces/dreddnafious/mini-agi-replication/code/harness/bench.py
-
curl -L -o bench.py https://huggingface.co/spaces/dreddnafious/mini-agi-replication/resolve/main/code/harness/bench.py
4.51 kB
| """Throughput and memory of the upstream training step at candidate scales. | |
| Drives the same per-chunk path as `train.py read`: peek/want experts, one | |
| FileReader.step (which re-forwards the whole window), backward, clip, AdamW over | |
| the trunk/pool split. Data is random bytes: cost does not depend on content. | |
| """ | |
| import argparse | |
| import copy | |
| import json | |
| import os | |
| import tempfile | |
| import time | |
| import numpy as np | |
| import torch | |
| import yaml | |
| from upstream import UPSTREAM, use_config | |
| BASE = yaml.safe_load(open(os.path.join(UPSTREAM, "config.yaml"))) | |
| SCALES = { | |
| # name: (model overrides, pool overrides, chunk, context) | |
| "xs": (dict(d_model=256, n_head=4, d_ff=704, max_steps=8, train_steps_mean=4.0, | |
| bptt_window=8, halt_prior=0.2), | |
| dict(experts=32, width=512, top_k=4, resident=16), 256, 512), | |
| "s": (dict(d_model=256, n_head=4, d_ff=704, max_steps=8, train_steps_mean=4.0, | |
| bptt_window=8, halt_prior=0.2), | |
| dict(experts=32, width=512, top_k=4, resident=16), 512, 1024), | |
| "m": (dict(d_model=384, n_head=6, d_ff=1024, max_steps=12, train_steps_mean=6.4, | |
| bptt_window=12, halt_prior=0.144), | |
| dict(experts=48, width=1024, top_k=6, resident=24), 512, 1024), | |
| "full": ({}, {}, 2048, 2048), | |
| } | |
| def make_config(scale): | |
| mo, po, chunk, ctx = SCALES[scale] | |
| c = copy.deepcopy(BASE) | |
| c["model"].update(mo) | |
| c["model"]["context_start"] = ctx | |
| c["model"]["context_end"] = ctx | |
| c["pool"].update(po) | |
| c["training"]["chunk"] = chunk | |
| return c, chunk, ctx | |
| def build(scale, device): | |
| from minagi.create import create | |
| from minagi.recur import RecurCoder # noqa: F401 | |
| import train as T | |
| c, chunk, ctx = make_config(scale) | |
| tmp = tempfile.mkdtemp(prefix=f"bench_{scale}_") | |
| cpath = os.path.join(tmp, "config.yaml") | |
| yaml.safe_dump(c, open(cpath, "w")) | |
| use_config(cpath) | |
| wdir = os.path.join(tmp, "weights") | |
| create(wdir, force=True, verbose=False) | |
| model, cfg, pool, man = T.build_paged(wdir, device) | |
| m = c["model"] | |
| cfg.train_steps_mean = float(m["train_steps_mean"]) | |
| cfg.min_steps = int(m["min_steps"]) | |
| cfg.bptt_window = min(int(m["bptt_window"]), cfg.max_steps) | |
| cfg.halt_prior = float(m["halt_prior"]) | |
| cfg.halt_thresh = float(m["halt_thresh"]) | |
| cfg.ponder_beta = float(m["ponder_beta"]) | |
| trunk, pool_ps = T._split_trunk_pool(model) | |
| lr = c["training"]["lr"] | |
| opt = torch.optim.AdamW( | |
| [{"params": trunk, "lr": lr * 0.1, "weight_decay": 0.1}, | |
| {"params": pool_ps, "lr": lr, "weight_decay": 0.1}], | |
| betas=(0.9, 0.95), fused=True) | |
| pool.attach_optimiser(opt) | |
| return model, cfg, pool, opt, chunk, ctx, trunk, pool_ps | |
| def run(scale, seconds): | |
| from minagi.precision import set_compute_dtype | |
| from minagi.stream import FileReader | |
| device = torch.device("cuda") | |
| set_compute_dtype("bf16") | |
| torch.manual_seed(0) | |
| model, cfg, pool, opt, chunk, ctx, trunk, pool_ps = build(scale, device) | |
| model.train() | |
| rng = np.random.default_rng(0) | |
| data = rng.integers(32, 127, size=4_000_000).astype(np.uint16) | |
| r = FileReader(model, data, "bench", chunk, ctx, device) | |
| torch.cuda.reset_peak_memory_stats() | |
| chars, t0, steps = 0, None, 0 | |
| while True: | |
| nxt = r.peek() | |
| model.want_experts(nxt) | |
| opt.zero_grad(set_to_none=True) | |
| loss = r.step(learn=True, aux_weight=cfg.pool_aux) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| opt.step() | |
| steps += 1 | |
| if steps == 5: | |
| torch.cuda.synchronize() | |
| t0 = time.time() | |
| elif steps > 5: | |
| chars += chunk | |
| if time.time() - t0 > seconds: | |
| break | |
| torch.cuda.synchronize() | |
| dt = time.time() - t0 | |
| out = {"scale": scale, "chunk": chunk, "context": ctx, | |
| "params_total_M": round(pool.n_params() / 1e6 + sum(p.numel() for p in trunk) / 1e6, 2), | |
| "trunk_M": round(sum(p.numel() for p in trunk) / 1e6, 2), | |
| "vram_M": round(pool.vram_params() / 1e6, 2), | |
| "char_per_s": round(chars / dt), "peak_GB": round(torch.cuda.max_memory_allocated() / 1e9, 2), | |
| "loss": round(float(loss), 3)} | |
| print(json.dumps(out), flush=True) | |
| if __name__ == "__main__": | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("scale", choices=list(SCALES)) | |
| ap.add_argument("--seconds", type=float, default=30) | |
| a = ap.parse_args() | |
| run(a.scale, a.seconds) | |