diff --git a/python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py b/python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py index 138934d36..a41af05d4 100644 --- a/python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py +++ b/python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py @@ -6,6 +6,10 @@ import torch import triton import triton.language as tl +from sglang.multimodal_gen.runtime.platforms import ( + current_platform, +) + @triton.jit def _ltx2_ada_values9_kernel( @@ -141,14 +145,17 @@ def ltx2_ada_values9( ) -> tuple[torch.Tensor, ...]: if timestep.ndim != 3: raise ValueError("timestep must have shape [B, S, 9 * D]") - if not timestep.is_cuda or timestep.dtype != torch.bfloat16: + if ( + not current_platform.tensor_on_device(timestep) + or timestep.dtype != torch.bfloat16 + ): raise ValueError("timestep must be a CUDA bfloat16 tensor") if not timestep.is_contiguous(): raise ValueError("timestep must be contiguous") if scale_shift_table.ndim != 2 or scale_shift_table.shape[0] != 9: raise ValueError("scale_shift_table must have shape [9, D]") if ( - not scale_shift_table.is_cuda + not current_platform.tensor_on_device(scale_shift_table) or scale_shift_table.dtype not in (torch.bfloat16, torch.float32) or scale_shift_table.stride(-1) != 1 ): diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py index 63e93cda3..c2bca4961 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py @@ -44,7 +44,7 @@ from sglang.multimodal_gen.runtime.distributed.communication_op import ( tensor_model_parallel_all_reduce, ) from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNormNoWeight +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, RMSNormNoWeight from sglang.multimodal_gen.runtime.layers.linear import ( ColumnParallelLinear, RowParallelLinear, @@ -237,12 +237,12 @@ def _ltx2_try_fused_ada_values9( if ( _LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED or get_tp_world_size() != 1 - or not timestep.is_cuda + or not current_platform.tensor_on_device(timestep) or timestep.dtype != torch.bfloat16 or timestep.ndim != 3 or int(timestep.shape[0]) != int(batch_size) or not timestep.is_contiguous() - or not scale_shift_table.is_cuda + or not current_platform.tensor_on_device(scale_shift_table) or scale_shift_table.dtype not in (torch.bfloat16, torch.float32) or scale_shift_table.ndim != 2 or int(scale_shift_table.shape[0]) != 9 @@ -811,8 +811,12 @@ class LTX2Attention(nn.Module): self.k_norm: nn.Module | None = None if self.qk_norm: if tp_size == 1: - self.q_norm = torch.nn.RMSNorm(self.inner_dim, eps=self.norm_eps) - self.k_norm = torch.nn.RMSNorm(self.inner_dim, eps=self.norm_eps) + if _is_npu: + self.q_norm = RMSNorm(self.inner_dim, eps=self.norm_eps) + self.k_norm = RMSNorm(self.inner_dim, eps=self.norm_eps) + else: + self.q_norm = torch.nn.RMSNorm(self.inner_dim, eps=self.norm_eps) + self.k_norm = torch.nn.RMSNorm(self.inner_dim, eps=self.norm_eps) else: self.q_norm = LTX2TPRMSNormAcrossHeads( full_hidden_size=self.inner_dim, @@ -1778,6 +1782,12 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): hf_config.get("rope_double_precision", arch.double_precision_rope) ) ) + if rope_double_precision and not current_platform.is_float64_supported(): + logger.warning( + "Current platform does not support float64. Falling back to float32." + ) + rope_double_precision = False + self.quantize_video_rope_coords_to_hidden_dtype = bool( hf_config.get("quantize_video_rope_coords_to_hidden_dtype", False) ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py index 162767f5d..7a0127ee6 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py @@ -39,6 +39,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( StageValidators as V, ) +from sglang.multimodal_gen.runtime.platforms import ( + current_platform, +) from sglang.multimodal_gen.runtime.server_args import ( ServerArgs, is_ltx2_two_stage_pipeline_name, @@ -563,7 +566,11 @@ class LTX2DenoisingStage(DenoisingStage): noise = torch.randn( reference_tensor.shape, generator=generator, - dtype=torch.float64, + dtype=( + torch.float32 + if not current_platform.is_float64_supported() + else torch.float64 + ), device=reference_tensor.device, ) noise = (noise - noise.mean()) / noise.std() diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index 8e24d7c39..ee4702660 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -453,6 +453,10 @@ class Platform: attention_cls_str = self.get_attn_backend_cls_str(*args, **kwargs) return resolve_obj_by_qualname(attention_cls_str) + def tensor_on_device(self, t: torch.Tensor) -> bool: + """Check if a tensor is on the current platform's device.""" + return t.is_cuda + class UnspecifiedPlatform(Platform): _enum = PlatformEnum.UNSPECIFIED diff --git a/python/sglang/multimodal_gen/runtime/platforms/npu.py b/python/sglang/multimodal_gen/runtime/platforms/npu.py index 27034bb69..478a0f2db 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/npu.py +++ b/python/sglang/multimodal_gen/runtime/platforms/npu.py @@ -2,6 +2,7 @@ # Adapted from vllm-ascend: https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/platform.py import os +from functools import lru_cache from typing import Any import torch @@ -40,6 +41,14 @@ class NPUPlatformBase(Platform): dispatch_key: str = "NPU" device_control_env_var: str = "ASCEND_RT_VISIBLE_DEVICES" + @classmethod + @lru_cache(maxsize=1) + def is_float64_supported(cls) -> bool: + return False + + def tensor_on_device(self, t: torch.Tensor) -> bool: + return t.is_npu + @classmethod def get_local_torch_device(cls) -> torch.device: return torch.device(f"npu:{envs.LOCAL_RANK}")