Skip to content

[DSv4.1] Speed up the FP4 indexer scorer by decoding keys without exp2 - #42861

Open
ormandj wants to merge 2 commits into
sgl-project:mainfrom
ormandj:dsv4-fp4-indexer-exp2-free
Open

ormandj wants to merge 2 commits into
sgl-project:mainfrom
ormandj:dsv4-fp4-indexer-exp2-free

Conversation

@ormandj

@ormandj ormandj commented Oct 7, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

_fp4_index_logits_tile (used by fp4_index_logits_decode and fp4_index_logits_paged) decodes every K element with tl.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_tile and added fp4_index_logits_paged. It kept both exp2 calls. 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 with tl.where instead of exp2(e - 1).
  • New _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 flushed exp2(-127) did; the quantizer never stores it (exponents are clamped to 1..254). A comment notes that _ue8m0_to_fp32 in dequant_k_cache.py decodes codes 0 and 255 differently (2^-127 and NaN, following torch.float8_e8m0fnu); the scorer only sees codes 1..254.
  • dequant_k_cache.py has a similar E2M1 decode with exp2. 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_exactly checks all 16 E2M1 and 256 E8M0 codes against these exact values. In all benchmark cases below, fp4_index_logits_paged returns 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), and test_fp4_paged_logits_replay[6-2-193-True]; and test_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_paged per 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 is fp4_indexer.py at 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.py is 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).

context main (us) this PR (us) speedup
16,384 13.8 11.2 1.23x
100,000 79.8 59.9 1.33x
400,000 283.3 219.0 1.29x
bench_fp4_index_scorer.py
"""Time fp4_index_logits_paged from two copies of fp4_indexer.py and check that they
return equal scores at every visible position.

  python bench_fp4_index_scorer.py BASE_fp4_indexer.py PR_fp4_indexer.py

Run it in an environment with SGLang installed (it provides the FP4 K cache writer).
Inputs follow test_fp4_indexer.py's _make_logits_case: exact dyadic Q, weights and
keys, 4 verify rows of one request, 32 heads, keys on shuffled 64-token pages. The
4 rows see the context minus 0 to 3 keys. Each version runs 20 calls in one CUDA
graph; after a warm-up replay, 10 replays are timed.
"""

import importlib.util
import sys

import torch

from sglang.kernels.ops.attention.dsv4.fp4_indexer import store_fp4_index_k_cache

HEAD_DIM, PAGE_SIZE, ROWS, HEADS = 128, 64, 4, 32


def load(path, name):
    spec = importlib.util.spec_from_file_location(name, path)
    module = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(module)
    return module


def make_case(width, seed=719):
    torch.manual_seed(seed)
    dev = "cuda"
    q = torch.randint(-8, 9, (ROWS, HEADS, HEAD_DIM), device=dev).to(torch.bfloat16) / 8
    w = torch.randint(-8, 9, (ROWS, HEADS), device=dev).to(torch.bfloat16) / 8
    padded = max(PAGE_SIZE, (width + PAGE_SIZE - 1) // PAGE_SIZE * PAGE_SIZE)
    pages = torch.randperm(padded // PAGE_SIZE, device=dev)
    slots = (pages[:, None] * PAGE_SIZE + torch.arange(PAGE_SIZE, device=dev)).flatten()
    keys = torch.randint(-8, 9, (padded, HEAD_DIM), device=dev).to(torch.bfloat16) / 8
    table = torch.zeros(
        padded // PAGE_SIZE, PAGE_SIZE * 68, device=dev, dtype=torch.uint8
    )
    store_fp4_index_k_cache(keys, table, slots, page_size=PAGE_SIZE, rne=True)
    cap = width + 64
    req = torch.zeros(ROWS, device=dev, dtype=torch.int32)
    req_table = torch.full((1, cap), -1, device=dev, dtype=torch.int32)
    req_table[0, :width] = slots[:width].to(torch.int32)
    lens = width - torch.arange(ROWS, device=dev, dtype=torch.int32)
    return q, w, req, req_table, lens, table, cap


def score(module, q, w, req, req_table, lens, table, cap):
    return module.fp4_index_logits_paged(
        q, w, req, req_table, lens, table, PAGE_SIZE, cap, 1, None
    )


def time_graph(fn, calls=20, replays=10):
    stream = torch.cuda.Stream()
    stream.wait_stream(torch.cuda.current_stream())
    with torch.cuda.stream(stream):
        fn()
        fn()
    torch.cuda.current_stream().wait_stream(stream)
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        for _ in range(calls):
            fn()
    graph.replay()
    torch.cuda.synchronize()
    start = torch.cuda.Event(enable_timing=True)
    end = torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(replays):
        graph.replay()
    end.record()
    torch.cuda.synchronize()
    return start.elapsed_time(end) * 1e3 / (calls * replays)


def main():
    base = load(sys.argv[1], "fp4_indexer_base")
    pr = load(sys.argv[2], "fp4_indexer_pr")
    for width in (16384, 100000, 400000):
        case = make_case(width)
        lens, cap = case[4], case[6]
        visible = torch.arange(cap, device="cuda")[None, :] < lens[:, None]
        equal = torch.equal(score(base, *case)[visible], score(pr, *case)[visible])
        t_base = time_graph(lambda: score(base, *case))
        t_pr = time_graph(lambda: score(pr, *case))
        print(
            f"context={width:6d} equal={equal} main {t_base:7.1f} us "
            f"PR {t_pr:7.1f} us speedup {t_base / t_pr:.2f}x",
            flush=True,
        )


if __name__ == "__main__":
    main()

Run from an SGLang checkout installed in the environment:

git show 9ddbba50dd:python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py > fp4_indexer_main.py
git show d03509e642:python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py > fp4_indexer_pr.py
python bench_fp4_index_scorer.py fp4_indexer_main.py fp4_indexer_pr.py

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

This branch has not been deployed

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant