Skip to content

fix: derive the SM120 fp8 mqa_logits swizzle mode from head_dim - #434

Open
aganhui wants to merge 2 commits into
deepseek-ai:nv_devfrom
aganhui:fix/sm120-fp8-mqa-swizzle
Open

aganhui wants to merge 2 commits into
deepseek-ai:nv_devfrom
aganhui:fix/sm120-fp8-mqa-swizzle

Conversation

@aganhui

@aganhui aganhui commented Sep 10, 2026 •

Copy link
Copy Markdown

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::copy sites 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 == 128 restriction — 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):

  • Dense, 384-config self-consistency scan: 189 failures → 0 (all failures were fp8 with D ∈ {32, 64}; zero at D=128, zero mxfp4)
  • Dense, full test_mqa_logits (20x bitwise + accuracy vs reference + bench): PASSED, all 384 configs, fp8 up to 411 TFLOPS
  • Paged, targeted fp8 repro (small configs; the full paged suite does not fit in 32 GB): accuracy diff vs reference D=32: 0.83 → 0.0008, D=64: 0.72 → 0.001, D=128 unchanged (~0.001)
  • Dense regression after adding the static asserts: fp8 D=32 accuracy diff 0.0006 (< 1e-3)

Also 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 invalid CuTeSwizzle.

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

… 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;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 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;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 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

Copy link
Copy Markdown
Collaborator

🤖 ds-review-bot Code Review

v6

修改使读取侧 swizzle 模式与 TMA 描述符及写入布局保持一致,并覆盖所有允许的 head_dim(32、64、128),未发现会破坏现有行为的问题。

v5

LGTM。该 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 通过的验证。建议合并。

v4f

The 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
Issues found: 🟡 1 warning | 🔵 1 suggestion
Inline comments posted: 2

…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.
@aganhui

aganhui commented Sep 10, 2026

Copy link
Copy Markdown
Author

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 head_dim == 128 assert belongs to the fp4 paged dispatch only (sm120_fp4_paged_mqa_logits), while sm120_fp8_paged_mqa_logits has no head_dim restriction and its host descriptors pass swizzle_mode = head_dim.

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 DG_STATIC_ASSERT pinning head_dim to {32, 64, 128} in both files. Dense regression after the asserts: fp8 D=32 accuracy diff 0.0006.

@aganhui

aganhui commented Sep 11, 2026

Copy link
Copy Markdown
Author

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 DG_STATIC_ASSERT guards and a longer explanatory comment. What #433 adds over both PRs is the evidence/documentation layer: the 189-config failure surface, the symptom forensics (nondeterministic in dense vs deterministically wrong in paged — only accuracy comparison exposes the latter), the cross-arch derivation analysis, and the connection to the misdiagnosed reports in #392.

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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants