[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:
Makcum888e
2026-02-08 10:45:30 +08:00
committed by GitHub
co-authored by dhx98 DHX98 ronnie_zheng DHX98 Yuhao Yang
parent 7b83659310
commit 00248d85c7
25 changed files with 476 additions and 30 deletions
@@ -16,7 +16,6 @@ import torch.distributed
from torch.cuda import synchronize
from torch.distributed import Backend, ProcessGroup
from sglang.multimodal_gen import envs
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import (
DeviceCommunicatorBase,
)
@@ -46,11 +45,7 @@ _group_name_counter: dict[str, int] = {}
def get_local_torch_device() -> torch.device:
"""Return the torch device for the current rank."""
return (
torch.device(f"cuda:{envs.LOCAL_RANK}")
if current_platform.is_cuda_alike()
else torch.device("mps")
)
return current_platform.get_local_torch_device()
def _get_unique_name(name: str) -> str:
@@ -190,8 +185,6 @@ class GroupCoordinator:
# TODO: fix it for other platforms
self.device = get_local_torch_device()
from sglang.multimodal_gen.runtime.platforms import current_platform
self.use_device_communicator = use_device_communicator
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
@@ -287,9 +280,6 @@ class GroupCoordinator:
@contextmanager
def graph_capture(self, graph_capture_context: GraphCaptureContext | None = None):
# Platform-aware graph capture
from sglang.multimodal_gen.runtime.platforms import current_platform
if current_platform.is_cuda_alike():
if graph_capture_context is None:
stream = torch.cuda.Stream()
@@ -248,7 +248,11 @@ def init_distributed_environment(
# For MPS and MUSA, don't pass device_id as it doesn't support device indices
extra_args = (
{}
if (current_platform.is_mps() or current_platform.is_musa())
if (
current_platform.is_mps()
or current_platform.is_musa()
or current_platform.is_npu()
)
else dict(device_id=device_id)
)
@@ -618,6 +622,7 @@ def maybe_init_distributed_environment_and_model_parallel(
local_rank=local_rank,
distributed_init_method=distributed_init_method,
device_id=device,
backend=current_platform.get_torch_distributed_backend_str(),
timeout=dist_timeout,
)
initialize_model_parallel(
@@ -14,8 +14,12 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
_is_cuda = current_platform.is_cuda()
_is_hip = current_platform.is_hip()
_is_npu = current_platform.is_npu()
if _is_cuda or _is_hip:
from sgl_kernel import silu_and_mul
if _is_npu:
import torch_npu
# TODO (will): remove this dependency
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
@@ -46,6 +50,10 @@ class SiluAndMul(CustomOp):
d = x.shape[-1] // 2
return F.silu(x[..., :d]) * x[..., d:]
def forward_npu(self, x: torch.Tensor) -> torch.Tensor:
out = torch_npu.npu_swiglu(x)
return out
@CustomOp.register("gelu_and_mul")
class GeluAndMul(CustomOp):
@@ -64,6 +64,11 @@ class CustomOp(nn.Module):
# PyTorch-native implementation.
return self.forward_native(*args, **kwargs)
def forward_npu(self, *args, **kwargs) -> Any:
# By default, we assume that NPU ops are compatible with the
# PyTorch-native implementation.
return self.forward_native(*args, **kwargs)
def dispatch_forward(self) -> Callable:
if _is_cuda:
return self.forward_cuda
@@ -12,9 +12,13 @@ import torch.nn.functional as F
from sglang.multimodal_gen.runtime.platforms import current_platform
_is_cuda = current_platform.is_cuda()
_is_npu = current_platform.is_npu()
if _is_cuda:
from sgl_kernel import fused_add_rmsnorm, rmsnorm
if _is_npu:
import torch_npu
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_tensor_model_parallel_rank,
@@ -28,11 +32,8 @@ from sglang.multimodal_gen.runtime.layers.triton_ops import (
rms_norm_fn,
triton_one_pass_rms_norm,
)
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
_is_cuda = current_platform.is_cuda()
# Copied and adapted from sglang
@CustomOp.register("rms_norm")
@@ -141,6 +142,18 @@ class RMSNorm(CustomOp):
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
return self.forward_native(x, residual)
def forward_npu(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
if residual is not None:
out, _, residual_out = torch_npu.npu_add_rms_norm(
residual, x, self.weight.data, self.variance_epsilon
)
return out, residual_out
return torch_npu.npu_rms_norm(x, self.weight.data, self.variance_epsilon)[0]
def forward_hip(
self,
x: torch.Tensor,
@@ -214,7 +227,7 @@ class LayerNorm(CustomOp):
x = x.view(-1, self.hidden_size)
return self.forward_triton(x).view(shape)
@torch.compile(backend="inductor")
@torch.compile(backend="inductor", disable=current_platform.is_npu())
def forward_native(
self,
x: torch.Tensor,
@@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.models.parameter import (
# yapf: enable
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@@ -152,7 +153,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
) -> torch.Tensor:
output = (
F.linear(x, layer.weight, bias)
if torch.cuda.is_available() or bias is None
if current_platform.is_amp_supported() or bias is None
else F.linear(x, layer.weight, bias.to(x.dtype))
) # NOTE: this line assumes that we are using amp when using cuda and is needed to account for the fact that amp isn't supported in mps
return output
@@ -8,6 +8,8 @@ import triton # type: ignore
import triton.language as tl # type: ignore
from torch import Tensor
from sglang.multimodal_gen.runtime.platforms import current_platform
@triton.autotune(
configs=[
@@ -524,8 +526,14 @@ def triton_autotune_configs():
max_threads_per_block = 1024
# Default to warp size 32 if not defined by device
warp_size = getattr(
torch.cuda.get_device_properties(torch.cuda.current_device()), "warp_size", 32
torch.get_device_module().get_device_properties(
torch.get_device_module().current_device()
),
"warp_size",
32,
)
if warp_size is None:
warp_size = 32
# Autotune for warp counts which are powers of 2 and do not exceed thread per block limit
return [
triton.Config({}, num_warps=warp_count)
@@ -820,7 +828,7 @@ def _layer_norm_fwd_impl(
BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N))
if N > BLOCK_N:
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
with torch.cuda.device(x.device.index):
with torch.get_device_module().device(x.device.index):
torch.library.wrap_triton(_layer_norm_fwd_1pass_kernel)[(M,)](
x,
out,
@@ -1166,3 +1174,31 @@ def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6
BLOCK_SIZE_SEQ=BLOCK_SIZE_SEQ,
)
return y
if current_platform.is_npu():
# TODO: remove this when triton ascend bug is fixed
def fuse_scale_shift_native(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
block_l: int = 128,
block_c: int = 128,
):
return x * (1 + scale) + shift
fuse_scale_shift_kernel = fuse_scale_shift_native
# TODO: remove this when triton ascend bug is fixed
def apply_rotary_embedding_native(
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
) -> torch.Tensor:
cos = cos.unsqueeze(-2).to(x.dtype)
sin = sin.unsqueeze(-2).to(x.dtype)
x1 = x[..., ::2]
x2 = x[..., 1::2]
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
return torch.stack((o1, o2), dim=-1).flatten(-2)
apply_rotary_embedding = apply_rotary_embedding_native
@@ -145,7 +145,11 @@ class VocabParallelEmbeddingShardIndices:
assert self.num_added_elements <= self.num_added_elements_padded
@torch.compile(dynamic=True, backend=current_platform.simple_compile_backend)
@torch.compile(
dynamic=True,
backend=current_platform.simple_compile_backend,
disable=current_platform.is_npu(),
)
def get_masked_input_and_mask(
input_: torch.Tensor,
org_vocab_start_index: int,
@@ -71,7 +71,7 @@ class GPUWorker:
def init_device_and_model(self) -> None:
"""Initialize the device and load the model."""
setproctitle(f"sgl_diffusion::scheduler_TP{self.local_rank}")
torch.cuda.set_device(self.local_rank)
torch.get_device_module().set_device(self.local_rank)
# Set environment variables for distributed initialization
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = str(self.master_port)
@@ -86,6 +86,7 @@ class GPUWorker:
ring_degree=self.server_args.ring_degree,
sp_size=self.server_args.sp_degree,
dp_size=self.server_args.dp_size,
distributed_init_method=f"tcp://127.0.0.1:{self.master_port}",
dist_timeout=self.server_args.dist_timeout,
)
@@ -160,7 +161,7 @@ class GPUWorker:
output_batch = None
try:
if self.rank == 0:
torch.cuda.reset_peak_memory_stats()
torch.get_device_module().reset_peak_memory_stats()
start_time = time.monotonic()
@@ -347,7 +348,8 @@ def run_scheduler_process(
"""
configure_logger(server_args)
globally_suppress_loggers()
set_cuda_arch()
if current_platform.is_cuda():
set_cuda_arch()
port_args = PortArgs.from_server_args(server_args)
@@ -854,7 +854,7 @@ class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
encoder_hidden_states = (
encoder_hidden_states.to(orig_dtype)
if current_platform.is_mps()
if not current_platform.is_amp_supported()
else encoder_hidden_states
) # cast to orig_dtype for MPS
@@ -264,7 +264,7 @@ class CLIPAttention(nn.Module):
key_states,
value_states,
attn_mask=attn_mask,
is_causal=True,
is_causal=attention_mask is None,
scale=self.scale,
)
attn_output = attn_output.transpose(1, 2)
@@ -1227,10 +1227,9 @@ class DenoisingStage(PipelineStage):
raw_latent_shape=batch.raw_latent_shape
)
else:
# attn_metadata can be None for SDPA attention backend
return None
assert attn_metadata is not None, "attn_metadata cannot be None"
return attn_metadata
def _predict_noise(
@@ -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)