Skip to content

dsv4.1: RoPE and FP4 packing kernels - #39656

Merged
hnyls2002 merged 39 commits into
mainfrom
dsv4.1-rope-fp4
Sep 16, 2026
Merged

hnyls2002 merged 39 commits into
mainfrom
dsv4.1-rope-fp4

Conversation

@hnyls2002

@hnyls2002 hnyls2002 commented Sep 15, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

  • Fused RoPE and FP4 packing kernels for the DeepSeek-V4.1 low-ratio indexer: keys and queries are rotated, fake-quantized to the per-32 UE8M0 FP4 grid and packed into the 68-byte index-K format (64 e2m1 bytes plus four block exponents) that the paged FP4 MQA logits kernel reads.
  • fp4_indexer_rope.cuh / fp4_indexer_rope.py: one warp per row. index_k_norm_rope_pack_store fuses RMSNorm, RoPE at the group's first position, both quantization stages and the paged cache store (slot 0 publishes nothing); index_q_rope_pack_weights does the query side and computes the indexer head weights in the same launch.
  • fp4_rope_fake_quant.py: Triton RoPE tail plus FP4 fake-quant for the torch-shaped path, per-32 UE8M0 by default and per-16 E4M3 for compressed KV, preserving the bf16 round-trip after RoPE and round-half-to-even.
  • fp4_indexer.py: _index_k_rope_pack_kernel / index_k_rope_pack, the Triton RoPE + fake-quant + pack + store kernel for the prefill path; _fp4_e2m1_code_rne and an rne flag on quantize_fp4_indexer_tensor / store_fp4_index_k_cache for the reference's round-to-nearest-even; INDEX_K_SLOT_BYTES names the slot size the store kernel already asserted.

Changes to existing kernels

  • fp4_indexer.py: rne defaults to False, so the existing C4 indexer path keeps its threshold rounding; the kernel gains an RNE constexpr branch and nothing else.

Verification

  • The CUDA kernels, the Triton pack kernel and the fake-quant + existing packer path agree byte for byte: 504/504 comparisons over rows in {1, 5, 64, 300}, input scales {1, 40, 1e-3} plus zero, subnormal and exact half-way rows, compress ratios {1, 2, 4} and int32/int64 positions, covering the packed payload, the block-exponent words, the paged cache pages with slot-0 suppression, and the fp32 head weights.
  • test_fp4_indexer.py guards the unchanged default rounding of the existing kernel.

CI States

Latest PR Test (Base): 🚫 Run #35160475205
Latest PR Test (Extra): ❌ Run #35160474900
Latest PR Test (AMD ROCm 10): ❌ Run #35160475128

@github-actions github-actions Bot added quant LLM Quantization jit-kernel labels Sep 15, 2026
@hnyls2002
hnyls2002 added this pull request to stack #39658 September 15, 2026 22:10
@hnyls2002
hnyls2002 removed this pull request from stack #39658 September 15, 2026 22:31
@hnyls2002
hnyls2002 added this pull request to stack #39667 September 15, 2026 22:32
@hnyls2002
hnyls2002 removed this pull request from stack #39667 September 15, 2026 22:43
@hnyls2002
hnyls2002 added this pull request to stack #39669 September 15, 2026 22:43
@hnyls2002
hnyls2002 removed this pull request from stack #39669 September 15, 2026 22:56
@hnyls2002
hnyls2002 requested a review from Ying1123 as a September 15, 2026 22:57
hnyls2002 and others added 2 commits September 16, 2026 03:09
flash_c2_decode_kernel loads the RoPE frequencies twice for the same
thread and the same position: once right after the norm weight, and
again inside the `tx >= kNopeThreads` branch that consumes them. `freq`
is never read between the two, so the first load is dead and the second
re-issues the same 8 bytes per lane.

Keep the early load -- it is issued before the softmax and the RMSNorm
reduction, so its latency overlaps that work -- and remove the reload
inside the branch.

Verified bit-exact on GB300 (sm_103) against the pre-change kernel:
identical sha256 of `out`, the paged cache and the pair-state ring for
V4, V41 and V41_FP4 decode plus a draft_len=4 verify batch, with the
JIT cache cleared between runs. A deliberate eps perturbation was used
as a negative control to confirm the harness recompiles and detects
kernel changes.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
@BBuf

