Fix draft extend cuda graph when spec_step=1 (#21709)

This commit is contained in:
Qiaolin Yu
2026-03-31 18:29:56 -07:00
committed by GitHub
parent e4c565f2f2
commit d8db3077ca
@@ -12,9 +12,9 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_r
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import (
EAGLEDraftNpuGraphRunner,
)
from sglang.srt.layers.attention.triton_backend import TritonMultiStepDraftBackend
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
from sglang.srt.layers.attention.trtllm_mla_backend import (
TRTLLMMLAMultiStepDraftBackend,
TRTLLMMLABackend,
)
from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
@@ -294,8 +294,8 @@ class EagleDraftWorker(BaseDraftWorker):
)
supports_cuda_draft_extend_graph = _is_cuda and (
isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend)
or isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend)
isinstance(self.draft_extend_attn_backend, TritonAttnBackend)
or isinstance(self.draft_extend_attn_backend, TRTLLMMLABackend)
)
# Capture extend
# TODO: support draft extend cuda graph for more attention backends