fix: align SMEM layout properly (128B aligned tmem + Q)

This commit is contained in:
2026-05-28 08:46:56 +00:00
parent 2a765be715
commit 9a51bfa578
2 changed files with 8 additions and 3 deletions

View File

@@ -26,9 +26,14 @@ fmha_qk_verify(
// 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)
extern __shared__ char sbuf[];
uint32_t* sTmemBase = (uint32_t*)sbuf;
bf16_t* sQ = (bf16_t*)(((uintptr_t)(sbuf + 4) + 127) & ~127);
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

View File

@@ -65,7 +65,7 @@ int main() {
cudaMemset(ds_out, 0, s_k*4);
// SMEM = 4 (tmem_base) + 128B align + 128*HD*2 (sQ) + 128*HD*2 (sK) + slack
int smem = 4 + 128 + 128 * HD * 2 + 128 * HD * 2 + 4096;
int smem = 256 + 128 * HD * 2 + 128 * HD * 2 + 4096;
dim3 grid(1, H, B);
dim3 block(NTHREADS);