diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py index 4f2c5ae2a..60650c8b3 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py @@ -74,9 +74,21 @@ def patch_model_npu( class NPUGraphRunner(CudaGraphRunner): """A NPUGraphRunner runs the forward pass of a model with npu graph and torch.compile.""" - def __init__(self, model_runner: ModelRunner): + def __init__( + self, + model_runner: ModelRunner, + *, + attn_backend=None, + speculative_num_steps: Optional[int] = None, + speculative_num_draft_tokens: Optional[int] = None, + ): sglang.srt.model_executor.cuda_graph_runner.patch_model = patch_model_npu - super().__init__(model_runner) + super().__init__( + model_runner, + attn_backend=attn_backend, + speculative_num_steps=speculative_num_steps, + speculative_num_draft_tokens=speculative_num_draft_tokens, + ) self.update_attr_name = None self.update_attr_type = None self.model_runner = model_runner diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index ae84f8226..35b00f40d 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -9,6 +9,7 @@ import torch from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import ( EAGLEDraftNpuGraphRunner, ) +from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner from sglang.srt.layers.dp_attention import get_attention_tp_group from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.moe.utils import ( @@ -376,7 +377,8 @@ class EAGLEWorker(TpModelWorker): target_graph_runner = None if not self.server_args.disable_cuda_graph: - target_graph_runner = CudaGraphRunner( + TargetGraphRunnerCls = NPUGraphRunner if _is_npu else CudaGraphRunner + target_graph_runner = TargetGraphRunnerCls( target_model_runner, attn_backend=target_attn_backend, speculative_num_steps=speculative_num_steps, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 8f3c97e5b..13c4dc52c 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -12,6 +12,7 @@ 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.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.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.layers.attention.triton_backend import TritonAttnBackend @@ -976,7 +977,8 @@ class EAGLEWorkerV2(BaseSpecWorker): target_graph_runner = None if not self.server_args.disable_cuda_graph: - target_graph_runner = CudaGraphRunner( + TargetGraphRunnerCls = NPUGraphRunner if _is_npu else CudaGraphRunner + target_graph_runner = TargetGraphRunnerCls( target_model_runner, attn_backend=target_attn_backend, speculative_num_steps=speculative_num_steps,