debug: skip hd_chunk>0 to isolate chunk0

This commit is contained in:
2026-05-30 06:59:01 +00:00
parent 25aeaca9ab
commit 9227b0e93f

View File

@@ -130,6 +130,15 @@ fmha_6warp_tma_multirow_multitile_kernel(FmhaTmaMultiRowMultiTileParams params)
const int n_sub_end = n_sub_start + N_NSUB_CHUNK;
int phase = 0; // reset mbarrier phase per hd_chunk
// DEBUG: skip hd_chunk > 0 to test just chunk 0
if (hd_chunk > 0) {
if (my_row_active) {
for (int d = 0; d < HD_CHUNK; d++)
o_head[my_row * HD + hd_chunk_start + d] = f32_to_bf16(0.0f);
}
continue;
}
// Zero accumulator
for (int i = tid; i < MAX_ROWS * HD_CHUNK; i += 192) sOacc[i] = 0.0f;
for (int i = tid; i < MAX_ROWS; i += 192) {
@@ -330,6 +339,13 @@ fmha_6warp_tma_multirow_multitile_kernel(FmhaTmaMultiRowMultiTileParams params)
// ---- Write chunk to GMEM ----
if (my_row_active) {
float inv_rs = 1.0f / sRunningSum[my_row];
// Debug: for first hd_chunk, first row, print some values
if (my_row == 0 && hd_chunk == 0) {
// Print first 8 O values and running stats
printf("hd_chunk=%d rmax=%f rsum=%f O[0..3]=%f %f %f %f\n",
hd_chunk, sRunningMax[0], sRunningSum[0],
sOacc[0] * inv_rs, sOacc[1] * inv_rs, sOacc[2] * inv_rs, sOacc[3] * inv_rs);
}
for (int d = 0; d < HD_CHUNK; d++) {
o_head[my_row * HD + hd_chunk_start + d] = f32_to_bf16(sOacc[my_row * HD_CHUNK + d] * inv_rs);
}