Conversation
…n SM120 (deepseek-ai#430) Port the deepseek-ai#343 tensormap drain to the SM120 k_grouped NT contiguous kernel (sm120_fp8_fp4_gemm_1d1d.cuh, copied from sm90_fp8_gemm_1d1d.cuh 12 days before deepseek-ai#343 fixed the same race on the SM90 side). Without the drain, the group-switch publishes GMEM tensormaps while TMA loads of the previous group may still be in flight, producing non-deterministic wrong results.
🤖 ds-review-bot Code Reviewv6该变更正确移植了 SM90 的 tensormap drain 逻辑,且同步位置与描述符更新流程匹配,未发现会破坏现有行为的问题。 v5本 MR(commit f115a07)将 #343 中修复 SM90 in-flight tensormap 竞态的 drain 逻辑,完整移植到 SM120 k_grouped NT contiguous 内核 sm120_fp8_fp4_gemm_1d1d.cuh 的 group-switch 路径。变更仅 7 行插入:在 smem 中的 tensormap 更新之后、发布到 GMEM(gmem_tm_a/gmem_tm_b 写入 + tensor_map_release_gpu/acquire_gpu)之前,插入 cute::tma_desc_commit_group() + cute::tma_desc_wait_group(),确保引用旧 tensormap 的在途 TMA load 全部完成;随后的单线程 _syncwarp(1u << lane_idx) 仅用于阻止 ptxas 将 GMEM store 重排到 wait 之前。经逐行比对,移植与已合入的 SM90 修复(sm90_fp8_gemm_1d1d.cuh 第 201–214 行)完全一致(仅变量名 smem_tm/gmem_tm_ 不同)。插入点选择正确:位于 kKGroupedConstantStride 常量步长与变量步长两个分支汇合之后,两条 smem tensormap 更新路径均被覆盖;该代码仅在 is_tma_leader(warp_idx == kNumMathWarps 且 lane_idx == 0)单线程执行,__syncwarp 掩码恒为 1u<<0,语义与 #343 中已确认有效的构造相同,无正确性问题。变更为纯插入、不触碰周边逻辑,回归面最小;PR 附带 RTX 5090 D 验证(修复前 13/14 次非确定性失败,修复后 bit-identical、100 次 soak 干净),性能代价(哨兵形状 +0.9%、典型 EP 形状 +0.26%)与 SM90 在 #343 中接受的权衡相同。结论:修复正确、范围最小、与上游先例一致,建议合入。 v4fApprove. Correct, minimal, faithful port of the #343 tensormap drain to the SM120 FP8/FP4 1D1D kernel (sm120_fp8_fp4_gemm_1d1d.cuh). This kernel was copied from sm90_fp8_gemm_1d1d.cuh before #343 landed there, so the SM120 K-grouped group-switch path still published per-CTA GMEM tensormaps (gmem_tm_a/b) while cp.async.bulk.tensor loads issued for the previous group could still have outstanding descriptor reads, producing the non-deterministic wrong results described in #430. The MR inserts cute::tma_desc_commit_group() + cute::tma_desc_wait_group() plus the ptxas-ordering __syncwarp immediately before the single GMEM publish site, byte-identical to the merged SM90 #343 fix. Verification performed on the source and vendored cute headers: gmem_tm_a/b have exactly one write site (lines 283-284) and are referenced only by TMA loads issued by the same single leader thread (warp kNumMathWarps, lane 0) that executes the drain, so the per-thread cp.async.bulk commit/wait scope is complete; cute::tma_desc_wait_group() maps to cp.async.bulk.wait_group.read 0 (third-party/cutlass include/cute/arch/copy_sm90_tma.hpp), which per the cute comment waits only until outstanding TMA descriptor reads are safe to modify (a descriptor-read drain, not a full-data drain, matching the small measured perf cost); lane_idx is in scope; includes and the CUTE_ARCH_TMA_SM90_ENABLED guard are satisfied for SM120. Edge cases are benign: the first group switch drains an empty group, and the drain is compiled out for non-KGroupedContiguous instantiations, so other gemm types get zero code/perf impact. No deadlock hazard: wait_group.read depends only on async-proxy descriptor fetch, not on consumer progress or the empty/full barrier pipeline. Empirical evidence in the description is strong (pre-fix non-determinism in 13/14 runs resolved to bit-identical runs, diff 0.000105; full k_grouped + m_grouped suites and 100x soak clean; disclosed perf cost +0.9% worst-case group-switch density / +0.26% typical EP, consistent with the trade-off SM90 accepted in #343). Could not re-run GPU validation in this environment. One non-blocking follow-up below for the sibling sm120_bf16_gemm.cuh kernel. Files reviewed: 1 📍 未定位到 diff 的评论🔵 suggestion 🔵 suggestion |
Fixes #430.
The SM120 1D1D kernel (
sm120_fp8_fp4_gemm_1d1d.cuh, introduced in #324) was copied fromsm90_fp8_gemm_1d1d.cuhtwelve days before #343 fixed the in-flight tensormap race there. Since the copy lives at a different file path, the #343 fix never reached it. This ports the same drain to the SM120 group-switch path so that in-flight TMA loads complete before the grouped tensormaps are published to GMEM.Validation on RTX 5090 D (SM120, CUDA 13.0, PyTorch 2.11.0+cu130, nv_dev @ 2642b32):
test_k_grouped_gemm_contiguousfails at the first shape in 13/14 runs; diff 0.0015–0.006, non-deterministic on identical inputsNote on
__syncwarp: same construct as the merged SM90 fix — for the prior discussion of its validity see #343.sm120_bf16_gemm.cuhhas the same un-drained pattern, but we could not reproduce failures there (full suite + 50 targeted stress runs pass), so this PR leaves it alone — details in the linked issue.Evidence and raw logs: https://gist.github.com/aganhui/3b374ffb3d6d13b52d5688780fafac28