Skip to content

mx.fast.gated_delta_update VJP: wrong gradient w.r.t. g (and ~1e-3 error on all inputs) on the NAX path for Hk=16 #4614

Description

@kevincaldwellgordon-bot

The fused VJP added in #4565 returns a badly wrong gradient for the decay input g when 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).

python mlx_issue_repro.py --T 128 --heads 16,48 --decay strong

Relative L2 error vs float64 (T=128, D=128, B=1, float32 inputs, h0=None):

Hk/Hv decay (mean log g) input fused VJP float32 ops recurrence
16/48 strong (−8.3) dg 8.8e+02 8.2e-07
16/48 strong dq / dk / dv / dbeta 1.8e-03 1.2e-07 – 2.2e-07
16/48 weak (−0.15) dg 2.2e-03 8.4e-07
16/16 strong (−8.5) dg 7.4e+02 8.1e-07
4/12 strong (−8.5) all 1.2e-07 – 8.9e-07 (= ops) same

Recurrence used for the references: the mlx_lm ops path. Per value head, with kv head h // (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 dg error is large enough to break training through a log_g parameterisation such as Qwen3.5/3.8's A_log/dt_bias. Through that chain we measured a da rel 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
"""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()

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions