"""Standalone repro for an upstream MLX issue: mx.fast.gated_delta_update VJP, gradient w.r.t. the decay g.
Only mlx + numpy, random inputs. Compares, per input gradient, against a float64 sequential reference (CPU):
fast mx.fast.gated_delta_update VJP (GPU, float32)
ops32 the same sequential recurrence in plain mx ops (GPU, float32): what float32 rounding alone costs
Recurrence (mlx_lm gated_delta_update ops path), per value head h, kv head h // (Hv // Hk):
S <- g_t * S ; S <- S + outer(beta_t * (v_t - S k_t), k_t) ; y_t = S q_t (S: [Dv, Dk], S0 = 0)
Usage: python mlx_issue_repro.py [--T 256] [--heads 16,48] [--D 128] [--decay strong|weak] [--seed 0]
"""
import argparse
import json
import mlx.core as mx
import numpy as np
def recurrence(q, k, v, g, beta, Hk, Hv):
B, T, _, D = q.shape
rep = Hv // Hk
qh, kh = mx.repeat(q, rep, axis=2), mx.repeat(k, rep, axis=2) # [B, T, Hv, D]
S = mx.zeros((B, Hv, v.shape[-1], D), dtype=q.dtype)
ys = []
for t in range(T):
S = S * g[:, t, :, None, None]
kv = (S * kh[:, t, :, None, :]).sum(-1) # [B, Hv, Dv]
delta = (v[:, t] - kv) * beta[:, t, :, None]
S = S + delta[..., None] * kh[:, t, :, None, :]
ys.append((S * qh[:, t, :, None, :]).sum(-1))
return mx.stack(ys, axis=1) # [B, T, Hv, Dv]
def rel(a, b):
a, b = np.asarray(a, np.float64).ravel(), np.asarray(b, np.float64).ravel()
return float(np.linalg.norm(a - b) / max(np.linalg.norm(b), 1e-300))
def main():
p = argparse.ArgumentParser()
p.add_argument("--T", type=int, default=256)
p.add_argument("--heads", default="16,48")
p.add_argument("--D", type=int, default=128)
p.add_argument("--decay", choices=("strong", "weak"), default="strong")
p.add_argument("--seed", type=int, default=0)
a = p.parse_args()
Hk, Hv = (int(x) for x in a.heads.split(","))
rng = np.random.default_rng(a.seed)
B, T, D = 1, a.T, a.D
def unit(x):
return x / np.linalg.norm(x, axis=-1, keepdims=True)
q = unit(rng.standard_normal((B, T, Hk, D))) * D ** -0.5
k = unit(rng.standard_normal((B, T, Hk, D)))
v = rng.standard_normal((B, T, Hv, D))
log_g = -np.exp(rng.normal(2.0 if a.decay == "strong" else -2.0, 0.5, (B, T, Hv)))
g = np.exp(log_g)
beta = 1 / (1 + np.exp(-rng.standard_normal((B, T, Hv))))
dy = rng.standard_normal((B, T, Hv, D))
names = ("dq", "dk", "dv", "dg", "dbeta")
def grads(fn, dtype, device):
ins = [mx.array(x.astype(dtype)) for x in (q, k, v, g, beta)]
with mx.stream(device):
_, vjps = mx.vjp(lambda *xs: [fn(*xs)], ins, [mx.array(dy.astype(dtype))])
mx.eval(vjps)
return [np.array(x) for x in vjps]
ref = grads(lambda *xs: recurrence(*xs, Hk, Hv), np.float64, mx.cpu)
ops = grads(lambda *xs: recurrence(*xs, Hk, Hv), np.float32, mx.gpu)
fast = grads(lambda q_, k_, v_, g_, b_: mx.fast.gated_delta_update(q_, k_, v_, g_, b_, None)[0],
np.float32, mx.gpu)
out = {"mlx": mx.__version__, "T": T, "Hk": Hk, "Hv": Hv, "D": D, "decay": a.decay,
"mean_log_g": float(log_g.mean()),
"rel_l2_vs_float64": {n: {"fast": rel(f, r), "ops32": rel(o, r)} for n, f, o, r in zip(names, fast, ops, ref)}}
print(json.dumps(out, indent=1))
if __name__ == "__main__":
main()
The fused VJP added in #4565 returns a badly wrong gradient for the decay input
gwhen the call takes the chunked NAX kernel (here Hk=16), and all other input gradients there are ~1e4× less accurate than the same recurrence written in plain float32 ops. Shapes that do not take that kernel (Hk=4, Hv=12) match the ops path exactly.Repro. Self-contained: mlx + numpy, random inputs. It compares each input gradient with a float64 sequential reference computed on the CPU. Script below (~90 lines).
Relative L2 error vs float64 (T=128, D=128, B=1, float32 inputs,
h0=None):Recurrence used for the references: the
mlx_lmops path. Per value head, with kv headh // (Hv/Hk),S = g_t*S; S += outer(beta_t*(v_t - S k_t), k_t); y_t = S q_t.Under strong decay (g ≈ 2e-4) the
dgerror is large enough to break training through alog_gparameterisation such as Qwen3.5/3.8'sA_log/dt_bias. Through that chain we measured adarel L2 of 5.7e-2 against 1.8e-3 for a chunked float32 implementation.Environment: mlx 0.32.4 (nightly "Release build" from commit 5c89fa1, which contains #4565), macOS 26.6, Apple M5 Pro.
mlx_issue_repro.py