From dfaf75b1a69e1dafaed89e78273f0cef1f452dba Mon Sep 17 00:00:00 2001 From: khalilzhk Date: Fri, 24 Jul 2026 21:13:21 +0800 Subject: [PATCH] [NPU] bugfix for extra device memory on Ascend (#30112) --- .../hardware_backend/npu/graph_runner/npu_cudagraph_backend.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py index a0339d8fb..105034c03 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_cudagraph_backend.py @@ -53,6 +53,7 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend): self._outputs: Dict[Any, Any] = {} self._pool = None self._device_module = cuda_graph_runner.device_module + self._device_id = self._device_module.current_device() self._tp_group = cuda_graph_runner.model_runner.tp_group self._capture_stream = None self._memory_saver_adapter: Optional[Any] = TorchMemorySaverAdapter.create( @@ -166,6 +167,7 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend): graph = self._graphs[shape_key] def _update(): + self._device_module.set_device(self._device_id) graph.update(cpu_update_input=cpu_update_input) thread = threading.Thread(target=_update)