[Spec] Support FlashInfer CUDA graph for EAGLE draft-extend (#28782)

This commit is contained in:
Liangsheng Yin
2026-06-21 14:46:25 -07:00
committed by GitHub
parent a4d0ff3def
commit 8e890391f5
4 changed files with 50 additions and 7 deletions
@@ -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,12 +1343,15 @@ 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
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
@@ -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:
@@ -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(
+8 -2
View File
@@ -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)
@@ -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 (