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 ( from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import (
EAGLEDraftNpuGraphRunner, 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 ( 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.dp_attention import get_attention_tp_group
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
@@ -294,8 +294,8 @@ class EagleDraftWorker(BaseDraftWorker):
) )
supports_cuda_draft_extend_graph = _is_cuda and ( supports_cuda_draft_extend_graph = _is_cuda and (
isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend) isinstance(self.draft_extend_attn_backend, TritonAttnBackend)
or isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend) or isinstance(self.draft_extend_attn_backend, TRTLLMMLABackend)
) )
# Capture extend # Capture extend
# TODO: support draft extend cuda graph for more attention backends # TODO: support draft extend cuda graph for more attention backends