Skip to content

[AMD][Perf] Split-KV flash-decode attention for EAGLE target-verify (Triton backend) - #27382

Merged
HaiShaw merged 8 commits into
sgl-project:mainfrom
ntgiang71096:feat/rocm-splitkv-verify
Jun 19, 2026
Merged

HaiShaw merged 8 commits into
sgl-project:mainfrom
ntgiang71096:feat/rocm-splitkv-verify

Conversation

@ntgiang71096

@ntgiang71096 ntgiang71096 commented Jun 5, 2026 •

Copy link
Copy Markdown
Contributor

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

  • New kernel python/sglang/srt/layers/attention/triton_ops/verify_splitkv.py — two Triton kernels: _verify_prefix_stage1 (split-KV over the shared prefix, applying the fp8 k_scale/v_scale dequant multipliers exactly as extend_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 as extend_attention_fwd and returns True if it ran, False otherwise.
  • Dispatch in TritonAttnBackend.forward_extend — on target-verify, try verify_splitkv_fwd(...); on True, return. Purely additive: the existing extend_attention_fwd call is unchanged and is the fallback.
  • can_handle() static-shape gate returns False (→ 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.
  • topk == 1 only — enabled via self.topk == 1 in the backend (the same condition aiter'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 tree custom_mask — matches the baseline exactly. topk>1 falls back to extend_attention_fwd.
  • Env flag SGLANG_ENABLE_SPLITKV_VERIFY (environ.py, default on) to opt out.
  • Scope: verify path only. Routing the spec-v2 draft-extend KV refill through the same kernel is a follow-up PR.

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

  • Bit-equivalent on the operative path. At topk=1 the tree mask is causal, so split-KV computes the same result as extend_attention_fwd up to bf16 reduction-order noise.
  • Unit test test/registered/attention/test_verify_splitkv.py — parity vs extend_attention_fwd across head_dim {128, 256}, GQA/MQA ratios, extend lengths, and fp8 KV scales, plus can_handle() fallback coverage. On gfx950 the max abs diff was ≈2e-3 (bf16).
  • GSM8K (end-to-end, lossless) — kernel-on vs kernel-off identical: Qwen3.6-35B-A3B-FP8 0.970 == 0.970; Qwen3.5-397B-A17B-FP8 0.980 (5-shot, 200 examples; thinking disabled for both models, 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:

ctx extend_attention_fwd split-KV verify speedup
1024 0.068 ms 0.029 ms 2.4×
4096 0.237 ms 0.046 ms 5.1×
8192 0.461 ms 0.048 ms 9.6×
16384 0.910 ms 0.082 ms 11.0×

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_VERIFY on vs off), radix cache disabled, standard random dataset, NEXTN spec (topk=1), MI350X, Qwen3.6-35B-A3B-FP8, median of 3 runs:

config output tok/s TPOT (ms) accept len
spec, verify kernel off (stock extend_attention_fwd) 855.7 7.6 2.87
spec, verify kernel on (this PR) 1117.2 5.8 2.85
Δ +30.6% −24% matched ✓

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:

# server — run twice, toggling the kernel flag (default on):
SGLANG_ENABLE_SPEC_V2=1 SGLANG_ENABLE_SPLITKV_VERIFY=1 \
python -m sglang.launch_server --model-path <hybrid-GDN model, head_dim 256> \
  --attention-backend triton --disable-radix-cache \
  --speculative-algorithm NEXTN --speculative-num-steps 3 \
  --speculative-eagle-topk 1 --speculative-num-draft-tokens 4 --context-length 16384
# bench:
python -m sglang.bench_serving --dataset-name random \
  --random-input-len 8192 --random-output-len 512 --num-prompts 32 --max-concurrency 8

Checklist

  • Code formatted with pre-commit (black / isort / ruff).
  • Unit tests added (test/registered/attention/test_verify_splitkv.py).
  • Documentation — n/a (internal backend kernel; the env var is documented inline in environ.py).
  • Accuracy and speed benchmark results provided (above).
  • Follows SGLang code style.

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

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Comment thread python/sglang/srt/layers/attention/triton_ops/verify_splitkv.py
Comment thread python/sglang/srt/layers/attention/triton_ops/verify_splitkv.py
Comment thread python/sglang/srt/layers/attention/triton_ops/verify_splitkv.py
Comment thread python/sglang/srt/layers/attention/triton_ops/verify_splitkv.py
Comment thread python/sglang/srt/layers/attention/triton_backend.py
@ntgiang71096

ntgiang71096 commented Jun 5, 2026 •

Copy link
Copy Markdown
Contributor Author

Addressed in 0bc7cce — _VK_CACHE no longer keys on the dynamic batch size:

  • VerifySplitKV scratch (att_out/att_lse) is sized by a stable max_bs (the backend passes req_to_token_pool.size), not bs.
  • The cache key drops bs, so a single instance is reused across all batch sizes — fixed buffer addresses (CUDA/HIP-graph safe) and no per-batch-size memory growth; grow_buffers() handles a larger max_bs if ever requested.
  • The kernel grid still uses the per-call bs (<= max_bs); only the first bs slices are touched (strides are independent of the leading dim).

Re-validated kernel numerics on MI350X (gfx950), including max_bs >> bs and the grow path — parity with extend_attention_fwd holds (max abs diff <= ~2e-3, bf16).

@gemini-code-assist

Copy link
Copy Markdown
Contributor

Thanks for the detailed update, @ntgiang71096. The refinements to VerifySplitKV—specifically using a stable max_bs for scratch buffer sizing and keying _VK_CACHE without bs—are excellent improvements for both memory stability and HIP/CUDA-graph compatibility. The parity results on MI350X are very promising. I have no further comments on this implementation.

@ntgiang71096

Copy link
Copy Markdown
Contributor Author

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

@HaiShaw

HaiShaw commented Jun 15, 2026

Copy link
Copy Markdown
Collaborator

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Jun 15, 2026
@ntgiang71096
ntgiang71096 force-pushed the feat/rocm-splitkv-verify branch from 93089c6 to 2fdb2d6 Compare June 15, 2026 12:20
@ntgiang71096

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@ntgiang71096

Copy link
Copy Markdown
Contributor Author

/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.
@ntgiang71096

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@ntgiang71096
ntgiang71096 force-pushed the feat/rocm-splitkv-verify branch from 526d000 to 372d63c Compare June 18, 2026 03:13

@HaiShaw HaiShaw 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.

@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.

@HaiShaw

HaiShaw commented Jun 18, 2026

Copy link
Copy Markdown
Collaborator

@kkHuang-amd @1am9trash @RolaoDenthu @Raiden-Makoto @raikonenfnu @yichiche let's triage for an implementation within aiter or DSL backend.

@HaiShaw

HaiShaw commented Jun 18, 2026

Copy link
Copy Markdown
Collaborator

@amd-bot ci-status

@amd-bot

amd-bot commented Jun 18, 2026

Copy link
Copy Markdown

@HaiShaw

CI Status for PR #27382

Merge 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 (SGLANG_ENABLE_SPLITKV_VERIFY=EnvBool(True), with no platform/ROCm guard) and passes the AMD-only Triton arg waves_per_eu, which NVIDIA's Triton rejects → scheduler crash during CUDA-graph capture. This is a real runtime regression of existing functionality, not just the new feature. PR CI is also incomplete: the bug failed base-a on NVIDIA, which fast-fail-cascaded and skipped all base-b/base-c stages — including the new kernel's own CUDA unit test.

Caution

The PR's changed code is exercised on NVIDIA and crashes there (base-a-test-1-gpu-small, test_basic_sanity_eagle3.py). Despite being described as a ROCm-only, correctness-neutral perf change, SGLANG_ENABLE_SPLITKV_VERIFY defaults to True for all backends and the kernel uses the AMD-specific waves_per_eu tuning hint. Because base-a failed, wait-for-base-a failed and all NVIDIA base-b/base-c stages were skipped (not tested) — including the PR's own CUDA unit test test/registered/attention/test_verify_splitkv.py (registered to base-b/1-gpu-small). AMD coverage is OK: that unit test ran and passed under stage-b-test-1-gpu-small-amd.

Changed files: verify_splitkv.py (+795), test_verify_splitkv.py (+239), bench_verify_splitkv.py (+138), triton_backend.py (+52), environ.py (+5)

Executed CI failure attribution: AMD: 1 failure (0 related) · Others: 4 failures (1 related) · NVIDIA base-b/base-c skipped by cascade (not tested) · 2 AMD MI35x stage-c jobs still queued.

Other Executed Failures

Job Test File Test Function Error Related? Why
base-a-test-1-gpu-small test/registered/core/test_basic_sanity_eagle3.py EAGLE3 server startup KeyError: 'Keyword argument waves_per_eu was specified but unrecognised' 🔴 Crash is in this PR's verify_splitkv.py:437 _verify_prefix_stage1[grid], reached via triton_backend.py:1122 verify_splitkv_fwd on the topk=1 verify path; NVIDIA Triton rejects waves_per_eu.
stage-b-test-1-gpu-xpu test/registered/xpu/test_xpu_embedding.py embedding server TimeoutError: Server failed to start 🟢 XPU embedding startup timeout; no EAGLE/Triton-verify path involved. Pre-existing/infra.
stage-b-test-1-npu-a2 NPU a2 suite N/A command terminated with exit code 255 🟢 NPU backend; does not use the Triton split-KV verify path.
multimodal-gen-test-2-npu-a3 multimodal_gen/.../test_server_2_npu.py test_diffusion_generation[flux_2/qwen_image/wan2_2…] diffusion testcase check failures 🟢 NPU diffusion (Wan/Flux/Qwen-Image) gen tests; unrelated to LLM EAGLE attention.

AMD Executed Failures

Job Test File Test Function Error Related? Why
stage-c-test-large-8-gpu-amd (…,3) test/registered/ops/test_aiter_allreduce_fusion_amd.py test_fused_ar_rms_residual_accuracy AssertionError: Residual accuracy check failed 🟢 aiter allreduce+RMSNorm fusion bit-exact residual check; no attention/verify code path. Pre-existing on AMD.

(*-finish, pr-test-finish, etc. are roll-up jobs reflecting the above; wait-for-base-a is the gate that cascaded. Not counted as independent failures.)

Details / what to do before merge

The root-cause for the 🔴 failure (verified in the log and diff):

  1. environ.py: SGLANG_ENABLE_SPLITKV_VERIFY = EnvBool(True) — on by default, no is_hip()/ROCm guard.
  2. triton_backend.py: self.use_verify_splitkv = envs.SGLANG_ENABLE_SPLITKV_VERIFY.get() and self.topk == 1 — gate is topk-only, not platform-aware. On NVIDIA + EAGLE topk=1 this enters verify_splitkv_fwd in forward_extend.
  3. verify_splitkv.py:437: _verify_prefix_stage1[grid](..., waves_per_eu=...) — waves_per_eu is an AMD/ROCm Triton hint; NVIDIA Triton raises KeyError, crashing the scheduler at CUDA-graph capture.

Suggested next steps for the author:

  • Gate the path to ROCm (e.g. only set use_verify_splitkv when is_hip()), or strip waves_per_eu for non-HIP Triton, so NVIDIA falls back to extend_attention_fwd.
  • After fixing, re-run full PR CI so the skipped NVIDIA base-b lane (which runs the new test_verify_splitkv.py CUDA test) and base-c actually execute — the green AMD signal alone does not cover the NVIDIA path this PR touches.
  • The XPU/NPU/AMD-allreduce failures (🟢) appear pre-existing/infra and unrelated to this PR; confirm against a base-branch run, but they should not block on this PR's account.

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.
@ntgiang71096

Copy link
Copy Markdown
Contributor Author

Thanks @HaiShaw 🙏 — addressed the NV crash in f7622c0: gated use_verify_splitkv on is_hip() and made the kernel NV-safe (the AMD-only waves_per_eu / matrix_instr_nonkdim are passed only on ROCm), so NVIDIA falls back to extend_attention_fwd. Scoped to where it's validated (gfx950). Happy to dig into the aiter/DSL direction as a separate follow-up.

@HaiShaw HaiShaw 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.

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

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.

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")

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.

please make sure adding to mi35x test group

@@ -0,0 +1,138 @@
"""Micro-benchmark: split-KV EAGLE-verify kernel vs extend_attention_fwd.

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.

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).
@ntgiang71096

Copy link
Copy Markdown
Contributor Author

Thanks @HaiShaw 🙏 — addressed all three in 2561dc3fb2:

  • triton_backend: gated use_verify_splitkv on is_gfx95_supported() instead of is_hip().
  • test_verify_splitkv: moved the AMD CI registration to the mi35x group (stage-b-test-1-gpu-small-amd-mi35x).
  • bench_verify_splitkv: added a docstring note + a runtime warning that it only applies to MI35x (gfx950).

PTAL.

Comment thread benchmark/kernels/verify_splitkv_triton/bench_verify_splitkv.py
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.
@HaiShaw
HaiShaw merged commit b36360d into sgl-project:main Jun 19, 2026
101 of 136 checks passed
@ntgiang71096

Copy link
Copy Markdown
Contributor Author

Thanks @HaiShaw for the thorough reviews and merges (this + #27793). Appreciate it!

sreerohi added a commit to sreerohi/miles_multi_turn_qwen_amd that referenced this pull request Aug 3, 2026
sreerohi added a commit to sreerohi/miles_latest that referenced this pull request Aug 10, 2026
sreerohi added a commit to sreerohi/miles_latest that referenced this pull request Aug 11, 2026
Chronostasys pushed a commit to MindLab-Research/sglang that referenced this pull request Aug 24, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants