From bd1309ba88d960b0fec37004c91ad05a15b29b4c Mon Sep 17 00:00:00 2001 From: biondizzle Date: Fri, 29 May 2026 18:40:11 +0000 Subject: [PATCH] =?UTF-8?q?test:=20minimal=20QK=20=E2=80=94=20128=20thread?= =?UTF-8?q?s,=20tid=3D=3D0=20MMA,=20match=20working=20gen=20kernel=20patte?= =?UTF-8?q?rn?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/unit/test_qk_minimal.cu | 45 ++++++++++++++++------------------- 1 file changed, 20 insertions(+), 25 deletions(-) diff --git a/tests/unit/test_qk_minimal.cu b/tests/unit/test_qk_minimal.cu index f260d802..646e26b7 100644 --- a/tests/unit/test_qk_minimal.cu +++ b/tests/unit/test_qk_minimal.cu @@ -27,7 +27,7 @@ constexpr int SK = 128; constexpr int NKT = HD / MMA_K_BF16; constexpr int CORES_MN = 16; // 128/8 -__global__ void __launch_bounds__(192) +__global__ void __launch_bounds__(128) test_qk_minimal_kernel(float* __restrict__ out_s, const bf16_t* __restrict__ q, const bf16_t* __restrict__ k, int T, int s_k) { @@ -36,9 +36,9 @@ test_qk_minimal_kernel(float* __restrict__ out_s, static constexpr int NUM_READS = SK / 8; const int tid = threadIdx.x; - const int wid = tid / 32; const int lane = tid % 32; - const bool is_mma_warp = (wid == 4); + // Simple 4-warp: warp 0 = load+softmax+MMA, warps 1-3 = softmax + // ALL 128 threads participate in loading and MMA is called by tid==0 // SMEM: sQ0 and sK0 are (128, 16) each extern __shared__ __align__(128) char sbuf[]; @@ -49,44 +49,39 @@ test_qk_minimal_kernel(float* __restrict__ out_s, off = (off + 127) & ~(size_t)127; bf16_t* sK0 = (bf16_t*)(sbuf + off); off += TILE_SZ * sizeof(bf16_t); - if (is_mma_warp) tmem_alloc(__cvta_generic_to_shared(sTmemBase), TMEM_N); + if (tid < 32) tmem_alloc(__cvta_generic_to_shared(sTmemBase), TMEM_N); __syncthreads(); uint32_t tb = *sTmemBase; for (int kt = 0; kt < NKT; kt++) { - // Load Q sub-tile - for (int i = tid; i < TILE_SZ; i += NTHREADS) sQ0[i] = 0; - for (int r = tid / 32; r < T; r += 6) { // one row per warp - if (r < T) { - for (int d = lane; d < MMA_K_BF16; d += 32) { - int full_d = kt * MMA_K_BF16 + d; - if (full_d < HD) { - int ck = d/8, lc = d%8, cm = r/8, lr = r%8; - sQ0[ck*CORES_MN*64 + cm*64 + lr*8 + lc] = q[r * HD + full_d]; - } - } + // Load Q sub-tile — all 128 threads participate + for (int i = tid; i < TILE_SZ; i += 128) sQ0[i] = 0; + for (int d = tid; d < MMA_K_BF16; d += 128) { + int full_d = kt * MMA_K_BF16 + d; + if (full_d < HD) { + // Q has 1 row (T=1), row 0: cm=0, lr=0 + int ck = d/8, lc = d%8; + sQ0[ck*CORES_MN*64 + 0*64 + 0*8 + lc] = q[full_d]; } } __syncthreads(); - // Load K sub-tile - for (int i = tid; i < TILE_SZ; i += NTHREADS) sK0[i] = 0; - for (int r = lane; r < s_k; r += 32) { + // Load K sub-tile — all 128 threads + for (int i = tid; i < TILE_SZ; i += 128) sK0[i] = 0; + for (int r = tid; r < s_k; r += 128) { for (int d = 0; d < MMA_K_BF16; d++) { int full_d = kt * MMA_K_BF16 + d; - if (full_d < HD) { - int ck = d/8, lc = d%8, cm = r/8, lr = r%8; - sK0[ck*CORES_MN*64 + cm*64 + lr*8 + lc] = k[r * HD + full_d]; - } + int ck = d/8, lc = d%8, cm = r/8, lr = r%8; + sK0[ck*CORES_MN*64 + cm*64 + lr*8 + lc] = k[r * HD + full_d]; } } __syncthreads(); - if (is_mma_warp) { + if (tid == 0) { uint32_t idesc = make_idesc(128, 128); uint64_t dq = make_umma_desc_kmajor_none(__cvta_generic_to_shared(sQ0), 128); uint64_t dk = make_umma_desc_kmajor_none(__cvta_generic_to_shared(sK0), 128); - if (tid == 128) umma_ss_f16(tb, dq, dk, idesc, kt > 0); + umma_ss_f16(tb, dq, dk, idesc, kt > 0); asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory"); } __syncthreads(); @@ -134,7 +129,7 @@ int main() { cudaMemcpy(d_k, h_k, SK * HD * sizeof(bf16_t), cudaMemcpyHostToDevice); int smem = 4 + 128 + 128*16*2 + 128*16*2 + 4096; - test_qk_minimal_kernel<<<1, 192, smem>>>(d_out, d_q, d_k, T, SK); + test_qk_minimal_kernel<<<1, 128, smem>>>(d_out, d_q, d_k, T, SK); cudaError_t err = cudaDeviceSynchronize(); if (err != cudaSuccess) { printf("CUDA ERROR: %s\n", cudaGetErrorString(err)); return 1; }