Conversation
🤖 ds-review-bot Code Reviewv6未发现本次变更新引入的明确缺陷。MX 清理前的全 CTA 同步与数学线程 barrier 相互独立,清理范围及 FP8、paged 路径语义保持不变。当前环境缺少 NVCC 和 GPU,未复跑 CUDA 测试或性能验证。 v5The MR moves MXFP8/MXFP4 MQA-logits cleanup in v4This MR relocates MXFP8/MXFP4 clean_logits cleanup out of the specialized cleaner warp's concurrent branch and runs it once after all producer/consumer work and TMEM release, synchronized on a dedicated user named barrier (ID 1). The existing grid-stride cleaner is factored into a local clean_logits lambda that both the concurrent FP8 path and the deferred MX path reuse, and the lambda's if constexpr (kCleanLogits) guard prevents the contiguous-only make_cleaner() from being instantiated for the paged scheduler. The cleanup ranges, in-coverage masks, metadata layout and public API are unchanged, and the MR reports bitwise-identical outputs together with 1.6-2.3x MX speedups on GB200/NVCC 13.3.33. The implementation looks correct: the deferred writes are disjoint from the math writes, barrier 1 is a separate hardware barrier from the math-only barrier 0 so the two full/partial CTA synchronizations cannot deadlock, and every warp role reaches the new barrier. My comments are about the toolchain-specific nature of the tuning, maintainability of the scheduler coupling, and the loss of a permanent regression guard rather than correctness. Files reviewed: 1 |
f2f2a0d to
7710a5b
Compare
On GB200 with NVCC 13.3.33, placing MXFP8/MXFP4 MQA-logits cleanup after the CTA's main computation improves
clean_logits=Trueperformance by about 1.6–2.3x for the H64/D128 cases below, with identical complete outputs. This PR proposes that scheduling change for the SM100 MX paths.Keep the existing inline cleaner block for FP8. For MXFP8/MXFP4, run the same cleaning loop at the end of the core, after the CTA's producer/consumer work and TMEM release, using a separate named barrier (ID 1) from the math-only barrier (ID 0). Both paths support
clean_logits=True; FP8 stays concurrent because deferring it caused a 5–8% slowdown in some exploratory cases. The two inline blocks preserve the desired code placement without introducing a helper or lambda. The cleanup ranges, in-coverage masks, metadata layout, and public API are unchanged.Performance
Baseline:
66081d4c9c7d7c44f13fea402e5b622aa0f409c2. Same GB200 GPU (152 SMs), NVCC 13.3.33, default DeepGEMM JIT flags targetingsm_100f, driver 580.167.08, PyTorch 2.13.0+cu130; PDL disabled.H=64, D=128, FP32 weights/output, uncompressed logits. Q/K/weights are all zero, MX scales are one, and each row uses
[0, KV-Q+row). These are CUDA Graph GPU timings of the completeclean_logits=Truecall, including all cleaning. Metadata is precomputed outside timing. Five calls per graph with distinct retained output buffers; 40 samples, rotating/reversing variant order, immediate warmup before each timed replay; median per-call latency. The timing runs use normal GPU clocks. Separate Nsight captures below use controlled base clocks.The
clean_logits=Falsecontrol stays within 0.2% between builds for these eight points. A further 18-shape screen (both metadata modes) includes H=8/32/64, Q=KV=8192, CP windows, and BF16 weights/output: MX latencies improve in every screened case; FP8 is within 0.2% of baseline. The improvement also appears with random inputs. The table uses an all-zero workload to reduce input-dependent power effects.The deferred schedule is an empirical tuning choice. When changing SM100 hardware or NVCC, remeasure both schedules with the complete-output benchmark below. Prefer concurrent cleanup if it matches or outperforms the deferred schedule within measurement variability on the target workloads; the full-CTA barrier is required by the deferred schedule, not by the MQA operation itself.
Timing data, including p10/p90 (microseconds)
Investigation notes
A few local ablations helped distinguish cleanup bandwidth costs from sensitivity to code placement. These use H64/D128, Q=2048, KV=32768:
clean_logits=False. The cleaner-loop unrolling and inline-PTX store variants showed similar timings to baseline.With the final patch, a controlled-clock Nsight Compute 2026.2 comparison for MXFP8, Q=2048, KV=32768, no metadata shows:
per_issue_activeratio)Exploratory PC samples place much of the additional wait around the math-side TMEM-result barrier check. Together, these results suggest that cleaner placement affects main-loop execution; memory traffic and resource allocation are largely unchanged, and neither binary has register spills. The precise low-level dependency remains to be investigated. The proposed change avoids this sensitivity on the tested configuration while retaining the existing output semantics.
Validation
test_paged_mqa_logitspasses all 4 sampled configurations, exercising the shared core with paged schedulers.test_mqa_logitsconfigurations and Compute Sanitizer synccheck/memcheck on 32 compact MX output configurations. Those broader checks were not rerun for this source cleanup.No permanent test-matrix or benchmark-file additions. Existing upstream tests can be run from a built checkout with:
Standalone performance reproducer (writes every timing sample to JSON)
Save as
pr-reproduce.pyand run it once with each built revision, using the same GPU and NVCC. Both builds and all JIT compilations must finish before timed sampling; the script warms every specialization before capture. It verifies the zero output and complete negative-infinity mask before timing.