"""Find the first tensor that goes non-finite in bf16 training. Variants: full (as trained) / nobal (balance_coef=0) / dense (mlp=dense) / noadam-embed (freeze embed+head only). """ import argparse import numpy as np import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim from mlx.utils import tree_flatten import configs from model import HOARD from train import TokenLoader, fp32_master, is_muon_param def finite(x): return bool(mx.all(mx.isfinite(x.astype(mx.float32))).item()) def scan(tree, tag): bad = [] for name, v in tree_flatten(tree): if not finite(v): bad.append(name) if bad: print(f" !! non-finite {tag}: {bad[:6]}{' …' if len(bad) > 6 else ''}") return bad def mstat(name, t): t32 = t.astype(mx.float32) print(f" {name}: max|{mx.abs(t32).max().item():.3e}| finite={bool(mx.all(mx.isfinite(t32)).item())}") def probe_mixer(mixer, xn): """Descend into GDNMixer on a failing input.""" from model import chunk_gated_delta_rule x = xn.astype(mx.float32) B, T, _ = x.shape qkv = mixer.qkv(x) arg = mixer.a(x) + mixer.dt_bias g = mixer._log_decay(x) beta = mx.sigmoid(mixer.b(x)) print(f" softplus arg: [{arg.min().item():.3e}, {arg.max().item():.3e}] " f"A_log max {mixer.A_log.max().item():.3e} g min {g.min().item():.3e}") xpad = mx.pad(qkv, [(0, 0), (mixer.K - 1, 0), (0, 0)]) qkv2 = nn.silu(mixer.conv(xpad)) q, k, v = mixer._split(qkv2, B, T) mstat("q", q); mstat("k", k); mstat("v", v) # manual chunk walk printing state growth qh, kh, vh = (t.transpose(0, 2, 1, 3) for t in (q, k, v)) gh, bh = g.transpose(0, 2, 1), beta.transpose(0, 2, 1) o, S = chunk_gated_delta_rule(qh, kh, vh, gh, bh, mixer.C) mstat("chunk o", o); mstat("final S", S) # per-chunk S trace, replicated math import model as M C = mixer.C f32 = mx.float32 qq = qh.astype(f32) * (mixer.dk ** -0.5); kk = kh.astype(f32); vv = vh.astype(f32) gg = gh.astype(f32); bb = bh.astype(f32) pad = (C - T % C) % C if pad: qq = mx.pad(qq, [(0,0),(0,0),(0,pad),(0,0)]); kk = mx.pad(kk, [(0,0),(0,0),(0,pad),(0,0)]) vv = mx.pad(vv, [(0,0),(0,0),(0,pad),(0,0)]); gg = mx.pad(gg, [(0,0),(0,0),(0,pad)]) bb = mx.pad(bb, [(0,0),(0,0),(0,pad)]) Tp = T + pad; nC = Tp // C qq = qq.reshape(B, mixer.H, nC, C, mixer.dk); kk = kk.reshape(B, mixer.H, nC, C, mixer.dk) vv = vv.reshape(B, mixer.H, nC, C, mixer.dv); gg = gg.reshape(B, mixer.H, nC, C) bb = bb.reshape(B, mixer.H, nC, C) G = mx.cumsum(gg, axis=-1) tril = mx.tril(mx.ones((C, C), dtype=mx.bool_)); strict = mx.tril(mx.ones((C, C), dtype=mx.bool_), k=-1) diff = G[..., :, None] - G[..., None, :] D = mx.exp(mx.where(tril, diff, -1e30)); Ds = mx.where(strict, D, 0.0) kb = kk * bb[..., None]; vb = vv * bb[..., None] A = (kb @ kk.swapaxes(-1, -2)) * Ds Tinv = M.inv_unit_lower(A) print(f" A max {mx.abs(A).max().item():.3e} Tinv max {mx.abs(Tinv).max().item():.3e} " f"D max {D.max().item():.3e}") W = Tinv @ (kb * mx.exp(G)[..., None]); U = Tinv @ vb S = mx.zeros((B, mixer.H, mixer.dk, mixer.dv), dtype=f32) for i in range(nC): qi, ki, Gi = qq[:, :, i], kk[:, :, i], G[:, :, i] v_new = U[:, :, i] - W[:, :, i] @ S g_last = Gi[..., -1] kdec = ki * mx.exp(g_last[..., None] - Gi)[..., None] S = S * mx.exp(g_last)[..., None, None] + kdec.swapaxes(-1, -2) @ v_new sm = mx.abs(S).max().item() if i % 4 == 0 or not np.isfinite(sm) or sm > 1e6: print(f" chunk {i:2d}: |S|max {sm:.3e} |v_new|max {mx.abs(v_new).max().item():.3e}") if not np.isfinite(sm): break def probe_hoard(mlp, xn): x32 = xn.reshape(-1, xn.shape[-1]) idx, gates, bal = mlp.route(x32) mstat("router gates", gates) print(f" bal {bal.item():.3e} idx range [{idx.min().item()}, {idx.max().item()}]") y = mlp(xn) mstat("hoard out", y) def probe_forward(model, x): """Re-run forward capturing intermediate norms; descend into first bad block.""" cfg = model.cfg e = model.embed(x) h = e print(f" embed: max|e| {mx.abs(e).max().item():.3e}") descended = False for l in range(cfg.n_loops): h = h + e + model.loop_embed[l] cell = model.cells[0] blocks = [("mixer", cell.mixer, cell.n1)] outs = [] for bname, mod, norm in blocks: xn = norm(h) o = mod(xn) outs.append((bname, o)) if not finite(o) and not descended: descended = True print(f" -> descending into {bname} at loop {l}") mstat("input h", h); mstat("normed", xn) probe_mixer(mod, xn) h = h + o p1x = cell.n2(h) p1 = cell.mlp1(p1x) if not finite(p1) and not descended: descended = True print(f" -> descending into mlp1 at loop {l}") probe_hoard(cell.mlp1, p1x) h = h + p1 parts = [f"{n} {mx.abs(o).max().item():.3e}" for n, o in outs] + [f"mlp1 {mx.abs(p1).max().item():.3e}"] if cfg.use_window_attn: at = cell.attn(cell.n3(h)) h = h + at p2x = cell.n4(h) p2 = cell.mlp2(p2x) if not finite(p2) and not descended: descended = True print(f" -> descending into mlp2 at loop {l}") probe_hoard(cell.mlp2, p2x) h = h + p2 parts += [f"attn {mx.abs(at).max().item():.3e}", f"mlp2 {mx.abs(p2).max().item():.3e}"] print(f" loop {l}: max|h| {mx.abs(h).max().item():.3e} | " + " | ".join(parts)) logits = model.logits(h) print(f" logits: max {mx.abs(logits).max().item():.3e}") def main(): ap = argparse.ArgumentParser() ap.add_argument("--variant", default="full", choices=["full", "nobal", "dense", "freeze_embed", "f32router"]) ap.add_argument("--steps", type=int, default=8) a = ap.parse_args() mx.random.seed(0) cfg = configs.get("hoard_small") if a.variant == "nobal": cfg.hoard_balance_coef = 0.0 if a.variant == "dense": cfg.mlp = "dense" model = HOARD(cfg) model.set_dtype(mx.bfloat16) loader = TokenLoader("data/shakespeare/train.bin", 1024, 32, 0) muon = fp32_master(optim.Muon)(learning_rate=0.02, weight_decay=0.0) adamw = fp32_master(optim.AdamW)(learning_rate=3e-3, betas=(0.9, 0.95), weight_decay=0.0) filters = [is_muon_param] if a.variant == "freeze_embed": zero = fp32_master(optim.SGD)(learning_rate=0.0) filters = [is_muon_param, lambda p, x: ("embed" in p or "head" in p)] opt = optim.MultiOptimizer([muon, zero, adamw], filters) else: opt = optim.MultiOptimizer([muon, adamw], filters) def loss_fn(model, batch): x, y = batch[:, :-1], batch[:, 1:] logits = model(x) ce = nn.losses.cross_entropy(logits.astype(mx.float32), y, reduction="mean") return ce + model.balance_loss(), ce grad_fn = nn.value_and_grad(model, loss_fn) for step in range(a.steps): batch = loader.next() (loss, ce), grads = grad_fn(model, batch) mx.eval(loss, grads) lv = loss.item() print(f"step {step}: loss {lv:.4f} ce {ce.item():.4f}") bad_g = scan(grads, "grads") if not np.isfinite(lv) or bad_g: print(" -> probing forward intermediates on this batch:") probe_forward(model, batch[:, :-1]) scan(model.parameters(), "params") break grads, _ = optim.clip_grad_norm(grads, 1.0) opt.update(model, grads) mx.eval(model.parameters(), opt.state) scan(model.parameters(), "params-after-update") if __name__ == "__main__": main()