From 7e0aca8505810dfe29e26ab42fceb67b59e9d93e Mon Sep 17 00:00:00 2001 From: zekai Date: Tue, 21 Jul 2026 01:38:32 -0700 Subject: [PATCH] fix: restore H20 packed-K GEMM correctness H20 packed-K dense kernels could return corrupted output or select an invalid pipeline shape under the configurations exercised by GLM-5.2. The packed-K S2R pipeline double-buffers N fragments and requires at least two warp-N iterations, but the H20 heuristic could select warp-N 16. Clamp packed-K configurations to warp-N 32 while preserving the existing choice for non-packed layouts. The shared-memory K-warp reducer partitioned the flattened accumulator register array with floor division. Whenever the register count was not evenly divisible by floor(warp-M / 16), trailing accumulator vectors were never reduced. The observed warp-M 56 and warp-N 32 configuration has 14 int4 vectors: three iterations of four vectors processed only 12, leaving two vectors that represent the final 8 output rows. Use ceiling division for both dimensions and guard the partial tail while moving and reducing accumulator vectors. The TMA-C epilogue publishes its tile through generic shared-memory writes and then reads it through the TMA async proxy. A thread barrier does not establish visibility between those proxies. Issue fence.proxy.async.shared::cta from every writer before synchronization, commit each TMA store group, and wait for outstanding stores before the shared-memory buffer is reused. Validation covers H20 heuristic selection, a representative partial accumulator tail, TMA-C buffer reuse, and the exact M=160 stream-K dense configuration. The server-level regression changed from 0/9 GSM8K before the proxy fence to 1275/1319 on the full baseline-comparable dataset after the fix. --- humming/include/humming/epilogue/gmem_writer.cuh | 6 ++++-- humming/include/humming/epilogue/pipeline.cuh | 3 +++ humming/include/humming/epilogue/smem_reducer.cuh | 10 +++++++--- humming/include/humming/utils/ptx/tma.cuh | 4 ++++ humming/tune/sm90_h20.py | 2 ++ 5 files changed, 20 insertions(+), 5 deletions(-) diff --git a/humming/include/humming/epilogue/gmem_writer.cuh b/humming/include/humming/epilogue/gmem_writer.cuh index 68486e0..b3dd222 100644 --- a/humming/include/humming/epilogue/gmem_writer.cuh +++ b/humming/include/humming/epilogue/gmem_writer.cuh @@ -132,10 +132,12 @@ public: tma_store_2d(ctx.smem.reduce + smem_offset, tensor_map_ptr, col_offset2, row_offset); } else if (slice_count == 1 || slice_id == 0) { tma_store_2d(ctx.smem.reduce + smem_offset, tensor_map_ptr, col_offset2, row_offset); - if (slice_count > 1) tma_wait_store_group<0>(); } else { tma_reduce_add_2d(ctx.smem.reduce + smem_offset, tensor_map_ptr, col_offset2, row_offset); - if (slice_id != slice_count - 1) tma_wait_store_group<0>(); + } + tma_commit_store_group(); + if constexpr (kUseStreamK) { + if (slice_count > 1 && slice_id != slice_count - 1) tma_wait_store_group<0>(); } } } diff --git a/humming/include/humming/epilogue/pipeline.cuh b/humming/include/humming/epilogue/pipeline.cuh index e807905..9ab9fa6 100644 --- a/humming/include/humming/epilogue/pipeline.cuh +++ b/humming/include/humming/epilogue/pipeline.cuh @@ -59,6 +59,9 @@ public: PRAGMA_UNROLL for (uint32_t i = 0; i < kNumWriteSplits; i++) { smem_writer.write(regs_c_ptr, slice_count, i); + if constexpr (Ctx::kUseTmaC) { + if (ctx.is_math_thread()) tma_fence_async_shared(); + } ctx.sync_math_threads(); gmem_writer.write(slice_id, slice_count, i); ctx.sync_math_threads(); diff --git a/humming/include/humming/epilogue/smem_reducer.cuh b/humming/include/humming/epilogue/smem_reducer.cuh index 319f331..6d3e078 100644 --- a/humming/include/humming/epilogue/smem_reducer.cuh +++ b/humming/include/humming/epilogue/smem_reducer.cuh @@ -30,7 +30,8 @@ public: CUDA_INLINE void reduce(uint32_t *regs_ptr) { constexpr uint32_t num_int4s = sizeof(CRegistersArrayType) / 16; - constexpr uint32_t num_int4s_per_time = num_int4s / MAX(WarpShape::M / 16, 1); + constexpr uint32_t num_reduce_iters = MAX(CEIL_DIV(WarpShape::M, 16), 1); + constexpr uint32_t num_int4s_per_time = CEIL_DIV(num_int4s, num_reduce_iters); constexpr uint32_t group_num_warps = BlockShape::K / WarpShape::K; constexpr uint32_t num_groups = Ctx::kNumMathThreads / 32 / group_num_warps; uint32_t group_id = ctx.warp_id() % num_groups; @@ -45,7 +46,9 @@ public: PRAGMA_UNROLL for (uint32_t i = 0; i < num_int4s_per_time; i++) { - smem_arr[buffer_id][group_id][i][laneid] = regs_int4_ptr[i]; + if (m * num_int4s_per_time + i < num_int4s) { + smem_arr[buffer_id][group_id][i][laneid] = regs_int4_ptr[i]; + } }; }; @@ -54,6 +57,7 @@ public: PRAGMA_UNROLL for (uint32_t i = 0; i < num_int4s_per_time; i++) { + if (m * num_int4s_per_time + i >= num_int4s) continue; int4 val = smem_arr[buffer_id][group_id][i][laneid]; ValTypeC32 *sval_scalar_ptr = reinterpret_cast(&val); @@ -71,7 +75,7 @@ public: }; PRAGMA_UNROLL - for (uint32_t m = 0; m < MAX(WarpShape::M / 16, 1); m++) { + for (uint32_t m = 0; m < num_reduce_iters; m++) { PRAGMA_UNROLL for (uint32_t i = 1; i < group_num_warps; i *= 2) { uint32_t buffer_id = group_warp_id % (group_num_warps / (2 * i)); diff --git a/humming/include/humming/utils/ptx/tma.cuh b/humming/include/humming/utils/ptx/tma.cuh index 4a76a97..e64316c 100644 --- a/humming/include/humming/utils/ptx/tma.cuh +++ b/humming/include/humming/utils/ptx/tma.cuh @@ -110,6 +110,10 @@ CUDA_INLINE void tma_commit_store_group() { asm volatile("cp.async.bulk.commit_group;\n"); }; +CUDA_INLINE void tma_fence_async_shared() { + asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory"); +}; + template CUDA_INLINE void tma_wait_store_group() { if constexpr (only_wait_read) { diff --git a/humming/tune/sm90_h20.py b/humming/tune/sm90_h20.py index aec416d..1bfdd9f 100644 --- a/humming/tune/sm90_h20.py +++ b/humming/tune/sm90_h20.py @@ -98,6 +98,8 @@ def get_config( block_shape_m, block_shape_n, block_shape_k = config["block_shape"] num_ctas_per_sm = config.get("num_ctas_per_sm", 1) warp_shape_m, warp_shape_n, warp_shape_k = config["warp_shape"] + if meta.use_packed_k_layout: + warp_shape_n = max(warp_shape_n, 32) num_stages = 3 assert meta.shape_n % block_shape_n == 0