File size: 4,506 Bytes
0662d8e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | """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)
|