From 6388b6cfb1d93c253714a408f1d66a093302acd7 Mon Sep 17 00:00:00 2001 From: Kurkur <102506892+litmei@users.noreply.github.com> Date: Mon, 14 Sep 2026 09:32:24 +0800 Subject: [PATCH] [feat] Add base NpuSRTPlatform implementation (#36472) --- python/sglang/srt/platforms/__init__.py | 8 ++ python/sglang/srt/platforms/npu.py | 89 ++++++++++++ .../unit/platforms/test_platform_interface.py | 129 ++++++++++++++++++ 3 files changed, 226 insertions(+) create mode 100644 python/sglang/srt/platforms/npu.py diff --git a/python/sglang/srt/platforms/__init__.py b/python/sglang/srt/platforms/__init__.py index cd740f253..b5b411ff5 100644 --- a/python/sglang/srt/platforms/__init__.py +++ b/python/sglang/srt/platforms/__init__.py @@ -21,6 +21,7 @@ from sglang.srt.environ import envs from sglang.srt.platforms.cpu import CpuSRTPlatform from sglang.srt.platforms.cuda import CudaSRTPlatform from sglang.srt.platforms.interface import SRTPlatform +from sglang.srt.platforms.npu import NPUSRTPlatform from sglang.srt.platforms.rocm import RocmSRTPlatform from sglang.srt.platforms.xpu import XpuSRTPlatform from sglang.srt.plugins import PLATFORM_PLUGINS_GROUP, load_plugins_by_group @@ -42,6 +43,10 @@ def _is_cpu_available() -> bool: return os.getenv("SGLANG_USE_CPU_ENGINE", "0") == "1" +def _is_npu_available() -> bool: + return hasattr(torch, "npu") and torch.npu.is_available() + + def _is_xpu_available() -> bool: return torch.xpu.is_available() @@ -127,6 +132,9 @@ def _resolve_platform() -> SRTPlatform: "No platform plugin detected. Using CUDA SRTPlatform defaults." ) return CudaSRTPlatform() + if _is_npu_available(): + logger.debug("No platform plugin detected. Using NPU SRTPlatform defaults.") + return NPUSRTPlatform() if _is_rocm_available(): logger.debug( "No platform plugin detected. Using ROCm SRTPlatform defaults." diff --git a/python/sglang/srt/platforms/npu.py b/python/sglang/srt/platforms/npu.py new file mode 100644 index 000000000..c1382ff76 --- /dev/null +++ b/python/sglang/srt/platforms/npu.py @@ -0,0 +1,89 @@ +"""NPU device operations for the SRT platform layer.""" + +from typing import Optional + +import torch + +from sglang.srt.platforms.device_mixin import ( + DeviceCapability, + DeviceMixin, + PlatformEnum, +) +from sglang.srt.platforms.interface import SRTPlatform + + +class NPUDeviceMixin(DeviceMixin): + """NPU implementation of the shared device operations.""" + + _enum: PlatformEnum = PlatformEnum.NPU + device_name: str = "npu" + device_type: str = "npu" + + def get_device_total_memory(self, device_id: int = 0) -> int: + return int(torch.npu.get_device_properties(device_id).total_memory) + + def get_current_memory_usage( + self, device: Optional["torch.device"] = None + ) -> float: + return float(torch.npu.max_memory_allocated(device)) + + def get_device(self, local_rank: int) -> "torch.device": + return torch.device("npu", local_rank) + + def set_device(self, device: "torch.device") -> None: + torch.npu.set_device(device) + + def get_device_name(self, device_id: int = 0) -> str: + return str(torch.npu.get_device_name(device_id)) + + def get_device_uuid(self, device_id: int = 0) -> str: + return str(torch.npu.get_device_properties(device_id).uuid) + + def get_device_capability(self, device_id: int = 0) -> DeviceCapability: + # The return value of torch_npu.npu.get_device_capability() is configured + # via the environment variable TORCH_NPU_DEVICE_CAPABILITY, which is only + # used for compatibility with native PyTorch and does not represent the + # actual capabilities of the NPU hardware + return DeviceCapability(0, 0) + + def empty_cache(self) -> None: + torch.npu.empty_cache() + + def synchronize(self) -> None: + torch.npu.synchronize() + + def get_available_memory(self, device_id: int = 0) -> tuple[int, int]: + return torch.npu.mem_get_info(device_id) + + def is_pin_memory_available(self, device=None) -> bool: + if device is not None and str(device) == "cpu": + return False + return True + + @classmethod + def seed_everything(cls, seed: int | None = None) -> None: + if seed is not None: + super().seed_everything(seed) + if hasattr(torch, "npu"): + torch.npu.manual_seed_all(seed) + + +class NPUSRTPlatform(NPUDeviceMixin, SRTPlatform): + """Default in-tree NPU SRT platform.""" + + def get_default_attention_backend(self) -> str: + return "ascend" + + def get_dispatch_key_name(self) -> str: + return "npu" + + def supports_fp8(self) -> bool: + # NPU quantization backends in hardware_backend/npu/quantization + return True + + def support_cuda_graph(self) -> bool: + # NPUGraphRunner in hardware_backend/npu/graph_runner + return True + + def support_piecewise_cuda_graph(self) -> bool: + return False diff --git a/test/registered/unit/platforms/test_platform_interface.py b/test/registered/unit/platforms/test_platform_interface.py index 35b1458df..5614dd7f0 100644 --- a/test/registered/unit/platforms/test_platform_interface.py +++ b/test/registered/unit/platforms/test_platform_interface.py @@ -19,6 +19,7 @@ from sglang.srt.platforms.device_mixin import ( PlatformEnum, ) from sglang.srt.platforms.interface import SRTPlatform +from sglang.srt.platforms.npu import NPUSRTPlatform from sglang.srt.platforms.rocm import RocmSRTPlatform from sglang.srt.platforms.xpu import XpuSRTPlatform from sglang.test.ci.ci_register import register_cpu_ci @@ -308,6 +309,134 @@ class TestXpuDeviceMixin(CustomTestCase): self.assertTrue(base.support_piecewise_cuda_graph()) +class TestNpuDeviceMixin(CustomTestCase): + """Tests for NPU device operation defaults.""" + + def setUp(self): + # torch.device("npu", ...) requires the "npu" device type, which + # torch_npu registers via the privateuse1 backend rename; CPU-only + # builds lack it. Register it per-test so only this suite carries + # the process-wide side effect. + try: + torch.utils.rename_privateuse1_backend("npu") + except Exception: + # Re-registration with a different name raises on some versions; + # real NPU machines may have already renamed the backend. + pass + super().setUp() + + def test_default_get_device_returns_npu_device(self): + base = NPUSRTPlatform() + self.assertEqual(base.get_device(2), torch.device("npu", 2)) + + def test_default_get_device_capability_reports_zero(self): + # torch_npu's get_device_capability is configured via the environment + # variable TORCH_NPU_DEVICE_CAPABILITY purely for native-PyTorch + # compatibility; it does not reflect the real NPU hardware. The + # platform therefore reports (0, 0) without consulting torch.npu. + base = NPUSRTPlatform() + mock_npu = MagicMock() + with patch.object(torch, "npu", mock_npu, create=True): + self.assertEqual(base.get_device_capability(1), DeviceCapability(0, 0)) + mock_npu.get_device_capability.assert_not_called() + + def test_memory_queries_delegate_to_torch_npu(self): + base = NPUSRTPlatform() + mock_npu = MagicMock() + mock_npu.get_device_properties.return_value.total_memory = 32 * 1024**3 + mock_npu.max_memory_allocated.return_value = 5 * 10**8 + mock_npu.mem_get_info.return_value = (10**9, 2 * 10**9) + with patch.object(torch, "npu", mock_npu, create=True): + self.assertEqual(base.get_device_total_memory(1), 32 * 1024**3) + mock_npu.get_device_properties.assert_called_once_with(1) + self.assertEqual(base.get_current_memory_usage(), 5 * 10**8) + mock_npu.max_memory_allocated.assert_called_once_with(None) + device = torch.device("npu", 0) + base.get_current_memory_usage(device) + mock_npu.max_memory_allocated.assert_called_with(device) + self.assertEqual(base.get_available_memory(2), (10**9, 2 * 10**9)) + mock_npu.mem_get_info.assert_called_once_with(2) + + def test_device_info_queries_delegate_to_torch_npu(self): + base = NPUSRTPlatform() + mock_npu = MagicMock() + mock_npu.get_device_name.return_value = "Ascend910B4" + mock_npu.get_device_properties.return_value.uuid = "npu-uuid-0" + with patch.object(torch, "npu", mock_npu, create=True): + self.assertEqual(base.get_device_name(1), "Ascend910B4") + mock_npu.get_device_name.assert_called_once_with(1) + self.assertEqual(base.get_device_uuid(1), "npu-uuid-0") + mock_npu.get_device_properties.assert_called_once_with(1) + + def test_device_state_ops_delegate_to_torch_npu(self): + base = NPUSRTPlatform() + mock_npu = MagicMock() + with patch.object(torch, "npu", mock_npu, create=True): + device = torch.device("npu", 3) + base.set_device(device) + mock_npu.set_device.assert_called_once_with(device) + base.empty_cache() + mock_npu.empty_cache.assert_called_once() + base.synchronize() + mock_npu.synchronize.assert_called_once() + + def test_pin_memory_available_for_npu_targets(self): + base = NPUSRTPlatform() + self.assertTrue(base.is_pin_memory_available()) + self.assertTrue(base.is_pin_memory_available(device="npu")) + self.assertTrue(base.is_pin_memory_available(device=torch.device("npu", 0))) + self.assertFalse(base.is_pin_memory_available(device="cpu")) + self.assertFalse(base.is_pin_memory_available(device=torch.device("cpu"))) + + def test_default_seed_everything_seeds_npu(self): + mock_npu = MagicMock() + with ( + patch.object(torch, "npu", mock_npu, create=True), + patch("torch.manual_seed") as mock_torch_seed, + patch("sglang.srt.platforms.device_mixin.np.random.seed") as mock_np_seed, + patch("sglang.srt.platforms.device_mixin.random.seed") as mock_random_seed, + ): + NPUSRTPlatform.seed_everything(123) + mock_random_seed.assert_called_once_with(123) + mock_np_seed.assert_called_once_with(123) + mock_torch_seed.assert_called_once_with(123) + mock_npu.manual_seed_all.assert_called_once_with(123) + + def test_seed_everything_none_seed_is_noop(self): + mock_npu = MagicMock() + with ( + patch.object(torch, "npu", mock_npu, create=True), + patch("torch.manual_seed") as mock_torch_seed, + patch("sglang.srt.platforms.device_mixin.np.random.seed") as mock_np_seed, + patch("sglang.srt.platforms.device_mixin.random.seed") as mock_random_seed, + ): + NPUSRTPlatform.seed_everything(None) + mock_random_seed.assert_not_called() + mock_np_seed.assert_not_called() + mock_torch_seed.assert_not_called() + mock_npu.manual_seed_all.assert_not_called() + + def test_npu_srt_platform_identity(self): + base = NPUSRTPlatform() + self.assertTrue(base.is_npu()) + self.assertFalse(base.is_cuda()) + self.assertFalse(base.is_cuda_alike()) + self.assertEqual(base.device_name, "npu") + self.assertEqual(base.device_type, "npu") + + def test_get_default_attention_backend_is_ascend(self): + self.assertEqual(NPUSRTPlatform().get_default_attention_backend(), "ascend") + + def test_get_dispatch_key_name_is_npu(self): + self.assertEqual(NPUSRTPlatform().get_dispatch_key_name(), "npu") + + def test_npu_srt_platform_capabilities(self): + base = NPUSRTPlatform() + self.assertTrue(base.supports_fp8()) + self.assertTrue(base.support_cuda_graph()) + self.assertFalse(base.support_piecewise_cuda_graph()) + + class TestCpuDeviceMixin(CustomTestCase): """Tests for CPU device operation defaults (covers both x86 and ARM)."""