Skip to content

[Spec] Enable split-KV EAGLE verify on CUDA for the Triton backend - #39316

Open
rluisr wants to merge 6 commits into
sgl-project:mainfrom
rluisr:triton-splitkv-verify-cuda
Open

rluisr wants to merge 6 commits into
sgl-project:mainfrom
rluisr:triton-splitkv-verify-cuda

Conversation

@rluisr

@rluisr rluisr commented Sep 13, 2026 •

Copy link
Copy Markdown

Motivation

verify_splitkv_fwd (the flash-decode style EAGLE target-verify kernel for the Triton backend) is currently gated behind is_gfx95_supported(), so it only runs on AMD MI35x. On CUDA the Triton backend's target-verify falls back to extend_attention_fwd, whose grid is (bs, head_num, cdiv(max_len_extend, BLOCK_M)). At topk=1 with 4 draft tokens that is one program per (seq, head) that walks the whole prefix serially in BLOCK_N steps — for a 24-head model, 24 programs vs. the num_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() inside verify_splitkv.py, and its own docstring says so), and test_verify_splitkv.py is registered on the CUDA CI lane. This PR simply drops the gfx95 gate. SGLANG_ENABLE_SPLITKV_VERIFY=0 remains 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 to is_hip(), later to is_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: remove is_gfx95_supported() from the use_verify_splitkv condition; keep the topk == 1 and env-var gates. Comment updated to explain why.
  • python/sglang/srt/environ.py: doc comment for SGLANG_ENABLE_SPLITKV_VERIFY now 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 that can_handle() already validates.

Accuracy Tests

test/registered/attention/test_verify_splitkv.py on CUDA (sm_80, Triton from the 0.5.19 wheel): 9 tests, OK (parity vs extend_attention_fwd across head dims / GQA ratios / prefix lengths / KV scales, plus the can_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:

prompt tokens no spec NEXTN s3/k1/d4 before NEXTN s3/k1/d4 after
512 69.6 209.4 217.5
8,192 65.0 94.7 162.8
32,768 53.4 38.9 131.3
131,072 31.1 10.1 57.7

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


CI States

Latest PR Test (Base): ❌ Run #35825241950
Latest PR Test (Extra): ❌ Run #35825241812
Latest PR Test (AMD ROCm 10): ❌ Run #35825241944

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

rluisr commented Sep 17, 2026

Copy link
Copy Markdown
Author

@zijiecode @yichiche Can you help review this PR?

This branch has not been deployed

No deployments
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