Conversation
… descriptors (deepseek-ai#433) The fp8 row is head_dim bytes, and the host TMA descriptors plus the tma::copy sites swizzle in head_dim-byte atoms (32B/64B/128B for D=32/64/128). The kernel readers hardcoded a 128-byte pattern, so for head_dim 32/64 every fragment load un-swizzled with the wrong XOR bits and consumed out-of-tile shared memory — silently wrong, run-to-run nondeterministic in-window logits. Derive the mode from kHeadDim like the SM90/SM100 siblings do.
| // the TMA writes are swizzled in head_dim-byte atoms (32B/64B/128B for D=32/64/128). | ||
| // A hardcoded 128 here only matches D=128; for D=32/64 the readers would un-swizzle | ||
| // with the wrong pattern and consume out-of-tile shared memory. | ||
| static constexpr uint32_t kSwizzleMode = kHeadDim; |
There was a problem hiding this comment.
🔵 suggestion: 可选(nit):建议在 kSwizzleMode 定义附近加一条 DG_STATIC_ASSERT(kHeadDim == 32 or kHeadDim == 64 or kHeadDim == 128, "..."),把合法 swizzle 取值固化在设备端,防止未来 dispatch 放宽 head_dim 时静默生成非法的 CuTeSwizzle(如 kHeadDim=256 会得到 Swizzle<4,4,3>,超出硬件 B128 上限)。目前主机侧已有 head_dim 断言兜底,不阻塞合并。
🤖 v5
| // the TMA writes are swizzled in head_dim-byte atoms (32B/64B/128B for D=32/64/128). | ||
| // A hardcoded 128 here only matches D=128; for D=32/64 the readers would un-swizzle | ||
| // with the wrong pattern and consume out-of-tile shared memory. | ||
| static constexpr uint32_t kSwizzleMode = kHeadDim; |
There was a problem hiding this comment.
🟡 warning: The same hard-coded kSwizzleMode = 128 remains in deep_gemm/include/deep_gemm/impls/sm120_fp8_paged_mqa_logits.cuh, but the fp8 paged path does not appear to be head_dim-128-only: fp8_fp4_paged_mqa_logits in csrc/apis/attention.hpp allows fp8 head_dim 32/64/128, sm120_fp8_paged_mqa_logits in csrc/jit_kernels/impls/sm120_mqa_logits.hpp passes head_dim as the TMA descriptor swizzle and does not assert head_dim == 128 (unlike its fp4 sibling), and tests/test_attention.py enumerates SM120 fp8 paged D=32/64/128. For D=32/64 the paged reader would hit the same XOR/row-address mismatch that this MR fixes in the dense kernel, while the PR description states both paged variants are head_dim-128-only. Could you confirm the paged fp8 variant is truly unaffected, or derive its kSwizzleMode from kHeadDim as well?
🤖 v4f
🤖 ds-review-bot Code Reviewv6修改使读取侧 swizzle 模式与 TMA 描述符及写入布局保持一致,并覆盖所有允许的 head_dim(32、64、128),未发现会破坏现有行为的问题。 v5LGTM。该 MR 将 SM120 fp8 mqa_logits 内核读取侧的 kSwizzleMode 从硬编码的 128 改为由 kHeadDim 派生,修复 #433。经核实:(1) 主机侧 TMA 描述符(csrc/jit_kernels/impls/sm120_mqa_logits.hpp:101-106)以 head_dim 字节为 swizzle 原子并断言 head_dim ∈ {32,64,128},内核内 tma::copy<kHeadDim, ..., kHeadDim> 写入与之一致;此前硬编码 128 仅在 D=128 时与写入模式匹配,D=32/64 时读取端用错误的 XOR 位 un-swizzle,会读取相邻行或尚未填充的流水线 stage,导致窗口内 logits 静默错误且不可复现——修复方向正确。(2) sm120_utils.cuh 的 SwizzleContext/CuTeSwizzle = Swizzle<ctz(bytes)-4,4,3> 完整支持 32B/64B/128B(B32/B64/B128),kSwizzleMode = kHeadDim 落在合法取值范围内。(3) kSwizzleAlignment = kHeadDim * 8 与 kSMEMKBytes = kHeadDim 原本就由 head_dim 派生,SMEM_*_SIZE_PER_STAGE % kSwizzleAlignment == 0 的静态断言在 D=32/64/128 下均成立,无需额外改动。(4) 与兄弟内核一致:SM90 用 to_swizzle_cute_type<kHeadDim>(),SM100 用 kHeadDim / kPackFactor。(5) 未受影响的内核确认无需改动:sm120_fp8_paged(kSwizzleMode=128,主机断言 head_dim==128)、sm120_fp4 两个变体(kSwizzleMode=64 = fp4 行字节数 head_dim/2,主机断言 head_dim==128 且描述符用 head_dim/2 swizzle),dense fp8 确为唯一允许多种行宽的实例化。改动最小、注释解释充分,作者提供了 RTX 5090 上 384 配置扫描(189 失败 → 0)和完整 test_mqa_logits 通过的验证。建议合并。 v4fThe dense-kernel change is correct and well-scoped. The host TMA descriptors and the tma::copy sites swizzle in head_dim-byte atoms (32B/64B/128B for D=32/64/128), so deriving kSwizzleMode from kHeadDim instead of hard-coding 128 makes SwizzleContext un-swizzle with exactly the XOR pattern TMA wrote. For D=128 the value is unchanged; for D=32/64 the previous 128-bit pattern XORed high column/row bits into the row-address bits and produced out-of-row fragment loads, matching the reported nondeterminism. The explanatory comment is accurate and the derivation mirrors the SM90 (to_swizzle_cute_type<kHeadDim>()) and SM100 (kHeadDim / kPackFactor) siblings. One follow-up is raised below about the fp8 paged variant, whose description claim looks inconsistent with the code. Files reviewed: 1 |
…c asserts (deepseek-ai#433) The fp8 paged dispatch has no head_dim == 128 restriction (unlike its fp4 sibling), and the host descriptors pass swizzle_mode = head_dim, so the paged kernel had the same read/write swizzle mismatch for D=32/64: on RTX 5090 D a targeted fp8 paged repro shows accuracy diff of 0.83 (D=32) / 0.72 (D=64) against the test's reference, while self-consistency passes — the paged pipeline's misreads land on deterministically-stale smem, so the wrong values are bit-stable across runs and only the accuracy comparison exposes them. Deriving kSwizzleMode from kHeadDim drops the diff to ~1e-3 (D=32: 0.000834, D=64: 0.000952; D=128 unchanged). Also add DG_STATIC_ASSERT pinning head_dim to {32, 64, 128} next to the derived constants (per review suggestion), so a future dispatch widening cannot silently instantiate an invalid CuTeSwizzle.
|
Addressed in 5825212 — and thank you for catching this, the original description's "both paged variants are head_dim-128-only" claim was wrong: the Empirically confirmed on the 5090 D with a targeted fp8 paged repro (small configs; the full paged suite exceeds 32 GB here): accuracy diff vs the test's reference was 0.83 (D=32) / 0.72 (D=64) while self-consistency passed — in the paged variant the misreads land on deterministically-stale shared memory, so the wrong values are bit-stable across runs and only the accuracy comparison exposes them. After deriving kSwizzleMode from kHeadDim: D=32: 0.0008, D=64: 0.001 (D=128 unchanged). Also added the suggested |
|
Cross-reference: leavelet had already identified this bug and prepared an essentially identical one-line fix in #379 (open since 7/15, bot-approved, covering both the dense and paged kernels from the start). We missed that open PR before filing — apologies for the duplication. Deltas of this branch over #379: the Happy to close this in favor of #379 if maintainers prefer; the derivation is identical and both validate cleanly on our 5090 D. Have posted our validation data to #379 to help it land. |
Fixes #433.
The SM120 fp8 MQA-logits readers un-swizzle shared memory with a hardcoded 128-byte pattern, while the host TMA descriptors and the
tma::copysites swizzle in head_dim-byte atoms (32B/64B/128B for D=32/64/128 — the fp8 row is head_dim bytes). For head_dim 32/64 the XOR bits don't match the write pattern, fragment loads address outside their row, and the kernel returns in-window logits that are silently wrong. D=128 is the only head_dim where the hardcoded value matches, which is why the main DSA path works.This derives the mode from
kHeadDim, matching the sibling kernels: SM90 (to_swizzle_cute_type<kHeadDim>()) and SM100 (kHeadDim / kPackFactor).Both fp8 variants are fixed (second commit, addressing the review finding): the fp8 paged dispatch has no
head_dim == 128restriction — only its fp4 sibling does — so the paged kernel had the same mismatch. Notably, in the paged variant the resulting garbage is deterministic (misreads land on deterministically-stale shared memory), so self-consistency passes and only the accuracy comparison against the reference exposes it.Validation on RTX 5090 D (CUDA 13.0, PyTorch 2.11.0+cu130, nv_dev @ 2642b32):
test_mqa_logits(20x bitwise + accuracy vs reference + bench): PASSED, all 384 configs, fp8 up to 411 TFLOPSAlso adds
DG_STATIC_ASSERT(kHeadDim ∈ {32, 64, 128})next to the derived constants (per review suggestion), so a future dispatch widening cannot silently instantiate an invalidCuTeSwizzle.Unaffected kernels (verified): dense fp4 and both fp4 paged variants are head_dim-128-only (
DG_HOST_ASSERT/DG_STATIC_ASSERT), where their constants already match. Dense fp8 and fp8 paged are the only instantiations whose dispatches allow multiple row widths, and both are now derived.Scan script and raw logs: https://gist.github.com/aganhui/f783450b212217284c062c3b46300649