Gala-598M-MLX / debug_nan.py
junafinity's picture
Gala-598M: weights, code, logs, report
76bbe95 verified
Raw History Blame Contribute Delete
7.96 kB
"""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()