Skip to content

Commit e3926e8

Browse files
committed
perf(v4): tile K-dim in KDA Path B bwd scratch arrays — ~30% faster at K=128
The per-timestep scratch arrays (S_t_col, S_prev, S_decayed, decay_arr) were sized [K] per thread. At K=128 that meant 4*128 = 512 floats of scratch + state[K] + dS[K] + inner_hist[T] = ~832 floats per thread, crushing register occupancy at K >= 128. Tile those 4 per-step scratch arrays to K_TILE=32 (or 16 for K%16 only). state[K] and dS[K] still must be full because they persist across the T scan and the per-timestep dinner_j requires the full-K sum_i k[i]*dS[i] reduction that cannot be tiled without a two-phase pre-pass. The savings on the four tiled arrays (4*K → 4*K_TILE) drop per-thread scratch from 512 to 128 floats at K=128, ~46% reduction in K-scratch register footprint. For K <= 64 the original layout already fit comfortably; the extra reads of state_hist + re-evaluation of exp(g) in the dk/dg passes are a net loss, so K_TILE = K (no-op) in that regime. Bench (B=1, T=64, H=HV=4, K=128, M3 Pro, best-of-3, mlx eval-only): V=32: baseline 9.87ms → tiled 7.61ms (1.30x) V=64: baseline 9.53ms → tiled 7.46ms (1.28x) V=128: baseline 11.64ms → tiled 8.25ms (1.41x) All 3 existing tests pass at original tolerances (atol=1e-4 rtol=1e-3); no tolerance widening was needed.
1 parent 8387052 commit e3926e8

1 file changed

Lines changed: 104 additions & 69 deletions

File tree

cppmega_v4/_tilelang/kda_path_b_bwd.py

