From 69bf20b09da53c61613b01319d1522159adb6c37 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Fri, 29 May 2026 18:43:44 +0000 Subject: [PATCH] fix: SMEM alignment in TMA K-only test --- tests/unit/test_tma_konly.cu | 21 +++++++++++++-------- 1 file changed, 13 insertions(+), 8 deletions(-) diff --git a/tests/unit/test_tma_konly.cu b/tests/unit/test_tma_konly.cu index 61f7a5cb..01dda2cb 100644 --- a/tests/unit/test_tma_konly.cu +++ b/tests/unit/test_tma_konly.cu @@ -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);