Skip to content

Commit 40a83ef

Browse files
committed
feat(v4): real MSL backward kernel for KDA Path B (8.8-15.9× speedup)
Replaces the fwd-fast/bwd-via-Path-A wrapper in kda_path_b_bwd.py with a hand-MSL Metal kernel mirroring the GDN Path B bwd design: - Snapshot strategy: forward replay writes per-step S_t[:,vj] into a device-memory workspace state_hist[B*HV, T+1, K, V]. Reverse pass reads S_t and S_{t-1} directly — no division by exp(g) (vectorized per-K, decay < 1, inverse-walk is numerically catastrophic — same blocker as mamba3_path_c's inverse-walk; snapshot fix mirrors that). - Per-thread register cache of inner_t[T] from forward so dv / dkth use exact values, not algebraic reconstruction. - One thread per (b, hv, vj) lane; one 32-lane simdgroup per (b, hv). simd_sum reduces over the vj axis for dq[i], dk[i], dbeta, ddecay[i]. - dq / dk / dv / dbeta written via atomic_fetch_add_explicit because multiple hv lanes share the same (b, t, h_idx) row of q/k when HV > H (KDA's GQA-style grouping); also dv accumulates from a single source but the atomic guards future tile splits. - dg is per-K (g has shape [B, T, HV, K]) — written directly per-i at each step (no accumulation across hv). KDA-specific algebra (vs GDN): - g vectorized: dg_t[i] = ddecay[i] * decay_t[i] per-K. - dbeta_t = sum_i k_i * (sum_j dS_t[i,j] * inner_t) — factored to pull k_i out of the j-reduction (saves one simd_sum vs naive). Fallback to mx.grad through naive_recurrent_kda for: - V > 32 (simdgroup width) - HV % H != 0 - shape inconsistencies (validated before calling kernel). Speedup vs Path A grad path (B=1, H=2, HV=4, K=16, V=16): T=64: PathA 8.12 ms → PathB-MSL 0.93 ms (8.8×) T=128: PathA 16.40 ms → PathB-MSL 1.43 ms (11.4×) T=256: PathA 31.23 ms → PathB-MSL 1.97 ms (15.9×) Tests: tests/v4/test_kda_path_b_bwd.py — all 3 pass at the original atol=1e-4 / rtol=1e-3 tolerance (forward parity, backward grad parity across (dq, dk, dv, dg, dbeta), finite-grads smoke). v4 suite: 285 passed / 2 skipped.
1 parent 025c761 commit 40a83ef

1 file changed

Lines changed: 360 additions & 8 deletions

File tree

Lines changed: 360 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,361 @@
1-
"""KDA Path B forward + backward — same fwd-fast / bwd-correct pattern as GDN."""
1+
"""KDA Path B forward + backward — fwd via fast Metal kernel, bwd via
2+
hand-MSL Metal kernel (real, fused recurrent backward).
3+
4+
Mirrors the GDN Path B bwd pattern in ``linear_attention_path_b_bwd.py``:
5+
forward replay snapshots S_t per j-column into a device-memory workspace
6+
(``state_hist[B*HV, T+1, K, V]``), and the reverse-time scan reads
7+
``S_t`` / ``S_{t-1}`` directly — never divides by ``decay = exp(g)``
8+
(``decay`` ≤ 1 per-K makes inverse-walk numerically catastrophic; this
9+
mirrors the mamba3_path_c switch from inverse-walk to cached snapshots).
10+
11+
KDA-specific bits vs GDN:
12+
- Gate ``g`` is per-K vector ``[B, T, HV, K]``, so ``dg`` is per-K.
13+
- ``v`` and the output share the V axis (not K); each thread owns one
14+
``vj`` column. Constraint ``V <= 32`` (= simd width) so ``simd_sum``
15+
over j fits one instruction.
16+
- ``q``/``k`` are ``[B, T, H, K]`` with HV groups expanded; head
17+
indexing is ``h_idx = hv_idx // (HV/H)``.
18+
- ``q`` is pre-scaled by ``1/sqrt(K)`` (FLA convention).
19+
20+
Backward algebra (derived from the forward in ``kda_path_b.py``):
21+
22+
Forward (with q' = q * scale, scale = 1/sqrt(K)):
23+
decay_t[i] = exp(g_t[i])
24+
S_decayed[i,j] = decay_t[i] * S_{t-1}[i,j]
25+
kth_t[j] = sum_i k_t[i] * S_decayed[i,j]
26+
inner_t[j] = v_t[j] - kth_t[j]
27+
S_t[i,j] = S_decayed[i,j] + beta_t * k_t[i] * inner_t[j]
28+
o_t[j] = sum_i q'_t[i] * S_t[i,j]
29+
30+
Backward (reverse t):
31+
dq'_i += sum_j dO[j] * S_t[i,j]
32+
dS_t[i,j] += dO[j] * q'_i
33+
dv[j] = dinner[j]
34+
dkth[j] = -dinner[j]
35+
dinner[j] = beta_t * sum_i dS_t[i,j] * k_t[i]
36+
dk_i (delta) += sum_j dS_t[i,j] * (beta_t * inner_t[j])
37+
dk_i (kth) += sum_j dkth[j] * S_decayed[i,j]
38+
dbeta_t += sum_{i,j} dS_t[i,j] * k_t[i] * inner_t[j]
39+
= sum_i k_t[i] * (sum_j dS_t[i,j] * inner_t[j])
40+
dS_decayed[i,j] = dS_t[i,j] + dkth[j] * k_t[i]
41+
ddecay[i] = sum_j dS_decayed[i,j] * S_{t-1}[i,j] (per-K)
42+
dS_{t-1}[i,j] = dS_decayed[i,j] * decay_t[i]
43+
dg_t[i] = ddecay[i] * decay_t[i] (per-K)
44+
45+
Falls back to ``mx.grad`` through ``naive_recurrent_kda`` for:
46+
- ``V > 32`` (multi-simdgroup not yet implemented)
47+
- ``initial_state`` provided
48+
- ``HV % H != 0``
49+
- any future shape outside the kernel's domain.
50+
"""
51+
52+
from __future__ import annotations
253

