Skip to content

Commit 61d235f

Browse files
committed
feat(v4): GDN Path E — vendor mlx-lm PR #1217 gated_delta_update
Verbatim copy of upstream gated_delta.py (MIT, Apple Inc.) plus a thin adapter that maps our (q, k, v, beta, g) signature onto upstream's (q, k, v, a, b, A_log, dt_bias): - g (log-decay, FLA convention) -> a = softplus_inverse(-g) so compute_g(A_log=0, dt_bias=0, a) = exp(g). Upstream parameterization only represents decay <= 1, so g is clamped at 0. - beta -> b = logit(beta) (upstream applies sigmoid(b)). - q pre-scaled by 1/sqrt(K) (FLA scales internally; upstream doesn't). - initial_state [B,H,K,V] transposed to upstream's [B,Hv,Dv,Dk]. Adapter falls back to upstream's pure-ops path when Dk%32!=0 or Dv%4!=0 (the Metal kernel's compile-time constraints). Path E now reports available=True in the dispatch table; auto_pick will select it ahead of Path A when the env doesn't pin a specific backend. Tests: 5 new in tests/v4/test_linear_attention_path_e.py — status, shape, parity vs Path A, env-forced dispatch, final-state return. Full v4 suite 124 passed / 10 skipped. Also: switch the deferred-fixture in test_benchmark_receipt.py from Path C to Path D — Path C's status reason varies by env (depends on whether tilelang imports), while Path D is uniformly unavailable on Apple Silicon (no Triton).
1 parent c119c35 commit 61d235f

5 files changed

Lines changed: 521 additions & 7 deletions

File tree

cppmega_v4/_tilelang/linear_attention_paths.py

Lines changed: 15 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -19,9 +19,14 @@
1919
- Path D: Triton frontend via ``tilelang.poc.triton_frontend.from_triton_kernel``
2020
on FLA's ``chunk_gated_delta_rule`` — scaffold; awaits frontend op
2121
coverage for the FLA kernel.
22-
- Path E: vendored mlx-lm ``gated_delta_update`` op (PR #1217) — scaffold;
23-
awaits cherry-pick + vendoring under
24-
``cppmega_v4/nn/_external/mlx_lm_gated_delta_update.py``.
22+
- Path E: vendored mlx-lm ``gated_delta_update`` op (PR #1217) —
23+
verbatim copy under
24+
``cppmega_v4/nn/_external/_mlx_lm_gated_delta_vendored.py`` with adapter
25+
``cppmega_v4/nn/_external/mlx_lm_gated_delta_update.py`` mapping our
26+
(q, k, v, beta, g) → upstream (q, k, v, a, b, A_log, dt_bias) by
27+
softplus_inverse(-g) for the gate and logit(beta) for the betas.
28+
Upstream Metal kernel needs Dk%32==0 & Dv%4==0; smaller dims fall back
29+
to the upstream ops path automatically.
2530
2631
When the user wants to validate against a specific backend they set
2732
``CPPMEGA_V4_KERNEL_PATH__LINEAR_ATTENTION=path_c`` (or the path of choice);
@@ -151,7 +156,13 @@ def _path_e_status() -> PathStatus:
151156
importlib.import_module(
152157
"cppmega_v4.nn._external.mlx_lm_gated_delta_update"
153158
)
154-
return PathStatus(path="path_e", available=True, reason="vendored mlx-lm op present")
159+
return PathStatus(
160+
path="path_e", available=True,
161+
reason=(
162+
"vendored mlx-lm PR #1217 gated_delta_update (Metal kernel "
163+
"for Dk%32==0 & Dv%4==0; ops fallback otherwise)"
164+
),
165+
)
155166
except Exception:
156167
return PathStatus(
157168
path="path_e",
Lines changed: 323 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,323 @@
1+
# Vendored verbatim from ml-explore/mlx-lm PR #1217
2+
# (mlx_lm/models/gated_delta.py)
3+
# Upstream license: MIT, Copyright © 2023 Apple Inc.
4+
#
5+
# Public surface used by cppmega_v4 Path E:
6+
# - gated_delta_update(...) -> (y, state)
7+
# - gated_delta_kernel(...) — low-level Metal kernel entry
8+
# - gated_delta_ops(...) — pure-MLX reference for prefill
9+
# - compute_g(A_log, a, dt_bias) — gate transform
10+
# No edits — pure copy. cppmega_v4/nn/path_e_adapter.py wraps our (g, beta)
11+
# API to this module's (a, b, A_log, dt_bias) API.
12+
13+
from functools import partial
14+
from typing import Optional, Tuple
15+
16+
import mlx.core as mx
17+
import mlx.nn as nn
18+
19+
20+
@partial(mx.compile, shapeless=True)
21+
def compute_g(A_log, a, dt_bias):
22+
return mx.exp(-mx.exp(A_log.astype(mx.float32)) * nn.softplus(a + dt_bias))
23+
24+
25+
def _make_gated_delta_kernel(has_mask=False, vectorized=False):
26+
if not mx.metal.is_available():
27+
return None
28+
mask_source = "mask[b_idx * T + t]" if has_mask else "true"
29+
30+
# Configure g indexing based on whether gating is vectorized
31+
if vectorized:
32+
g_comment = "// g: [B, T, Hv, Dk]"
33+
g_setup = "auto g_ = g + (b_idx * T * Hv + hv_idx) * Dk;"
34+
g_access = "g_[s_idx]"
35+
g_advance = "g_ += Hv * Dk;"
36+
else:
37+
g_comment = "// g: [B, T, Hv]"
38+
g_setup = "auto g_ = g + b_idx * T * Hv;"
39+
g_access = "g_[hv_idx]"
40+
g_advance = "g_ += Hv;"
41+
42+
source = f"""
43+
auto n = thread_position_in_grid.z;
44+
auto b_idx = n / Hv;
45+
auto hv_idx = n % Hv;
46+
auto hk_idx = hv_idx / (Hv / Hk);
47+
constexpr int n_per_t = Dk / 32;
48+
49+
// q, k: [B, T, Hk, Dk]
50+
auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk;
51+
auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk;
52+
53+
// v, y: [B, T, Hv, Dv]
54+
auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv;
55+
y += b_idx * T * Hv * Dv + hv_idx * Dv;
56+
57+
auto dk_idx = thread_position_in_threadgroup.x;
58+
auto dv_idx = thread_position_in_grid.y;
59+
60+
// state_in, state_out: [B, Hv, Dv, Dk]
61+
auto i_state = state_in + (n * Dv + dv_idx) * Dk;
62+
auto o_state = state_out + (n * Dv + dv_idx) * Dk;
63+
64+
float state[n_per_t];
65+
for (int i = 0; i < n_per_t; ++i) {{
66+
auto s_idx = n_per_t * dk_idx + i;
67+
state[i] = static_cast<float>(i_state[s_idx]);
68+
}}
69+
70+
{g_comment}
71+
{g_setup}
72+
auto beta_ = beta + b_idx * T * Hv;
73+
74+
for (int t = 0; t < T; ++t) {{
75+
if ({mask_source}) {{
76+
float kv_mem = 0.0f;
77+
for (int i = 0; i < n_per_t; ++i) {{
78+
auto s_idx = n_per_t * dk_idx + i;
79+
state[i] = state[i] * {g_access};
80+
kv_mem += state[i] * k_[s_idx];
81+
}}
82+
kv_mem = simd_sum(kv_mem);
83+
84+
auto delta = (v_[dv_idx] - kv_mem) * beta_[hv_idx];
85+
86+
float out = 0.0f;
87+
for (int i = 0; i < n_per_t; ++i) {{
88+
auto s_idx = n_per_t * dk_idx + i;
89+
state[i] = state[i] + k_[s_idx] * delta;
90+
out += state[i] * q_[s_idx];
91+
}}
92+
out = simd_sum(out);
93+
if (thread_index_in_simdgroup == 0) {{
94+
y[dv_idx] = static_cast<InT>(out);
95+
}}
96+
}} else {{
97+
y[dv_idx] = static_cast<InT>(0);
98+
}}
99+
// Increment data pointers to next time step
100+
q_ += Hk * Dk;
101+
k_ += Hk * Dk;
102+
v_ += Hv * Dv;
103+
y += Hv * Dv;
104+
{g_advance}
105+
beta_ += Hv;
106+
}}
107+
for (int i = 0; i < n_per_t; ++i) {{
108+
auto s_idx = n_per_t * dk_idx + i;
109+
o_state[s_idx] = static_cast<StT>(state[i]);
110+
}}
111+
"""
112+
inputs = ["q", "k", "v", "g", "beta", "state_in", "T"]
113+
if has_mask:
114+
inputs.append("mask")
115+
116+
suffix = ""
117+
if vectorized:
118+
suffix += "_vec"
119+
if has_mask:
120+
suffix += "_mask"
121+
122+
return mx.fast.metal_kernel(
123+
name=f"gated_delta_step{suffix}",
124+
input_names=inputs,
125+
output_names=["y", "state_out"],
126+
source=source,
127+
)
128+
129+
130+
_gated_delta_kernel = _make_gated_delta_kernel(has_mask=False, vectorized=False)
131+
_gated_delta_kernel_masked = _make_gated_delta_kernel(has_mask=True, vectorized=False)
132+
_gated_delta_kernel_vec = _make_gated_delta_kernel(has_mask=False, vectorized=True)
133+
_gated_delta_kernel_vec_masked = _make_gated_delta_kernel(
134+
has_mask=True, vectorized=True
135+
)
136+
137+
138+
@mx.compile
139+
def _gated_delta_step_ops(
140+
q: mx.array,
141+
k: mx.array,
142+
v: mx.array,
143+
g: mx.array,
144+
beta: mx.array,
145+
state: mx.array,
146+
mask: Optional[mx.array] = None,
147+
) -> Tuple[mx.array, mx.array]:
148+
"""
149+
Ops-based reference implementation for a single recurrent step.
150+
151+
Shapes:
152+
- q, k: [B, H, Dk]
153+
- v: [B, H, Dv]
154+
- g: [B, H] or [B, H, Dk]
155+
- beta: [B, H]
156+
- state: [B, H, Dv, Dk]
157+
Returns:
158+
- y: [B, H, Dv]
159+
- new_state: [B, H, Dv, Dk]
160+
"""
161+
162+
# Decay
163+
old_state = state
164+
if g.ndim == 2:
165+
decay = g[..., None, None]
166+
elif g.ndim == 3:
167+
decay = g[..., None, :]
168+
else:
169+
raise ValueError(f"Unsupported gating shape {g.shape}")
170+
state = state * decay
171+
kv_mem = (state * k[..., None, :]).sum(axis=-1) # [B, H, Dv]
172+
delta = (v - kv_mem) * beta[..., None] # [B, H, Dv]
173+
state = state + k[..., None, :] * delta[..., None]
174+
# Output projection along key dim with q
175+
y = (state * q[..., None, :]).sum(axis=-1) # [B, H, Dv]
176+
177+
if mask is not None:
178+
mask = mx.expand_dims(mask, axis=(1, 2, 3))
179+
state = mx.where(mask, state, old_state)
180+
return y.astype(q.dtype), state
181+
182+
183+
def gated_delta_kernel(
184+
q: mx.array,
185+
k: mx.array,
186+
v: mx.array,
187+
g: mx.array,
188+
beta: mx.array,
189+
state: mx.array,
190+
mask: Optional[mx.array] = None,
191+
) -> Tuple[mx.array, mx.array]:
192+
B, T, Hk, Dk = k.shape
193+
Hv, Dv = v.shape[2:]
194+
input_type = q.dtype
195+
state_type = state.dtype
196+
if g.ndim == 4:
197+
kernel = _gated_delta_kernel_vec
198+
inputs = [q, k, v, g, beta, state, T]
199+
if mask is not None:
200+
kernel = _gated_delta_kernel_vec_masked
201+
inputs.append(mask)
202+
else:
203+
kernel = _gated_delta_kernel
204+
inputs = [q, k, v, g, beta, state, T]
205+
if mask is not None:
206+
kernel = _gated_delta_kernel_masked
207+
inputs.append(mask)
208+
209+
return kernel(
210+
inputs=inputs,
211+
template=[
212+
("InT", input_type),
213+
("StT", state_type),
214+
("Dk", Dk),
215+
("Dv", Dv),
216+
("Hk", Hk),
217+
("Hv", Hv),
218+
],
219+
grid=(32, Dv, B * Hv),
220+
threadgroup=(32, 4, 1),
221+
output_shapes=[(B, T, Hv, Dv), state.shape],
222+
output_dtypes=[input_type, state_type],
223+
)
224+
225+
226+
def gated_delta_ops(
227+
q: mx.array,
228+
k: mx.array,
229+
v: mx.array,
230+
g: mx.array,
231+
beta: mx.array,
232+
state: Optional[mx.array] = None,
233+
mask: Optional[mx.array] = None,
234+
) -> Tuple[mx.array, mx.array]:
235+
"""
236+
Ops-based reference implementation for prompt prefill (sequential loop).
237+
Supports both scalar and vectorized gating.
238+
239+
Shapes:
240+
- q, k: [B, T, Hk, Dk]
241+
- v: [B, T, Hv, Dv]
242+
- g: [B, T, Hv] (scalar) or [B, T, Hv, Dk] (vectorized)
243+
- beta: [B, T, Hv]
244+
- state: [B, Hv, Dv, Dk]
245+
Returns:
246+
- y: [B, T, Hv, Dv]
247+
- state: [B, Hv, Dv, Dk]
248+
"""
249+
B, T, Hk, Dk = q.shape
250+
Hv, Dv = v.shape[-2:]
251+
if state is None:
252+
state = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32)
253+
254+
if (repeat_factor := Hv // Hk) > 1:
255+
q = mx.repeat(q, repeat_factor, -2)
256+
k = mx.repeat(k, repeat_factor, -2)
257+
258+
ys = []
259+
for t in range(T):
260+
y, state = _gated_delta_step_ops(
261+
q[:, t],
262+
k[:, t],
263+
v[:, t],
264+
g[:, t],
265+
beta[:, t],
266+
state,
267+
None if mask is None else mask[:, t],
268+
)
269+
ys.append(y)
270+
y = mx.stack(ys, axis=1)
271+
return y, state
272+
273+
274+
def gated_delta_update(
275+
q: mx.array,
276+
k: mx.array,
277+
v: mx.array,
278+
a: mx.array,
279+
b: mx.array,
280+
A_log: mx.array,
281+
dt_bias: mx.array,
282+
state: Optional[mx.array] = None,
283+
mask: Optional[mx.array] = None,
284+
use_kernel: bool = True,
285+
training: bool = False,
286+
) -> Tuple[mx.array, mx.array]:
287+
if training:
288+
# Chunked VJP path with O(T/chunk) autodiff graph — fits T≥2048
289+
# on 36 GB Apple Silicon where the Python-ops path OOMs.
290+
# Metal backward kernel is 8–11× faster than the Python reference,
291+
# but only handles the unmasked GPU path with Dk%32==0 and Dv%4==0.
292+
Dk = q.shape[-1]
293+
Dv = v.shape[-1]
294+
can_use_metal = (
295+
mx.metal.is_available()
296+
and mx.default_device() == mx.gpu
297+
and mask is None
298+
and Dk % 32 == 0
299+
and Dv % 4 == 0
300+
)
301+
if can_use_metal:
302+
try:
303+
from .gated_delta_vjp_metal import gated_delta_update_vjp_metal
304+
305+
return gated_delta_update_vjp_metal(
306+
q, k, v, a, b, A_log, dt_bias, state, mask
307+
)
308+
except ImportError:
309+
pass
310+
from .gated_delta_vjp import gated_delta_update_vjp
311+
312+
return gated_delta_update_vjp(q, k, v, a, b, A_log, dt_bias, state, mask)
313+
314+
beta = mx.sigmoid(b)
315+
g = compute_g(A_log, a, dt_bias)
316+
if state is None:
317+
B, _, Hk, Dk = q.shape
318+
Hv, Dv = v.shape[-2:]
319+
state = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32)
320+
321+
if not use_kernel or mx.default_device() != mx.gpu or not mx.metal.is_available():
322+
return gated_delta_ops(q, k, v, g, beta, state, mask)
323+
return gated_delta_kernel(q, k, v, g, beta, state, mask)

0 commit comments

Comments
 (0)