Skip to content

Commit e8378a5

Browse files
committed
perf(v4): KDA Path B bwd — multi-simdgroup V via shared-mem; raise cap to 256
Previously the KDA bwd kernel hard-capped V at 32 (single simdgroup) and fell back to mx.grad through naive_recurrent_kda for V > 32. This commit adds the multi-simdgroup path on V (matching the GDN bwd revision) and raises the cap to V <= 256. Within a single threadgroup, cross-simdgroup reductions now go through a batched threadgroup tile (tg_vec[kdim * n_simd] + tg_scalar0[n_simd]) instead of atomic_fetch_add. Inter-HV-group races on dq/dk/dbeta (different threadgroups touching the same shared-K cells through the HV expansion) still need atomic adds — those live in device memory across threadgroups — but each threadgroup now contributes ONE reduced partial per output cell instead of N_simd serialised partials. Other cleanups: - dv write no longer uses atomic_fetch_add (unique per (b,t,hv,vj), no race) — direct store. - dg per-K direct store (unique per (b,t,hv_idx,i), no race). - Final reduce uses a stride loop (oi += tg_size) so kdim > tg_size cases are handled correctly. Shared memory budget: kdim * n_simd * 4 + n_simd * 4 bytes = 256 * 8 * 4 + 32 = 8224 bytes per threadgroup at the cap (V=256, 8 simdgroups), well under the 32KB M-series limit. Single-simdgroup fast path (V <= 32) unchanged. Tests stay green at original tolerances. New native support at V=64 gives ~1.7× speedup vs the previous Path A fallback.
1 parent 4635186 commit e8378a5

1 file changed

Lines changed: 130 additions & 30 deletions

File tree

cppmega_v4/_tilelang/kda_path_b_bwd.py

Lines changed: 130 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)