BBuf commented Sep 16, 2026

Copy link
Copy Markdown
Collaborator

Read fp4_rope.cuh closely; the tiling argument is well made and the tricky bits hold up:

  • clear_negative_zero is correct — (packed | packed>>1 | packed>>2) & 0x11111111 sets bit 4k iff any magnitude bit of nibble k is set, and no shift is large enough to bleed across nibbles, so the sign survives exactly when the magnitude is non-zero. Matching the reference packer's sign = (x < 0) & (idx != 0) rather than the hardware's -0 is the right call and the comment says why.
  • index_pack_exponent reproduces torch_quant.ceil_pow2's biased exponent ((bits >> 23) + (mantissa != 0)), and deliberately keeps amax / 6.0f and a post-divide floor rather than reusing fp4::block_scale — the "the two stages therefore require separate scales" comment is the kind of thing that saves a future reader an afternoon.
  • index_scale_word being valid only on the lower half-warp is documented; worth keeping that comment glued to it.
  • Using libdevice in the Triton siblings is consistent with getting bit-exactness against torch rather than approximating it.

The one gap: both sides of an oracle are in this PR and nothing compares them

This PR ships the fused CUDA path and Triton implementations of the same math:

CUDA Triton
flash_index_k_kernel / index_rope_quant_pack (fp4_rope.cuh) _rope_fake_quant_pack_indexer_kernel (rope_pack_indexer.py)
the RoPE + fake-quant half of the above _rope_tail_fake_quant_fp4_kernel (rope_fake_quant_fp4.py)

and rope_pack_indexer.py's docstring states the invariant that must hold between them:

Keep both quantization stages: the indexer packer has a different scale floor from fake_quant_fp4, so directly packing the first stage is not equivalent.

No test in the PR exercises index_k_norm_rope_pack_store, index_q_rope_pack, index_q_rope_pack_weights, rope_tail_fake_quant_fp4 or rope_fake_quant_pack_indexer. A byte-comparison of the two implementations over a handful of shapes would pin the packed payload, the four block exponents, and the -0 handling in one assert — and it is the only thing that will catch the two drifting, since they are edited independently and both are "the reference" depending on which file you are in.

Two specific things I would want a test to hold, because they are stated as reachability arguments in comments rather than checked:

  • index_pack_exponent: "Neither bound is reachable for finite fp32 inputs" — the min(max(exponent, 1), 254) clamp.
  • index_rope_quant_pack: "inv_scale_ue8m0 instead of the reference's division [...] except at exponent 254, which needs a block absmax above 6 * 2^126 and so cannot come from a finite float."

Both are true as far as I can tell, but they are the kind of invariant that a later change to the amax floor quietly invalidates.

Smaller

kFp4RopeWarpsPerCTA = 4 is justified by "B200 decode measurements, where occupancy has little effect". The stack targets GB300 too — was 4 re-checked on sm_103, or is it carried over? Not asking for a re-tune, just whether the comment should say B200-and-GB300 or B200-only.

(Review by Claude Opus 5, run by @BBuf.)

Base automatically changed from dsv4.1-candidate to main September 16, 2026 22:15
@hnyls2002

Copy link
Copy Markdown
Collaborator Author

/rerun-test test/registered/kernels/ops/attention/test_fp4_indexer.py

@github-actions

github-actions Bot commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

Results for /rerun-test test/registered/kernels/ops/attention/test_fp4_indexer.py:

🚀 1-gpu-h100 (1 test): ✅ View workflow run

cd test/ && python3 registered/kernels/ops/attention/test_fp4_indexer.py

@hnyls2002

Copy link
Copy Markdown
Collaborator Author

/tag-and-rerun-ci

2 similar comments
@hnyls2002

Copy link
Copy Markdown
Collaborator Author

/tag-and-rerun-ci

@hnyls2002

Copy link
Copy Markdown
Collaborator Author

/tag-and-rerun-ci

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

Labels

jit-kernel quant LLM Quantization run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants