diff --git a/tests/unit/test_umma_qk.cu b/tests/unit/test_umma_qk.cu index 5544d03b..ab220b04 100644 --- a/tests/unit/test_umma_qk.cu +++ b/tests/unit/test_umma_qk.cu @@ -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();