Repository navigation
dsv4.1: RoPE and FP4 packing kernels - #39656
Conversation
0ab8faf to
3186760
Compare
5230624 to
8dba862
Compare
3186760 to
cd871a9
Compare
8dba862 to
ce26eb1
Compare
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>
|
Read
The one gap: both sides of an oracle are in this PR and nothing compares themThis PR ships the fused CUDA path and Triton implementations of the same math:
and
No test in the PR exercises Two specific things I would want a test to hold, because they are stated as reachability arguments in comments rather than checked:
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
(Review by Claude Opus 5, run by @BBuf.) |
…_quant next to the kernels; align compress kernel names
…dexer; share slot and floor constants
|
/rerun-test test/registered/kernels/ops/attention/test_fp4_indexer.py |
|
Results for 🚀 |
|
/tag-and-rerun-ci |
2 similar comments
|
/tag-and-rerun-ci |
|
/tag-and-rerun-ci |
Summary
fp4_indexer_rope.cuh/fp4_indexer_rope.py: one warp per row.index_k_norm_rope_pack_storefuses RMSNorm, RoPE at the group's first position, both quantization stages and the paged cache store (slot 0 publishes nothing);index_q_rope_pack_weightsdoes 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_rneand anrneflag onquantize_fp4_indexer_tensor/store_fp4_index_k_cachefor the reference's round-to-nearest-even;INDEX_K_SLOT_BYTESnames the slot size the store kernel already asserted.Changes to existing kernels
fp4_indexer.py:rnedefaults toFalse, so the existing C4 indexer path keeps its threshold rounding; the kernel gains anRNEconstexpr branch and nothing else.Verification
test_fp4_indexer.pyguards 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