@@ -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