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