test: single PV K-tile debug

This commit is contained in:
2026-05-28 13:43:24 +00:00
parent 3a40ed6d69
commit 482328160a

View File

@@ -137,12 +137,12 @@ test_fmha_ts(const bf16_t* q, const bf16_t* k, const bf16_t* v,
// MMA: M=128, N=16. K=16 per MMA call, 8 K-tiles total.
uint32_t idesc_pv = make_idesc(BLOCK_MN, HD);
for (int kt = 0; kt < VKT; kt++) {
for (int kt = 0; kt < 1; kt++) { // DEBUG: only first PV K-tile
bf16_t* sv = sV_base + kt * V_TILE_SZ;
uint64_t dv = make_umma_desc_kmajor_none(__cvta_generic_to_shared(sv), MMA_K_BF16);
uint32_t tmem_a = tb + kt * MMA_K_BF16; // P's K-tile columns [16*kt, 16*kt+15]
if (tid == 0) umma_ts_f16(tb_o, tmem_a, dv, idesc_pv, kt > 0);
if (tid == 0) umma_ts_f16(tb_o, tmem_a, dv, idesc_pv, false); // no accumulate for first tile
asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");
__syncthreads();
}