Skip to content

[DSv4.1] Score prefill consumer index layers on candidate blocks with DeepGEMM - #40352

Merged
BBuf merged 2 commits into
sgl-project:mainfrom
yuan-luo:dsv41-indexer-fused-topk
Sep 22, 2026
Merged

BBuf merged 2 commits into
sgl-project:mainfrom
yuan-luo:dsv41-indexer-fused-topk

Conversation

@yuan-luo

@yuan-luo yuan-luo commented Sep 19, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

DeepSeek-V4.1's hierarchical sparse indexer picks attention positions in two levels on the dense fp4 prefill path (DeepseekV4AttnBackend._low_ratio_index_topk_dense). The first Full-mode index layer (layer 20, the candidate source) keeps the best candidate_topk_blocks = 2048 blocks of candidate_block_size = 8 compressed positions per query row and publishes them. The later index layers (24 / 28 / 32 / 36, the consumers) only run their top-512 over those candidates.

Today both levels are torch glue over the full [tokens, context] fp32 score matrix that fp8_fp4_mqa_logits writes, and every consumer computes that matrix again. For a 16384-token chunk at 262144 context the matrix is 17 GB per index layer:

  • the source masks the unreachable tail with an int64 broadcast compare plus masked_fill_, then runs select_candidate_blocks in 16 row chunks (padded copy, block amax, torch.topk, scatter_, repeat_interleave, torch.cat) into a [tokens, context] bool position mask of 4.3 GB;
  • each consumer scores all 262K columns (16 ms), builds ~mask (another 4.3 GB), masked_fill_s its 17 GB of scores to -inf, runs the ragged top-k over the full row, and drops the -inf picks with mask_topk_scores. It uses 16K of the 262K columns it computed.

On 4x B200 (TP4 / EP4, one 262144-token prompt) the indexer is about 240 ms of the last 544 ms chunk, and it is the only part of prefill that grows with context (+16 ms per chunk for every 16K tokens; the sparse attention itself is flat at 57 ms per chunk).

The decode path does not have this problem: DeepGemmCandidateIndexer publishes a block table from the source layer's scores and the consumers score only their blocks with DeepGEMM's paged sparse logits kernel.

This PR reduces 256K input's TTFT with -22%. (6.75 -> 5.23)
For the backends other than SM100 will be tracked in #40574.

Related issue: #42170.

Design

shot-mu9wxkse

The source layer scores its whole context and keeps the 2048 best-scoring blocks of 8 positions per query row (1 in 16 at full context). Before, every consumer scored the whole context again and masked 15 of every 16 columns away; now it scores only the kept blocks, straight from its paged index-K pool, and picks its top-512 among them. Per chunk that turns four 17 GB dense matrices plus their masks into four 0.5 GB sparse rows.

No kernel changes. Everything the two levels need already exists for the decode path; the PR puts the prefill side behind one interface and wires the DeepGEMM implementation to prefill rows.

The interface. CandidateIndexer (dsv4/candidate_indexer.py) has three prefill methods. publish_prefill(inputs) returns what the source layer's consumers will select from; select_prefill(published, inputs, out_positions) is a consumer's top-k over that, written as flattened-K columns like the dense top-k writes them; prefill_tail(published, tail_lens) restricts what was published to the last rows of each request, for the late layers that run on the tail only. PrefillIndexerInputs carries a chunk's operands (fp4 query and head weights, compressed lengths, request starts, rows and lengths per request, the index-K pool view, the KV page table) plus a memoized dense_scores(), so an implementation that scores sparsely never computes the [tokens, context] matrix and one that needs it computes it once. Two implementations: DeepGemmCandidateIndexer (SM100), below, and MaskCandidateIndexer, the position masks that used to be inline in the backend, kept for Hopper and for the CP layout, whose local rows are not the page table's rows. make_candidate_indexer picks. The backend's dense prefill path only calls the three methods: a consumer calls select_prefill, the source computes its dense scores, calls publish_prefill and runs its own top-512; it never looks at what was published.

