[diffusion] [NPU] Optimize LTX-2/2.3 inference performance for NPU (#34722)
Co-authored-by: Elizaveta Martirosian <you@example.com> Co-authored-by: ronnie_zheng <zl19940307@163.com> Co-authored-by: root <root@localhost.localdomain> Co-authored-by: Makcum888e <79456407+Makcum888e@users.noreply.github.com> Co-authored-by: mickqian <mickqian@users.noreply.github.com>
This commit is contained in:
co-authored by
Elizaveta Martirosian
ronnie_zheng
root
Makcum888e
mickqian
parent
2b1c4e4c85
commit
0027af2eac
@@ -6,6 +6,10 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import (
|
||||||
|
current_platform,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _ltx2_ada_values9_kernel(
|
def _ltx2_ada_values9_kernel(
|
||||||
@@ -141,14 +145,17 @@ def ltx2_ada_values9(
|
|||||||
) -> tuple[torch.Tensor, ...]:
|
) -> tuple[torch.Tensor, ...]:
|
||||||
if timestep.ndim != 3:
|
if timestep.ndim != 3:
|
||||||
raise ValueError("timestep must have shape [B, S, 9 * D]")
|
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")
|
raise ValueError("timestep must be a CUDA bfloat16 tensor")
|
||||||
if not timestep.is_contiguous():
|
if not timestep.is_contiguous():
|
||||||
raise ValueError("timestep must be contiguous")
|
raise ValueError("timestep must be contiguous")
|
||||||
if scale_shift_table.ndim != 2 or scale_shift_table.shape[0] != 9:
|
if scale_shift_table.ndim != 2 or scale_shift_table.shape[0] != 9:
|
||||||
raise ValueError("scale_shift_table must have shape [9, D]")
|
raise ValueError("scale_shift_table must have shape [9, D]")
|
||||||
if (
|
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.dtype not in (torch.bfloat16, torch.float32)
|
||||||
or scale_shift_table.stride(-1) != 1
|
or scale_shift_table.stride(-1) != 1
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -44,7 +44,7 @@ from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
|||||||
tensor_model_parallel_all_reduce,
|
tensor_model_parallel_all_reduce,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention
|
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 (
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
ColumnParallelLinear,
|
ColumnParallelLinear,
|
||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
@@ -237,12 +237,12 @@ def _ltx2_try_fused_ada_values9(
|
|||||||
if (
|
if (
|
||||||
_LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED
|
_LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED
|
||||||
or get_tp_world_size() != 1
|
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.dtype != torch.bfloat16
|
||||||
or timestep.ndim != 3
|
or timestep.ndim != 3
|
||||||
or int(timestep.shape[0]) != int(batch_size)
|
or int(timestep.shape[0]) != int(batch_size)
|
||||||
or not timestep.is_contiguous()
|
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.dtype not in (torch.bfloat16, torch.float32)
|
||||||
or scale_shift_table.ndim != 2
|
or scale_shift_table.ndim != 2
|
||||||
or int(scale_shift_table.shape[0]) != 9
|
or int(scale_shift_table.shape[0]) != 9
|
||||||
@@ -811,8 +811,12 @@ class LTX2Attention(nn.Module):
|
|||||||
self.k_norm: nn.Module | None = None
|
self.k_norm: nn.Module | None = None
|
||||||
if self.qk_norm:
|
if self.qk_norm:
|
||||||
if tp_size == 1:
|
if tp_size == 1:
|
||||||
self.q_norm = torch.nn.RMSNorm(self.inner_dim, eps=self.norm_eps)
|
if _is_npu:
|
||||||
self.k_norm = torch.nn.RMSNorm(self.inner_dim, eps=self.norm_eps)
|
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:
|
else:
|
||||||
self.q_norm = LTX2TPRMSNormAcrossHeads(
|
self.q_norm = LTX2TPRMSNormAcrossHeads(
|
||||||
full_hidden_size=self.inner_dim,
|
full_hidden_size=self.inner_dim,
|
||||||
@@ -1778,6 +1782,12 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
hf_config.get("rope_double_precision", arch.double_precision_rope)
|
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(
|
self.quantize_video_rope_coords_to_hidden_dtype = bool(
|
||||||
hf_config.get("quantize_video_rope_coords_to_hidden_dtype", False)
|
hf_config.get("quantize_video_rope_coords_to_hidden_dtype", False)
|
||||||
)
|
)
|
||||||
|
|||||||
+8
-1
@@ -39,6 +39,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
|||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||||
StageValidators as V,
|
StageValidators as V,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import (
|
||||||
|
current_platform,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import (
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
is_ltx2_two_stage_pipeline_name,
|
is_ltx2_two_stage_pipeline_name,
|
||||||
@@ -563,7 +566,11 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
noise = torch.randn(
|
noise = torch.randn(
|
||||||
reference_tensor.shape,
|
reference_tensor.shape,
|
||||||
generator=generator,
|
generator=generator,
|
||||||
dtype=torch.float64,
|
dtype=(
|
||||||
|
torch.float32
|
||||||
|
if not current_platform.is_float64_supported()
|
||||||
|
else torch.float64
|
||||||
|
),
|
||||||
device=reference_tensor.device,
|
device=reference_tensor.device,
|
||||||
)
|
)
|
||||||
noise = (noise - noise.mean()) / noise.std()
|
noise = (noise - noise.mean()) / noise.std()
|
||||||
|
|||||||
@@ -453,6 +453,10 @@ class Platform:
|
|||||||
attention_cls_str = self.get_attn_backend_cls_str(*args, **kwargs)
|
attention_cls_str = self.get_attn_backend_cls_str(*args, **kwargs)
|
||||||
return resolve_obj_by_qualname(attention_cls_str)
|
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):
|
class UnspecifiedPlatform(Platform):
|
||||||
_enum = PlatformEnum.UNSPECIFIED
|
_enum = PlatformEnum.UNSPECIFIED
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
# Adapted from vllm-ascend: https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/platform.py
|
# Adapted from vllm-ascend: https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/platform.py
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
from functools import lru_cache
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -40,6 +41,14 @@ class NPUPlatformBase(Platform):
|
|||||||
dispatch_key: str = "NPU"
|
dispatch_key: str = "NPU"
|
||||||
device_control_env_var: str = "ASCEND_RT_VISIBLE_DEVICES"
|
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
|
@classmethod
|
||||||
def get_local_torch_device(cls) -> torch.device:
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
return torch.device(f"npu:{envs.LOCAL_RANK}")
|
return torch.device(f"npu:{envs.LOCAL_RANK}")
|
||||||
|
|||||||
Reference in New Issue
Block a user