Repository navigation
Conversation
ormandj
requested review from
BBuf,
DarkSharpness,
HaiShaw,
HydraQYH,
celve and
yuan-luo
as code owners
October 7, 2026 02:52
ormandj
force-pushed
the
dsv4-fp4-indexer-exp2-free
branch
from
October 7, 2026 05:42
f42ef09 to
d03509e
Compare
This was referenced Oct 7, 2026
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Motivation
_fp4_index_logits_tile(used byfp4_index_logits_decodeandfp4_index_logits_paged) decodes every K element withtl.exp2: one for the E2M1 exponent and one for the E8M0 scale, which is evaluated over the whole[BLOCK_L, HALF_D]tile. The E2M1 exponent factor (1, 2 or 4) and the E8M0 scale are powers of two, which can be produced exactly with selects and integer ops, without the special-function unit.#41251 (merged 2026-10-03) moved this K decode into the shared
_fp4_index_logits_tileand addedfp4_index_logits_paged. It kept bothexp2calls. This PR is based on it, updates the decoding helpers and the tile's scale call, and adds exhaustive decode coverage.Modifications
_e2m1_decode: the power of two (1, 2 or 4) is chosen withtl.whereinstead ofexp2(e - 1)._e8m0_decode:2^(code - 127)is built by writing the code into the FP32 exponent field (code << 23, bitcast). Code 0 decodes to 0, as the flushedexp2(-127)did; the quantizer never stores it (exponents are clamped to 1..254). A comment notes that_ue8m0_to_fp32indequant_k_cache.pydecodes codes 0 and 255 differently (2^-127 and NaN, followingtorch.float8_e8m0fnu); the scorer only sees codes 1..254.dequant_k_cache.pyhas a similar E2M1 decode withexp2. It is not on the scorer path and is left unchanged.test_fp4_scorer_decodes_every_code_exactly: runs both helpers over all 16 E2M1 and 256 E8M0 codes and compares the results with the exact values.Accuracy Tests
Exact comparisons. The new helpers decode every E2M1 code to its exact value and every stored E8M0 code (1-254) to exactly 2^(code-127); codes 0 and 255 are never stored, and decode to 0 and inf. The new test
test_fp4_scorer_decodes_every_code_exactlychecks all 16 E2M1 and 256 E8M0 codes against these exact values. In all benchmark cases below,fp4_index_logits_pagedreturns numerically equal scores (torch.equal) at every visible position before and after this change. On SM120 (RTX PRO 6000), at the head's parent (now b4c0ea5; the head only adds a comment), these tests passed:test_fp4_scorer_decodes_every_code_exactly; the existing logits tests that compare with the torch reference at zero tolerance:test_fp4_logits_visibility_boundaries(36),test_fp4_logits_noncontiguous_queries_and_weights(4),test_fp4_logits_graph_replay_updates_visibility_and_mapping(2), andtest_fp4_paged_logits_replay[6-2-193-True]; andtest_fp4_logits_invisible_tiles_ignore_nonfinite_queries(3), which checks that invisible positions stay -inf.Speed Tests and Profiling
Hardware: one RTX PRO 6000 Blackwell Max-Q Workstation Edition (SM120, 188 SMs), power limit 250 W, used only by this benchmark. Software: torch 2.14.1+cu130, Triton 3.8.0 built from release/3.8.x (c01b6774b) plus a cherry-pick of triton-lang/triton#11940. This change uses no API added by that cherry-pick, but the timings below were taken with the patched compiler; I have not measured stock Triton 3.8.0.
Method:
fp4_index_logits_pagedper call, no candidate mask, on the dyadic inputs of the existing test (_make_logits_case): 4 verify rows of one request, 32 heads. The context is the number of keys the request can see; the 4 rows see the context minus 0 to 3 keys. For each version, 20 calls are captured in one CUDA graph; after one warm-up replay, the graph is replayed 10 times under timing; the table shows time per call, the median of 3 runs (the runs agreed within 0.4 us). main isfp4_indexer.pyat this PR's base and this PR is the file at its head. The timings were taken with the script below before this branch was rebased onto current main (base 4ab720e, head f42ef09);fp4_indexer.pyis byte-identical at those commits and at the current base 9ddbba5 and head d03509e. Outputs of the two versions were equal at every visible position in every run. Reproducer:bench_fp4_index_scorer.py(below).bench_fp4_index_scorer.py
Run from an SGLang checkout installed in the environment:
Checklist
Developed with AI assistance.
CI States
Latest PR Test (Base): ❌ Run #37577689832
Latest PR Test (Extra): ❌ Run #37577689587
Latest PR Test (AMD ROCm 10): ❌ Run #37577689805