@@ -94,12 +94,29 @@ def _kda_backward_kernel(
9494 hv , vdim = v .shape [2 ], v .shape [- 1 ]
9595 if hv % h != 0 :
9696 raise ValueError (f"HV ({ hv } ) must be divisible by H ({ h } )" )
97- if vdim > _SIMD_WIDTH :
97+ if vdim > _SIMD_WIDTH * 8 :
98+ # Cap at 8 simdgroups (V up to 256). Beyond that per-thread register
99+ # pressure (state[K], dS[K], S_t_col[K], S_prev[K], S_decayed[K],
100+ # decay_arr[K]) gets prohibitive for the K dimension too.
98101 raise ValueError (
99- f"Real-MSL KDA bwd currently requires V<=32 (got { vdim } ); "
102+ f"Real-MSL KDA bwd currently requires V<=256 (got { vdim } ); "
100103 f"caller should fall back to Path A grad path"
101104 )
102105 group = hv // h
106+ # Multi-simdgroup path when V > 32: pad threadgroup to a 32-multiple so
107+ # simd_sum still works (32 lanes per simdgroup). Cross-simdgroup
108+ # reductions go through threadgroup-shared-memory tiles (the previous
109+ # revision used atomic_fetch_add inside the threadgroup which serialised
110+ # at V >= 64). Inter-HV-group dq/dk/dbeta races (different threadgroups
111+ # touching the same (b,t,h_idx)/(b,t,hv) cells through HV expansion) still
112+ # need atomic_fetch_add — those live in device memory.
113+ use_shared = vdim > _SIMD_WIDTH
114+ tg_size = ((vdim + _SIMD_WIDTH - 1 ) // _SIMD_WIDTH ) * _SIMD_WIDTH
115+ n_simd = tg_size // _SIMD_WIDTH
116+ # Shared-memory: one [kdim, n_simd] vector tile (reused across the three
117+ # K-vector reductions: dq, dk, ddecay) + a few [n_simd] scalar tiles.
118+ # Worst case kdim=256, n_simd=8 → 256*8*4 = 8192 bytes + ~64 bytes scalars.
119+ shared_bytes = (kdim * n_simd + 2 * n_simd ) * 4 if use_shared else 0
103120
104121 scale = kdim ** - 0.5
105122 q_f = q .astype (mx .float32 ).reshape (- 1 )
@@ -109,15 +126,25 @@ def _kda_backward_kernel(
109126 beta_f = beta .astype (mx .float32 ).reshape (- 1 )
110127 dy_f = dy .astype (mx .float32 ).reshape (- 1 )
111128
129+ shared_decls = (
130+ f"""
131+ threadgroup float tg_vec[{ kdim * n_simd } ]; // batched per-i partials
132+ threadgroup float tg_scalar0[{ n_simd } ]; // scalar reduction A
133+ """ if use_shared else ""
134+ )
135+
112136 source = f"""
113137 uint tid_in_tg = thread_position_in_threadgroup.x;
114138 uint bhv = threadgroup_position_in_grid.x;
115139 uint vj = tid_in_tg;
116140 bool active = (vj < { vdim } u) && (bhv < { b * hv } u);
141+ uint simd_id = tid_in_tg / 32u;
142+ uint lane = tid_in_tg & 31u;
117143
118144 uint bb = bhv / { hv } u;
119145 uint hv_idx = bhv % { hv } u;
120146 uint h_idx = hv_idx / { group } u;
147+ { shared_decls }
121148
122149 // Per-thread registers:
123150 // state[K], dS[K], inner_hist[T]
@@ -217,15 +244,35 @@ def _kda_backward_kernel(
217244 float q_i_scaled = q[qk_base + i] * { scale } f;
218245 float contrib = active ? (dY_j * S_t_col[i]) : 0.0f;
219246 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- }}
247+ { (
248+ f'''if (lane == 0u) tg_vec[i * { n_simd } + simd_id] = dq_i_sum;'''
249+ if use_shared else
250+ f'''if (active && vj == 0u) {{
251+ atomic_fetch_add_explicit(
252+ (device atomic_float*)&dq[qk_base + i],
253+ dq_i_sum * { scale } f,
254+ memory_order_relaxed
255+ );
256+ }}'''
257+ )}
227258 dS[i] += dY_j * q_i_scaled;
228259 }}
260+ { (
261+ f'''threadgroup_barrier(metal::mem_flags::mem_threadgroup);
262+ // Reduce per-i across simdgroups, then single atomic_fetch_add
263+ // to handle the inter-HV-group race on dq[(b,t,h_idx,i)].
264+ for (int oi = (int)tid_in_tg; oi < { kdim } ; oi += { tg_size } ) {{
265+ float total = 0.0f;
266+ for (int s = 0; s < { n_simd } ; s++) total += tg_vec[oi * { n_simd } + s];
267+ atomic_fetch_add_explicit(
268+ (device atomic_float*)&dq[qk_base + oi],
269+ total * { scale } f,
270+ memory_order_relaxed
271+ );
272+ }}
273+ threadgroup_barrier(metal::mem_flags::mem_threadgroup);'''
274+ if use_shared else ""
275+ )}
229276
230277 // ---- (2) S_t = S_decayed + beta * k * inner ----
231278 // dinner[j] = beta_t * sum_i dS_t[i, j] * k_i (per-j scalar)
@@ -238,17 +285,37 @@ def _kda_backward_kernel(
238285 }}
239286 float dinner_j = beta_t * sum_k_dS;
240287
241- // dv[j] = dinner_j
288+ // dv[j] = dinner_j — unique per (b,t,hv,vj), no race.
242289 if (active) {{
290+ dv[v_idx] = dinner_j;
291+ }}
292+
293+ // dbeta_t = sum_i k_i * (sum_j dS_t[i,j] * inner_t)
294+ { (
295+ f'''// Stash per-simdgroup term_sum[i] into tg_vec, then one
296+ // thread combines across simdgroups and i.
297+ for (int i = 0; i < { kdim } ; i++) {{
298+ float term = active ? (dS[i] * inner_t) : 0.0f;
299+ float term_sum = simd_sum(term);
300+ if (lane == 0u) tg_vec[i * { n_simd } + simd_id] = term_sum;
301+ }}
302+ threadgroup_barrier(metal::mem_flags::mem_threadgroup);
303+ if (tid_in_tg == 0u) {{
304+ float dbeta_partial = 0.0f;
305+ for (int i = 0; i < { kdim } ; i++) {{
306+ float total_i = 0.0f;
307+ for (int s = 0; s < { n_simd } ; s++) total_i += tg_vec[i * { n_simd } + s];
308+ dbeta_partial += k[qk_base + i] * total_i;
309+ }}
243310 atomic_fetch_add_explicit(
244- (device atomic_float*)&dv[v_idx ],
245- dinner_j ,
311+ (device atomic_float*)&dbeta[beta_idx ],
312+ dbeta_partial ,
246313 memory_order_relaxed
247314 );
248315 }}
249-
250- // dbeta_t = sum_i k_i * (sum_j dS_t[i,j] * inner_t)
251- float dbeta_partial = 0.0f;
316+ threadgroup_barrier(metal::mem_flags::mem_threadgroup);'''
317+ if use_shared else
318+ f''' float dbeta_partial = 0.0f;
252319 for (int i = 0; i < { kdim } ; i++) {{
253320 float term = active ? (dS[i] * inner_t) : 0.0f;
254321 float term_sum = simd_sum(term); // sum over j
@@ -262,7 +329,8 @@ def _kda_backward_kernel(
262329 dbeta_partial,
263330 memory_order_relaxed
264331 );
265- }}
332+ }}'''
333+ )}
266334
267335 // dk_i (delta) += sum_j dS_t[i, j] * (beta_t * inner_t)
268336 // Note inner_t is a per-j scalar; cannot factor it out of simd_sum.
@@ -272,14 +340,32 @@ def _kda_backward_kernel(
272340 float dk_delta = active ? (dS[i] * beta_t * inner_t) : 0.0f;
273341 float dk_kth = active ? (dkth_j * S_decayed[i]) : 0.0f;
274342 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- }}
343+ { (
344+ f'''if (lane == 0u) tg_vec[i * { n_simd } + simd_id] = dk_i_sum;'''
345+ if use_shared else
346+ f'''if (active && vj == 0u) {{
347+ atomic_fetch_add_explicit(
348+ (device atomic_float*)&dk[qk_base + i],
349+ dk_i_sum,
350+ memory_order_relaxed
351+ );
352+ }}'''
353+ )}
282354 }}
355+ { (
356+ f'''threadgroup_barrier(metal::mem_flags::mem_threadgroup);
357+ for (int oi = (int)tid_in_tg; oi < { kdim } ; oi += { tg_size } ) {{
358+ float total = 0.0f;
359+ for (int s = 0; s < { n_simd } ; s++) total += tg_vec[oi * { n_simd } + s];
360+ atomic_fetch_add_explicit(
361+ (device atomic_float*)&dk[qk_base + oi],
362+ total,
363+ memory_order_relaxed
364+ );
365+ }}
366+ threadgroup_barrier(metal::mem_flags::mem_threadgroup);'''
367+ if use_shared else ""
368+ )}
283369
284370 // ---- (3) dS_decayed[i, j] = dS_t[i, j] + dkth[j] * k_i ----
285371 for (int i = 0; i < { kdim } ; i++) {{
@@ -293,12 +379,26 @@ def _kda_backward_kernel(
293379 // dg[b, t, hv_idx, i] is the per-K dg write at this step.
294380 for (int i = 0; i < { kdim } ; i++) {{
295381 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- }}
382+ float ddecay_simd = simd_sum(contrib);
383+ { (
384+ f'''if (lane == 0u) tg_vec[i * { n_simd } + simd_id] = ddecay_simd;'''
385+ if use_shared else
386+ f'''if (active && vj == 0u) {{
387+ dg[g_base + i] = ddecay_simd * decay_arr[i];
388+ }}'''
389+ )}
300390 dS[i] = dS[i] * decay_arr[i];
301391 }}
392+ { (
393+ f'''threadgroup_barrier(metal::mem_flags::mem_threadgroup);
394+ for (int oi = (int)tid_in_tg; oi < { kdim } ; oi += { tg_size } ) {{
395+ float total = 0.0f;
396+ for (int s = 0; s < { n_simd } ; s++) total += tg_vec[oi * { n_simd } + s];
397+ dg[g_base + oi] = total * decay_arr[oi];
398+ }}
399+ threadgroup_barrier(metal::mem_flags::mem_threadgroup);'''
400+ if use_shared else ""
401+ )}
302402 }}
303403 """
304404
@@ -310,8 +410,8 @@ def _kda_backward_kernel(
310410 source = source ,
311411 )
312412
313- grid = (_SIMD_WIDTH * b * hv , 1 , 1 )
314- threadgroup = (_SIMD_WIDTH , 1 , 1 )
413+ grid = (tg_size * b * hv , 1 , 1 )
414+ threadgroup = (tg_size , 1 , 1 )
315415
316416 dq_flat , dk_flat , dv_flat , dg_flat , dbeta_flat , _ = kernel (
317417 inputs = [q_f , k_f , v_f , g_f , beta_f , dy_f ],
@@ -372,7 +472,7 @@ def _kda_apply_path_b_vjp(primals, cotangent, output):
372472 and g .shape == (* v .shape [:3 ], k .shape [- 1 ])
373473 and beta .shape == v .shape [:3 ]
374474 and hv % h == 0
375- and vdim <= _SIMD_WIDTH
475+ and vdim <= _SIMD_WIDTH * 8
376476 )
377477 if not bwd_ok :
378478 return _path_a_grad_fallback (primals , cotangent )
0 commit comments