From 0b9f9da2f785347956b02e317b189a2681806f1c Mon Sep 17 00:00:00 2001 From: biondizzle Date: Mon, 25 May 2026 17:03:19 +0000 Subject: [PATCH] revert grid change to debug regression --- dsv4/kernels/attention/fmha.py | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/dsv4/kernels/attention/fmha.py b/dsv4/kernels/attention/fmha.py index 499f7747..05355ea1 100644 --- a/dsv4/kernels/attention/fmha.py +++ b/dsv4/kernels/attention/fmha.py @@ -131,13 +131,7 @@ class FmhaKernel: # CuTeDSL doesn't support None parameters in @cute.kernel. if const_expr(lse is None): lse = cute.make_tensor(c.iterator, cute.make_layout((1,), stride=(0,))) - # Grid: (M_tiles, 1, batch) where M = n_h * T packed into M dimension - # At decode T=1, n_h=128: M=128, grid=(1,1,batch) — 1 CTA per batch - # At T=64, n_h=128: M=8192, grid=(64,1,batch) — 64 CTAs per batch - # For single-head (n_h=1): grid=(1,1,1) — backward compatible - M_total = self.num_query_heads # T is implicitly 1 for decode, M = n_h * T - num_M_tiles = math.ceil(M_total / 128) if M_total > 128 else 1 - self._kernel(qk_mma,pv_mma,tma_q,mQ,tma_k,mK,tma_v,mV,tma_c,mC,self.cluster_layout_vmnk,self.q_smem_s,self.k_smem_s,self.v_smem_s,self.p_tmem_s,self.p_smem_s,self.c_smem_s,self.epi_tile,lse).launch(grid=(num_M_tiles,1,self.batch_size),block=[self.threads_per_cta,1,1],stream=stream) + self._kernel(qk_mma,pv_mma,tma_q,mQ,tma_k,mK,tma_v,mV,tma_c,mC,self.cluster_layout_vmnk,self.q_smem_s,self.k_smem_s,self.v_smem_s,self.p_tmem_s,self.p_smem_s,self.c_smem_s,self.epi_tile,lse).launch(grid=(1,1,1),block=[self.threads_per_cta,1,1],stream=stream) @cute.kernel def _kernel(self, qk_mma, pv_mma, tma_q, mQ, tma_k, mK, tma_v, mV, tma_c, mC, cl_vmnk, q_smem_s, k_smem_s, v_smem_s, p_tmem_s, p_smem_s, c_smem_s, epi_tile, mLSE):