test: add 8KB padding after sQ to prevent MMA read overrun
This commit is contained in:
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user