fix: SMEM alignment in TMA K-only test
This commit is contained in:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user