Skip to content

[SM100] Paged MQA logits: fused coarse score histogram for LiteTopK decode - #94

Open
Heisenberg-Yin wants to merge 2 commits into
sgl-project:devfrom
Heisenberg-Yin:litetopk-decode
Open

Heisenberg-Yin wants to merge 2 commits into
sgl-project:devfrom
Heisenberg-Yin:litetopk-decode

Conversation

@Heisenberg-Yin

@Heisenberg-Yin Heisenberg-Yin commented Sep 29, 2026 •

Copy link
Copy Markdown

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=None on fp8_fp4_paged_mqa_logits and the legacy fp8_paged_mqa_logits, exposed through pybind, tvm-ffi and the deep_gemm / sgl_deep_gemm Python 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 needs next_n == 1 and regular calls need next_n <= 4.

  • Bins are in descending order, so bin 0 holds the largest scores. Only live (col < context_len of 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_histogram in tests/test_attention.py is 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. With next_n = 4 the 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 < 4 use BLOCK_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 = 6 and varlen compile to the same SASS as dev; next_n = 1..3 differ 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; varlen next_n = 1, regular next_n = 1..4, fine / unit / mixed score distributions). test_paged_mqa_logits also runs next_n = 2, 3 for FP8 / H32 / D128, so the token Q tile is checked against the reference for every next_n it applies to.

Producer cost

This PR (with the histogram) vs the dev base (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 + j keys. 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,400 tokens), 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 example B = 128, L = 1M would 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).

layer next_n with histogram: median / max extra (us) without histogram: median extra (us)
0 1 +0.05 / +0.98 -0.19
0 2 +0.18 / +0.89 -0.12
0 3 +0.40 / +1.42 +0.00
0 4 +0.79 / +1.41 +0.03
38 3 +0.56 / +1.96 +0.00
38 4 +1.09 / +2.66 +0.04

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, 2 slightly faster than dev in 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, 4 are 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

B \ L 8K 16K 32K 64K 128K 256K 512K ~1M
1 5.26 / +0.00 / -0.17 5.36 / +0.09 / -0.06 5.88 / +0.10 / -0.11 6.97 / +0.01 / -0.20 8.79 / +0.04 / -0.23 12.51 / -0.02 / -0.18 19.67 / -0.05 / -0.26 32.83 / -0.04 / -0.18
2 5.38 / +0.08 / -0.13 5.86 / +0.08 / -0.09 7.02 / +0.06 / -0.14 8.78 / +0.04 / -0.19 12.54 / +0.01 / -0.14 19.52 / -0.02 / -0.20 32.94 / +0.02 / -0.20 58.03 / -0.02 / -0.26 *
4 5.90 / +0.06 / -0.11 6.96 / +0.03 / -0.20 8.77 / +0.03 / -0.24 12.56 / +0.01 / -0.20 19.50 / +0.11 / -0.15 32.76 / +0.09 / -0.16 57.89 / -0.03 / -0.21 * 100.24 / +0.10 / +0.03 *
8 6.94 / +0.04 / -0.18 8.71 / +0.07 / -0.24 12.68 / +0.03 / -0.21 19.50 / +0.01 / -0.22 32.61 / -0.07 / -0.16 57.74 / +0.04 / -0.22 * 100.88 / +0.05 / -0.19 * n/a
16 8.96 / +0.05 / -0.23 12.82 / -0.00 / -0.21 19.20 / +0.07 / -0.23 32.58 / +0.07 / -0.24 58.20 / +0.10 / -0.09 * n/a n/a n/a
32 12.94 / +0.04 / -0.26 19.82 / +0.07 / -0.25 32.46 / +0.02 / -0.19 58.12 / +0.06 / -0.10 * 100.52 / +0.42 / +0.10 * n/a n/a n/a
64 19.55 / +0.07 / -0.18 32.88 / +0.04 / -0.15 58.01 / +0.14 / -0.17 * 100.64 / +0.46 / +0.09 * n/a n/a n/a n/a
128 32.73 / +0.13 / -0.15 58.03 / +0.14 / -0.20 * 100.99 / +0.18 / -0.20 * 182.64 / +0.98 / +0.80 * n/a n/a n/a n/a

