dreddnafious's picture
Revision 2: step density matched; headline replicates
0662d8e verified
Raw History Blame Contribute Delete
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)