test: add 8KB padding after sQ to prevent MMA read overrun

This commit is contained in:
2026-05-28 09:43:17 +00:00
parent 764ed01d6f
commit bcc6ed114d

View File

@@ -32,7 +32,8 @@ test_umma_qk_hd16(
extern __shared__ char sbuf[];
uint32_t* sTmemBase = (uint32_t*)sbuf;
bf16_t* sQ = (bf16_t*)(((uintptr_t)(sbuf + 4) + 15) & ~(uintptr_t)15);
bf16_t* sK = sQ + 128 * 16;
// Add 8KB padding after sQ to prevent MMA from reading into sK
bf16_t* sK = sQ + 128 * 16 + 4096; // 4096 BF16 = 8KB padding
float* sQ_row = (float*)(sK + 128 * 16);
for (int d = tid; d < 16; d += NTHREADS)
@@ -54,9 +55,9 @@ test_umma_qk_hd16(
// Descriptors
uint32_t sQ_smem = __cvta_generic_to_shared(sQ);
uint32_t sK_smem = __cvta_generic_to_shared(sK);
uint64_t desc_q = make_umma_desc_kmajor_none(sQ_smem, 64); // Try M=64
uint64_t desc_k = make_umma_desc_kmajor_none(sK_smem, 64); // Try M=64
uint32_t idesc = make_idesc(64, 128); // M=64, N=128
uint64_t desc_q = make_umma_desc_kmajor_none(sQ_smem, 128);
uint64_t desc_k = make_umma_desc_kmajor_none(sK_smem, 128);
uint32_t idesc = make_idesc(128, 128); // M=128, N=128
// Verify SMEM Q and K by reading back row 0
if (tid == 0) {
@@ -133,7 +134,7 @@ int main() {
cudaMemset(d_s_out, 0, 256 * sizeof(float));
cudaMemset(d_s_scalar, 0, SK * sizeof(float));
int smem = (4 + 16 + 128*16*2 + 128*16*2 + 16*4 + 256 + 127) & ~127;
int smem = (4 + 16 + 128*16*2 + 4096*2 + 128*16*2 + 16*4 + 256 + 127) & ~127;
test_umma_qk_hd16<<<1, NTHREADS, smem>>>(d_q, d_k, d_s_out, d_s_scalar, SCALE);
cudaError_t err = cudaDeviceSynchronize();