layer 0, next_n = 2

B \ L 8K 16K 32K 64K 128K 256K 512K ~1M
1 5.14 / +0.13 / -0.04 5.33 / +0.16 / -0.01 5.92 / +0.12 / -0.06 7.02 / +0.15 / -0.07 8.91 / +0.18 / -0.10 12.74 / +0.06 / -0.15 19.78 / +0.14 / -0.17 32.87 / +0.10 / -0.12
2 5.39 / +0.14 / -0.06 5.88 / +0.12 / -0.03 7.11 / +0.20 / -0.09 8.97 / +0.55 / -0.12 12.68 / +0.16 / -0.18 19.76 / +0.12 / -0.16 33.05 / +0.11 / -0.10 58.34 / +0.20 / -0.02 *
4 5.96 / +0.12 / -0.08 7.00 / +0.11 / -0.11 9.02 / +0.23 / -0.09 12.76 / +0.08 / -0.18 19.40 / +0.16 / -0.14 32.68 / +0.16 / -0.09 58.38 / +0.12 / -0.10 * 102.03 / +0.50 / +0.23 *
8 7.10 / +0.19 / -0.09 9.01 / +0.18 / -0.14 12.70 / +0.09 / -0.15 19.56 / +0.21 / -0.10 33.20 / +0.16 / -0.12 58.36 / +0.03 / -0.18 * 102.62 / +0.42 / +0.09 * n/a
16 9.15 / +0.28 / -0.16 13.18 / +0.14 / -0.14 19.83 / +0.29 / -0.18 33.08 / +0.42 / -0.15 58.35 / +0.33 / -0.07 * n/a n/a n/a
32 13.02 / +0.13 / -0.17 19.66 / +0.28 / -0.13 33.01 / +0.15 / -0.19 58.68 / +0.18 / -0.14 * 102.03 / +0.67 / +0.28 * n/a n/a n/a
64 19.85 / +0.29 / -0.14 33.40 / +0.26 / -0.12 58.65 / +0.19 / -0.20 * 102.66 / +0.61 / +0.26 * n/a n/a n/a n/a
128 33.39 / +0.24 / -0.17 59.06 / +0.24 / -0.19 * 103.42 / +0.66 / +0.17 * 185.58 / +0.89 / +0.63 * n/a n/a n/a n/a

layer 0, next_n = 3

B \ L 8K 16K 32K 64K 128K 256K 512K ~1M
1 5.19 / +0.23 / +0.02 5.46 / +0.23 / +0.05 6.00 / +0.26 / +0.05 7.23 / +0.28 / +0.00 9.21 / +0.34 / +0.02 12.88 / +0.26 / -0.02 19.82 / +0.36 / -0.02 33.05 / +0.32 / +0.02
2 5.40 / +0.21 / +0.02 5.95 / +0.24 / +0.00 7.22 / +0.25 / -0.04 9.27 / +0.33 / +0.00 13.04 / +0.43 / -0.04 19.83 / +0.39 / -0.02 33.05 / +0.51 / -0.00 58.30 / +0.42 / +0.02 *
4 6.02 / +0.20 / -0.01 7.13 / +0.32 / +0.02 9.16 / +0.35 / +0.01 12.89 / +0.32 / -0.03 19.76 / +0.35 / -0.02 32.97 / +0.60 / -0.02 58.88 / +0.59 / -0.00 * 104.13 / +0.88 / +0.33 *
8 7.14 / +0.57 / +0.00 9.22 / +0.55 / -0.03 12.91 / +0.30 / -0.02 19.93 / +0.60 / +0.00 33.05 / +0.40 / -0.05 58.76 / +0.37 / +0.03 * 104.05 / +0.88 / +0.14 * n/a
16 9.30 / +0.33 / -0.00 13.13 / +0.48 / -0.03 20.32 / +0.72 / -0.04 33.46 / +0.40 / +0.01 58.88 / +0.72 / +0.07 * n/a n/a n/a
32 13.36 / +0.41 / -0.05 20.02 / +0.86 / -0.06 33.52 / +0.34 / -0.05 58.87 / +0.40 / -0.08 * 104.23 / +0.88 / +0.22 * n/a n/a n/a
64 20.18 / +0.74 / -0.06 33.70 / +0.42 / -0.08 59.10 / +0.48 / -0.04 * 104.85 / +1.12 / +0.28 * n/a n/a n/a n/a
128 33.58 / +0.71 / +0.00 59.05 / +0.58 / +0.03 * 104.99 / +1.22 / +0.12 * 190.71 / +1.42 / +0.12 * n/a n/a n/a n/a

