From 67f02681c9a3dd578c8d987f48796e17c1f98aa5 Mon Sep 17 00:00:00 2001 From: Hubert Lu <55214931+hubertlu-tw@users.noreply.github.com> Date: Wed, 11 Mar 2026 17:01:01 -0700 Subject: [PATCH] [AMD] Support speculative decoding v2 for aiter backend on ROCm/HIP (#17450) Co-authored-by: kkHuang-amd Co-authored-by: HaiShaw --- .../srt/layers/attention/aiter_backend.py | 328 +++++++++++++++++- .../sglang/srt/speculative/eagle_info_v2.py | 2 +- .../sglang/srt/speculative/eagle_worker_v2.py | 27 +- .../amd/test_deepseek_r1_mxfp4_8gpu.py | 8 +- 4 files changed, 341 insertions(+), 24 deletions(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 2f36a088e..554477f10 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -106,6 +106,7 @@ class AiterAttnBackend(AttentionBackend): model_runner: ModelRunner, skip_prefill: bool = False, kv_indptr_buf: Optional[torch.Tensor] = None, + topk: int = 1, ): super().__init__() # Lazy import to avoid the initialization of cuda context @@ -123,6 +124,7 @@ class AiterAttnBackend(AttentionBackend): self.is_multimodal = model_runner.model_config.is_multimodal self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.speculative_num_steps = model_runner.server_args.speculative_num_steps + self.topk = topk self.num_head = ( model_runner.model_config.num_attention_heads // get_attention_tp_size() ) @@ -171,6 +173,7 @@ class AiterAttnBackend(AttentionBackend): self.mask_indptr = torch.zeros( (max_bs + 1,), dtype=torch.int64, device=model_runner.device ) + self._kv_indices_scratch: Optional[torch.Tensor] = None # Create prefill indices updater if not skip_prefill: @@ -432,6 +435,74 @@ class AiterAttnBackend(AttentionBackend): is_causal=is_causal, ) + def _resolve_v2_num_draft_tokens( + self, + extend_seq_lens: Optional[torch.Tensor] = None, + extend_seq_lens_cpu: Optional[list[int]] = None, + ) -> int: + """Resolve fixed per-request extend length for DRAFT_EXTEND_V2.""" + num_draft_tokens = self.num_draft_tokens + if num_draft_tokens is None: + if extend_seq_lens is not None and extend_seq_lens.numel() > 0: + # Avoid list scans in hot path when tensor lengths are already available. + num_draft_tokens = int(extend_seq_lens[0].item()) + elif extend_seq_lens_cpu: + num_draft_tokens = max(extend_seq_lens_cpu) + else: + raise ValueError( + "DRAFT_EXTEND_V2 requires speculative_num_draft_tokens or " + "non-empty extend_seq_lens/extend_seq_lens_cpu." + ) + + num_draft_tokens = int(num_draft_tokens) + if extend_seq_lens is not None and extend_seq_lens.numel() > 0: + if not torch.all(extend_seq_lens == num_draft_tokens): + raise ValueError( + "DRAFT_EXTEND_V2 expects fixed extend length per request; got " + f"extend_seq_lens={extend_seq_lens}, expected all == {num_draft_tokens}." + ) + if extend_seq_lens_cpu and any( + x != num_draft_tokens for x in extend_seq_lens_cpu + ): + raise ValueError( + "DRAFT_EXTEND_V2 expects fixed extend length per request; got " + f"{extend_seq_lens_cpu}, expected all == {num_draft_tokens}." + ) + return num_draft_tokens + + def _get_kv_indices_scratch( + self, required_tokens: int, device: torch.device + ) -> torch.Tensor: + if ( + self._kv_indices_scratch is None + or self._kv_indices_scratch.device != device + or self._kv_indices_scratch.numel() < required_tokens + ): + self._kv_indices_scratch = torch.empty( + required_tokens, dtype=torch.int32, device=device + ) + return self._kv_indices_scratch[:required_tokens] + + def _set_uniform_qo_indptr( + self, bs: int, tokens_per_req: int, device: torch.device + ) -> torch.Tensor: + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[: bs + 1] = torch.arange( + 0, + bs * tokens_per_req + 1, + step=tokens_per_req, + dtype=torch.int32, + device=device, + ) + return qo_indptr + + def _ensure_spec_v2_topk_supported(self): + if self.topk > 1: + raise NotImplementedError( + "AiterAttnBackend SPEC_V2 path currently supports topk <= 1 only. " + f"Got topk={self.topk}." + ) + def mla_fp8_prefill_attn( self, q: torch.Tensor, @@ -508,7 +579,7 @@ class AiterAttnBackend(AttentionBackend): return output def init_forward_metadata(self, forward_batch: ForwardBatch): - """Init auxiliary variables for triton attention backend.""" + """Init auxiliary variables for aiter attention backend.""" bs = forward_batch.batch_size kv_indptr = self.kv_indptr @@ -531,8 +602,8 @@ class AiterAttnBackend(AttentionBackend): if spec_info is None or forward_batch.forward_mode.is_idle(): kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) kv_indptr = kv_indptr[: bs + 1] - kv_indices = torch.empty( - forward_batch.seq_lens_sum, dtype=torch.int32, device=self.device + kv_indices = self._get_kv_indices_scratch( + forward_batch.seq_lens_sum, forward_batch.seq_lens.device ) create_flashinfer_kv_indices_triton[(bs,)]( self.req_to_token, @@ -598,7 +669,97 @@ class AiterAttnBackend(AttentionBackend): run_graph=False, ) + elif forward_batch.forward_mode.is_draft_extend_v2(): + # EAGLE V2: DRAFT_EXTEND_V2 mode - extend draft KV cache with all predicted tokens + self._ensure_spec_v2_topk_supported() + if self.use_mla: + device = forward_batch.seq_lens.device + num_draft_tokens = self._resolve_v2_num_draft_tokens( + extend_seq_lens=forward_batch.extend_seq_lens + ) + qo_indptr = self._set_uniform_qo_indptr(bs, num_draft_tokens, device) + + kv_indptr = self.kv_indptr[: bs + 1] + kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) + + kv_indices = self._get_kv_indices_scratch( + forward_batch.seq_lens_sum, device + ) + + create_flashinfer_kv_indices_triton[(bs,)]( + self.req_to_token, + forward_batch.req_pool_indices, + forward_batch.seq_lens, + kv_indptr, + None, + kv_indices, + self.req_to_token.stride(0), + ) + + if _use_mla_ps_kernel: + max_seqlen_qo = num_draft_tokens + ( + work_metadata, + work_indptr, + work_info_set, + reduce_indptr, + reduce_final_map, + reduce_partial_map, + ) = self.make_mla_decode_meta_data_buffer(max_seqlen_qo, bs) + + num_kv_splits = self.max_split_per_batch + + self.make_mla_meta_data( + qo_indptr, + kv_indptr, + self.kv_last_page_len[:bs], + work_metadata, + work_info_set, + work_indptr, + reduce_indptr, + reduce_final_map, + reduce_partial_map, + max_seqlen_qo, + fast_mode=fast_mode, + max_split_per_batch=num_kv_splits, + intra_batch_mode=intra_batch_mode, + ) + + self.forward_metadata = ForwardMetadata( + kv_indptr, + kv_indices, + qo_indptr, + self.kv_last_page_len[:bs], + num_draft_tokens, + forward_batch.seq_lens_cpu.max().item(), + work_metadata=work_metadata, + work_info_set=work_info_set, + work_indptr=work_indptr, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + num_kv_splits=num_kv_splits, + run_graph=False, + ) + else: + self.indices_updater_prefill.update( + forward_batch.req_pool_indices, + forward_batch.seq_lens, + forward_batch.seq_lens_sum, + prefix_lens=None, + encoder_lens=forward_batch.encoder_lens, + spec_info=forward_batch.spec_info, + ) + self.forward_metadata = ForwardMetadata( + self.indices_updater_prefill.kv_indptr, + self.indices_updater_prefill.kv_indices, + None, + None, + self.indices_updater_prefill.max_q_len, + self.indices_updater_prefill.max_kv_len, + ) elif forward_batch.forward_mode.is_draft_extend(): + # EAGLE V1: DRAFT_EXTEND mode - uses spec_info.accept_length if self.use_mla: kv_indices, kv_indptr, qo_indptr, custom_mask = ( spec_info.generate_attn_arg_prefill( @@ -686,20 +847,19 @@ class AiterAttnBackend(AttentionBackend): kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs device = forward_batch.seq_lens.device - qo_indptr = torch.arange( + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[: bs + 1] = torch.arange( 0, (1 + bs) * draft_num, step=draft_num, dtype=torch.int32, device=device, ) - kv_indptr = self.kv_indptr + kv_indptr = self.kv_indptr[: bs + 1] kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] - kv_indices = torch.empty( + kv_indices = self._get_kv_indices_scratch( kv_lens_sum, - dtype=torch.int32, - device=device, + device, ) create_flashinfer_kv_indices_triton[(bs,)]( self.req_to_token, @@ -1040,7 +1200,6 @@ class AiterAttnBackend(AttentionBackend): reduce_final_map=reduce_final_map, reduce_partial_map=reduce_partial_map, num_kv_splits=num_kv_splits, - # num_kv_splits_indptr=num_kv_splits_indptr, ) elif forward_mode.is_target_verify(): @@ -1134,7 +1293,70 @@ class AiterAttnBackend(AttentionBackend): mask_indptr=mask_indptr, max_extend_len=max_q_len, ) + elif forward_mode.is_draft_extend_v2(): + # EAGLE V2: Uses fixed num_draft_tokens per batch + self._ensure_spec_v2_topk_supported() + num_tokens_per_bs = self._resolve_v2_num_draft_tokens() + qo_indptr = self._set_uniform_qo_indptr(bs, num_tokens_per_bs, self.device) + kv_indptr = self.kv_indptr[: bs + 1] + kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) + kv_indices = self.cuda_graph_kv_indices + create_flashinfer_kv_indices_triton[(bs,)]( + self.req_to_token, + req_pool_indices, + seq_lens, + kv_indptr, + None, + kv_indices, + self.req_to_token.stride(0), + ) + kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] + max_q_len = num_tokens_per_bs + + if _use_mla_ps_kernel: + num_kv_splits = self.max_split_per_batch + + self.make_mla_meta_data( + qo_indptr, + kv_indptr, + kv_last_page_len, + self.work_metadata, + self.work_info_set, + self.work_indptr, + self.reduce_indptr, + self.reduce_final_map, + self.reduce_partial_map, + max_q_len, + fast_mode=fast_mode, + max_split_per_batch=num_kv_splits, + intra_batch_mode=intra_batch_mode, + ) + + work_metadata = self.work_metadata + work_info_set = self.work_info_set + work_indptr = self.work_indptr + + reduce_indptr = self.reduce_indptr + reduce_final_map = self.reduce_final_map + reduce_partial_map = self.reduce_partial_map + + self.forward_metadata = ForwardMetadata( + kv_indptr, + kv_indices, + qo_indptr, + kv_last_page_len, + max_q_len, + kv_indptr[-1].item(), + work_metadata=work_metadata, + work_info_set=work_info_set, + work_indptr=work_indptr, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + num_kv_splits=num_kv_splits, + ) elif forward_mode.is_draft_extend(): + # EAGLE V1: Uses speculative_num_steps + 1 num_tokens_per_bs = self.speculative_num_steps + 1 qo_indptr = self.qo_indptr[: bs + 1] qo_indptr[: bs + 1] = torch.arange( @@ -1314,7 +1536,6 @@ class AiterAttnBackend(AttentionBackend): reduce_final_map=reduce_final_map, reduce_partial_map=reduce_partial_map, num_kv_splits=num_kv_splits, - # num_kv_splits_indptr=num_kv_splits_indptr, ) elif forward_mode.is_target_verify(): @@ -1408,8 +1629,78 @@ class AiterAttnBackend(AttentionBackend): mask_indptr=mask_indptr, max_extend_len=max_q_len, ) + elif forward_mode.is_draft_extend_v2(): + # EAGLE V2: Fixed num_draft_tokens per batch + self._ensure_spec_v2_topk_supported() + seq_lens = seq_lens[:bs] + num_tokens_per_bs = self._resolve_v2_num_draft_tokens() + extend_lens = torch.full( + (bs,), num_tokens_per_bs, dtype=torch.int32, device=seq_lens.device + ) + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[1 : bs + 1] = torch.cumsum(extend_lens, dim=0) + kv_indptr = self.kv_indptr[: bs + 1] + kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) + kv_indices = self.cuda_graph_kv_indices + create_flashinfer_kv_indices_triton[(bs,)]( + self.req_to_token, + req_pool_indices, + seq_lens, + kv_indptr, + None, + kv_indices, + self.req_to_token.stride(0), + ) + + kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] + max_q_len = num_tokens_per_bs + + if _use_mla_ps_kernel: + + num_kv_splits = self.max_split_per_batch + + self.make_mla_meta_data( + qo_indptr, + kv_indptr, + kv_last_page_len, + self.work_metadata, + self.work_info_set, + self.work_indptr, + self.reduce_indptr, + self.reduce_final_map, + self.reduce_partial_map, + max_q_len, + fast_mode=fast_mode, + max_split_per_batch=num_kv_splits, + intra_batch_mode=intra_batch_mode, + ) + + work_metadata = self.work_metadata + work_info_set = self.work_info_set + work_indptr = self.work_indptr + + reduce_indptr = self.reduce_indptr + reduce_final_map = self.reduce_final_map + reduce_partial_map = self.reduce_partial_map + + self.forward_metadata = ForwardMetadata( + kv_indptr, + kv_indices, + qo_indptr, + kv_last_page_len, + max_q_len, + kv_indptr[-1].item(), + work_metadata=work_metadata, + work_info_set=work_info_set, + work_indptr=work_indptr, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + num_kv_splits=num_kv_splits, + ) elif forward_mode.is_draft_extend(): + # EAGLE V1: Uses spec_info.accept_length num_tokens_per_bs = self.speculative_num_steps + 1 seq_lens = seq_lens[:bs] accept_lens = spec_info.accept_length[:bs] @@ -1481,6 +1772,14 @@ class AiterAttnBackend(AttentionBackend): def get_cuda_graph_seq_len_fill_value(self): return 1 + def update_verify_buffers_to_fill_after_draft( + self, spec_info: SpecInput, cuda_graph_bs: Optional[int] + ): + # AITER verify path does not require post-draft buffer patching currently. + # This override prevents overlap-plan stream mode from failing with the + # base class NotImplementedError. + pass + def forward_extend( self, q: torch.Tensor, @@ -1528,6 +1827,7 @@ class AiterAttnBackend(AttentionBackend): forward_batch.forward_mode.is_extend() and not forward_batch.forward_mode.is_target_verify() and not forward_batch.forward_mode.is_draft_extend() + and not forward_batch.forward_mode.is_draft_extend_v2() ): extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) if kv_indices.shape[0] == 0 or extend_no_prefix: @@ -1680,7 +1980,10 @@ class AiterAttnBackend(AttentionBackend): num_kv_splits=num_kv_splits, ) return o - elif forward_batch.forward_mode.is_draft_extend(): + elif ( + forward_batch.forward_mode.is_draft_extend() + or forward_batch.forward_mode.is_draft_extend_v2() + ): work_metadata = self.forward_metadata.work_metadata work_indptr = self.forward_metadata.work_indptr @@ -2156,6 +2459,7 @@ class AiterMultiStepDraftBackend: model_runner, skip_prefill=True, kv_indptr_buf=self.kv_indptr[i], + topk=topk, ) ) self.max_context_len = self.attn_backends[0].max_context_len diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index d96c89da5..1224fbd33 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -310,7 +310,7 @@ class EagleVerifyInputV2Mixin: accept_length = torch.empty((bs,), dtype=torch.int32, device=device) # Sample tokens - if sampling_info.is_all_greedy or _is_npu: + if sampling_info.is_all_greedy or _is_npu or _is_hip: target_predict = torch.argmax(next_token_logits, dim=-1) target_predict = target_predict.reshape(bs, self.draft_token_num) predict, accept_index, accept_length = verify_tree_greedy_func( diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 24003af5f..248e7015c 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -59,6 +59,7 @@ from sglang.srt.utils.common import ( fast_topk, get_available_gpu_memory, is_cuda, + is_hip, is_npu, next_power_of_2, ) @@ -66,6 +67,7 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions _is_npu = is_npu() _is_cuda = is_cuda() +_is_hip = is_hip() logger = logging.getLogger(__name__) @@ -280,18 +282,27 @@ class EagleDraftWorker(BaseDraftWorker): "npu": EAGLEDraftExtendNpuGraphRunner, "cuda": EAGLEDraftExtendCudaGraphRunner, } + supports_hip_aiter_draft_extend_graph = False + if _is_hip: + # Keep import local so non-HIP environments do not require aiter. + from sglang.srt.layers.attention.aiter_backend import ( + AiterMultiStepDraftBackend, + ) + + supports_hip_aiter_draft_extend_graph = isinstance( + self.draft_attn_backend, AiterMultiStepDraftBackend + ) + + supports_cuda_draft_extend_graph = _is_cuda and ( + isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend) + or isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend) + ) # Capture extend # TODO: support draft extend cuda graph for more attention backends if self.draft_extend_attn_backend and ( _is_npu - or ( - _is_cuda - and isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend) - ) - or ( - _is_cuda - and isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend) - ) + or supports_cuda_draft_extend_graph + or supports_hip_aiter_draft_extend_graph ): tic = time.perf_counter() before_mem = get_available_gpu_memory(self.device, self.gpu_id) diff --git a/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py b/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py index 1851079ff..28249e706 100644 --- a/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py +++ b/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py @@ -1,9 +1,9 @@ -import os import unittest from types import SimpleNamespace import requests +from sglang.srt.environ import envs from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import register_amd_ci from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k @@ -87,6 +87,10 @@ class TestDeepseekR1MXFP4MTP(CustomTestCase): def setUpClass(cls): cls.model = DEEPSEEK_R1_MODEL_PATH cls.base_url = DEFAULT_URL_FOR_TEST + + envs.SGLANG_ENABLE_SPEC_V2.set(True) + envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.set(True) + other_args = [ "--tp", "8", @@ -113,8 +117,6 @@ class TestDeepseekR1MXFP4MTP(CustomTestCase): @classmethod def tearDownClass(cls): kill_process_tree(cls.process.pid) - if "SGLANG_ENABLE_SPEC_V2" in os.environ: - del os.environ["SGLANG_ENABLE_SPEC_V2"] def test_a_gsm8k( self,