Lines changed: 104 additions & 69 deletions
Original file line numberDiff line numberDiff line change
@@ -95,13 +95,27 @@ def _kda_backward_kernel(
9595
if hv % h != 0:
9696
raise ValueError(f"HV ({hv}) must be divisible by H ({h})")
9797
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.
98+
# Cap at 8 simdgroups (V up to 256).
10199
raise ValueError(
102100
f"Real-MSL KDA bwd currently requires V<=256 (got {vdim}); "
103101
f"caller should fall back to Path A grad path"
104102
)
103+
# K-tiling: the per-timestep scratch arrays (S_t_col, S_prev, S_decayed,
104+
# decay_arr) used to be sized [K] per thread, which at K=128 + state[K] +
105+
# dS[K] + inner_hist[T] crushed occupancy. We tile those 4 scratch arrays
106+
# to K_TILE elements. state[K] and dS[K] still must be full because they
107+
# persist across the T scan and the per-step dinner_j reduction requires
108+
# the full-K k·dS sum. For small K (<=64) the original layout already fits
109+
# comfortably in registers and the recomputation overhead (re-reading
110+
# state_hist + re-doing exp(g) in dk/dg passes) is a net loss, so we keep
111+
# K_TILE = K in that regime.
112+
if kdim > 64 and kdim % 32 == 0:
113+
k_tile = 32
114+
elif kdim > 64 and kdim % 16 == 0:
115+
k_tile = 16
116+
else:
117+
k_tile = kdim # no tiling — original behavior
118+
n_k_tiles = (kdim + k_tile - 1) // k_tile
105119
group = hv // h
106120
# Multi-simdgroup path when V > 32: pad threadgroup to a 32-multiple so
107121
# simd_sum still works (32 lanes per simdgroup). Cross-simdgroup
@@ -214,48 +228,48 @@ def _kda_backward_kernel(
214228
float dY_j = active ? dy[v_idx] : 0.0f;
215229
float inner_t = inner_hist[ti];
216230
217-
// Read S_t (after step ti) and S_{{t-1}} (after step ti-1).
218-
float S_t_col[{kdim}];
219-
float S_prev[{kdim}];
220-
for (int i = 0; i < {kdim}; i++) {{
221-
int idx_t = bhv * hist_bh_stride + (ti + 1) * hist_ti_stride + i * {vdim} + (int)vj;
222-
int idx_tm1 = bhv * hist_bh_stride + ti * hist_ti_stride + i * {vdim} + (int)vj;
223-
S_t_col[i] = active ? state_hist[idx_t] : 0.0f;
224-
S_prev[i] = active ? state_hist[idx_tm1] : 0.0f;
225-
}}
226-
227-
// S_decayed[i, vj] = decay_t[i] * S_prev[i, vj]
228-
float decay_arr[{kdim}];
229-
float S_decayed[{kdim}];
230-
for (int i = 0; i < {kdim}; i++) {{
231-
decay_arr[i] = exp(g[g_base + i]);
232-
S_decayed[i] = decay_arr[i] * S_prev[i];
233-
}}
231+
// K-tiled scratch: process K in chunks of K_TILE to keep
232+
// S_t_col / S_prev / S_decayed / decay_arr small in registers.
233+
float S_t_col[{k_tile}];
234+
float S_prev[{k_tile}];
235+
float decay_arr[{k_tile}];
236+
float S_decayed[{k_tile}];
234237
235238
// ---- (1) o_t[j] = sum_i q'_i * S_t[i, j] ----
236239
// dq'_i = sum_j dY[j] * S_t[i, j] (reduce over j == vj axis)
237240
// dS_t[i, j] += dY[j] * q'_i
238-
// dq writes per-K to (b, t, h_idx, i); groups inside HV share q/k,
239-
// so multiple hv lanes in the same (b, t, h_idx) would race.
240-
// We restrict the dq write to the first hv in each group via
241-
// (hv_idx % group == 0) and multiply by group_size? No — we need
242-
// the SUM of grads from all hv in the group. Use atomic_fetch_add.
243-
for (int i = 0; i < {kdim}; i++) {{
244-
float q_i_scaled = q[qk_base + i] * {scale}f;
245-
float contrib = active ? (dY_j * S_t_col[i]) : 0.0f;
246-
float dq_i_sum = simd_sum(contrib);
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-
)}
258-
dS[i] += dY_j * q_i_scaled;
241+
// Tiled across K; dS[i] (persistent) updated in-place.
242+
for (int kt0 = 0; kt0 < {kdim}; kt0 += {k_tile}) {{
243+
for (int ii = 0; ii < {k_tile}; ii++) {{
244+
int i = kt0 + ii;
245+
int idx_t = bhv * hist_bh_stride + (ti + 1) * hist_ti_stride + i * {vdim} + (int)vj;
246+
int idx_tm1 = bhv * hist_bh_stride + ti * hist_ti_stride + i * {vdim} + (int)vj;
247+
S_t_col[ii] = active ? state_hist[idx_t] : 0.0f;
248+
S_prev[ii] = active ? state_hist[idx_tm1] : 0.0f;
249+
decay_arr[ii] = exp(g[g_base + i]);
250+
S_decayed[ii] = decay_arr[ii] * S_prev[ii];
251+
}}
252+
for (int ii = 0; ii < {k_tile}; ii++) {{
253+
int i = kt0 + ii;
254+
float q_i_scaled = q[qk_base + i] * {scale}f;
255+
float contrib = active ? (dY_j * S_t_col[ii]) : 0.0f;
256+
float dq_i_sum = simd_sum(contrib);
257+
{(
258+
f'''if (lane == 0u) tg_vec[i * {n_simd} + simd_id] = dq_i_sum;'''
259+
if use_shared else
260+
f'''if (active && vj == 0u) {{
261+
atomic_fetch_add_explicit(
262+
(device atomic_float*)&dq[qk_base + i],
263+
dq_i_sum * {scale}f,
264+
memory_order_relaxed
265+
);
266+
}}'''
267+
)}
268+
dS[i] += dY_j * q_i_scaled;
269+
}}
270+
// Persist this tile's S_decayed/decay_arr for later passes
271+
// via re-derivation. To avoid restoring in registers, we
272+
// simply recompute them in the dk/dg passes below.
259273
}}
260274
{(
261275
f'''threadgroup_barrier(metal::mem_flags::mem_threadgroup);
@@ -335,22 +349,33 @@ def _kda_backward_kernel(
335349
// dk_i (delta) += sum_j dS_t[i, j] * (beta_t * inner_t)
336350
// Note inner_t is a per-j scalar; cannot factor it out of simd_sum.
337351
// dk_i (kth) += sum_j dkth[j] * S_decayed[i, j]; dkth[j] = -dinner_j
352+
// Tiled across K; recompute S_decayed locally.
338353
float dkth_j = -dinner_j;
339-
for (int i = 0; i < {kdim}; i++) {{
340-
float dk_delta = active ? (dS[i] * beta_t * inner_t) : 0.0f;
341-
float dk_kth = active ? (dkth_j * S_decayed[i]) : 0.0f;
342-
float dk_i_sum = simd_sum(dk_delta + dk_kth);
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-
)}
354+
for (int kt0 = 0; kt0 < {kdim}; kt0 += {k_tile}) {{
355+
for (int ii = 0; ii < {k_tile}; ii++) {{
356+
int i = kt0 + ii;
357+
int idx_tm1 = bhv * hist_bh_stride + ti * hist_ti_stride + i * {vdim} + (int)vj;
358+
S_prev[ii] = active ? state_hist[idx_tm1] : 0.0f;
359+
decay_arr[ii] = exp(g[g_base + i]);
360+
S_decayed[ii] = decay_arr[ii] * S_prev[ii];
361+
}}
362+
for (int ii = 0; ii < {k_tile}; ii++) {{
363+
int i = kt0 + ii;
364+
float dk_delta = active ? (dS[i] * beta_t * inner_t) : 0.0f;
365+
float dk_kth = active ? (dkth_j * S_decayed[ii]) : 0.0f;
366+
float dk_i_sum = simd_sum(dk_delta + dk_kth);
367+
{(
368+
f'''if (lane == 0u) tg_vec[i * {n_simd} + simd_id] = dk_i_sum;'''
369+
if use_shared else
370+
f'''if (active && vj == 0u) {{
371+
atomic_fetch_add_explicit(
372+
(device atomic_float*)&dk[qk_base + i],
373+
dk_i_sum,
374+
memory_order_relaxed
375+
);
376+
}}'''
377+
)}
378+
}}
354379
}}
355380
{(
356381
f'''threadgroup_barrier(metal::mem_flags::mem_threadgroup);
@@ -376,25 +401,35 @@ def _kda_backward_kernel(
376401
// ddecay[i] = sum_j dS_decayed[i, j] * S_{{t-1}}[i, j] (per-i)
377402
// dg_t[i] = ddecay[i] * decay[i] (per-K)
378403
// dS_{{t-1}}[i, j] = dS_decayed[i, j] * decay[i]
379-
// dg[b, t, hv_idx, i] is the per-K dg write at this step.
380-
for (int i = 0; i < {kdim}; i++) {{
381-
float contrib = active ? (dS[i] * S_prev[i]) : 0.0f;
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-
)}
390-
dS[i] = dS[i] * decay_arr[i];
404+
// Tiled across K; recompute S_prev + decay locally.
405+
for (int kt0 = 0; kt0 < {kdim}; kt0 += {k_tile}) {{
406+
for (int ii = 0; ii < {k_tile}; ii++) {{
407+
int i = kt0 + ii;
408+
int idx_tm1 = bhv * hist_bh_stride + ti * hist_ti_stride + i * {vdim} + (int)vj;
409+
S_prev[ii] = active ? state_hist[idx_tm1] : 0.0f;
410+
decay_arr[ii] = exp(g[g_base + i]);
411+
}}
412+
for (int ii = 0; ii < {k_tile}; ii++) {{
413+
int i = kt0 + ii;
414+
float contrib = active ? (dS[i] * S_prev[ii]) : 0.0f;
415+
float ddecay_simd = simd_sum(contrib);
416+
{(
417+
f'''if (lane == 0u) tg_vec[i * {n_simd} + simd_id] = ddecay_simd;'''
418+
if use_shared else
419+
f'''if (active && vj == 0u) {{
420+
dg[g_base + i] = ddecay_simd * decay_arr[ii];
421+
}}'''
422+
)}
423+
dS[i] = dS[i] * decay_arr[ii];
424+
}}
391425
}}
392426
{(
393427
f'''threadgroup_barrier(metal::mem_flags::mem_threadgroup);
394428
for (int oi = (int)tid_in_tg; oi < {kdim}; oi += {tg_size}) {{
395429
float total = 0.0f;
396430
for (int s = 0; s < {n_simd}; s++) total += tg_vec[oi * {n_simd} + s];
397-
dg[g_base + oi] = total * decay_arr[oi];
431+
// decay_arr is K_TILE-sized; recompute from g for the final write.
432+
dg[g_base + oi] = total * exp(g[g_base + oi]);
398433
}}
399434
threadgroup_barrier(metal::mem_flags::mem_threadgroup);'''
400435
if use_shared else ""

0 commit comments

Comments
 (0)