[diffusion] platform: support WAN/FLUX/Qwen-Image/Qwen-Image-edit on Ascend (#13662)
Co-authored-by: dhx98 <haox.dai@gmail.com> Co-authored-by: DHX98 <haoxiand@andrew.cmu.edu> Co-authored-by: ronnie_zheng <zl19940307@163.com> Co-authored-by: DHX98 <DHX98@noreply.gitcode.com> Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com>
This commit is contained in:
co-authored by
dhx98
DHX98
ronnie_zheng
DHX98
Yuhao Yang
parent
7b83659310
commit
00248d85c7
@@ -101,6 +101,24 @@ def rocm_platform_plugin() -> str | None:
|
||||
)
|
||||
|
||||
|
||||
def npu_platform_plugin() -> str | None:
|
||||
is_npu = False
|
||||
|
||||
try:
|
||||
import torch
|
||||
|
||||
if torch.npu.is_available():
|
||||
is_npu = True
|
||||
logger.info("NPU is available")
|
||||
except Exception as e:
|
||||
logger.info("NPU detection failed: %s", e)
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.platforms.npu.NPUPlatformBase"
|
||||
if is_npu
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
def musa_platform_plugin() -> str | None:
|
||||
is_musa = False
|
||||
|
||||
@@ -125,6 +143,7 @@ builtin_platform_plugins = {
|
||||
"rocm": rocm_platform_plugin,
|
||||
"mps": mps_platform_plugin,
|
||||
"cpu": cpu_platform_plugin,
|
||||
"npu": npu_platform_plugin,
|
||||
"musa": musa_platform_plugin,
|
||||
}
|
||||
|
||||
@@ -148,6 +167,11 @@ def resolve_current_platform_cls_qualname() -> str:
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to NPU
|
||||
platform_cls_qualname = npu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to MUSA
|
||||
platform_cls_qualname = musa_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
|
||||
@@ -15,6 +15,7 @@ import psutil
|
||||
import torch
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
from sglang.multimodal_gen import envs
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
AttentionBackendEnum,
|
||||
DeviceCapability,
|
||||
@@ -74,6 +75,10 @@ class CudaPlatformBase(Platform):
|
||||
dispatch_key: str = "CUDA"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_local_torch_device(cls) -> torch.device:
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -47,6 +47,7 @@ class PlatformEnum(enum.Enum):
|
||||
TPU = enum.auto()
|
||||
CPU = enum.auto()
|
||||
MPS = enum.auto()
|
||||
NPU = enum.auto()
|
||||
MUSA = enum.auto()
|
||||
OOT = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
@@ -99,6 +100,10 @@ class Platform:
|
||||
def is_cuda(self) -> bool:
|
||||
return self.is_cuda_static()
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_npu(self) -> bool:
|
||||
return self._enum == PlatformEnum.NPU
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_rocm(self) -> bool:
|
||||
return self.is_rocm_static()
|
||||
@@ -175,6 +180,15 @@ class Platform:
|
||||
def is_hip(self) -> bool:
|
||||
return self.is_rocm()
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=1)
|
||||
def is_amp_supported(cls) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_local_torch_device(cls) -> torch.device:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls_str(
|
||||
cls,
|
||||
@@ -236,6 +250,8 @@ class Platform:
|
||||
def get_device(self, local_rank: int) -> torch.device:
|
||||
if self.is_cuda() or self.is_rocm():
|
||||
return torch.device("cuda", local_rank)
|
||||
elif self.is_npu():
|
||||
return torch.device("npu", local_rank)
|
||||
elif self.is_musa():
|
||||
return torch.device("musa", local_rank)
|
||||
elif self.is_mps():
|
||||
@@ -247,6 +263,8 @@ class Platform:
|
||||
def get_torch_distributed_backend_str(self) -> str:
|
||||
if self.is_cuda_alike():
|
||||
return "nccl"
|
||||
elif self.is_npu():
|
||||
return "hccl"
|
||||
elif self.is_musa():
|
||||
return "mccl"
|
||||
elif self.is_mps():
|
||||
|
||||
@@ -26,6 +26,15 @@ class MpsPlatform(Platform):
|
||||
dispatch_key: str = "MPS"
|
||||
device_control_env_var: str = "MPS_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=1)
|
||||
def is_amp_supported(cls) -> bool:
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def get_local_torch_device(cls) -> torch.device:
|
||||
return torch.device("mps")
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm-ascend: https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/platform.py
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen import envs
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
AttentionBackendEnum,
|
||||
DeviceCapability,
|
||||
Platform,
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def device_id_to_physical_device_id(device_id: int) -> int:
|
||||
if "ASCEND_RT_VISIBLE_DEVICES" in os.environ:
|
||||
device_ids = os.environ["ASCEND_RT_VISIBLE_DEVICES"].split(",")
|
||||
if device_ids == [""]:
|
||||
msg = (
|
||||
"ASCEND_RT_VISIBLE_DEVICES is set to empty string, which means"
|
||||
" NPU support is disabled"
|
||||
)
|
||||
raise RuntimeError(msg)
|
||||
physical_device_id = device_ids[device_id]
|
||||
return int(physical_device_id)
|
||||
else:
|
||||
return device_id
|
||||
|
||||
|
||||
class NPUPlatformBase(Platform):
|
||||
_enum = PlatformEnum.NPU
|
||||
device_name: str = "npu"
|
||||
device_type: str = "npu"
|
||||
dispatch_key: str = "NPU"
|
||||
device_control_env_var: str = "ASCEND_RT_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_local_torch_device(cls) -> torch.device:
|
||||
return torch.device(f"npu:{envs.LOCAL_RANK}")
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return str(torch.npu.get_device_name(device_id))
|
||||
|
||||
@classmethod
|
||||
def get_device_total_memory(cls, device_id: int = 0) -> int:
|
||||
device_props = torch.npu.get_device_properties(device_id)
|
||||
return int(device_props.total_memory)
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
if enforce_eager:
|
||||
logger.warning(
|
||||
"To see benefits of async output processing, enable NPU "
|
||||
"graph. Since, enforce-eager is enabled, async output "
|
||||
"processor cannot be used"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool:
|
||||
logger.exception(
|
||||
"NVLink detection not possible, as context support was"
|
||||
" not found. Assuming no NVLink available."
|
||||
)
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def get_available_gpu_memory(
|
||||
cls,
|
||||
device_id: int = 0,
|
||||
distributed: bool = False,
|
||||
empty_cache: bool = True,
|
||||
cpu_group: Any = None,
|
||||
) -> float:
|
||||
if empty_cache:
|
||||
torch.npu.empty_cache()
|
||||
|
||||
free_gpu_memory, _ = torch.npu.mem_get_info(device_id)
|
||||
|
||||
if distributed:
|
||||
import torch.distributed as dist
|
||||
|
||||
tensor = torch.tensor(free_gpu_memory, dtype=torch.float32, device="npu")
|
||||
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=cpu_group)
|
||||
free_gpu_memory = float(tensor.item())
|
||||
|
||||
return free_gpu_memory / (1 << 30)
|
||||
|
||||
@classmethod
|
||||
def log_warnings(cls) -> None:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(
|
||||
cls, device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
torch.npu.reset_peak_memory_stats(device)
|
||||
return float(torch.npu.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls_str(
|
||||
cls,
|
||||
selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
) -> str:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return (
|
||||
"sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
|
||||
@@ -11,6 +11,7 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.multimodal_gen.envs as envs
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
AttentionBackendEnum,
|
||||
DeviceCapability,
|
||||
@@ -30,6 +31,10 @@ class RocmPlatform(Platform):
|
||||
dispatch_key: str = "CUDA"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
def get_local_torch_device(cls) -> torch.device:
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability:
|
||||
major, minor = torch.cuda.get_device_capability(device_id)
|
||||
|
||||
Reference in New Issue
Block a user