Text Generation
MLX
English
apple-silicon
pretrained-from-scratch
gated-deltanet
linear-attention
product-key-memory
long-context
Instructions to use junafinity/Gala-598M-MLX with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use junafinity/Gala-598M-MLX with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # if on a CUDA device, also pip install mlx[cuda] # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("junafinity/Gala-598M-MLX") prompt = "Once upon a time in" text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- MLX LM
How to use junafinity/Gala-598M-MLX with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Generate some text mlx_lm.generate --model "junafinity/Gala-598M-MLX" --prompt "Once upon a time"
- Atomic Chat
Download debug_nan.py from junafinity/Gala-598M-MLX: direct link, hf CLI and curl.
- Browser
- Download file 7.96 kB
-
https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/debug_nan.py
- Command line
-
hf download hf://junafinity/Gala-598M-MLX/debug_nan.py
-
curl -L -o debug_nan.py https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/debug_nan.py
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() | |