[feat] Add base NpuSRTPlatform implementation (#36472)
This commit is contained in:
@@ -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."
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user