Repository navigation
[AMD][Perf] Split-KV flash-decode attention for EAGLE target-verify (Triton backend) - #27382
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces a new split-KV (flash-decode) attention kernel (verify_splitkv_fwd) for EAGLE speculative verification when topk == 1 on the Triton backend (ROCm), along with a micro-benchmark and correctness tests. The review feedback highlights a critical memory optimization opportunity: caching VerifySplitKV instances using the dynamic batch size bs can lead to substantial GPU memory bloat and potential Out-Of-Memory (OOM) errors. The reviewer recommends refactoring the cache mechanism to use a stable max_bs (derived from the server's request pool size) and dynamically growing the scratch buffers when necessary.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
|
Addressed in 0bc7cce —
Re-validated kernel numerics on MI350X (gfx950), including |
|
Thanks for the detailed update, @ntgiang71096. The refinements to |
|
Hi maintainers 👋 — this one's ready whenever someone has a moment. Would you mind adding the run-ci label so the CUDA lane can run? It's a small, self-contained change: a new Triton split-KV verify kernel + dispatch, gated behind SGLANG_ENABLE_SPLITKV_VERIFY and topk==1 (falls back to the existing path otherwise), accuracy-neutral, with unit + fallback tests under test/registered/attention/. Happy to adjust anything — just let me know. Thanks! cc @HaiShaw |
|
/tag-and-rerun-ci |
93089c6 to
2fdb2d6
Compare
|
/rerun-failed-ci |
|
/rerun-failed-ci |
On the Triton backend, EAGLE target-verify runs through the prefill extend_attention_fwd (serial over the prefix KV per (seq, head)); with only a few draft-token queries this is bandwidth-starved at long context. This adds a split-KV (flash-decode) verify kernel that splits the prefix KV across parallel programs and merges the partials via log-sum-exp, then handles the small causal draft-draft block -- recovering HBM bandwidth on the verify path. - New kernel: triton_ops/verify_splitkv.py (prefix split-KV + fused combine). - Additive dispatch in TritonAttnBackend.forward_extend; verify_splitkv_fwd's can_handle() falls back to extend_attention_fwd for any unsupported case, so correctness is never at risk. Enabled only at speculative topk == 1 (the EAGLE tree reduces to pure causal), matching aiter's unified-verify condition. - Env flag SGLANG_ENABLE_SPLITKV_VERIFY (environ.py, default on). - Tests: test/registered/attention/test_verify_splitkv.py -- numerical parity vs extend_attention_fwd (head dims, GQA ratios, extend lens, KV scales) plus can_handle() fallback coverage. - Benchmark: benchmark/kernels/verify_splitkv_triton/. Draft-extend routing and topk>1 are out of scope (follow-ups).
Addresses review on sgl-project#27382: _VK_CACHE was keyed on the dynamic batch size, allocating a separate fp32 scratch buffer (att_out/att_lse) per bs -> GPU memory bloat / OOM risk. Now the buffers are sized by a stable max_bs (the backend passes req_to_token_pool.size), the cache key drops bs, and grow_buffers() handles a larger max_bs if ever requested -> a single reused instance, fixed buffer addresses (CUDA/HIP-graph safe), no per-batch-size growth. The kernel grid still uses the per-call bs (<= max_bs); only the first bs slices are touched. Re-validated numerics on gfx950 (incl. max_bs>>bs and the grow path).
…string - test_verify_splitkv.py: register_amd_ci suite stage-b-test-1-gpu-small-amd-mi35x -> stage-b-test-1-gpu-small-amd (the -mi35x suffix is not a valid suite; now matches sibling attention tests). Fixes base-a-test-cpu 'Tests registered to invalid suites'. - environ.py: drop extraneous f-prefix on a placeholder-less string (ruff F541, surfaced by ruff 0.15.1 now that this PR touches environ.py).
verify_splitkv.py is an imported library module, not an executable script, but carried a '#!/usr/bin/env python3' shebang -> pre-commit check-shebang-scripts-are-executable failed (lint). Drop the shebang.
|
/rerun-failed-ci |
526d000 to
372d63c
Compare
HaiShaw
left a comment
There was a problem hiding this comment.
@ntgiang71096 design is nice! There are places mentioning AMD gfx950, though not much code is AMD specific (correct me if wrong), therefore I guess you developed over INSTINCT GPUs, and test on NV isn't conducted. So you can either gate this to AMD specific (is_hip, _gfx94, _gfx95), or test on NV as well to move forward.
|
@kkHuang-amd @1am9trash @RolaoDenthu @Raiden-Makoto @raikonenfnu @yichiche let's triage for an implementation within aiter or DSL backend. |
|
@amd-bot ci-status |
CI Status for PR #27382Merge verdict: 🔴 Do not merge. This PR breaks EAGLE topk=1 target-verify on NVIDIA (and any non-ROCm Triton backend). The new split-KV verify kernel is enabled by default ( Caution The PR's changed code is exercised on NVIDIA and crashes there ( Changed files: Executed CI failure attribution: AMD: 1 failure (0 related) · Others: 4 failures (1 related) · NVIDIA Other Executed Failures
AMD Executed Failures
( Details / what to do before mergeThe root-cause for the 🔴 failure (verified in the log and diff):
Suggested next steps for the author:
Generated by amd-bot using Claude Code CLI |
Review (HaiShaw / amd-bot on sgl-project#27382): the kernel was enabled by default on the Triton backend for all platforms, but passes AMD/CDNA-only Triton launch hints (waves_per_eu, matrix_instr_nonkdim) that NVIDIA's Triton rejects -> crash during CUDA-graph capture on NV EAGLE topk=1 target-verify. The path was developed and validated only on AMD gfx950. - triton_backend: gate use_verify_splitkv on is_hip() so the path runs only on ROCm and falls back to extend_attention_fwd on NVIDIA/other. - verify_splitkv: pass waves_per_eu / matrix_instr_nonkdim only on ROCm via _AMD_LAUNCH_KWARGS, keeping the kernel NV-safe (the numerics test can still run on the CUDA CI lane). AMD launch is unchanged.
|
Thanks @HaiShaw 🙏 — addressed the NV crash in |
HaiShaw
left a comment
There was a problem hiding this comment.
Some to address, thanks.
| # NVIDIA's Triton rejects those kwargs and the path is unvalidated on NV, | ||
| # so restrict it to ROCm and fall back to extend_attention_fwd elsewhere. | ||
| self.use_verify_splitkv = ( | ||
| is_hip() and envs.SGLANG_ENABLE_SPLITKV_VERIFY.get() and self.topk == 1 |
There was a problem hiding this comment.
please use gfx95x check here instead of hip.
| from sglang.test.test_utils import CustomTestCase | ||
|
|
||
| register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-small") | ||
| register_amd_ci(est_time=30, suite="stage-b-test-1-gpu-small-amd") |
There was a problem hiding this comment.
please make sure adding to mi35x test group
| @@ -0,0 +1,138 @@ | |||
| """Micro-benchmark: split-KV EAGLE-verify kernel vs extend_attention_fwd. | |||
There was a problem hiding this comment.
please signal to user that this bench only applies to mi35x at run time and/or comment.
Addresses HaiShaw review on sgl-project#27382: - triton_backend: gate use_verify_splitkv on is_gfx95_supported() instead of is_hip() -- the kernel's block config + CDNA launch hints are tuned/validated only on gfx950 (MI350X/CDNA4), so restrict the dispatch to gfx95. - test_verify_splitkv: register the AMD CI to the mi35x group (stage-b-test-1-gpu-small-amd-mi35x) so it runs on the validated hardware. - bench_verify_splitkv: docstring note + runtime warning that the benchmark is only meaningful on MI35x (gfx950).
|
Thanks @HaiShaw 🙏 — addressed all three in
PTAL. |
Per HaiShaw review: the micro-benchmark is meaningful only on MI35x (gfx950); exit with a clear message on other hardware instead of warning and continuing.
Motivation
On the Triton attention backend, EAGLE/NEXTN target-verify is dispatched through the prefill
extend_attention_fwd, which loops the prefix KV serially per(sequence, head)with no KV-split. In the verify shape — a few draft-token queries (speculative_num_draft_tokens, e.g. 4) attending over a long KV cache — this has almost no query parallelism and walks the long KV serially. On MI350X (gfx950) at ctx 16k the verify attention reads about 536 MB of K+V in 0.79 ms, i.e. ≈680 GB/s — ≈8% of the 8 TB/s HBM peak: it is badly bandwidth-starved, and the cost grows with context until it cancels the accept-length benefit of speculative decoding. This is a large part of why spec decoding showed little/negative speedup on AMD at long context (cf. #23123).The fix is the standard flash-decode / split-KV technique the decode path already uses, adapted to the multi-query verify shape: split the long shared prefix across parallel programs and merge with a log-sum-exp reduction, then handle the small causal draft-draft block. The Triton backend has no split-KV verify path today — verify falls back to the prefill kernel.
Modifications
python/sglang/srt/layers/attention/triton_ops/verify_splitkv.py— two Triton kernels:_verify_prefix_stage1(split-KV over the shared prefix, applying the fp8k_scale/v_scaledequant multipliers exactly asextend_attention_fwd) and_verify_combine_stage2(log-sum-exp merge of the prefix splits + the small causal draft-draft block).verify_splitkv_fwd(...)takes the same positional args asextend_attention_fwdand returnsTrueif it ran,Falseotherwise.TritonAttnBackend.forward_extend— on target-verify, tryverify_splitkv_fwd(...); onTrue, return. Purely additive: the existingextend_attention_fwdcall is unchanged and is the fallback.can_handle()static-shape gate returnsFalse(→ fallback) for anything it can't serve bit-equivalently (non-causal, sinks, sliding-window, logit-cap, xai-temperature, ragged extend). It reads no tensor values, so it is HIP/CUDA-graph-capture safe.self.topk == 1in the backend (the same conditionaiter's unified-verify uses). At topk=1 the EAGLE tree reduces to a pure causal chain, so the kernel — which computes pure-causal attention and ignores the treecustom_mask— matches the baseline exactly. topk>1 falls back toextend_attention_fwd.SGLANG_ENABLE_SPLITKV_VERIFY(environ.py, default on) to opt out.Files:
verify_splitkv.py(new, kernel),triton_backend.py(+dispatch),environ.py(+flag),test/registered/attention/test_verify_splitkv.py(new),benchmark/kernels/verify_splitkv_triton/(new).Accuracy Tests
extend_attention_fwdup to bf16 reduction-order noise.test/registered/attention/test_verify_splitkv.py— parity vsextend_attention_fwdacross head_dim {128, 256}, GQA/MQA ratios, extend lengths, and fp8 KV scales, pluscan_handle()fallback coverage. On gfx950 the max abs diff was ≈2e-3 (bf16).enable_thinking=false).Speed Tests and Profiling
Kernel-level —
benchmark/kernels/verify_splitkv_triton/bench_verify_splitkv.py, MI350X (gfx950), bf16, verify shape (H_Q=16, H_KV=2, head_dim=256, 4 draft tokens), batch=1 to isolate per-sequence scaling:The baseline scales roughly linearly with context (serial prefix loop) while split-KV stays nearly flat, so the kernel-level speedup grows with context. (At production batch sizes the baseline gains occupancy from the batch, so the end-to-end gain is smaller — see below.)
End-to-end — isolated A/B with the verify kernel the only variable (
SGLANG_ENABLE_SPLITKV_VERIFYon vs off), radix cache disabled, standardrandomdataset, NEXTN spec (topk=1), MI350X, Qwen3.6-35B-A3B-FP8, median of 3 runs:extend_attention_fwd)Accept length is unchanged across arms (the kernel changes speed, not what is accepted) — confirming a clean isolation. Measured at 8k context; the per-kernel speedup (and the e2e gain) grows with context (see the kernel table above). Reproduce with stock sglang:
Checklist
test/registered/attention/test_verify_splitkv.py).environ.py).Related issues: #23123 (MTP / spec decoding shows little speedup on AMD at long context).
CI States
Latest PR Test (Base): ❌ Run #27795826802
Latest PR Test (Extra): ❌ Run #27795826686