layer 0, next_n = 4

B \ L 8K 16K 32K 64K 128K 256K 512K ~1M
1 5.36 / +0.31 / +0.02 5.59 / +0.32 / +0.04 6.17 / +0.30 / +0.04 7.38 / +0.39 / +0.03 9.48 / +0.46 / +0.05 13.25 / +0.62 / +0.04 20.19 / +0.57 / +0.04 33.32 / +0.62 / +0.07
2 5.58 / +0.30 / +0.05 6.07 / +0.32 / +0.04 7.46 / +0.38 / +0.04 9.62 / +0.98 / +0.02 13.46 / +1.13 / +0.04 20.24 / +0.52 / +0.01 33.30 / +0.64 / +0.00 58.93 / +0.66 / +0.02 *
4 6.16 / +0.30 / +0.03 7.42 / +0.37 / +0.02 9.62 / +0.55 / +0.02 13.37 / +0.58 / +0.05 20.50 / +0.90 / +0.05 33.76 / +0.64 / +0.05 59.52 / +0.76 / -0.03 * 106.38 / +1.10 / -0.14 *
8 7.34 / +0.79 / +0.06 9.61 / +0.71 / +0.01 13.48 / +0.94 / +0.06 20.49 / +0.63 / +0.03 33.98 / +1.27 / +0.04 59.17 / +0.79 / -0.01 * 106.76 / +1.14 / +0.02 * n/a
16 9.84 / +0.81 / +0.07 13.38 / +0.88 / +0.03 20.58 / +1.08 / +0.06 33.78 / +0.78 / -0.04 59.32 / +0.99 / -0.00 * n/a n/a n/a
32 13.77 / +0.94 / -0.01 20.62 / +0.97 / -0.00 33.92 / +0.76 / -0.02 59.41 / +0.96 / +0.06 * 106.94 / +1.27 / +0.02 * n/a n/a n/a
64 20.75 / +1.11 / +0.03 34.14 / +1.06 / +0.02 59.69 / +0.84 / +0.02 * 107.77 / +1.18 / +0.09 * n/a n/a n/a n/a
128 33.96 / +1.07 / +0.07 59.70 / +1.05 / +0.14 * 108.35 / +1.38 / -0.05 * 196.58 / +1.41 / +0.27 * n/a n/a n/a n/a

layer 38, next_n = 3