354
import mlx.core as mx
455

56+
from cppmega_v4._tilelang._kernel_cache import get_or_build_kernel
557
from cppmega_v4._tilelang.kda_path_b import kda_forward_path_b
658
from cppmega_v4.nn._external.fla_naive_kda import naive_recurrent_kda
759

860

61+
_SIMD_WIDTH = 32
62+
63+
64+
def _kda_backward_kernel(
65+
q: mx.array,
66+
k: mx.array,
67+
v: mx.array,
68+
g: mx.array,
69+
beta: mx.array,
70+
dy: mx.array,
71+
) -> tuple[mx.array, mx.array, mx.array, mx.array, mx.array]:
72+
"""Real Metal backward for the KDA recurrence.
73+
74+
Returns float32 grads ``(dq, dk, dv, dg, dbeta)``. Shapes match inputs:
75+
dq/dk: [B, T, H, K]
76+
dv: [B, T, HV, V]
77+
dg: [B, T, HV, K]
78+
dbeta: [B, T, HV]
79+
"""
80+
if q.ndim != 4 or k.shape != q.shape:
81+
raise ValueError(
82+
f"q/k must match shape [B, T, H, K]; got q={q.shape}, k={k.shape}"
83+
)
84+
if v.ndim != 4 or v.shape[:2] != q.shape[:2]:
85+
raise ValueError(f"v must be [B, T, HV, V]; got v={v.shape}")
86+
if g.shape != (*v.shape[:3], k.shape[-1]):
87+
raise ValueError(f"g must be [B, T, HV, K]; got g={g.shape}")
88+
if beta.shape != v.shape[:3]:
89+
raise ValueError(f"beta must be [B, T, HV]; got beta={beta.shape}")
90+
if dy.shape != v.shape:
91+
raise ValueError(f"dy must match v shape; got dy={dy.shape}, v={v.shape}")
92+
93+
b, t, h, kdim = q.shape
94+
hv, vdim = v.shape[2], v.shape[-1]
95+
if hv % h != 0:
96+
raise ValueError(f"HV ({hv}) must be divisible by H ({h})")
97+
if vdim > _SIMD_WIDTH:
98+
raise ValueError(
99+
f"Real-MSL KDA bwd currently requires V<=32 (got {vdim}); "
100+
f"caller should fall back to Path A grad path"
101+
)
102+
group = hv // h
103+
104+
scale = kdim ** -0.5
105+
q_f = q.astype(mx.float32).reshape(-1)
106+
k_f = k.astype(mx.float32).reshape(-1)
107+
v_f = v.astype(mx.float32).reshape(-1)
108+
g_f = g.astype(mx.float32).reshape(-1)
109+
beta_f = beta.astype(mx.float32).reshape(-1)
110+
dy_f = dy.astype(mx.float32).reshape(-1)
111+
112+
source = f"""
113+
uint tid_in_tg = thread_position_in_threadgroup.x;
114+
uint bhv = threadgroup_position_in_grid.x;
115+
uint vj = tid_in_tg;
116+
bool active = (vj < {vdim}u) && (bhv < {b * hv}u);
117+
118+
uint bb = bhv / {hv}u;
119+
uint hv_idx = bhv % {hv}u;
120+
uint h_idx = hv_idx / {group}u;
121+
122+
// Per-thread registers:
123+
// state[K], dS[K], inner_hist[T]
124+
// Device-memory workspace (extra kernel output):
125+
// state_hist[B*HV, T+1, K, V]
126+
float state[{kdim}];
127+
float dS[{kdim}];
128+
float inner_hist[{max(t, 1)}];
129+
for (int i = 0; i < {kdim}; i++) {{ state[i] = 0.0f; dS[i] = 0.0f; }}
130+
131+
int hist_bh_stride = {(t + 1) * kdim * vdim};
132+
int hist_ti_stride = {kdim * vdim};
133+
134+
// Snapshot t=0 = zero state for this column.
135+
for (int i = 0; i < {kdim}; i++) {{
136+
int idx = bhv * hist_bh_stride + 0 * hist_ti_stride + i * {vdim} + (int)vj;
137+
if (active) state_hist[idx] = 0.0f;
138+
}}
139+
140+
// ============================================================
141+
// Forward replay: rebuild final state, capture inner[t] per
142+
// (vj), snapshot S_t[:,vj] into device memory at every step.
143+
// ============================================================
144+
for (int ti = 0; ti < {t}; ti++) {{
145+
int g_base = ((bb * {t} + ti) * {hv} + hv_idx) * {kdim};
146+
int beta_idx = (bb * {t} + ti) * {hv} + hv_idx;
147+
int qk_base = ((bb * {t} + ti) * {h} + h_idx) * {kdim};
148+
int v_idx = ((bb * {t} + ti) * {hv} + hv_idx) * {vdim} + (int)vj;
149+
150+
float beta_t = beta[beta_idx];
151+
float v_j = active ? v[v_idx] : 0.0f;
152+
153+
// Per-K decay + interleaved KS reduction.
154+
float kth_j = 0.0f;
155+
for (int i = 0; i < {kdim}; i++) {{
156+
float decay_i = exp(g[g_base + i]);
157+
state[i] *= decay_i;
158+
kth_j += k[qk_base + i] * state[i];
159+
}}
160+
float inner_j = v_j - kth_j;
161+
inner_hist[ti] = inner_j;
162+
163+
// Rank-1 outer add: S[i, vj] += beta * k[i] * inner_j.
164+
for (int i = 0; i < {kdim}; i++) {{
165+
state[i] += beta_t * k[qk_base + i] * inner_j;
166+
}}
167+
168+
// Snapshot S_t at slot (ti+1) for this thread's column.
169+
for (int i = 0; i < {kdim}; i++) {{
170+
int idx = bhv * hist_bh_stride + (ti + 1) * hist_ti_stride + i * {vdim} + (int)vj;
171+
if (active) state_hist[idx] = state[i];
172+
}}
173+
}}
174+
175+
// ============================================================
176+
// Reverse-time backward scan.
177+
// ============================================================
178+
for (int rr = 0; rr < {t}; rr++) {{
179+
int ti = {t} - 1 - rr;
180+
int g_base = ((bb * {t} + ti) * {hv} + hv_idx) * {kdim};
181+
int beta_idx = (bb * {t} + ti) * {hv} + hv_idx;
182+
int qk_base = ((bb * {t} + ti) * {h} + h_idx) * {kdim};
183+
int v_idx = ((bb * {t} + ti) * {hv} + hv_idx) * {vdim} + (int)vj;
184+
185+
float beta_t = beta[beta_idx];
186+
float v_j = active ? v[v_idx] : 0.0f;
187+
float dY_j = active ? dy[v_idx] : 0.0f;
188+
float inner_t = inner_hist[ti];
189+
190+
// Read S_t (after step ti) and S_{{t-1}} (after step ti-1).
191+
float S_t_col[{kdim}];
192+
float S_prev[{kdim}];
193+
for (int i = 0; i < {kdim}; i++) {{
194+
int idx_t = bhv * hist_bh_stride + (ti + 1) * hist_ti_stride + i * {vdim} + (int)vj;
195+
int idx_tm1 = bhv * hist_bh_stride + ti * hist_ti_stride + i * {vdim} + (int)vj;
196+
S_t_col[i] = active ? state_hist[idx_t] : 0.0f;
197+
S_prev[i] = active ? state_hist[idx_tm1] : 0.0f;
198+
}}
199+
200+
// S_decayed[i, vj] = decay_t[i] * S_prev[i, vj]
201+
float decay_arr[{kdim}];
202+
float S_decayed[{kdim}];
203+
for (int i = 0; i < {kdim}; i++) {{
204+
decay_arr[i] = exp(g[g_base + i]);
205+
S_decayed[i] = decay_arr[i] * S_prev[i];
206+
}}
207+
208+
// ---- (1) o_t[j] = sum_i q'_i * S_t[i, j] ----
209+
// dq'_i = sum_j dY[j] * S_t[i, j] (reduce over j == vj axis)
210+
// dS_t[i, j] += dY[j] * q'_i
211+
// dq writes per-K to (b, t, h_idx, i); groups inside HV share q/k,
212+
// so multiple hv lanes in the same (b, t, h_idx) would race.
213+
// We restrict the dq write to the first hv in each group via
214+
// (hv_idx % group == 0) and multiply by group_size? No — we need
215+
// the SUM of grads from all hv in the group. Use atomic_fetch_add.
216+
for (int i = 0; i < {kdim}; i++) {{
217+
float q_i_scaled = q[qk_base + i] * {scale}f;
218+
float contrib = active ? (dY_j * S_t_col[i]) : 0.0f;
219+
float dq_i_sum = simd_sum(contrib);
220+
if (active && vj == 0u) {{
221+
atomic_fetch_add_explicit(
222+
(device atomic_float*)&dq[qk_base + i],
223+
dq_i_sum * {scale}f,
224+
memory_order_relaxed
225+
);
226+
}}
227+
dS[i] += dY_j * q_i_scaled;
228+
}}
229+
230+
// ---- (2) S_t = S_decayed + beta * k * inner ----
231+
// dinner[j] = beta_t * sum_i dS_t[i, j] * k_i (per-j scalar)
232+
// dk_i (delta) += sum_j dS_t[i, j] * (beta_t * inner_t)
233+
// dbeta_t += sum_{{i,j}} dS_t[i, j] * k_i * inner_t
234+
// = sum_i k_i * (sum_j dS_t[i, j] * inner_t)
235+
float sum_k_dS = 0.0f;
236+
for (int i = 0; i < {kdim}; i++) {{
237+
sum_k_dS += k[qk_base + i] * dS[i];
238+
}}
239+
float dinner_j = beta_t * sum_k_dS;
240+
241+
// dv[j] = dinner_j
242+
if (active) {{
243+
atomic_fetch_add_explicit(
244+
(device atomic_float*)&dv[v_idx],
245+
dinner_j,
246+
memory_order_relaxed
247+
);
248+
}}
249+
250+
// dbeta_t = sum_i k_i * (sum_j dS_t[i,j] * inner_t)
251+
float dbeta_partial = 0.0f;
252+
for (int i = 0; i < {kdim}; i++) {{
253+
float term = active ? (dS[i] * inner_t) : 0.0f;
254+
float term_sum = simd_sum(term); // sum over j
255+
if (vj == 0u) {{
256+
dbeta_partial += k[qk_base + i] * term_sum;
257+
}}
258+
}}
259+
if (active && vj == 0u) {{
260+
atomic_fetch_add_explicit(
261+
(device atomic_float*)&dbeta[beta_idx],
262+
dbeta_partial,
263+
memory_order_relaxed
264+
);
265+
}}
266+
267+
// dk_i (delta) += sum_j dS_t[i, j] * (beta_t * inner_t)
268+
// Note inner_t is a per-j scalar; cannot factor it out of simd_sum.
269+
// dk_i (kth) += sum_j dkth[j] * S_decayed[i, j]; dkth[j] = -dinner_j
270+
float dkth_j = -dinner_j;
271+
for (int i = 0; i < {kdim}; i++) {{
272+
float dk_delta = active ? (dS[i] * beta_t * inner_t) : 0.0f;
273+
float dk_kth = active ? (dkth_j * S_decayed[i]) : 0.0f;
274+
float dk_i_sum = simd_sum(dk_delta + dk_kth);
275+
if (active && vj == 0u) {{
276+
atomic_fetch_add_explicit(
277+
(device atomic_float*)&dk[qk_base + i],
278+
dk_i_sum,
279+
memory_order_relaxed
280+
);
281+
}}
282+
}}
283+
284+
// ---- (3) dS_decayed[i, j] = dS_t[i, j] + dkth[j] * k_i ----
285+
for (int i = 0; i < {kdim}; i++) {{
286+
dS[i] = dS[i] + dkth_j * k[qk_base + i];
287+
}}
288+
289+
// ---- (4) S_decayed[i, j] = decay[i] * S_{{t-1}}[i, j] ----
290+
// ddecay[i] = sum_j dS_decayed[i, j] * S_{{t-1}}[i, j] (per-i)
291+
// dg_t[i] = ddecay[i] * decay[i] (per-K)
292+
// dS_{{t-1}}[i, j] = dS_decayed[i, j] * decay[i]
293+
// dg[b, t, hv_idx, i] is the per-K dg write at this step.
294+
for (int i = 0; i < {kdim}; i++) {{
295+
float contrib = active ? (dS[i] * S_prev[i]) : 0.0f;
296+
float ddecay_i = simd_sum(contrib);
297+
if (active && vj == 0u) {{
298+
dg[g_base + i] = ddecay_i * decay_arr[i];
299+
}}
300+
dS[i] = dS[i] * decay_arr[i];
301+
}}
302+
}}
303+
"""
304+
305+
name = f"v4_kda_bwd_{b}_{t}_{h}_{hv}_{kdim}_{vdim}"
306+
kernel = get_or_build_kernel(
307+
name=name,
308+
input_names=["q", "k", "v", "g", "beta", "dy"],
309+
output_names=["dq", "dk", "dv", "dg", "dbeta", "state_hist"],
310+
source=source,
311+
)
312+
313+
grid = (_SIMD_WIDTH * b * hv, 1, 1)
314+
threadgroup = (_SIMD_WIDTH, 1, 1)
315+
316+
dq_flat, dk_flat, dv_flat, dg_flat, dbeta_flat, _ = kernel(
317+
inputs=[q_f, k_f, v_f, g_f, beta_f, dy_f],
318+
output_shapes=[
319+
(b * t * h * kdim,),
320+
(b * t * h * kdim,),
321+
(b * t * hv * vdim,),
322+
(b * t * hv * kdim,),
323+
(b * t * hv,),
324+
(b * hv * (t + 1) * kdim * vdim,),
325+
],
326+
output_dtypes=[mx.float32] * 6,
327+
grid=grid,
328+
threadgroup=threadgroup,
329+
init_value=0.0,
330+
)
331+
return (
332+
dq_flat.reshape(b, t, h, kdim),
333+
dk_flat.reshape(b, t, h, kdim),
334+
dv_flat.reshape(b, t, hv, vdim),
335+
dg_flat.reshape(b, t, hv, kdim),
336+
dbeta_flat.reshape(b, t, hv),
337+
)
338+
339+
340+
def _path_a_grad_fallback(primals, cotangent):
341+
q, k, v, g, beta = primals
342+
343+
def _loss(q_, k_, v_, g_, beta_):
344+
y, _ = naive_recurrent_kda(q_, k_, v_, g_, beta_)
345+
return (y * cotangent).sum()
346+
347+
return mx.grad(_loss, argnums=(0, 1, 2, 3, 4))(q, k, v, g, beta)
348+
349+
9350
@mx.custom_function
10351
def kda_apply_path_b(
11352
q: mx.array, k: mx.array, v: mx.array, g: mx.array, beta: mx.array,
12353
) -> mx.array:
13-
"""Forward via fast Path B Metal kernel; backward via Path A reference grad."""
354+
"""Forward via fast Path B Metal kernel; backward via real Metal kernel.
355+
356+
Falls back to ``mx.grad`` through ``naive_recurrent_kda`` for shapes
357+
outside the kernel's domain (``V > 32``, ``HV % H != 0``, etc.).
358+
"""
14359
y, _ = kda_forward_path_b(q, k, v, g, beta, output_final_state=False)
15360
return y
16361

@@ -19,12 +364,19 @@ def kda_apply_path_b(
19364
def _kda_apply_path_b_vjp(primals, cotangent, output):
20365
del output
21366
q, k, v, g, beta = primals
22-
23-
def _loss_proxy(q_, k_, v_, g_, beta_):
24-
y, _ = naive_recurrent_kda(q_, k_, v_, g_, beta_)
25-
return (y * cotangent).sum()
26-
27-
return mx.grad(_loss_proxy, argnums=(0, 1, 2, 3, 4))(q, k, v, g, beta)
367+
vdim = v.shape[-1]
368+
hv = v.shape[2]
369+
h = q.shape[2]
370+
bwd_ok = (
371+
v.ndim == 4
372+
and g.shape == (*v.shape[:3], k.shape[-1])
373+
and beta.shape == v.shape[:3]
374+
and hv % h == 0
375+
and vdim <= _SIMD_WIDTH
376+
)
377+
if not bwd_ok:
378+
return _path_a_grad_fallback(primals, cotangent)
379+
return _kda_backward_kernel(q, k, v, g, beta, cotangent)
28380

29381

30382
__all__ = ["kda_apply_path_b"]

0 commit comments

Comments
 (0)