fix: SMEM alignment in TMA K-only test

This commit is contained in:
2026-05-29 18:43:44 +00:00
parent 2c0ee69aea
commit 69bf20b09d

View File

@@ -45,15 +45,20 @@ fmha_tma_konly_kernel(
const int wid = tid / 32;
const int lane = tid % 32;
// SMEM
// SMEM — carefully aligned
extern __shared__ __align__(128) char sbuf[];
uint32_t* sTmemBase = (uint32_t*)sbuf;
bf16_t* sQ0 = (bf16_t*)(((uintptr_t)(sbuf + 4) + 15) & ~(uintptr_t)15);
bf16_t* sK0 = sQ0 + TILE_SZ;
// TMA staging buffer after sK0
bf16_t* sTmaBuf = (bf16_t*)(((uintptr_t)(sK0 + TILE_SZ) + 127) & ~(uintptr_t)127);
// mbarrier after TMA buffer
uint64_t* sMbar = (uint64_t*)(((uintptr_t)(sTmaBuf + TILE_SZ) + 15) & ~(uintptr_t)15);
size_t off = 0;
uint32_t* sTmemBase = (uint32_t*)sbuf; off = 4;
// 16-byte align for Q0
off = (off + 15) & ~(size_t)15;
bf16_t* sQ0 = (bf16_t*)(sbuf + off); off += TILE_SZ * sizeof(bf16_t);
bf16_t* sK0 = (bf16_t*)(sbuf + off); off += TILE_SZ * sizeof(bf16_t);
// 128-byte align for TMA buffer
off = (off + 127) & ~(size_t)127;
bf16_t* sTmaBuf = (bf16_t*)(sbuf + off); off += TILE_SZ * sizeof(bf16_t);
// 16-byte align for mbarrier
off = (off + 15) & ~(size_t)15;
uint64_t* sMbar = (uint64_t*)(sbuf + off); off += 8;
// Init
if (wid == 1) tmem_alloc(__cvta_generic_to_shared(sTmemBase), TMEM_N);