Repository navigation
Conversation
verify_splitkv_fwd was gated to gfx95 (AMD MI35x). On CUDA the Triton backend's target-verify therefore falls back to extend_attention_fwd, whose grid is (bs, head_num, cdiv(draft_tokens, BLOCK_M)) - i.e. one program per (seq, head) that walks the entire prefix serially with BLOCK_N. At topk=1 with 4 draft tokens that is 24 programs for a 24-head model, roughly 1/8 the parallelism of the split-KV decode kernel, so verify cost grows linearly with prefix length and speculative decoding becomes slower than no-spec beyond ~32k context. The kernel is already NV-safe (HIP-only launch kwargs are gated on is_hip()), and its numerics test is registered on the CUDA CI lane. Drop the gfx95 gate; SGLANG_ENABLE_SPLITKV_VERIFY=0 remains as the opt-out. Measured on Qwen3.5-27B (hybrid GDN, 16 full-attn layers, GQA 24/4, head_dim 256), NEXTN steps=3 topk=1 draft=4, sm_80, 1 GPU, batch 1, decode tok/s at 128 output tokens: prompt no-spec spec (before) spec (after) 512 69.6 209.4 217.5 8192 65.0 94.7 162.8 32768 53.4 38.9 131.3 131072 31.1 10.1 57.7 Accept length was 3.0-3.5 in both spec runs; only the verify kernel changed. test_verify_splitkv passes on CUDA (9 tests).
rluisr
requested review from
Fridge003,
HaiShaw,
Qiaolin-Yu,
hebiao064,
ispobock and
merrymercy
as code owners
September 13, 2026 20:05
5 tasks done
Author
|
@zijiecode @yichiche Can you help review this PR? |
This branch has not been deployed
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
verify_splitkv_fwd(the flash-decode style EAGLE target-verify kernel for the Triton backend) is currently gated behindis_gfx95_supported(), so it only runs on AMD MI35x. On CUDA the Triton backend's target-verify falls back toextend_attention_fwd, whose grid is(bs, head_num, cdiv(max_len_extend, BLOCK_M)). Attopk=1with 4 draft tokens that is one program per (seq, head) that walks the whole prefix serially inBLOCK_Nsteps — for a 24-head model, 24 programs vs. thenum_kv_splits=8× 24 = 192 the regular decode kernel uses for the same prefix.The consequence is that verify cost grows linearly with prefix length and speculative decoding becomes slower than no-spec beyond ~32k context, even with a healthy accept length. On a 27B hybrid-GDN model at 128k we measured decode dropping from 31 tok/s (no spec) to 10 tok/s (NEXTN), which is easy to misattribute to the SSM state handling.
The kernel itself is already NV-safe (HIP-only launch kwargs are gated on
is_hip()insideverify_splitkv.py, and its own docstring says so), andtest_verify_splitkv.pyis registered on the CUDA CI lane. This PR simply drops the gfx95 gate.SGLANG_ENABLE_SPLITKV_VERIFY=0remains as the opt-out.Why the gate is gfx95-only today
The gate is historical, not technical. In #27382 the kernel originally shipped with no platform gate; the NVIDIA CI lane then failed (
test_basic_sanity_eagle3.py,KeyError: Keyword argument waves_per_eu was specified but unrecognised) because the AMD-only Triton launch hints were passed unconditionally. The author fixed that in f7622c0 by making the kwargs HIP-conditional — and at the same time narrowed the dispatch tois_hip(), later tois_gfx95_supported()at review request ("scoped to where it's validated"). So the kernel has been NV-safe since before merge, and its numerics test has been running on the CUDA CI lane the whole time; only the dispatch was left closed. #35521 makes the same observation for gfx942 ("the kernel is arch-neutral Triton"). This PR supplies the CUDA validation that was missing at the time.Modifications
python/sglang/srt/layers/attention/triton_backend.py: removeis_gfx95_supported()from theuse_verify_splitkvcondition; keep thetopk == 1and env-var gates. Comment updated to explain why.python/sglang/srt/environ.py: doc comment forSGLANG_ENABLE_SPLITKV_VERIFYnow says ROCm and CUDA._should_use_verify_shared_kv(the grouped-head MLA / single-KV-head kernel) is left gfx95-only; this PR only touches the per-head split-KV path thatcan_handle()already validates.Accuracy Tests
test/registered/attention/test_verify_splitkv.pyon CUDA (sm_80, Triton from the 0.5.19 wheel): 9 tests, OK (parity vsextend_attention_fwdacross head dims / GQA ratios / prefix lengths / KV scales, plus thecan_handle()rejection cases).End-to-end: served a Qwen3.5-27B-class checkpoint with
--speculative-algorithm NEXTN --speculative-num-steps 3 --speculative-eagle-topk 1 --speculative-num-draft-tokens 4 --attention-backend triton; accept length stayed at 3.0–3.5 before and after (only the verify kernel changed), tool-call parsing and reasoning parsing unchanged.Speed Tests and Profiling
Qwen3.5-27B (hybrid GDN: 48 linear-attn + 16 full-attn layers, GQA 24/4, head_dim 256), compressed-tensors W4A16, 1× sm_80 GPU (64 GB),
--mem-fraction-static 0.80,--attention-backend triton --mamba-backend triton, batch 1, cold prefix cache, 128 output tokens. Decode tok/s:Before this change spec decoding is a net loss past 32k; after it is +86% over no-spec at 128k. Prefill / TTFT is unaffected by this PR (verify is a decode-side kernel).
Concurrency 4 at 32k (sum of per-stream decode tok/s) is flat vs. no-spec (156 vs 172), as expected once the GPU is saturated; the win is at low concurrency / long context.
Checklist
test_verify_splitkv.pyalready covers the kernel on the CUDA lane; no new kernel code.)CI States
Latest PR Test (Base): ❌ Run #35825241950
Latest PR Test (Extra): ❌ Run #35825241812
Latest PR Test (AMD ROCm 10): ❌ Run #35825241944