Level one, publish_prefill on DeepGEMM (source layer, once per chunk). From the dense scores the source needs anyway, amax8_varlen writes one key per block of 8 positions below the row's compress_len, newest block forced to +inf so it is always kept: one read of the matrix, 4 bytes written per block. The plain ragged top-k over the keys (k = 2048) returns each row's block ids; rows with at most 2048 blocks take its trivial path and keep everything. sort_candidate_blocks sorts them ascending in place, INT32_MAX padded, and derives the pool slots, and get_paged_sparse_mqa_logits_metadata builds the DeepGEMM schedule for the rows. The result is a PrefillSparseBlockTable: one row per query token, with the compress_lens / page_table / request_ids it was built from kept attached, because prefill_tail has to rebuild the schedule for the tail rows (it is per row set and cannot be sliced). This replaces the compare, the masked_fill_, the row-chunk loop and the position mask; the scores are read once instead of five times.

Level two, select_prefill on DeepGEMM (each consumer layer). fp8_fp4_paged_sparse_mqa_logits scores the consumer's queries against the source layer's index-K pool through the schedule, reading only the 2048 published blocks of each row, and writes bf16 [tokens, 2048 x 8] in ascending block order. topk_transform_bf16_small takes the top-512 of the first valid_len columns and applies its page transform with the block table as the page table (column i -> blocks[row, i // 8] * 8 + i % 8), which gives request-relative compressed positions, -1 past min(512, valid_len); adding the request start makes them the flattened-K columns of the interface, and the rest of the path (sort ascending, slot gather, raw indices) is unchanged. A consumer never computes, gathers or reads a dense score row, and no per-consumer mask exists. The consumers now score with bf16 weights and get bf16 scores, exactly what the decode path does for the same layers; the selection agrees with the mask implementation on 97-98% of positions, the rest are boundary picks a few bf16 ulps apart.

Tail rows. The tail metadata exists when the source publishes, so _publish_prefill writes the full table on the forward metadata and, through prefill_tail, the tail rows' table on the tail metadata. enter_late_layer_tail has nothing to cut any more; the torch prefill path's inline masks take the same route.

Memory. A published table outlives its chunk (the late layers and the overlap scheduler read it during the next one). In the shared caching-allocator pool its tensors (blocks, schedule, page table, about 700 MB) were carved from the free remainder of the block that held the source layer's 15 GiB dense scores, and that block could then neither be released nor reused for the next chunk's 16 GiB request: a flush / cold / cold / warm sequence at 262K ended in an OOM with 27 GiB reserved but unusable, while the plain version passed by luck. The DeepGEMM indexer allocates its tables from its own torch.cuda.MemPool, whose block sizes depend only on the row count, so it settles after one chunk and never touches the score blocks. (Putting the dense scores themselves in a private pool does not work: a private pool's cached blocks are not released under memory pressure, and the per-chunk sizes grow.)

The protocol, and what the DeepGEMM kernels do with one query row from the source layer's dense scores down to the consumer's positions:

shot-mu9wzqx5

Modifications

python/sglang/srt/layers/attention/dsv4/candidate_indexer.py — CandidateIndexer (publish_prefill / select_prefill / prefill_tail), PrefillIndexerInputs (msgspec.Struct), expand_index_page_table (moved from the backend), cut_request_masks, and MaskCandidateIndexer: its publish_prefill is the block selection that was in the backend's _publish_or_consume_candidates (tail masked to -inf, row-chunked select_candidate_blocks), select_prefill masks the dense scores, runs topk_transform_ragged_v2 and drops the -inf picks with mask_topk_scores, prefill_tail slices the masks. make_candidate_indexer returns the DeepGEMM indexer on SM100 (with the mask implementation for prefill under CP) and the mask indexer on Hopper.

python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py — DeepGemmCandidateIndexer(CandidateIndexer): publish_prefill (candidate_row_lens -> amax8_varlen -> topk_transform_ragged_v2 over the keys -> _prefill_table: sort_candidate_blocks, build_sparse_indexer_schedule, a CUDA event), prefill_tail (the table restricted to the tail rows, schedule rebuilt), select_prefill (sparse_logits + topk_transform_bf16_small, then the request start), the PrefillSparseBlockTable they share, and the MemPool the tables come from. The TODO(dark) for publish / select prefill is retired.

python/sglang/srt/layers/attention/deepseek_v4_backend.py — _low_ratio_index_topk_dense builds the PrefillIndexerInputs (_prefill_indexer_inputs; the dense scores are padded to 8 columns, amax8_varlen wants 32-byte aligned blocks where the top-k only needed 4). A consumer calls select_prefill and skips the index-K gather, the dense logits and the ragged top-k; the source computes its dense scores, calls publish_prefill and runs its own top-k. _publish_prefill writes the full table and, with prefill_tail, the late-layer tail's; _publish_prefill_masks does the same for the torch prefill path's inline masks. Removed: the candidate cut in enter_late_layer_tail, _publish_or_consume_candidates, the empty-batch mask publish (an empty chunk publishes None). The decode paths and the torch prefill path's own publish / consume are not touched.

test/registered/kernels/ops/attention/test_dsv41_prefill_sparse_indexer.py (SM100, one GPU) — publish_prefill returns the same blocks as select_candidate_blocks on the -inf-masked scores, block for block (ragged lengths, an empty row, block counts above and below 2048); the DeepGEMM select_prefill agrees with MaskCandidateIndexer.select_prefill on the same scores on at least 95% of the positions, picks inside its blocks, and every disagreement is within 2^-5 relative of the selection floor; prefill_tail carries the tail rows' blocks and selects what the full table selects for them. All pass on B200.

Performance

DeepSeek-V4.1-Flash on 4x B200, TP4 / EP4, flashinfer_mxfp4 MoE, chunked_prefill_size=16384, one 262144-token prompt plus 1024 output tokens, no speculative decoding. main (567d5925f) on GPUs 4-7 and this PR on GPUs 0-3 of the same box, requests alternating between the two. Warm numbers are after a 256K request on the same server; "cold" is the first request after /flush_cache, which empties the caching allocator.

main this PR
TTFT (262144 tokens, bs=1), warm 6.75 / 6.90 / 7.07 / 7.17 s 5.55 / 5.59 / 5.61 / 5.72 s (−18%)
TTFT, cold after /flush_cache 10.8 / 13.0 s 7.7 / 8.7 / 9.6 s
TTFT, 200K natural-text prompt (sglang sources + a question) 4.83 s 4.17 s
decode at 256K context 3.8 ms / token 3.8 ms / token
prefill wall / GPU-busy (torch profiler, 16 chunks) 6.77 / 6.55 s 5.47 / 5.30 s
last chunk (262K context) 544 ms 376 ms
last chunk, indexer kernels dense logits x8 104.9, ragged top-k x8 32.3, masked_fill 28.6, block amax 27.4, int64 compare 16.3, scatter / repeat 9.6, bitwise_not 8.7, torch.topk 7.0, copies 8.0 ms dense logits x4 40.6, ragged top-k x5 13.5, schedule 13.9, sparse logits x4 5.8, amax8 2.7, small top-k x4 1.2, sort 0.5 ms
four consumer layers, last chunk 4 x (16 logits + 7.7 mask + 5 top-k + 0.3) = 116 ms 4 x (1.45 + 0.3) = 7 ms
chunk growth per 16K tokens of context +16 ms +4.8 ms

Accuracy

GPQA Diamond (198 questions, single sample) through sgl-eval run gpqa, both servers with --reasoning-parser deepseek-v41, thinking on, reasoning_effort=max, temperature 1.0 / top_p 0.95, max_tokens 65536, seed 1; main on GPUs 4-7 and this PR on GPUs 0-3 of the same box, run in parallel:

main (567d5925f) this PR
score 87.88% 91.41%
truncated at max_tokens 3.54% (7) 1.01% (2)
errors 0 0
generated tokens / throughput 2.2M / 3426 tok/s 2.2M / 3560 tok/s

Checklist


CI States

Latest PR Test (Base): 🚫 Run #35706914278
Latest PR Test (Extra): ⚠️ Not run on latest push -- push again to dispatch.
Latest PR Test (AMD ROCm 10): ❌ Run #35706914128

@yuan-luo yuan-luo added run-ci CI: run the baseline test suite on this PR run-ci-extra CI: also run the extra suite (requires run-ci) labels Sep 19, 2026
@yuan-luo
yuan-luo force-pushed the dsv41-indexer-fused-topk branch from 57aadcc to 428c65f Compare September 19, 2026 13:01
@DarkSharpness

Copy link
Copy Markdown
Collaborator

This is not a good solution. I would recommend take an approach like this:

#pragma once
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh>
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include <cstdint>
#include <limits>
namespace sglang {
/// Level-one keys of the two-level indexer: the max of each kBlockTokens-score
/// block, the row's newest block forced to +inf. Contract: BlockAmaxKernel.
struct BlockAmaxConfig {
using DType = float;
static constexpr uint32_t kBlockTokens = 8; // scores per key
static constexpr uint32_t kBlockSize = 512;
static constexpr uint32_t kNumItems = 2; // keys per thread
static constexpr uint32_t kOccupancy = 4;
static constexpr uint32_t kKeysPerCTA = kBlockSize * kNumItems;
// One block is 32 B: a single load on Blackwell, two 16 B loads before it.
static constexpr uint32_t kVecSize = device::kMaxVecBytes / sizeof(DType);
static constexpr uint32_t kVecsPerBlock = kBlockTokens / kVecSize;
static_assert(kVecsPerBlock * kVecSize == kBlockTokens);
using vec_t = device::AlignedVector<DType, kVecSize>;
};
struct BlockAmaxParams {
const BlockAmaxConfig::DType* __restrict__ scores;
BlockAmaxConfig::DType* __restrict__ amax_scores;
const int32_t* __restrict__ seq_len;
int64_t stride_scores; // in elements
int64_t stride_amax_scores; // in elements
uint32_t topk; // rows with <= topk blocks are skipped, 0 = never skip
};
/// grid = (rows, ceil(max_keys / kKeysPerCTA)); a CTA owns kKeysPerCTA consecutive
/// keys of one row, a thread kNumItems keys kBlockSize apart (coalesced loads).
template <bool kUsePDL>
__global__ __launch_bounds__(BlockAmaxConfig::kBlockSize, BlockAmaxConfig::kOccupancy) //
void amax8_varlen_kernel(const __grid_constant__ BlockAmaxParams params) {
using namespace device;
using C = BlockAmaxConfig;
using T = typename C::DType;
using vec_t = typename C::vec_t;
const auto bx = blockIdx.x;
const auto by = blockIdx.y;
const auto tx = threadIdx.x;
PDLWaitPrimary<kUsePDL>(); // seq_len and scores are the previous kernels' outputs
const auto seq_len = static_cast<uint32_t>(params.seq_len[bx]);
const auto num_keys = (seq_len + C::kBlockTokens - 1) / C::kBlockTokens;
const auto first_key = by * C::kKeysPerCTA;
if (num_keys <= params.topk || first_key >= num_keys) {
return PDLTriggerSecondary<kUsePDL>();
}
const auto* __restrict__ in = params.scores + bx * params.stride_scores;
auto* __restrict__ out = params.amax_scores + bx * params.stride_amax_scores;
vec_t vec[C::kNumItems][C::kVecsPerBlock];
#pragma unroll
for (uint32_t i = 0; i < C::kNumItems; ++i) {
const auto idx = first_key + tx + i * C::kBlockSize;
if (idx < num_keys) {
#pragma unroll
for (uint32_t v = 0; v < C::kVecsPerBlock; ++v) {
vec[i][v].load(in, idx * C::kVecsPerBlock + v);
}
}
}
// The dependent grid may start its prologue now; its griddepcontrol.wait still
// covers every store below (it waits for this grid to complete).
PDLTriggerSecondary<kUsePDL>();
#pragma unroll
for (uint32_t i = 0; i < C::kNumItems; ++i) {
const auto idx = first_key + tx + i * C::kBlockSize;
if (idx < num_keys) {
T key = vec[i][0][0];
#pragma unroll
for (uint32_t v = 0; v < C::kVecsPerBlock; ++v) {
#pragma unroll
for (uint32_t j = 0; j < C::kVecSize; ++j) {
key = fmaxf(key, vec[i][v][j]); // a NaN score is ignored, torch.amax would propagate it
}
}
out[idx] = idx + 1 == num_keys ? std::numeric_limits<T>::infinity() : key;
}
}
}
/// Host entry: `amax_scores[b, i] = max(scores[b, 8 i : 8 i + 8])` for
/// `i < ceil(seq_len[b] / 8)`, the last of them +inf; rows with at most `topk`
/// blocks untouched. `scores` rows must stay 32 B aligned (stride % 8 == 0).
/// The grid covers `amax_scores`' width, so the caller sizes it for the longest
/// row: `seq_len[b] <= 8 * amax_scores.shape[1]` for every row (not checked).
template <bool kPDL>
struct BlockAmaxKernel {
static void amax8_varlen(
const tvm::ffi::TensorView scores,
const tvm::ffi::TensorView seq_lens,
const tvm::ffi::TensorView amax_scores,
const uint32_t topk) {
using namespace host;
using C = BlockAmaxConfig;
auto B = SymbolicSize{"batch_size"};
auto L = SymbolicSize{"max_seq_len"};
auto S = SymbolicSize{"stride_scores"};
auto K = SymbolicSize{"max_keys"};
auto O = SymbolicSize{"stride_amax_scores"};
auto device_ = SymbolicDevice{};
device_.set_options<kDLGPU>();
TensorMatcher({B, L}) // scores
.with_strides({S, 1})
.with_dtype<typename C::DType>()
.with_device(device_)
.verify(scores);
TensorMatcher({B}) // seq_lens
.with_dtype<int32_t>()
.with_device(device_)
.verify(seq_lens);
TensorMatcher({B, K}) // amax_scores
.with_strides({O, 1})
.with_dtype<typename C::DType>()
.with_device(device_)
.verify(amax_scores);
RuntimeCheck(S.unwrap() % C::kBlockTokens == 0, "stride_scores must keep every block 32 B aligned");
RuntimeCheck(
reinterpret_cast<uintptr_t>(scores.data_ptr()) % (C::kBlockTokens * sizeof(typename C::DType)) == 0,
"scores must be 32 B aligned");
RuntimeCheck(K.unwrap() > 0, "amax_scores must hold at least one key per row");
const auto max_keys = K.unwrap(); // ceil(longest row / 8), sized by the caller
const auto params = BlockAmaxParams{
.scores = static_cast<const typename C::DType*>(scores.data_ptr()),
.amax_scores = static_cast<typename C::DType*>(amax_scores.data_ptr()),
.seq_len = static_cast<const int32_t*>(seq_lens.data_ptr()),
.stride_scores = S.unwrap(),
.stride_amax_scores = O.unwrap(),
.topk = topk,
};
const auto grid = dim3(
static_cast<uint32_t>(B.unwrap()),
static_cast<uint32_t>(div_ceil(max_keys, static_cast<int64_t>(C::kKeysPerCTA))));
LaunchKernel(grid, C::kBlockSize, device_.unwrap())
.config({.use_pdl = kPDL})
.launch(amax8_varlen_kernel<kPDL>, params);
}
};
} // namespace sglang

The main reason here is:

  1. Avoid too much intrusive modification into topk kernel.
  2. Avoid applying mask (and reading the score) twice. We only need to condense the score once. Also topk over smaller number of scores gives better performance.
  3. Thie is not a good implementation either. We should integrate DG's sparse indexer, which only score against the selected blocks.

@yuan-luo
yuan-luo force-pushed the dsv41-indexer-fused-topk branch 2 times, most recently from 665a462 to fa2b999 Compare September 20, 2026 11:22
@yuan-luo

Copy link
Copy Markdown
Collaborator Author

This is not a good solution. I would recommend take an approach like this:

#pragma once
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh>
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include <cstdint>
#include <limits>
namespace sglang {
/// Level-one keys of the two-level indexer: the max of each kBlockTokens-score
/// block, the row's newest block forced to +inf. Contract: BlockAmaxKernel.
struct BlockAmaxConfig {
using DType = float;
static constexpr uint32_t kBlockTokens = 8; // scores per key
static constexpr uint32_t kBlockSize = 512;
static constexpr uint32_t kNumItems = 2; // keys per thread
static constexpr uint32_t kOccupancy = 4;
static constexpr uint32_t kKeysPerCTA = kBlockSize * kNumItems;
// One block is 32 B: a single load on Blackwell, two 16 B loads before it.
static constexpr uint32_t kVecSize = device::kMaxVecBytes / sizeof(DType);
static constexpr uint32_t kVecsPerBlock = kBlockTokens / kVecSize;
static_assert(kVecsPerBlock * kVecSize == kBlockTokens);
using vec_t = device::AlignedVector<DType, kVecSize>;
};
struct BlockAmaxParams {
const BlockAmaxConfig::DType* __restrict__ scores;
BlockAmaxConfig::DType* __restrict__ amax_scores;
const int32_t* __restrict__ seq_len;
int64_t stride_scores; // in elements
int64_t stride_amax_scores; // in elements
uint32_t topk; // rows with <= topk blocks are skipped, 0 = never skip
};
/// grid = (rows, ceil(max_keys / kKeysPerCTA)); a CTA owns kKeysPerCTA consecutive
/// keys of one row, a thread kNumItems keys kBlockSize apart (coalesced loads).
template <bool kUsePDL>
__global__ __launch_bounds__(BlockAmaxConfig::kBlockSize, BlockAmaxConfig::kOccupancy) //
void amax8_varlen_kernel(const __grid_constant__ BlockAmaxParams params) {
using namespace device;
using C = BlockAmaxConfig;
using T = typename C::DType;
using vec_t = typename C::vec_t;
const auto bx = blockIdx.x;
const auto by = blockIdx.y;
const auto tx = threadIdx.x;
PDLWaitPrimary<kUsePDL>(); // seq_len and scores are the previous kernels' outputs
const auto seq_len = static_cast<uint32_t>(params.seq_len[bx]);
const auto num_keys = (seq_len + C::kBlockTokens - 1) / C::kBlockTokens;
const auto first_key = by * C::kKeysPerCTA;
if (num_keys <= params.topk || first_key >= num_keys) {
return PDLTriggerSecondary<kUsePDL>();
}
const auto* __restrict__ in = params.scores + bx * params.stride_scores;
auto* __restrict__ out = params.amax_scores + bx * params.stride_amax_scores;
vec_t vec[C::kNumItems][C::kVecsPerBlock];
#pragma unroll
for (uint32_t i = 0; i < C::kNumItems; ++i) {
const auto idx = first_key + tx + i * C::kBlockSize;
if (idx < num_keys) {
#pragma unroll
for (uint32_t v = 0; v < C::kVecsPerBlock; ++v) {
vec[i][v].load(in, idx * C::kVecsPerBlock + v);
}
}
}
// The dependent grid may start its prologue now; its griddepcontrol.wait still
// covers every store below (it waits for this grid to complete).
PDLTriggerSecondary<kUsePDL>();
#pragma unroll
for (uint32_t i = 0; i < C::kNumItems; ++i) {
const auto idx = first_key + tx + i * C::kBlockSize;
if (idx < num_keys) {
T key = vec[i][0][0];
#pragma unroll
for (uint32_t v = 0; v < C::kVecsPerBlock; ++v) {
#pragma unroll
for (uint32_t j = 0; j < C::kVecSize; ++j) {
key = fmaxf(key, vec[i][v][j]); // a NaN score is ignored, torch.amax would propagate it
}
}
out[idx] = idx + 1 == num_keys ? std::numeric_limits<T>::infinity() : key;
}
}
}
/// Host entry: `amax_scores[b, i] = max(scores[b, 8 i : 8 i + 8])` for
/// `i < ceil(seq_len[b] / 8)`, the last of them +inf; rows with at most `topk`
/// blocks untouched. `scores` rows must stay 32 B aligned (stride % 8 == 0).
/// The grid covers `amax_scores`' width, so the caller sizes it for the longest
/// row: `seq_len[b] <= 8 * amax_scores.shape[1]` for every row (not checked).
template <bool kPDL>
struct BlockAmaxKernel {
static void amax8_varlen(
const tvm::ffi::TensorView scores,
const tvm::ffi::TensorView seq_lens,
const tvm::ffi::TensorView amax_scores,
const uint32_t topk) {
using namespace host;
using C = BlockAmaxConfig;
auto B = SymbolicSize{"batch_size"};
auto L = SymbolicSize{"max_seq_len"};
auto S = SymbolicSize{"stride_scores"};
auto K = SymbolicSize{"max_keys"};
auto O = SymbolicSize{"stride_amax_scores"};
auto device_ = SymbolicDevice{};
device_.set_options<kDLGPU>();
TensorMatcher({B, L}) // scores
.with_strides({S, 1})
.with_dtype<typename C::DType>()
.with_device(device_)
.verify(scores);
TensorMatcher({B}) // seq_lens
.with_dtype<int32_t>()
.with_device(device_)
.verify(seq_lens);
TensorMatcher({B, K}) // amax_scores
.with_strides({O, 1})
.with_dtype<typename C::DType>()
.with_device(device_)
.verify(amax_scores);
RuntimeCheck(S.unwrap() % C::kBlockTokens == 0, "stride_scores must keep every block 32 B aligned");
RuntimeCheck(
reinterpret_cast<uintptr_t>(scores.data_ptr()) % (C::kBlockTokens * sizeof(typename C::DType)) == 0,
"scores must be 32 B aligned");
RuntimeCheck(K.unwrap() > 0, "amax_scores must hold at least one key per row");
const auto max_keys = K.unwrap(); // ceil(longest row / 8), sized by the caller
const auto params = BlockAmaxParams{
.scores = static_cast<const typename C::DType*>(scores.data_ptr()),
.amax_scores = static_cast<typename C::DType*>(amax_scores.data_ptr()),
.seq_len = static_cast<const int32_t*>(seq_lens.data_ptr()),
.stride_scores = S.unwrap(),
.stride_amax_scores = O.unwrap(),
.topk = topk,
};
const auto grid = dim3(
static_cast<uint32_t>(B.unwrap()),
static_cast<uint32_t>(div_ceil(max_keys, static_cast<int64_t>(C::kKeysPerCTA))));
LaunchKernel(grid, C::kBlockSize, device_.unwrap())
.config({.use_pdl = kPDL})
.launch(amax8_varlen_kernel<kPDL>, params);
}
};
} // namespace sglang

The main reason here is:

  1. Avoid too much intrusive modification into topk kernel.
  2. Avoid applying mask (and reading the score) twice. We only need to condense the score once. Also topk over smaller number of scores gives better performance.
  3. Thie is not a good implementation either. We should integrate DG's sparse indexer, which only score against the selected blocks.

@DarkSharpness Thanks. I agreed on your three points and reworked the PR with new design.

  1. The top-k kernel is untouched.
  2. The consumer index layers (24 / 28 / 32 / 36) now go through DeepGEMM's sparse indexer exactly the way the decode path does.
  3. The position-mask path stays for Hopper and for the CP layout, gated in _prefill_sparse_indexer.

Comment on lines +1775 to +1782
elif isinstance(full_masks, PrefillSparseBlockTable):
rows, start = [], 0
for n, t in zip(full_masks.rows_per_request, tail_lens_cpu):
rows.append(torch.arange(start + n - t, start + n))
start += n
tail_metadata.candidate_metadata = self.candidate_indexer.prefill_rows(
full_masks, torch.cat(rows).to(full_masks.blocks.device), tail_lens_cpu
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we avoid the if here? I guess we should make publish_prefill kind of a generic approach (update the interface for candidate indexer, and implement that for all backend).

@yuan-luo yuan-luo Sep 20, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@DarkSharpness Your suggestion is excellent. I've refactored the PR and now the candidate indexer is a protocol now. I added you as a co-author, as your comments sharpen this PR quite a lot.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

For all the other backends, I'll create an issue and work on new follow up PRs.

@yuan-luo yuan-luo changed the title [DSv4.1] Fuse the dense prefill indexer's candidate masks into the top-k [DSv4.1] Score prefill consumer index layers on candidate blocks with DeepGEMM Sep 20, 2026
@yuan-luo

yuan-luo commented Sep 20, 2026 •

Copy link
Copy Markdown
Collaborator Author

Updated PR description based on the new design. TTFT drops 18% for 256K input.

Latest update: TTFT drops 22% for 256k input.

luoyuan.luo and others added 2 commits September 22, 2026 16:49
…s with DeepGEMM

The dense prefill indexer's consumer layers (24 / 28 / 32 / 36) scored the
whole context again and masked 15 of every 16 columns away. They now go
through DeepGEMM's paged sparse indexer as the decode path does: the
candidate source publishes a block table from its dense scores (block keys
by amax8_varlen in the same tiled pass as its own top-k, ragged top-k over
the keys, sort_candidate_blocks, DeepGEMM's schedule) and a consumer scores
only the published blocks straight from the index-K pool
(fp8_fp4_paged_sparse_mqa_logits + topk_transform_bf16_small). No dense
row is computed or read by a consumer.

The backend talks to one protocol, CandidateIndexer: publish_prefill,
select_prefill and prefill_tail over a PrefillIndexerInputs.
DeepGemmCandidateIndexer implements it with the block table;
DenseCandidateIndexer wraps the tiled dense_prefill_topk (block ids per
request, consumed through masks) for prefill under CP, whose local rows are
not the page table's rows. The source publishes the tail rows' table at
publish time, so enter_late_layer_tail cuts nothing.

4x B200 TP4/EP4, one 262144-token prompt: TTFT 6.9 -> 5.6 s, the last chunk
544 -> 376 ms, the four consumers 116 -> 7 ms per chunk, decode unchanged.
GPQA-diamond 91.4% vs 87.9% for main (single sample, noise).

Co-authored-by: DarkSharpness <76582120+DarkSharpness@users.noreply.github.com>
Register it on the B200 pool it needs (it skipped on the H100 runner),
run it as a CustomTestCase with subtests like its neighbours, slice the
tail inputs with one helper, and name the numeric tolerances.
@yuan-luo
yuan-luo force-pushed the dsv41-indexer-fused-topk branch from 4e41a5a to 04525c0 Compare September 22, 2026 08:49

@DarkSharpness DarkSharpness left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Let's get this merged first. We need many more clean up of inside the v4 backend and need better abstraction.

@BBuf
BBuf merged commit c79510c into sgl-project:main Sep 22, 2026
178 of 202 checks passed
@yuan-luo yuan-luo added the release-highlight Candidate PR for release note highlight label Sep 22, 2026
kevin-mii added a commit to kevin-mii/sglang that referenced this pull request Sep 23, 2026
Main moved kernel tests under test/registered/kernels/ops (sgl-project#39966), renamed
wo_a_bf16.py to wo_a.py (sgl-project#39957) and moved the low-ratio page-table expansion
into dsv4/candidate_indexer.py (sgl-project#40352). Place the remaining AMD suites in the
plural tree and follow the two renames.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
@YJR722 YJR722 mentioned this pull request Sep 23, 2026
4 of 5 tasks
@yuan-luo yuan-luo mentioned this pull request Oct 7, 2026
10 of 41 tasks
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

deepseek jit-kernel release-highlight Candidate PR for release note highlight run-ci CI: run the baseline test suite on this PR run-ci-extra CI: also run the extra suite (requires run-ci)

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants