debug: clean QK verify with scalar sanity + MMA result
This commit is contained in:
@@ -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<uint64_t>(sQ_smem >> 4) & 0x3FFF); // start_address
|
||||
desc_q |= (static_cast<uint64_t>(1) << 46); // version
|
||||
// Everything else = 0 (no strides, no swizzle)
|
||||
desc_q |= (static_cast<uint64_t>(sQ_smem >> 4) & 0x3FFF);
|
||||
desc_q |= (static_cast<uint64_t>(16) & 0x3FFF) << 16;
|
||||
desc_q |= (static_cast<uint64_t>(128) & 0x3FFF) << 32;
|
||||
desc_q |= (static_cast<uint64_t>(1) << 46);
|
||||
|
||||
// K-major NONE for (128, 64) BF16:
|
||||
// LBO=16, SBO=32
|
||||
uint64_t desc_k = 0;
|
||||
desc_k |= (static_cast<uint64_t>(sK_smem >> 4) & 0x3FFF);
|
||||
desc_k |= (static_cast<uint64_t>(16) & 0x3FFF) << 16;
|
||||
desc_k |= (static_cast<uint64_t>(32) & 0x3FFF) << 32;
|
||||
desc_k |= (static_cast<uint64_t>(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();
|
||||
|
||||
|
||||
@@ -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<s_k;j++) {
|
||||
|
||||
Reference in New Issue
Block a user