From 8e890391f585f5aa7f04c6620e3632197f531c57 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sun, 21 Jun 2026 14:46:25 -0700 Subject: [PATCH] [Spec] Support FlashInfer CUDA graph for EAGLE draft-extend (#28782) --- .../layers/attention/flashinfer_backend.py | 30 +++++++++++++++++-- .../eagle_draft_extend_cuda_graph_runner.py | 2 ++ python/sglang/srt/speculative/eagle_info.py | 10 +++++-- .../sglang/srt/speculative/eagle_worker_v2.py | 15 ++++++++-- 4 files changed, 50 insertions(+), 7 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 30cd64e52..e8a919a84 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -572,6 +572,18 @@ class FlashInferAttnBackend(AttentionBackend): encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, spec_info=None, ) + elif forward_mode.is_draft_extend_v2(): + self.indices_updater_prefill.update( + req_pool_indices[:bs], + seq_lens[:bs], + seq_lens_cpu[:bs] if seq_lens_cpu is not None else None, + seq_lens_sum, + prefix_lens=None, + prefill_wrappers=self.draft_extend_cuda_graph_metadata[bs], + use_ragged=False, + encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, + spec_info=spec_info, + ) else: raise ValueError("Invalid forward mode") @@ -678,6 +690,7 @@ class FlashInferAttnBackend(AttentionBackend): fixed_split_size=self.prefill_split_tile_size, multi_item_params=multi_item_params, cross_attention_custom_mask=forward_batch.cross_attention_custom_mask, + extend_prefix_lens_cpu=forward_batch.extend_prefix_lens_cpu, ) self.forward_metadata = PrefillMetadata( self.prefill_wrappers_paged, @@ -797,6 +810,11 @@ class FlashInferAttnBackend(AttentionBackend): self.forward_metadata = PrefillMetadata( prefill_wrappers, forward_mode.is_dllm_extend(), False ) + elif forward_mode.is_draft_extend_v2(): + # Draft-extend: causal paged prefill over the full sequence (no mask). + prefill_wrappers = self._create_prefill_wrappers(bs, use_custom_mask=False) + self.draft_extend_cuda_graph_metadata[bs] = prefill_wrappers + self.forward_metadata = PrefillMetadata(prefill_wrappers, False, False) else: raise ValueError(f"Invalid mode: {forward_mode=}") @@ -1306,6 +1324,7 @@ class FlashInferIndicesUpdaterPrefill: fixed_split_size: Optional[int] = None, multi_item_params: Optional[MultiItemScoringParams] = None, cross_attention_custom_mask: Optional[torch.Tensor] = None, + extend_prefix_lens_cpu: Optional[List[int]] = None, ): # Keep the signature for type checking. It will be assigned during runtime. raise NotImplementedError() @@ -1324,13 +1343,16 @@ class FlashInferIndicesUpdaterPrefill: fixed_split_size: Optional[int] = None, multi_item_params: Optional[MultiItemScoringParams] = None, cross_attention_custom_mask: Optional[torch.Tensor] = None, + extend_prefix_lens_cpu: Optional[List[int]] = None, ): if use_ragged: assert prefix_lens is not None - # TODO: remove this device sync, we can use forward_batch.extend_prefix_lens_cpu - # and forward_batch.extend_seq_lens_cpu paged_kernel_lens = prefix_lens - paged_kernel_lens_sum = paged_kernel_lens.sum().item() + if extend_prefix_lens_cpu is not None: + # Host-known prefix lens; avoids a per-step D2H sync. + paged_kernel_lens_sum = sum(extend_prefix_lens_cpu) + else: + paged_kernel_lens_sum = paged_kernel_lens.sum().item() else: paged_kernel_lens = seq_lens paged_kernel_lens_sum = seq_lens_sum @@ -1366,6 +1388,7 @@ class FlashInferIndicesUpdaterPrefill: fixed_split_size: Optional[int] = None, multi_item_params: Optional[MultiItemScoringParams] = None, cross_attention_custom_mask: Optional[torch.Tensor] = None, + extend_prefix_lens_cpu: Optional[List[int]] = None, ): if prefix_lens is None: num_accept_tokens = getattr(spec_info, "num_accept_tokens", None) @@ -1486,6 +1509,7 @@ class FlashInferIndicesUpdaterPrefill: fixed_split_size: Optional[int] = None, multi_item_params: Optional[MultiItemScoringParams] = None, cross_attention_custom_mask: Optional[torch.Tensor] = None, + extend_prefix_lens_cpu: Optional[List[int]] = None, ): for wrapper_id in range(2): if wrapper_id == 0: diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 37e9d0a04..a50d75eed 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -340,6 +340,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): hidden_states=hidden_states, num_correct_drafts=num_correct_drafts, num_accept_tokens=num_accept_tokens, + # Padded tree width per req; drives the constant qo layout. + num_tokens_per_req=self.num_tokens_per_bs, ) forward_batch = ForwardBatch( diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 4c7eb2b6e..6fc815bc7 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -418,8 +418,14 @@ class EagleDraftExtendInput(SpecInput): ): device = req_pool_indices.device bs = self.num_correct_drafts.numel() - qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=device) - qo_indptr[1:] = torch.cumsum(self.num_accept_tokens, dim=0) + # Constant num_tokens_per_req qo layout (required for cuda-graph capture). + qo_indptr = torch.arange( + 0, + (bs + 1) * self.num_tokens_per_req, + step=self.num_tokens_per_req, + dtype=torch.int32, + device=device, + ) cum_kv_seq_len = torch.zeros((bs + 1,), dtype=torch.int32, device=device) cum_kv_seq_len[1:] = torch.cumsum(paged_kernel_lens, dim=0) diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index a97490960..463a3fe3a 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -14,6 +14,7 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner i ) from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner from sglang.srt.kv_canary.runner.canary_manager import context_tuple +from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.layers.attention.triton_backend import TritonAttnBackend from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend @@ -407,8 +408,9 @@ class EagleDraftWorker(EagleDraftWorkerBase): self.draft_attn_backend, AiterMultiStepDraftBackend ) - supports_cuda_draft_extend_graph = (_is_cuda or _is_musa) and isinstance( - self.draft_extend_attn_backend, + draft_extend_backend = self.draft_extend_attn_backend + graph_supported_backend = isinstance( + draft_extend_backend, ( TritonAttnBackend, TRTLLMMLABackend, @@ -416,6 +418,15 @@ class EagleDraftWorker(EagleDraftWorkerBase): TokenspeedMLABackend, ), ) + # FlashInfer draft-extend graph does not support a reduced draft vocab + # (speculative_token_map / FR-Spec); fall back to eager in that case. + flashinfer_graph_supported = ( + isinstance(draft_extend_backend, FlashInferAttnBackend) + and self.server_args.speculative_token_map is None + ) + supports_cuda_draft_extend_graph = (_is_cuda or _is_musa) and ( + graph_supported_backend or flashinfer_graph_supported + ) # Capture extend # TODO: support draft extend cuda graph for more attention backends if self.draft_extend_attn_backend and (