test: single PV K-tile debug
This commit is contained in:
@@ -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();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user