B \ L 8K 16K 32K 64K 128K 256K 512K ~1M
1 5.19 / +0.27 / -0.00 5.40 / +0.28 / +0.06 5.99 / +0.25 / +0.05 7.23 / +0.34 / +0.02 9.26 / +0.44 / -0.00 12.82 / +0.48 / +0.02 19.72 / +0.49 / +0.04 33.05 / +0.41 / +0.04
2 5.42 / +0.27 / +0.01 5.98 / +0.26 / +0.02 7.22 / +0.35 / +0.01 9.26 / +0.44 / -0.01 12.91 / +0.56 / -0.00 19.86 / +0.51 / +0.02 33.00 / +0.61 / -0.02 58.27 / +0.53 / +0.02 *
4 6.03 / +0.25 / +0.00 7.15 / +0.34 / -0.04 9.28 / +0.41 / -0.04 12.92 / +0.43 / -0.06 19.85 / +0.55 / +0.08 32.93 / +0.73 / +0.07 58.76 / +0.65 / +0.07 * 104.02 / +1.37 / +0.24 *
8 7.08 / +0.50 / -0.03 9.17 / +0.56 / -0.03 12.92 / +0.60 / +0.01 19.89 / +0.72 / -0.00 33.10 / +0.58 / -0.01 58.82 / +0.56 / -0.05 * 104.44 / +1.50 / +0.19 * n/a
16 9.46 / +0.63 / +0.00 13.04 / +0.62 / -0.04 20.14 / +1.08 / +0.03 33.38 / +0.46 / -0.04 58.73 / +0.81 / +0.07 * n/a n/a n/a
32 13.24 / +0.56 / -0.02 20.06 / +1.00 / +0.00 33.45 / +0.44 / -0.05 58.79 / +0.63 / -0.03 * 104.09 / +1.13 / +0.13 * n/a n/a n/a
64 20.30 / +0.88 / +0.01 33.67 / +0.67 / -0.07 59.16 / +0.75 / -0.01 * 104.83 / +1.51 / +0.29 * n/a n/a n/a n/a
128 33.63 / +0.83 / -0.06 59.10 / +0.70 / -0.01 * 104.60 / +1.31 / +0.09 * 190.02 / +1.96 / +0.40 * n/a n/a n/a n/a

layer 38, next_n = 4

B \ L 8K 16K 32K 64K 128K 256K 512K ~1M
1 5.34 / +0.36 / +0.04 5.55 / +0.34 / +0.03 6.16 / +0.32 / -0.00 7.43 / +0.42 / +0.02 9.54 / +0.60 / +0.05 13.25 / +0.94 / +0.02 20.33 / +0.76 / +0.02 33.31 / +0.89 / +0.02
2 5.57 / +0.31 / +0.02 6.11 / +0.34 / +0.02 7.47 / +0.42 / +0.03 9.64 / +1.11 / +0.02 13.42 / +1.48 / +0.02 20.23 / +0.76 / -0.00 33.31 / +1.05 / +0.06 58.90 / +1.09 / +0.01 *
4 6.20 / +0.34 / +0.04 7.42 / +0.46 / +0.02 9.65 / +0.72 / +0.02 13.38 / +0.96 / +0.00 20.60 / +1.22 / +0.06 33.74 / +0.90 / +0.10 59.54 / +1.24 / +0.07 * 106.89 / +1.66 / +0.05 *
8 7.39 / +0.83 / +0.04 9.62 / +0.88 / +0.04 13.41 / +1.34 / +0.05 20.34 / +0.98 / +0.04 33.89 / +1.43 / +0.05 59.09 / +1.22 / +0.05 * 106.68 / +1.95 / +0.05 * n/a
16 9.74 / +0.93 / +0.06 13.36 / +1.04 / +0.05 20.67 / +1.32 / +0.04 33.70 / +1.09 / +0.11 59.20 / +1.46 / +0.10 * n/a n/a n/a
32 13.61 / +1.22 / +0.02 20.44 / +1.16 / +0.05 33.97 / +1.14 / -0.04 59.36 / +1.36 / +0.08 * 107.14 / +1.80 / +0.03 * n/a n/a n/a
64 20.63 / +1.44 / +0.07 34.20 / +1.38 / -0.02 59.69 / +1.28 / +0.02 * 107.93 / +2.07 / +0.14 * n/a n/a n/a n/a
128 34.04 / +1.41 / +0.04 59.79 / +1.23 / +0.08 * 108.04 / +2.16 / +0.08 * 196.57 / +2.66 / +0.05 * n/a n/a n/a n/a

…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>
@Heisenberg-Yin

Heisenberg-Yin commented Sep 29, 2026 •

Copy link
Copy Markdown
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.

@Heisenberg-Yin Heisenberg-Yin changed the title [SM100] Paged MQA logits: fused coarse score histogram for LiteTopK decode + per-next_n Q tile [SM100] Paged MQA logits: fused coarse score histogram for LiteTopK decode Sep 29, 2026
- 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
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant