[SM100] Paged MQA logits: fused coarse score histogram for LiteTopK decode - #94
Open
Heisenberg-Yin wants to merge 2 commits into
Open
Heisenberg-Yin wants to merge 2 commits into
Heisenberg-Yin wants to merge 2 commits into
Conversation
…n Q tile Add an optional `histogram` output to `fp8_fp4_paged_mqa_logits` / `fp8_paged_mqa_logits`. When given a zeroed int32 [B * next_n, 1024] tensor, the SM100 FP8 H32/D128 paged kernel accumulates every live, non-NaN FP32 score of each row into 1024 descending coarse bins while the scores are still in registers, so a downstream top-k can pick its threshold bin without another pass over the logits. - Bins: FP16-RN (exponent + 4 mantissa bits) below |x| = 16, unit-width bins from 16 to 223, bin 0 holds the largest scores. Counts accumulate across calls; the caller owns zero-at-entry / consumer reset. - Verify steps (next_n > 1) keep one shared-memory slot per token and publish with packed 64-bit atomics when the CTA moves to the next request; next_n = 4 drops the Q ring to 2 stages to fit the slots. - Token Q tile: for non-varlen FP8 H32/D128 calls with next_n < 4 the Q block is sized to next_n (UMMA_N = next_n * 32) instead of computing empty tokens. Applies with or without the histogram. Ported onto dev (post public release 26/09): the smem for the bins sits after the statically checked `MQALogitsSharedStorage`, the generic `get_mqa_logits_smem_size` covers the token tile, and the argument is plumbed through pybind, tvm-ffi and the Python wrappers. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Author
|
@DarkSharpness could you take a look? This is the producer side for LiteTopK decode. The paged MQA logits kernel emits the coarse histogram that the SGLang top-k v2 streaming and cluster paths currently build in their first pass over the logits. |
- sgl_deep_gemm: add `histogram=None` to `fp8_paged_mqa_logits` and `fp8_fp4_paged_mqa_logits`, matching the deep_gemm wrappers. - test_paged_mqa_logits: also run next_n = 2 / 3 for FP8 H32/D128 on SM100, covering the per-next_n token Q tile without the histogram. - test_paged_mqa_logits_histogram: the varlen branch rebinds batch_size to the row count; iterate over num_requests and reset batch_size per case so it no longer leaks into the next shape. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Heisenberg-Yin
marked this pull request as ready for review
September 29, 2026 08:33
Heisenberg-Yin
added a commit
to Heisenberg-Yin/sglang
that referenced
this pull request
Sep 30, 2026
- select.cuh + coarse_bins.cuh: exact FP32 top-2048 selector driven by DeepGEMM's coarse score histogram, compiled on first use through load_jit. Covers decode (next_n = 1) and speculative verify (next_n <= 4). A stale histogram, or a crossing bin larger than the candidate capacity, falls back to an exact whole-row radix select instead of raising a device error. Ties go to the lower physical slot. - fused.py: FusedDecodePlan(batch, next_n, device, candidate_capacity=8192) calls deep_gemm.fp8_fp4_paged_mqa_logits(histogram=...) directly. - Drop the CuTe DSL selector, its AOT build and manifest, and the vendored DeepGEMM patch and private build: the histogram argument now comes from sgl-project/DeepGEMM#94. - smoke_fused.py: decode and verify, eager and CUDA graph replay, plus a stale histogram and a candidate overflow. - Add kernels/experimental/.clang-format, identical to kernels/jit's. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
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
In DSA decode, the indexer runs paged MQA logits and then a top-k over every
[row, context_len]score row. For long contexts the top-k's first full pass over the logits exists only to build a coarse histogram and find the threshold bin. The logits kernel already has every score in registers, so this PR lets it emit that histogram as a by-product. A downstream top-k (LiteTopK decode on the SGLang side) can then skip its histogram pass over the logits and go straight to collecting.What this adds
histogram=Noneonfp8_fp4_paged_mqa_logitsand the legacyfp8_paged_mqa_logits, exposed through pybind, tvm-ffi and thedeep_gemm/sgl_deep_gemmPython wrappers. Without it, none of the histogram code runs.Contract:
int32 [B * next_n, 1024], contiguous, 8-byte aligned, zero at entry. Counts accumulate across calls, so the consumer resets them. Supported configs: SM100 FP8 (non-MX), H = 32, D = 128, FP32 weights and logits,clean_logits=False. Varlen needsnext_n == 1and regular calls neednext_n <= 4.Bins are in descending order, so bin 0 holds the largest scores. Only live (
col < context_lenof that row), non-NaN scores are counted.|x| < 16: FP16-RN exponent plus 4 mantissa bits. This is the usual 10-bit FP16 coarse key, reversed.16 <= |x| < 223: unit-width bins, lower-inclusive for negatives.|x| >= 223, including ±inf: one saturating bin at each end.ref_coarse_histogramintests/test_attention.pyis the reference.Verify steps (
next_n > 1) keep one 4 KiB shared-memory slot per token and flush it with packed 64-bit global atomics when the CTA moves to the next request. Withnext_n = 4the Q ring drops from 3 to 2 stages so the slots fit.Token Q tile (applies with or without the histogram): regular FP8 H32/D128 calls with FP32 weights and
next_n < 4useBLOCK_Q = next_n(UMMA_N = 32 * next_n) instead of padding the Q block to 4 tokens.Without the histogram,
next_n = 4,next_n = 6and varlen compile to the same SASS asdev;next_n = 1..3differ only by the token Q tile.Tests: new
test_paged_mqa_logits_histogram(histogram equals the PyTorch reference in every bin, logits bitwise identical with and without the histogram, counts accumulate across calls; varlennext_n = 1, regularnext_n = 1..4, fine / unit / mixed score distributions).test_paged_mqa_logitsalso runsnext_n = 2, 3for FP8 / H32 / D128, so the token Q tile is checked against the reference for everynext_nit applies to.Producer cost
This PR (with the histogram) vs the
devbase (no histogram), paired in one process.Method: B200, real GLM-5.2 indexer capture (layers 0 and 38), token j of a request sees
L + jkeys. Cold L2 (1 GiB memset before every call), CUPTI kernel time, 4 independent processes x 150 randomized paired blocks per cell. Measured: every cell whose KV fits one rank's KV pool (B x (L + next_n) <= 1,550,400tokens), plus larger batches beyond that pool, marked*(B2 ~1M; B4 512K / ~1M; B8 256K / 512K; B16 128K; B32 64K / 128K; B64 32K / 64K; B128 16K / 32K / 64K), measured in separate runs with the same method. The remaining cells are n/a (for exampleB = 128, L = 1Mwould need ~134M tokens, ~6.4 TB of GLM-5.2 KV). SM clock 1965 MHz in every cell, except short SW power-cap events in the largest cells, mostly B128 64K (lowest sampled clock 1710 MHz).A/A (dev vs dev): |extra| median 0.016 us, max 0.212 us over 294 cells.
Without the histogram the token Q tile makes
next_n = 1, 2slightly faster thandevin most cells; in the largest cells (about 100 us and up) it is at par or slightly slower instead, at most +0.80 us (0.4 % of a 183 us kernel, B = 128, 64K).next_n = 3, 4are unchanged.Per-cell results, all measured cells. Each cell:
dev us / +extra with histogram / +extra without histogram(us; extra = this PR - dev, mean of the 4 per-process paired medians).layer 0, next_n = 1
layer 0, next_n = 2
layer 0, next_n = 3
layer 0, next_n = 4
layer 38, next_n = 3
layer 38, next_n = 4