From e5ba0ca119887d917169b14ad5d0466454b503af Mon Sep 17 00:00:00 2001 From: biondizzle Date: Thu, 28 May 2026 08:53:35 +0000 Subject: [PATCH] debug: clean QK verify with scalar sanity + MMA result --- dsv4/kernels/attention/fmha_qk_verify.cuh | 114 ++++++++-------------- tests/unit/test_qk_mma.cu | 5 + 2 files changed, 45 insertions(+), 74 deletions(-) diff --git a/dsv4/kernels/attention/fmha_qk_verify.cuh b/dsv4/kernels/attention/fmha_qk_verify.cuh index 3d4628d6..cb6e49f9 100644 --- a/dsv4/kernels/attention/fmha_qk_verify.cuh +++ b/dsv4/kernels/attention/fmha_qk_verify.cuh @@ -1,6 +1,5 @@ /** * DSV4 FMHA — QK GEMM verification with SWIZZLE_NONE UMMA layout. - * Uses simple row-major SMEM (no swizzle) with the NONE descriptor. */ #pragma once @@ -24,44 +23,46 @@ fmha_qk_verify( const bf16_t* qh = q + batch*bstride_q + head*HD; const bf16_t* kb = k + batch*bstride_kv; - // SMEM: sQ (128×HD BF16 row-major) + sK (128×HD BF16 row-major) + tmem_base - // Must be 16-byte aligned for UMMA - // SMEM layout: must be 128-byte aligned for UMMA descriptors - // [0..127] padding for alignment - // [128..131] tmem_base (4 bytes) - // [132..255] padding for Q alignment - // [256..] sQ (128*HD*2 bytes) + sK (128*HD*2 bytes) + // SMEM layout (256B aligned for UMMA): + // [0..127] padding + // [128..131] tmem_base + // [132..255] padding + // [256..256+128*HD*2) sQ (128×HD BF16 row-major) + // [256+128*HD*2..) sK (128×HD BF16 row-major) extern __shared__ char sbuf[]; uint32_t* sTmemBase = (uint32_t*)(sbuf + 128); bf16_t* sQ = (bf16_t*)(sbuf + 256); bf16_t* sK = sQ + 128 * HD; - // Load Q: (1, HD) padded to (128, HD) with zeros - for (int i = tid; i < 128 * HD; i += NTHREADS) sQ[i] = 0; - for (int d = tid; d < HD; d += NTHREADS) sQ[d] = qh[d]; - - // Load K: (min(128, s_k), HD) padded to (128, HD) + // Load Q and K to SMEM int kv_len = min(128, s_k); for (int i = tid; i < 128 * HD; i += NTHREADS) { int r = i / HD, c = i % HD; + sQ[i] = (r == 0 && c < HD) ? qh[c] : 0; sK[i] = (r < kv_len) ? kb[r * HD + c] : 0; } __syncthreads(); - // SKIP TMEM — just test SMEM loads and scalar QK - // No TMEM alloc, no MMA - /* - // TMEM alloc for S: 128 columns + // Sanity check: scalar QK dot product + if (tid == 0) { + float dot = 0; + for (int d = 0; d < HD; d++) { + dot += bf16_to_f32(sQ[d]) * bf16_to_f32(sK[d]); + } + s_out[0] = dot * scale; // row 0, col 0 (scalar reference) + s_out[1] = bf16_to_f32(sQ[0]); // Q[0,0] + s_out[2] = bf16_to_f32(sK[0]); // K[0,0] + } + + // TMEM alloc if (wid == 0) { uint32_t smem_ptr = __cvta_generic_to_shared(sTmemBase); tmem_alloc(smem_ptr, 128); } __syncthreads(); uint32_t tmem_base = *sTmemBase; - */ - uint32_t tmem_base = 0; // dummy - // Zero TMEM S + // Zero TMEM if (wid == 0) { for (int col = 0; col < 128; col++) { tmem_store(tmem_base + col, 0, 0, 0, 0); @@ -70,48 +71,27 @@ fmha_qk_verify( } __syncthreads(); - // UMMA descriptors + // UMMA descriptors (SWIZZLE_NONE with proper strides) uint32_t sQ_smem = __cvta_generic_to_shared(sQ); uint32_t sK_smem = __cvta_generic_to_shared(sK); + // MN-major NONE for (128, 64) BF16: + // LBO=16 (uint128_t), SBO=128 (uint128_t) uint64_t desc_q = 0; - desc_q |= (static_cast(sQ_smem >> 4) & 0x3FFF); // start_address - desc_q |= (static_cast(1) << 46); // version - // Everything else = 0 (no strides, no swizzle) + desc_q |= (static_cast(sQ_smem >> 4) & 0x3FFF); + desc_q |= (static_cast(16) & 0x3FFF) << 16; + desc_q |= (static_cast(128) & 0x3FFF) << 32; + desc_q |= (static_cast(1) << 46); + // K-major NONE for (128, 64) BF16: + // LBO=16, SBO=32 uint64_t desc_k = 0; desc_k |= (static_cast(sK_smem >> 4) & 0x3FFF); + desc_k |= (static_cast(16) & 0x3FFF) << 16; + desc_k |= (static_cast(32) & 0x3FFF) << 32; desc_k |= (static_cast(1) << 46); - // Quick test: verify SMEM data was loaded correctly - // Write Q[0,0..3] * K[0,0..3] dot product (scalar) to s_out[0] as sanity check - if (tid == 0) { - // Read first few Q values directly from SMEM - float q0 = bf16_to_f32(sQ[0]); - float q1 = bf16_to_f32(sQ[1]); - float k0 = bf16_to_f32(sK[0]); - float k1 = bf16_to_f32(sK[1]); - - float dot = 0; - for (int d = 0; d < HD; d++) { - dot += bf16_to_f32(sQ[d]) * bf16_to_f32(sK[d]); - } - s_out[0] = dot * scale; - s_out[1] = q0; // first Q value - s_out[2] = k0; // first K value - s_out[3] = (float)(sQ_smem & 0xFFFF); // low 16 bits of SMEM address - } - __syncthreads(); - // TEMPORARILY SKIP MMA — just verify SMEM loads - if (wid == 0) tmem_dealloc(tmem_base, 128); - return; // ALL threads return — we're just testing SMEM loads - // MMA: called by ONE lane (elect_one_sync pattern) - if (tid == 0) printf("[qk] tmem_base=%u sQ_smem=0x%x sK_smem=0x%x\n", tmem_base, sQ_smem, sK_smem); - __syncthreads(); - - // Try: ALL 32 lanes of warp 0 call MMA (not just lane 0) - // The CUTLASS code uses elect_one_sync which selects lane 0, - // but maybe the PTX requires all lanes to participate? + // QK GEMM if (wid == 0) { umma_ss_f16(tmem_base, desc_q, desc_k, false); } @@ -119,28 +99,14 @@ fmha_qk_verify( if (wid == 0 && lane == 0) tmem_fence_store(); __syncthreads(); - // Read S row 0 from TMEM — try multiple column offsets + // Read S row 0 from TMEM if (tid == 0) { - printf("[qk] tmem_base=%u, reading S...\n", tmem_base); - for (int col = 0; col < 4; col++) { - uint32_t u0, u1, u2, u3; - tmem_load(tmem_base + col, u0, u1, u2, u3); - printf("[qk] S[0,%d] raw: %f %f %f %f\n", col, - u32_to_f32(u0), u32_to_f32(u1), u32_to_f32(u2), u32_to_f32(u3)); - } - } - __syncthreads(); - - // Also read using the output array (same as before) - if (wid == 0) { - for (int col = lane; col < 128; col += WARP) { - uint32_t u0, u1, u2, u3; - tmem_load(tmem_base + col, u0, u1, u2, u3); - if (lane * 4 + 0 < kv_len) s_out[lane * 4 + 0] = u32_to_f32(u0) * scale; - if (lane * 4 + 1 < kv_len) s_out[lane * 4 + 1] = u32_to_f32(u1) * scale; - if (lane * 4 + 2 < kv_len) s_out[lane * 4 + 2] = u32_to_f32(u2) * scale; - if (lane * 4 + 3 < kv_len) s_out[lane * 4 + 3] = u32_to_f32(u3) * scale; - } + uint32_t u0, u1, u2, u3; + tmem_load(tmem_base + 0, u0, u1, u2, u3); + s_out[3] = u32_to_f32(u0) * scale; // MMA result: S[0,0] + s_out[4] = u32_to_f32(u1) * scale; // S[0,1] + s_out[5] = u32_to_f32(u2) * scale; // S[0,2] + s_out[6] = u32_to_f32(u3) * scale; // S[0,3] } __syncthreads(); diff --git a/tests/unit/test_qk_mma.cu b/tests/unit/test_qk_mma.cu index a862fad8..7ba3b7ec 100644 --- a/tests/unit/test_qk_mma.cu +++ b/tests/unit/test_qk_mma.cu @@ -87,6 +87,11 @@ int main() { cudaMemcpy(hs_gpu, ds_out, s_k*4, cudaMemcpyDeviceToHost); // Compare + // Compare — but also print raw values for debugging + printf("Scalar QK[0]: %.6f (ref: %.6f)\n", hs_gpu[0], hs_ref[0]); + printf("MMA QK[0]: %.6f\n", hs_gpu[3]); + printf("sQ[0]: %.6f sK[0]: %.6f\n", hs_gpu[1], hs_gpu[2]); + float max_diff = 0; int nan_count = 0; for(int j=0;j