Repository navigation
[DSv4.1] Score prefill consumer index layers on candidate blocks with DeepGEMM - #40352
Conversation
57aadcc to
428c65f
Compare
|
This is not a good solution. I would recommend take an approach like this: sglang/python/sglang/kernels/jit/csrc/deepseek_v4/block_amax.cuh Lines 1 to 155 in 76f9213 The main reason here is:
|
665a462 to
fa2b999
Compare
@DarkSharpness Thanks. I agreed on your three points and reworked the PR with new design.
|
| 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 | ||
| ) |
There was a problem hiding this comment.
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).
There was a problem hiding this comment.
@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.
There was a problem hiding this comment.
For all the other backends, I'll create an issue and work on new follow up PRs.
|
Updated PR description based on the new design. TTFT drops 18% for 256K input. Latest update: TTFT drops 22% for 256k input. |
14eb8d5 to
cecc114
Compare
…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.
4e41a5a to
04525c0
Compare
DarkSharpness
left a comment
There was a problem hiding this comment.
Let's get this merged first. We need many more clean up of inside the v4 backend and need better abstraction.
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>
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 bestcandidate_topk_blocks = 2048blocks ofcandidate_block_size = 8compressed 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 thatfp8_fp4_mqa_logitswrites, and every consumer computes that matrix again. For a 16384-token chunk at 262144 context the matrix is 17 GB per index layer:masked_fill_, then runsselect_candidate_blocksin 16 row chunks (padded copy, blockamax,torch.topk,scatter_,repeat_interleave,torch.cat) into a[tokens, context]bool position mask of 4.3 GB;~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-infpicks withmask_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:
DeepGemmCandidateIndexerpublishes 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
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.PrefillIndexerInputscarries 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 memoizeddense_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, andMaskCandidateIndexer, 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_indexerpicks. The backend's dense prefill path only calls the three methods: a consumer callsselect_prefill, the source computes its dense scores, callspublish_prefilland runs its own top-512; it never looks at what was published.Level one,
publish_prefillon DeepGEMM (source layer, once per chunk). From the dense scores the source needs anyway,amax8_varlenwrites one key per block of 8 positions below the row'scompress_len, newest block forced to+infso 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_blockssorts them ascending in place,INT32_MAXpadded, and derives the pool slots, andget_paged_sparse_mqa_logits_metadatabuilds the DeepGEMM schedule for the rows. The result is aPrefillSparseBlockTable: one row per query token, with thecompress_lens/page_table/request_idsit was built from kept attached, becauseprefill_tailhas to rebuild the schedule for the tail rows (it is per row set and cannot be sliced). This replaces the compare, themasked_fill_, the row-chunk loop and the position mask; the scores are read once instead of five times.Level two,
select_prefillon DeepGEMM (each consumer layer).fp8_fp4_paged_sparse_mqa_logitsscores 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_smalltakes the top-512 of the firstvalid_lencolumns 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,-1pastmin(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_prefillwrites the full table on the forward metadata and, throughprefill_tail, the tail rows' table on the tail metadata.enter_late_layer_tailhas 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:
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, andMaskCandidateIndexer: itspublish_prefillis the block selection that was in the backend's_publish_or_consume_candidates(tail masked to-inf, row-chunkedselect_candidate_blocks),select_prefillmasks the dense scores, runstopk_transform_ragged_v2and drops the-infpicks withmask_topk_scores,prefill_tailslices the masks.make_candidate_indexerreturns 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_v2over 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), thePrefillSparseBlockTablethey share, and theMemPoolthe tables come from. TheTODO(dark)for publish / select prefill is retired.python/sglang/srt/layers/attention/deepseek_v4_backend.py—_low_ratio_index_topk_densebuilds thePrefillIndexerInputs(_prefill_indexer_inputs; the dense scores are padded to 8 columns,amax8_varlenwants 32-byte aligned blocks where the top-k only needed 4). A consumer callsselect_prefilland skips the index-K gather, the dense logits and the ragged top-k; the source computes its dense scores, callspublish_prefilland runs its own top-k._publish_prefillwrites the full table and, withprefill_tail, the late-layer tail's;_publish_prefill_masksdoes the same for the torch prefill path's inline masks. Removed: the candidate cut inenter_late_layer_tail,_publish_or_consume_candidates, the empty-batch mask publish (an empty chunk publishesNone). 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_prefillreturns the same blocks asselect_candidate_blockson the-inf-masked scores, block for block (ragged lengths, an empty row, block counts above and below 2048); the DeepGEMMselect_prefillagrees withMaskCandidateIndexer.select_prefillon the same scores on at least 95% of the positions, picks inside its blocks, and every disagreement is within2^-5relative of the selection floor;prefill_tailcarries 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_mxfp4MoE,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./flush_cacheAccuracy
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:567d5925f)max_tokensChecklist
CI States
Latest PR Test (Base): 🚫 Run #35706914278⚠️ Not run on latest push -- push again to dispatch.
Latest PR Test (Extra):
Latest PR Test (AMD ROCm 10): ❌ Run #35706914128