[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,
|
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
|
||||||
spec_info=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:
|
else:
|
||||||
raise ValueError("Invalid forward mode")
|
raise ValueError("Invalid forward mode")
|
||||||
|
|
||||||
@@ -678,6 +690,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
fixed_split_size=self.prefill_split_tile_size,
|
fixed_split_size=self.prefill_split_tile_size,
|
||||||
multi_item_params=multi_item_params,
|
multi_item_params=multi_item_params,
|
||||||
cross_attention_custom_mask=forward_batch.cross_attention_custom_mask,
|
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.forward_metadata = PrefillMetadata(
|
||||||
self.prefill_wrappers_paged,
|
self.prefill_wrappers_paged,
|
||||||
@@ -797,6 +810,11 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
self.forward_metadata = PrefillMetadata(
|
self.forward_metadata = PrefillMetadata(
|
||||||
prefill_wrappers, forward_mode.is_dllm_extend(), False
|
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:
|
else:
|
||||||
raise ValueError(f"Invalid mode: {forward_mode=}")
|
raise ValueError(f"Invalid mode: {forward_mode=}")
|
||||||
|
|
||||||
@@ -1306,6 +1324,7 @@ class FlashInferIndicesUpdaterPrefill:
|
|||||||
fixed_split_size: Optional[int] = None,
|
fixed_split_size: Optional[int] = None,
|
||||||
multi_item_params: Optional[MultiItemScoringParams] = None,
|
multi_item_params: Optional[MultiItemScoringParams] = None,
|
||||||
cross_attention_custom_mask: Optional[torch.Tensor] = 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.
|
# Keep the signature for type checking. It will be assigned during runtime.
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
@@ -1324,13 +1343,16 @@ class FlashInferIndicesUpdaterPrefill:
|
|||||||
fixed_split_size: Optional[int] = None,
|
fixed_split_size: Optional[int] = None,
|
||||||
multi_item_params: Optional[MultiItemScoringParams] = None,
|
multi_item_params: Optional[MultiItemScoringParams] = None,
|
||||||
cross_attention_custom_mask: Optional[torch.Tensor] = None,
|
cross_attention_custom_mask: Optional[torch.Tensor] = None,
|
||||||
|
extend_prefix_lens_cpu: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
if use_ragged:
|
if use_ragged:
|
||||||
assert prefix_lens is not None
|
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 = 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:
|
else:
|
||||||
paged_kernel_lens = seq_lens
|
paged_kernel_lens = seq_lens
|
||||||
paged_kernel_lens_sum = seq_lens_sum
|
paged_kernel_lens_sum = seq_lens_sum
|
||||||
@@ -1366,6 +1388,7 @@ class FlashInferIndicesUpdaterPrefill:
|
|||||||
fixed_split_size: Optional[int] = None,
|
fixed_split_size: Optional[int] = None,
|
||||||
multi_item_params: Optional[MultiItemScoringParams] = None,
|
multi_item_params: Optional[MultiItemScoringParams] = None,
|
||||||
cross_attention_custom_mask: Optional[torch.Tensor] = None,
|
cross_attention_custom_mask: Optional[torch.Tensor] = None,
|
||||||
|
extend_prefix_lens_cpu: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
if prefix_lens is None:
|
if prefix_lens is None:
|
||||||
num_accept_tokens = getattr(spec_info, "num_accept_tokens", None)
|
num_accept_tokens = getattr(spec_info, "num_accept_tokens", None)
|
||||||
@@ -1486,6 +1509,7 @@ class FlashInferIndicesUpdaterPrefill:
|
|||||||
fixed_split_size: Optional[int] = None,
|
fixed_split_size: Optional[int] = None,
|
||||||
multi_item_params: Optional[MultiItemScoringParams] = None,
|
multi_item_params: Optional[MultiItemScoringParams] = None,
|
||||||
cross_attention_custom_mask: Optional[torch.Tensor] = None,
|
cross_attention_custom_mask: Optional[torch.Tensor] = None,
|
||||||
|
extend_prefix_lens_cpu: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
for wrapper_id in range(2):
|
for wrapper_id in range(2):
|
||||||
if wrapper_id == 0:
|
if wrapper_id == 0:
|
||||||
|
|||||||
@@ -340,6 +340,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
num_correct_drafts=num_correct_drafts,
|
num_correct_drafts=num_correct_drafts,
|
||||||
num_accept_tokens=num_accept_tokens,
|
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(
|
forward_batch = ForwardBatch(
|
||||||
|
|||||||
@@ -418,8 +418,14 @@ class EagleDraftExtendInput(SpecInput):
|
|||||||
):
|
):
|
||||||
device = req_pool_indices.device
|
device = req_pool_indices.device
|
||||||
bs = self.num_correct_drafts.numel()
|
bs = self.num_correct_drafts.numel()
|
||||||
qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
# Constant num_tokens_per_req qo layout (required for cuda-graph capture).
|
||||||
qo_indptr[1:] = torch.cumsum(self.num_accept_tokens, dim=0)
|
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 = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
||||||
cum_kv_seq_len[1:] = torch.cumsum(paged_kernel_lens, dim=0)
|
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.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
||||||
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
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.tokenspeed_mla_backend import TokenspeedMLABackend
|
||||||
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
|
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
|
||||||
from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend
|
from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend
|
||||||
@@ -407,8 +408,9 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.draft_attn_backend, AiterMultiStepDraftBackend
|
self.draft_attn_backend, AiterMultiStepDraftBackend
|
||||||
)
|
)
|
||||||
|
|
||||||
supports_cuda_draft_extend_graph = (_is_cuda or _is_musa) and isinstance(
|
draft_extend_backend = self.draft_extend_attn_backend
|
||||||
self.draft_extend_attn_backend,
|
graph_supported_backend = isinstance(
|
||||||
|
draft_extend_backend,
|
||||||
(
|
(
|
||||||
TritonAttnBackend,
|
TritonAttnBackend,
|
||||||
TRTLLMMLABackend,
|
TRTLLMMLABackend,
|
||||||
@@ -416,6 +418,15 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
TokenspeedMLABackend,
|
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
|
# Capture extend
|
||||||
# TODO: support draft extend cuda graph for more attention backends
|
# TODO: support draft extend cuda graph for more attention backends
|
||||||
if self.draft_extend_attn_backend and (
|
if self.draft_extend_attn_backend and (
|
||||||
|
|||||||
Reference in New Issue
Block a user