[Spec] Support FlashInfer CUDA graph for EAGLE draft-extend (#28782)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 (
|
||||
|
||||
Reference in New Issue
Block a user