From eec794bce0808ae26cc1dcb84a56b65d2df82af5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?H=E1=BB=93=20S=E1=BB=B9=20Th=E1=BA=BF?= <95731191+hsthe29@users.noreply.github.com> Date: Sun, 23 Aug 2026 04:22:43 +0700 Subject: [PATCH] [AMD][Spec] Fix aiter GQA packing + split-KV routing in NEXTN spec attention (verify & draft_extend) (#30105) --- python/sglang/srt/environ.py | 8 +- .../srt/layers/attention/aiter_backend.py | 74 +++++++++++++++++-- .../test_aiter_unified_draft_extend_env.py | 40 ++++++++++ 3 files changed, 115 insertions(+), 7 deletions(-) create mode 100644 test/manual/test_aiter_unified_draft_extend_env.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 66a988b64..060e2a0e3 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -513,7 +513,13 @@ class Envs: # HND KV layout folds (page, head) into one paged index for per-kv-head sparse # page tables (DP attn); paged backends like trtllm_mha consume it directly. SGLANG_USE_HND_KVCACHE = EnvBool(False) - # Size the KV pool after CUDA-graph capture. + + # Attention (aiter, ROCm): route NEXTN spec draft_extend (EAGLE-v2 KV + # catch-up) through aiter unified_attention (GQA-packed + split-KV) instead + # of the occupancy-starved mha_batch_prefill FMHA. Independent kill-switch + # for the new path; pairs with SGLANG_AITER_UNIFIED_VERIFY. Default on. + SGLANG_AITER_UNIFIED_DRAFT_EXTEND = EnvBool(True) + # size the KV pool after CUDA-graph capture SGLANG_ENABLE_POST_CAPTURE_KV_SIZING = EnvBool(False) # =================================================================== diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 3943a5812..994cfe108 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -70,6 +70,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( scaled_fp8_quant, ) from sglang.srt.configs.model_config import AttentionArch +from sglang.srt.environ import envs from sglang.srt.layers.attention.aiter_utils import ( forward_decode_vectorized_5d, forward_extend_vectorized_5d, @@ -2285,12 +2286,15 @@ class AiterAttnBackend(AttentionBackend): v_unified = v_cache.view( -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim ) - if layer.tp_k_head_num == 1 and layer.tp_q_head_num > 1: - # Qwen3.5 can replicate one KV head across multiple TP ranks. - # Present the local KV head as per-Q-head stride-0 views so - # target_verify uses the same local head mapping as the model. - k_unified = k_unified.expand(-1, -1, layer.tp_q_head_num, -1) - v_unified = v_unified.expand(-1, -1, layer.tp_q_head_num, -1) + # GQA-packing fix: do NOT expand the single KV head to + # tp_q_head_num. Passing K/V with the true kv-head count (exactly + # like forward_decode) lets unified_attention derive + # num_queries_per_kv = tp_q_head_num (the GQA group) and pack all Q + # heads against one KV load. The old stride-0 .expand() made the + # wrapper see num_kv_heads=tp_q_head_num -> num_queries_per_kv=1 -> + # full MHA tiling (~7x more KV traffic at long context; trace + # signature num_query_heads_16/num_queries_per_kv_1). GQA head + # mapping here is identical to the proven decode path. # The seq_lens + draft_num add has to run INSIDE the graph # region; a host-side pre-add would allocate a new tensor @@ -2339,6 +2343,64 @@ class AiterAttnBackend(AttentionBackend): ) return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + # draft_extend (EAGLE-v2 KV catch-up) is decode-shaped: short Q + # (accepted tokens) against long paged KV, GQA. The default + # mha_batch_prefill_func FMHA has no split-KV, so at short Q it is + # occupancy-starved (~0.2% HBM BW, ~400x over the memory floor). + # Route it through unified_attention (GQA-packed + split-KV), exactly + # like target_verify, so it runs near the memory floor. + if ( + self._use_unified_verify + and forward_batch.forward_mode.is_draft_extend_v2() + and envs.SGLANG_AITER_UNIFIED_DRAFT_EXTEND.get() + ): + bs = forward_batch.batch_size + if layer.qk_head_dim != layer.v_head_dim: + o = q.new_empty( + (q.shape[0], layer.tp_q_head_num * layer.v_head_dim) + ) + else: + o = torch.empty_like(q) + k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id) + page_table, swa_page_table = self._build_unified_page_table_from_spec( + self.forward_metadata, bs + ) + pt = page_table + de_window = (-1, -1) + if ( + layer.sliding_window_size is not None + and layer.sliding_window_size > -1 + ): + de_window = (layer.sliding_window_size - 1, 0) + if swa_page_table is not None: + pt = swa_page_table + kv_indptr = self.forward_metadata.kv_indptr + seqused_k = (kv_indptr[1 : bs + 1] - kv_indptr[:bs]).to(torch.int32) + unified_attention( + q=q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), + k=k_cache.view( + -1, self.page_size, layer.tp_k_head_num, layer.qk_head_dim + ), + v=v_cache.view( + -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim + ), + out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim), + cu_seqlens_q=self.qo_indptr[: bs + 1], + seqused_k=seqused_k, + max_seqlen_q=self.forward_metadata.max_q_len, + max_seqlen_k=pt.shape[1] * self.page_size, + softmax_scale=layer.scaling, + causal=True, + window_size=de_window, + block_table=pt, + softcap=layer.logit_cap, + q_descale=None, + k_descale=k_descale, + v_descale=v_descale, + sinks=sinks, + ) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + bs0 = forward_batch.batch_size + 1 q_descale = None diff --git a/test/manual/test_aiter_unified_draft_extend_env.py b/test/manual/test_aiter_unified_draft_extend_env.py new file mode 100644 index 000000000..84a45ede3 --- /dev/null +++ b/test/manual/test_aiter_unified_draft_extend_env.py @@ -0,0 +1,40 @@ +""" +Unit test for the aiter NEXTN spec-decode draft_extend routing flag. + +aiter_backend.py routes EAGLE-v2 draft_extend (KV catch-up) through +unified_attention (GQA-packed + split-KV) instead of the occupancy-starved +mha_batch_prefill FMHA, gated behind SGLANG_AITER_UNIFIED_DRAFT_EXTEND. This +guards the flag's registration and default so an accidental rename or default +flip is caught in CI. The kernel path itself requires ROCm/gfx950 and is +covered by the accuracy/speed evidence in the PR. + +Run: + python test/manual/test_aiter_unified_draft_extend_env.py +""" + +import os +import sys +import unittest + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../../python")) + + +class TestAiterUnifiedDraftExtendEnv(unittest.TestCase): + def setUp(self): + from sglang.srt.environ import envs + + self.flag = envs.SGLANG_AITER_UNIFIED_DRAFT_EXTEND + + def test_registered_and_default_on(self): + # Registered as an EnvBool defaulting to True (path enabled in prod). + self.assertTrue(self.flag.get()) + + def test_override_toggles_and_restores(self): + self.assertTrue(self.flag.get()) + with self.flag.override(False): + self.assertFalse(self.flag.get()) + self.assertTrue(self.flag.get()) + + +if __name__ == "__main__": + unittest.main()