fix: align SMEM layout properly (128B aligned tmem + Q)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user