diff --git a/.github/workflows/pr-test-npu.yml b/.github/workflows/pr-test-npu.yml index 0361d078f..703680d00 100644 --- a/.github/workflows/pr-test-npu.yml +++ b/.github/workflows/pr-test-npu.yml @@ -64,7 +64,7 @@ jobs: multimodal_gen: - "python/sglang/multimodal_gen/**" - "python/pyproject_npu.toml" - - "scripts/ci/npu_ci_install_dependency.sh" + - "scripts/ci/npu/npu_ci_install_dependency.sh" - ".github/workflows/pr-test-npu.yml" # ==================== PR Gate ==================== # @@ -241,3 +241,42 @@ jobs: run: | cd test/srt python3 run_suite.py --suite per-commit-16-npu-a3 --timeout-per-file 3600 --auto-partition-id ${{ matrix.part }} --auto-partition-size 2 + + multimodal-gen-test-1-npu-a3: + needs: [check-changes, pr-gate] + if: needs.check-changes.outputs.multimodal_gen == 'true' + runs-on: linux-aarch64-a3-16 + container: + image: swr.cn-southwest-2.myhuaweicloud.com/base_image/ascend-ci/cann:8.3.rc2-a3-ubuntu22.04-py3.11 + steps: + - name: Checkout code + uses: actions/checkout@v4 + + - name: Install dependencies + run: | + # speed up by using infra cache services + CACHING_URL="cache-service.nginx-pypi-cache.svc.cluster.local" + sed -Ei "s@(ports|archive).ubuntu.com@${CACHING_URL}:8081@g" /etc/apt/sources.list + pip config set global.index-url http://${CACHING_URL}/pypi/simple + pip config set global.extra-index-url "https://pypi.tuna.tsinghua.edu.cn/simple" + pip config set global.trusted-host "${CACHING_URL} pypi.tuna.tsinghua.edu.cn" + + bash scripts/ci/npu/npu_ci_install_dependency.sh a3 + # copy required file from our daily cache + cp ~/.cache/modelscope/hub/datasets/otavia/ShareGPT_Vicuna_unfiltered/ShareGPT_V3_unfiltered_cleaned_split.json /tmp + # copy download through proxy + curl -o /tmp/test.jsonl -L https://gh-proxy.test.osinfra.cn/https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl + + - name: Run test + timeout-minutes: 60 + env: + SGLANG_USE_MODELSCOPE: true + SGLANG_IS_IN_CI: true + HF_ENDPOINT: https://hf-mirror.com + TORCH_EXTENSIONS_DIR: /tmp/torch_extensions + PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True" + STREAMS_PER_DEVICE: 32 + run: | + export PATH="/usr/local/Ascend/8.3.RC1/compiler/bishengir/bin:${PATH}" + cd python + python3 sglang/multimodal_gen/test/run_suite.py --suite 1-npu diff --git a/python/pyproject_npu.toml b/python/pyproject_npu.toml index 70403f8ee..26ab1a6a3 100644 --- a/python/pyproject_npu.toml +++ b/python/pyproject_npu.toml @@ -77,7 +77,8 @@ diffusion = [ "moviepy>=2.0.0", "opencv-python==4.10.0.84", "remote-pdb", - "cache-dit==1.1.8" + "cache-dit==1.2.1", + "addict" ] tracing = [ diff --git a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py index d9915fd8c..cabd056ff 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py +++ b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py @@ -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() diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index 24d1dca77..7bc7d8b4e 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/layers/activation.py b/python/sglang/multimodal_gen/runtime/layers/activation.py index 668a5b2f1..43b325950 100644 --- a/python/sglang/multimodal_gen/runtime/layers/activation.py +++ b/python/sglang/multimodal_gen/runtime/layers/activation.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/layers/custom_op.py b/python/sglang/multimodal_gen/runtime/layers/custom_op.py index 373261183..fd7355a3f 100644 --- a/python/sglang/multimodal_gen/runtime/layers/custom_op.py +++ b/python/sglang/multimodal_gen/runtime/layers/custom_op.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py index f855795f2..0f76516b8 100644 --- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py +++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/layers/linear.py b/python/sglang/multimodal_gen/runtime/layers/linear.py index 709642352..73a668ff1 100644 --- a/python/sglang/multimodal_gen/runtime/layers/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/linear.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/triton_ops.py b/python/sglang/multimodal_gen/runtime/layers/triton_ops.py index 931e9a473..13d613934 100644 --- a/python/sglang/multimodal_gen/runtime/layers/triton_ops.py +++ b/python/sglang/multimodal_gen/runtime/layers/triton_ops.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py b/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py index 0227351ef..e32f662ac 100644 --- a/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py +++ b/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 2cefd7761..4673a798b 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py index f468ab8ad..7d9d5d117 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/clip.py b/python/sglang/multimodal_gen/runtime/models/encoders/clip.py index deac46cf0..4b85a1260 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/clip.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/clip.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 13da5f65a..c5dabf507 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/platforms/__init__.py b/python/sglang/multimodal_gen/runtime/platforms/__init__.py index 88d2f0155..55aa8ad5d 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/__init__.py +++ b/python/sglang/multimodal_gen/runtime/platforms/__init__.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index 6b4530a68..cf368f453 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index 7b78168ab..93bde4220 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -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(): diff --git a/python/sglang/multimodal_gen/runtime/platforms/mps.py b/python/sglang/multimodal_gen/runtime/platforms/mps.py index 2208d70ce..bb9116a4d 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/mps.py +++ b/python/sglang/multimodal_gen/runtime/platforms/mps.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/platforms/npu.py b/python/sglang/multimodal_gen/runtime/platforms/npu.py new file mode 100644 index 000000000..4b15e55a6 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/platforms/npu.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/platforms/rocm.py b/python/sglang/multimodal_gen/runtime/platforms/rocm.py index 24d1bd95a..02eb2ffa4 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/rocm.py +++ b/python/sglang/multimodal_gen/runtime/platforms/rocm.py @@ -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) diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index d5823d30c..6610a4cfb 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -38,6 +38,15 @@ SUITES = { ], } +suites_ascend = { + "1-npu": [ + "ascend/test_server_1_npu.py", + # add new 1-npu test files here + ] +} + +SUITES.update(suites_ascend) + def parse_args(): parser = argparse.ArgumentParser(description="Run multimodal_gen test suite") diff --git a/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json b/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json new file mode 100644 index 000000000..12d05de79 --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json @@ -0,0 +1,76 @@ +{ + "metadata": { + "model": "Diffusion Server", + "hardware": "CI A2 64GB pool", + "description": "Reference numbers captured from the CI diffusion server baseline run" + }, + "scenarios": { + "wan2_1_t2v_1.3b_1_npu": { + "stages_ms": { + "InputValidationStage": 0.1, + "TextEncodingStage": 1609.27, + "ConditioningStage": 0.02, + "TimestepPreparationStage": 3.46, + "LatentPreparationStage": 0.39, + "DenoisingStage": 26324.0, + "DecodingStage": 817.68, + "per_frame_generation": null + }, + "denoise_step_ms": { + "0": 195.27, + "1": 329.05, + "2": 545.43, + "3": 541.3, + "4": 537.07, + "5": 537.21, + "6": 537.19, + "7": 537.19, + "8": 537.27, + "9": 537.05, + "10": 537.02, + "11": 537.11, + "12": 537.42, + "13": 537.2, + "14": 537.16, + "15": 537.11, + "16": 537.14, + "17": 537.19, + "18": 537.1, + "19": 537.0, + "20": 537.26, + "21": 537.18, + "22": 537.16, + "23": 537.24, + "24": 537.15, + "25": 537.14, + "26": 536.99, + "27": 537.19, + "28": 537.22, + "29": 537.23, + "30": 537.06, + "31": 537.06, + "32": 537.18, + "33": 537.07, + "34": 537.19, + "35": 537.28, + "36": 537.17, + "37": 537.38, + "38": 537.31, + "39": 537.25, + "40": 537.28, + "41": 537.26, + "42": 537.1, + "43": 537.19, + "44": 537.19, + "45": 537.31, + "46": 537.19, + "47": 537.16, + "48": 537.23, + "49": 532.91 + }, + "expected_e2e_ms": 28769.9, + "expected_avg_denoise_ms": 526.34, + "expected_median_denoise_ms": 537.19 + } + } +} diff --git a/python/sglang/multimodal_gen/test/server/ascend/test_server_1_npu.py b/python/sglang/multimodal_gen/test/server/ascend/test_server_1_npu.py new file mode 100644 index 000000000..3be09a899 --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/ascend/test_server_1_npu.py @@ -0,0 +1,29 @@ +""" +Config-driven diffusion performance test with pytest parametrization. + + +If the actual run is significantly better than the baseline, the improved cases with their updated baseline will be printed +""" + +from __future__ import annotations + +import pytest + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.test.server.ascend.testcase_configs_npu import ONE_NPU_CASES +from sglang.multimodal_gen.test.server.test_server_common import ( # noqa: F401 + DiffusionServerBase, + diffusion_server, +) +from sglang.multimodal_gen.test.server.testcase_configs import DiffusionTestCase + +logger = init_logger(__name__) + + +class TestDiffusionServerOneNpu(DiffusionServerBase): + """Performance tests for 1-NPU diffusion cases.""" + + @pytest.fixture(params=ONE_NPU_CASES, ids=lambda c: c.id) + def case(self, request) -> DiffusionTestCase: + """Provide a DiffusionTestCase for each 1-NPU test.""" + return request.param diff --git a/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py b/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py new file mode 100644 index 000000000..96211ecb9 --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py @@ -0,0 +1,22 @@ +from sglang.multimodal_gen.test.server.testcase_configs import ( + T2V_PROMPT, + DiffusionSamplingParams, + DiffusionServerArgs, + DiffusionTestCase, +) + +ONE_NPU_CASES: list[DiffusionTestCase] = [ + # === Text to Video (T2V) === + DiffusionTestCase( + "wan2_1_t2v_1.3b_1_npu", + DiffusionServerArgs( + model_path="/root/.cache/modelscope/hub/models/Wan-AI/Wan2.1-T2V-1.3B-Diffusers", + modality="video", + warmup=0, + custom_validator="video", + ), + DiffusionSamplingParams( + prompt=T2V_PROMPT, + ), + ), +] diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 07d12d5ad..a73bf307c 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -132,6 +132,24 @@ class BaselineConfig: ), ) + def update(self, path: Path): + """Load baseline configuration from JSON file.""" + with path.open("r", encoding="utf-8") as fh: + data = json.load(fh) + + scenarios_new = {} + for name, cfg in data["scenarios"].items(): + scenarios_new[name] = ScenarioConfig( + stages_ms=cfg["stages_ms"], + denoise_step_ms={int(k): v for k, v in cfg["denoise_step_ms"].items()}, + expected_e2e_ms=float(cfg["expected_e2e_ms"]), + expected_avg_denoise_ms=float(cfg["expected_avg_denoise_ms"]), + expected_median_denoise_ms=float(cfg["expected_median_denoise_ms"]), + ) + + self.scenarios.update(scenarios_new) + return self + @dataclass(frozen=True) class DiffusionServerArgs: @@ -729,4 +747,6 @@ TWO_GPU_CASES_B = [ ] # Load global configuration -BASELINE_CONFIG = BaselineConfig.load(Path(__file__).with_name("perf_baselines.json")) +BASELINE_CONFIG = BaselineConfig.load( + Path(__file__).with_name("perf_baselines.json") +).update(Path(__file__).parent / "ascend" / "perf_baselines_npu.json")