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
354import mlx .core as mx
455
56+ from cppmega_v4 ._tilelang ._kernel_cache import get_or_build_kernel
557from cppmega_v4 ._tilelang .kda_path_b import kda_forward_path_b
658from 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
10351def 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(
19364def _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