[diffusion] model: support two stage pipeline of LTX-2 (#20707)

Co-authored-by: daiweitao <dwti614707404@163.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
Co-authored-by: GMI Xiao Jin <xiao.j@gmicloud.ai>
This commit is contained in:
Prozac614
2026-04-04 09:37:28 +08:00
committed by GitHub
co-authored by daiweitao Mick GMI Xiao Jin
parent 95cdbce34f
commit db3d4f4b76
41 changed files with 2201 additions and 647 deletions
+15 -12
View File
@@ -33,23 +33,26 @@ default parameters when initializing and generating videos.
| TurboWan2.1 T2V 14B | `IPostYellow/TurboWan2.1-T2V-14B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ | | TurboWan2.1 T2V 14B | `IPostYellow/TurboWan2.1-T2V-14B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
| TurboWan2.1 T2V 14B 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ | | TurboWan2.1 T2V 14B 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
| TurboWan2.2 I2V A14B | `IPostYellow/TurboWan2.2-I2V-A14B-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ | | TurboWan2.2 I2V A14B | `IPostYellow/TurboWan2.2-I2V-A14B-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
| LTX-2 | `Lightricks/LTX-2` | 1536×1024 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
**Note**: **Note**:
1.Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue. 1. Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
2.SageSLA Based on SpargeAttn. Install it first with `pip install git+https://github.com/thu-ml/SpargeAttn.git --no-build-isolation` 2. SageSLA is based on SpargeAttn. Install it first with `pip install git+https://github.com/thu-ml/SpargeAttn.git --no-build-isolation`
3. LTX-2 two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`. The spatial upsampler and distilled LoRA are auto-resolved from the model snapshot by default, and can still be overridden with `--spatial-upsampler-path` and `--distilled-lora-path`.
### Image Generation Models ### Image Generation Models
| Model Name | HuggingFace Model ID | Resolutions | | Model Name | HuggingFace Model ID |
|:-----------------|:----------------------------------------|:---------------| |:---------------------|:------------------------------------|
| FLUX.1-dev | `black-forest-labs/FLUX.1-dev` | Any resolution | | FLUX.1-dev | `black-forest-labs/FLUX.1-dev` |
| FLUX.2-dev | `black-forest-labs/FLUX.2-dev` | Any resolution | | FLUX.2-dev | `black-forest-labs/FLUX.2-dev` |
| FLUX.2-Klein | `black-forest-labs/FLUX.2-klein-4B` | Any resolution | | FLUX.2-Klein | `black-forest-labs/FLUX.2-klein-4B` |
| Z-Image-Turbo | `Tongyi-MAI/Z-Image-Turbo` | Any resolution | | Z-Image-Turbo | `Tongyi-MAI/Z-Image-Turbo` |
| GLM-Image | `zai-org/GLM-Image` | Any resolution | | GLM-Image | `zai-org/GLM-Image` |
| Qwen Image | `Qwen/Qwen-Image` | Any resolution | | Qwen Image | `Qwen/Qwen-Image` |
| Qwen Image 2512 | `Qwen/Qwen-Image-2512` | Any resolution | | Qwen Image 2512 | `Qwen/Qwen-Image-2512` |
| Qwen Image Edit | `Qwen/Qwen-Image-Edit` | Any resolution | | Qwen Image Edit | `Qwen/Qwen-Image-Edit` |
| Qwen Image Edit 2511 | `Qwen/Qwen-Image-Edit-2511` |
## Verified LoRA Examples ## Verified LoRA Examples
+4 -6
View File
@@ -27,12 +27,12 @@ def _is_overlay_diffusion_model(model_path: str) -> bool:
def _is_registered_diffusion_model(model_path: str) -> bool: def _is_registered_diffusion_model(model_path: str) -> bool:
try: try:
# if diffusion dependencies are not installed from sglang.multimodal_gen.registry import has_registered_diffusion_model_path
from sglang.multimodal_gen.registry import get_model_info
except ImportError: except ImportError:
# if diffusion dependencies are not installed
return False return False
return get_model_info(model_path, backend="sglang") is not None return has_registered_diffusion_model_path(model_path)
def _is_diffusers_model_dir(model_dir: str) -> bool: def _is_diffusers_model_dir(model_dir: str) -> bool:
@@ -93,9 +93,7 @@ def get_is_diffusion_model(model_path: str) -> bool:
return _is_diffusers_model_dir(os.path.dirname(file_path)) return _is_diffusers_model_dir(os.path.dirname(file_path))
except Exception as e: except Exception as e:
logger.debug("Failed to auto-detect diffusion model for %s: %s", model_path, e) logger.debug("Failed to auto-detect diffusion model for %s: %s", model_path, e)
# For gated repos, file download fails but model card is still accessible. return False
# Check library_name from HF metadata as a fallback.
return _is_gated_diffusion_repo(model_path)
def get_model_path(extra_argv): def get_model_path(extra_argv):
+6 -1
View File
@@ -9,7 +9,7 @@ SGLang diffusion features an end-to-end unified pipeline for accelerating diffus
## Key Features ## Key Features
SGLang Diffusion has the following features: SGLang Diffusion has the following features:
- Broad model support: Wan series, FastWan series, Hunyuan, Qwen-Image, Qwen-Image-Edit, Flux, Z-Image, GLM-Image - Broad model support: Wan series, FastWan series, Hunyuan, LTX-2, Qwen-Image, Qwen-Image-Edit, Flux, Z-Image, GLM-Image
- Fast inference speed: enpowered by highly optimized kernel from sgl-kernel and efficient scheduler loop - Fast inference speed: enpowered by highly optimized kernel from sgl-kernel and efficient scheduler loop
- Ease of use: OpenAI-compatible api, CLI, and python sdk support - Ease of use: OpenAI-compatible api, CLI, and python sdk support
- Multi-platform support: - Multi-platform support:
@@ -76,6 +76,11 @@ sglang generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--save-output --save-output
``` ```
For LTX-2 two-stage generation, use `--pipeline-class-name LTX2TwoStagePipeline`. The
spatial upsampler and distilled LoRA are auto-resolved from the same model snapshot by
default, and can still be overridden with `--spatial-upsampler-path` and
`--distilled-lora-path` when needed.
### LoRA support ### LoRA support
Apply LoRA adapters via `--lora-path`: Apply LoRA adapters via `--lora-path`:
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAECon
@dataclass @dataclass
class LTXAudioVAEArchConfig(VAEArchConfig): class LTXAudioVAEArchConfig(VAEArchConfig):
# Architecture params # Architecture params
temporal_compression_ratio: int = 4
causality_axis: str = "height" causality_axis: str = "height"
attn_resolutions: Optional[Tuple[int, ...]] = None attn_resolutions: Optional[Tuple[int, ...]] = None
base_channels: int = 128 base_channels: int = 128
@@ -20,6 +21,7 @@ class LTXAudioVAEArchConfig(VAEArchConfig):
mid_block_add_attention: bool = False mid_block_add_attention: bool = False
sample_rate: int = 16000 sample_rate: int = 16000
mel_hop_length: int = 160 mel_hop_length: int = 160
mel_compression_ratio: int = 4
is_causal: bool = True is_causal: bool = True
mel_bins: Optional[int] = 64 mel_bins: Optional[int] = 64
double_z: bool = True double_z: bool = True
@@ -11,6 +11,8 @@ class LTXVideoVAEArchConfig(VAEArchConfig):
in_channels: int = 3 in_channels: int = 3
latent_channels: int = 128 latent_channels: int = 128
out_channels: int = 3 out_channels: int = 3
temporal_compression_ratio: int = 8
spatial_compression_ratio: int = 32
block_out_channels: List[int] = field( block_out_channels: List[int] = field(
default_factory=lambda: [256, 512, 1024, 2048] default_factory=lambda: [256, 512, 1024, 2048]
) )
@@ -169,6 +169,7 @@ class PipelineConfig:
# controls the timestep embedding generation # controls the timestep embedding generation
should_use_guidance: bool = True should_use_guidance: bool = True
embedded_cfg_scale: float = 6.0 embedded_cfg_scale: float = 6.0
generator_device: str | None = None
flow_shift: float | None = None flow_shift: float | None = None
disable_autocast: bool = False disable_autocast: bool = False
@@ -420,6 +421,9 @@ class PipelineConfig:
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {} return {}
def _unpad_and_unpack_latents(self, latents, audio_latents, batch, vae, audio_vae):
raise NotImplementedError("not yet implemented")
@staticmethod @staticmethod
def add_cli_args( def add_cli_args(
parser: FlexibleArgumentParser, prefix: str = "" parser: FlexibleArgumentParser, prefix: str = ""
@@ -1,6 +1,6 @@
import dataclasses import dataclasses
from dataclasses import field from dataclasses import field
from typing import Callable from typing import Callable, Optional
import torch import torch
@@ -11,6 +11,7 @@ from sglang.multimodal_gen.configs.models.encoders import (
) )
from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config
from sglang.multimodal_gen.configs.models.vaes.ltx_audio import LTXAudioVAEConfig from sglang.multimodal_gen.configs.models.vaes.ltx_audio import LTXAudioVAEConfig
from sglang.multimodal_gen.configs.models.vaes.ltx_video import LTXVideoVAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import ( from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType, ModelTaskType,
PipelineConfig, PipelineConfig,
@@ -93,8 +94,11 @@ def pack_text_embeds(
def _gemma_postprocess_func( def _gemma_postprocess_func(
outputs: BaseEncoderOutput, text_inputs: dict outputs: BaseEncoderOutput,
text_inputs: dict,
pipeline_config: Optional["LTX2PipelineConfig"] = None,
) -> torch.Tensor: ) -> torch.Tensor:
_ = pipeline_config
# LTX-2 requires all hidden states concatenated for the connector # LTX-2 requires all hidden states concatenated for the connector
if hasattr(outputs, "hidden_states") and outputs.hidden_states is not None: if hasattr(outputs, "hidden_states") and outputs.hidden_states is not None:
# outputs.hidden_states is a tuple of tensors # outputs.hidden_states is a tuple of tensors
@@ -116,6 +120,7 @@ class LTX2PipelineConfig(PipelineConfig):
task_type: ModelTaskType = ModelTaskType.TI2V task_type: ModelTaskType = ModelTaskType.TI2V
skip_input_image_preprocess: bool = True skip_input_image_preprocess: bool = True
generator_device: str = "cpu"
dit_config: LTX2Config = field(default_factory=LTX2Config) dit_config: LTX2Config = field(default_factory=LTX2Config)
# Model architecture # Model architecture
@@ -125,35 +130,23 @@ class LTX2PipelineConfig(PipelineConfig):
patch_size_t: int = 1 patch_size_t: int = 1
# Audio VAE configuration # Audio VAE configuration
vae_config: LTXVideoVAEConfig = field(default_factory=LTXVideoVAEConfig)
audio_vae_config: LTXAudioVAEConfig = field(default_factory=LTXAudioVAEConfig) audio_vae_config: LTXAudioVAEConfig = field(default_factory=LTXAudioVAEConfig)
audio_vae_precision: str = "fp32" audio_vae_precision: str = "fp32"
audio_vae_temporal_compression_ratio: int = 4
audio_vae_mel_compression_ratio: int = 4
@property @property
def vae_scale_factor(self): def vae_scale_factor(self):
return getattr(self.vae_config.arch_config, "spatial_compression_ratio", 32) return self.vae_config.arch_config.spatial_compression_ratio
@property @property
def vae_temporal_compression(self): def vae_temporal_compression(self):
return getattr(self.vae_config.arch_config, "temporal_compression_ratio", 8) return self.vae_config.arch_config.temporal_compression_ratio
def prepare_latent_shape(self, batch, batch_size, num_frames): def prepare_latent_shape(self, batch, batch_size, num_frames):
"""Return packed latent shape [B, seq, C] directly.""" """Return unpacked latent shape [B, C, F, H, W]."""
height = batch.height // self.vae_scale_factor height = batch.height // self.vae_scale_factor
width = batch.width // self.vae_scale_factor width = batch.width // self.vae_scale_factor
return (batch_size, self.in_channels, num_frames, height, width)
post_patch_num_frames = num_frames // self.patch_size_t
post_patch_height = height // self.patch_size
post_patch_width = width // self.patch_size
seq_len = post_patch_num_frames * post_patch_height * post_patch_width
num_channels = (
self.in_channels * self.patch_size_t * self.patch_size * self.patch_size
)
shape = (batch_size, seq_len, num_channels)
return shape
def prepare_audio_latent_shape(self, batch, batch_size, num_frames): def prepare_audio_latent_shape(self, batch, batch_size, num_frames):
# Adapted from diffusers pipeline prepare_audio_latents # Adapted from diffusers pipeline prepare_audio_latents
@@ -161,7 +154,9 @@ class LTX2PipelineConfig(PipelineConfig):
sample_rate = self.audio_vae_config.arch_config.sample_rate sample_rate = self.audio_vae_config.arch_config.sample_rate
hop_length = self.audio_vae_config.arch_config.mel_hop_length hop_length = self.audio_vae_config.arch_config.mel_hop_length
temporal_compression = self.audio_vae_temporal_compression_ratio temporal_compression = (
self.audio_vae_config.arch_config.temporal_compression_ratio
)
latents_per_second = ( latents_per_second = (
float(sample_rate) / float(hop_length) / float(temporal_compression) float(sample_rate) / float(hop_length) / float(temporal_compression)
@@ -169,15 +164,13 @@ class LTX2PipelineConfig(PipelineConfig):
latent_length = round(duration_s * latents_per_second) latent_length = round(duration_s * latents_per_second)
num_mel_bins = self.audio_vae_config.arch_config.mel_bins num_mel_bins = self.audio_vae_config.arch_config.mel_bins
mel_compression_ratio = self.audio_vae_mel_compression_ratio mel_compression_ratio = self.audio_vae_config.arch_config.mel_compression_ratio
latent_mel_bins = num_mel_bins // mel_compression_ratio latent_mel_bins = num_mel_bins // mel_compression_ratio
# Default to 8 # Default to 8
num_channels_latents = self.audio_vae_config.arch_config.latent_channels num_channels_latents = self.audio_vae_config.arch_config.latent_channels
shape = (batch_size, latent_length, num_channels_latents * latent_mel_bins) return (batch_size, num_channels_latents, latent_length, latent_mel_bins)
return shape
# Text encoding stage (Gemma) # Text encoding stage (Gemma)
# LTX-2 needs separate contexts for video/audio streams. We model this as # LTX-2 needs separate contexts for video/audio streams. We model this as
@@ -221,6 +214,7 @@ class LTX2PipelineConfig(PipelineConfig):
padding="max_length", padding="max_length",
max_length=max_sequence_length, max_length=max_sequence_length,
truncation=True, truncation=True,
add_special_tokens=True,
return_tensors="pt", return_tensors="pt",
) )
return text_inputs return text_inputs
@@ -524,7 +518,9 @@ class LTX2PipelineConfig(PipelineConfig):
sample_rate = self.audio_vae_config.arch_config.sample_rate sample_rate = self.audio_vae_config.arch_config.sample_rate
hop_length = self.audio_vae_config.arch_config.mel_hop_length hop_length = self.audio_vae_config.arch_config.mel_hop_length
temporal_compression = self.audio_vae_temporal_compression_ratio temporal_compression = (
self.audio_vae_config.arch_config.temporal_compression_ratio
)
duration_s = num_frames / batch.fps duration_s = num_frames / batch.fps
latents_per_second = ( latents_per_second = (
@@ -533,43 +529,9 @@ class LTX2PipelineConfig(PipelineConfig):
audio_num_frames = round(duration_s * latents_per_second) audio_num_frames = round(duration_s * latents_per_second)
num_mel_bins = self.audio_vae_config.arch_config.mel_bins num_mel_bins = self.audio_vae_config.arch_config.mel_bins
mel_compression_ratio = self.audio_vae_mel_compression_ratio mel_compression_ratio = self.audio_vae_config.arch_config.mel_compression_ratio
latent_mel_bins = num_mel_bins // mel_compression_ratio latent_mel_bins = num_mel_bins // mel_compression_ratio
audio_latents_mean = getattr(audio_vae, "latents_mean", None)
audio_latents_std = getattr(audio_vae, "latents_std", None)
if (
isinstance(audio_latents_mean, torch.Tensor)
and isinstance(audio_latents_std, torch.Tensor)
and audio_latents_mean.numel() == audio_latents_std.numel()
):
audio_latents_mean = audio_latents_mean.to(
device=audio_latents.device, dtype=audio_latents.dtype
)
audio_latents_std = audio_latents_std.to(
device=audio_latents.device, dtype=audio_latents.dtype
)
if audio_latents.ndim == 3:
if audio_latents.shape[-1] != audio_latents_mean.numel():
raise ValueError(
f"audio_latents last dim {audio_latents.shape[-1]} "
f"does not match audio_vae stats {audio_latents_mean.numel()}"
)
audio_latents = audio_latents * audio_latents_std.view(
1, 1, -1
) + audio_latents_mean.view(1, 1, -1)
elif audio_latents.ndim == 2:
if audio_latents.shape[-1] != audio_latents_mean.numel():
raise ValueError(
f"audio_latents last dim {audio_latents.shape[-1]} "
f"does not match audio_vae stats {audio_latents_mean.numel()}"
)
audio_latents = audio_latents * audio_latents_std.view(
1, -1
) + audio_latents_mean.view(1, -1)
else:
audio_latents = audio_latents * audio_latents_std + audio_latents_mean
audio_latents = self._unpack_audio_latents( audio_latents = self._unpack_audio_latents(
audio_latents, audio_num_frames, num_mel_bins=latent_mel_bins audio_latents, audio_num_frames, num_mel_bins=latent_mel_bins
) )
@@ -10,6 +10,7 @@ class LTX2SamplingParams(SamplingParams):
# Match the reference defaults used by ltx-pipelines (one-stage). # Match the reference defaults used by ltx-pipelines (one-stage).
# See: LTX-2/packages/ltx-pipelines/src/ltx_pipelines/utils/constants.py # See: LTX-2/packages/ltx-pipelines/src/ltx_pipelines/utils/constants.py
seed: int = 10 seed: int = 10
generator_device: str = "cpu"
# Video parameters # Video parameters
height: int = 512 height: int = 512
@@ -128,7 +128,7 @@ class SamplingParams:
# Batch info # Batch info
num_outputs_per_prompt: int = 1 num_outputs_per_prompt: int = 1
seed: int = 42 seed: int = 42
generator_device: str = "cuda" # Device for random generator: "cuda" or "cpu" generator_device: str | None = None # None means use the pipeline/model default
# Original dimensions (before VAE scaling) # Original dimensions (before VAE scaling)
num_frames: int = 1 # Default for image models num_frames: int = 1 # Default for image models
@@ -685,7 +685,7 @@ class SamplingParams:
"--generator-device", "--generator-device",
type=str, type=str,
choices=["cuda", "musa", "cpu"], choices=["cuda", "musa", "cpu"],
help="Device for random generator (cuda, musa or cpu). Default: cuda", help="Device for random generator (cuda, musa or cpu). Default: use the model-specific setting.",
) )
add_argument( add_argument(
"--num-frames", "--num-frames",
+26
View File
@@ -264,6 +264,29 @@ def _normalize_hf_cache_path(path: str) -> str:
return os.path.normpath(path).lower().replace("\\", "/") return os.path.normpath(path).lower().replace("\\", "/")
def has_registered_diffusion_model_path(model_path: str) -> bool:
all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(), key=len, reverse=True)
if model_path in _MODEL_HF_PATH_TO_NAME:
return True
model_short_name = get_model_short_name(model_path.lower())
for registered_model_hf_id in all_model_hf_paths:
registered_model_name = get_model_short_name(registered_model_hf_id.lower())
if registered_model_name in model_short_name:
return True
normalized_model_path = _normalize_hf_cache_path(model_path)
for registered_model_hf_id in all_model_hf_paths:
cache_repo_fragment = (
f"models--{registered_model_hf_id.lower().replace('/', '--')}"
)
if cache_repo_fragment in normalized_model_path:
return True
return False
@lru_cache(maxsize=1) @lru_cache(maxsize=1)
def _get_config_info( def _get_config_info(
model_path: str, model_id: Optional[str] = None model_path: str, model_id: Optional[str] = None
@@ -570,6 +593,9 @@ def _register_configs():
register_configs( register_configs(
sampling_param_cls=LTX2SamplingParams, sampling_param_cls=LTX2SamplingParams,
pipeline_config_cls=LTX2PipelineConfig, pipeline_config_cls=LTX2PipelineConfig,
hf_model_paths=[
"Lightricks/LTX-2",
],
model_detectors=[ model_detectors=[
lambda path: "ltx" in path.lower() and "video" in path.lower(), lambda path: "ltx" in path.lower() and "video" in path.lower(),
lambda path: "ltx-2" in path.lower(), lambda path: "ltx-2" in path.lower(),
@@ -270,7 +270,6 @@ def _maybe_mux_audio_into_mp4(
sample_rate=selected_sr, sample_rate=selected_sr,
ffmpeg_exe=ffmpeg_exe, ffmpeg_exe=ffmpeg_exe,
) )
logger.info(f"Merged video saved to {CYAN}{save_file_path}{RESET}")
except Exception as e: except Exception as e:
logger.warning( logger.warning(
"Failed to mux audio into mp4 (saved silent video): %s", "Failed to mux audio into mp4 (saved silent video): %s",
@@ -266,6 +266,7 @@ class LocalAttention(nn.Module):
q: torch.Tensor, q: torch.Tensor,
k: torch.Tensor, k: torch.Tensor,
v: torch.Tensor, v: torch.Tensor,
attn_mask: torch.Tensor | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Apply local attention between query, key and value tensors. Apply local attention between query, key and value tensors.
@@ -284,6 +285,35 @@ class LocalAttention(nn.Module):
forward_context: ForwardContext = get_forward_context() forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata ctx_attn_metadata = forward_context.attn_metadata
if attn_mask is not None:
q_ = q.transpose(1, 2)
k_ = k.transpose(1, 2)
v_ = v.transpose(1, 2)
if torch.is_floating_point(attn_mask):
mask = attn_mask.to(dtype=q_.dtype, device=q_.device)
if mask.dim() == 2:
mask = mask[:, None, None, :]
elif mask.dim() == 3:
mask = mask[:, None, :, :]
else:
mask = attn_mask.to(dtype=q_.dtype, device=q_.device)
if mask.dim() == 2:
mask = mask[:, None, None, :]
elif mask.dim() == 3:
mask = mask[:, None, :, :]
mask = (mask - 1.0) * torch.finfo(q_.dtype).max
return torch.nn.functional.scaled_dot_product_attention(
q_,
k_,
v_,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
scale=self.softmax_scale,
).transpose(1, 2)
output = self.attn_impl.forward(q, k, v, attn_metadata=ctx_attn_metadata) output = self.attn_impl.forward(q, k, v, attn_metadata=ctx_attn_metadata)
return output return output
@@ -84,18 +84,22 @@ class BaseLayerWithLoRA(nn.Module):
# TODO: Support multiple LoRA adapters when use not merged mode # TODO: Support multiple LoRA adapters when use not merged mode
if not self.merged and not self.disable_lora: if not self.merged and not self.disable_lora:
lora_A_sliced = self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)) lora_dtype = lora_A.dtype
lora_B_sliced = self.slice_lora_b_weights(lora_B.to(x, non_blocking=True)) x_lora = x.to(dtype=lora_dtype)
delta = x @ lora_A_sliced.T @ lora_B_sliced.T lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(device=x.device, non_blocking=True)
)
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(device=x.device, non_blocking=True)
)
delta = x_lora @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank: if self.lora_alpha != self.lora_rank:
delta = delta * ( delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore ) # type: ignore
delta = delta * self.strength delta = delta * self.strength
if delta.dim() > 2:
delta = delta.reshape(-1, delta.shape[-1])
out, output_bias = self.base_layer(x) out, output_bias = self.base_layer(x)
return out + delta, output_bias return out + delta.to(dtype=out.dtype), output_bias
else: else:
out, output_bias = self.base_layer(x) out, output_bias = self.base_layer(x)
return out, output_bias return out, output_bias
@@ -113,6 +117,7 @@ class BaseLayerWithLoRA(nn.Module):
lora_path: str | None = None, lora_path: str | None = None,
strength: float = 1.0, strength: float = 1.0,
clear_existing: bool = False, clear_existing: bool = False,
merge_weights: bool = True,
) -> None: ) -> None:
""" """
Set LoRA weights. Supports multiple LoRA adapters. Set LoRA weights. Supports multiple LoRA adapters.
@@ -149,7 +154,10 @@ class BaseLayerWithLoRA(nn.Module):
self.strength = strength self.strength = strength
self.disable_lora = False self.disable_lora = False
self.merge_lora_weights() if merge_weights:
self.merge_lora_weights()
elif self.merged:
self.unmerge_lora_weights()
@torch.no_grad() @torch.no_grad()
def _merge_lora_into_data( def _merge_lora_into_data(
@@ -309,11 +317,34 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
super().__init__(base_layer, lora_rank, lora_alpha) super().__init__(base_layer, lora_rank, lora_alpha)
def forward(self, input_: torch.Tensor) -> torch.Tensor: def forward(self, input_: torch.Tensor) -> torch.Tensor:
# duplicate the logic in ColumnParallelLinear lora_A = self.lora_A
lora_B = self.lora_B
if isinstance(self.lora_B, DTensor):
lora_B = self.lora_B.to_local()
lora_A = self.lora_A.to_local()
bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None
output_parallel = self.base_layer.quant_method.apply( output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_, bias self.base_layer, input_, bias
) )
if not self.merged and not self.disable_lora:
lora_dtype = lora_A.dtype
input_lora = input_.to(dtype=lora_dtype)
lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(device=input_.device, non_blocking=True)
)
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(device=input_.device, non_blocking=True)
)
delta_parallel = input_lora @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank:
delta_parallel = delta_parallel * (
self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore
delta_parallel = delta_parallel * self.strength
output_parallel = output_parallel + delta_parallel.to(
dtype=output_parallel.dtype
)
if self.base_layer.gather_output: if self.base_layer.gather_output:
output = tensor_model_parallel_all_gather(output_parallel) output = tensor_model_parallel_all_gather(output_parallel)
else: else:
@@ -399,7 +430,12 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
super().__init__(base_layer, lora_rank, lora_alpha) super().__init__(base_layer, lora_rank, lora_alpha)
def forward(self, input_: torch.Tensor): def forward(self, input_: torch.Tensor):
# duplicate the logic in RowParallelLinear lora_A = self.lora_A
lora_B = self.lora_B
if isinstance(self.lora_B, DTensor):
lora_B = self.lora_B.to_local()
lora_A = self.lora_A.to_local()
if self.base_layer.input_is_parallel: if self.base_layer.input_is_parallel:
input_parallel = input_ input_parallel = input_
else: else:
@@ -411,6 +447,24 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
output_parallel = self.base_layer.quant_method.apply( output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_parallel self.base_layer, input_parallel
) )
if not self.merged and not self.disable_lora:
lora_dtype = lora_A.dtype
input_parallel_lora = input_parallel.to(dtype=lora_dtype)
lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(device=input_parallel.device, non_blocking=True)
)
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(device=input_parallel.device, non_blocking=True)
)
delta_parallel = input_parallel_lora @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank:
delta_parallel = delta_parallel * (
self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore
delta_parallel = delta_parallel * self.strength
output_parallel = output_parallel + delta_parallel.to(
dtype=output_parallel.dtype
)
if self.base_layer.reduce_results and self.base_layer.tp_size > 1: if self.base_layer.reduce_results and self.base_layer.tp_size > 1:
output_ = tensor_model_parallel_all_reduce(output_parallel) output_ = tensor_model_parallel_all_reduce(output_parallel)
@@ -466,19 +520,23 @@ class LinearWithLoRA(BaseLayerWithLoRA):
# TODO: Support multiple LoRA adapters when use not merged mode # TODO: Support multiple LoRA adapters when use not merged mode
if not self.merged and not self.disable_lora: if not self.merged and not self.disable_lora:
lora_A_sliced = self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)) lora_dtype = lora_A.dtype
lora_B_sliced = self.slice_lora_b_weights(lora_B.to(x, non_blocking=True)) x_lora = x.to(dtype=lora_dtype)
delta = x @ lora_A_sliced.T @ lora_B_sliced.T lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(device=x.device, non_blocking=True)
)
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(device=x.device, non_blocking=True)
)
delta = x_lora @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank: if self.lora_alpha != self.lora_rank:
delta = delta * ( delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore ) # type: ignore
delta = delta * self.strength delta = delta * self.strength
if delta.dim() > 2:
delta = delta.reshape(-1, delta.shape[-1])
# nn.Linear.forward() returns a single tensor, not a tuple # nn.Linear.forward() returns a single tensor, not a tuple
out = self.base_layer(x) out = self.base_layer(x)
return out + delta return out + delta.to(dtype=out.dtype)
else: else:
# nn.Linear.forward() returns a single tensor # nn.Linear.forward() returns a single tensor
out = self.base_layer(x) out = self.base_layer(x)
@@ -15,6 +15,7 @@ from torch import nn
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.configs.models import ModelConfig
from sglang.multimodal_gen.configs.pipeline_configs.flux import Flux2PipelineConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.loader.utils import (
_normalize_component_type, _normalize_component_type,
@@ -294,9 +295,15 @@ class TokenizerLoader(ComponentLoader):
def load_customized( def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str self, component_model_path: str, server_args: ServerArgs, component_name: str
) -> Any: ) -> Any:
# Flux.2 aligns to the tokenizer defaults from the original baseline.
# TODO: abstract this
if isinstance(server_args.pipeline_config, Flux2PipelineConfig):
return AutoTokenizer.from_pretrained(component_model_path)
return AutoTokenizer.from_pretrained( return AutoTokenizer.from_pretrained(
component_model_path, component_model_path,
padding_size="right", padding_side="right",
use_fast=True,
) )
@@ -0,0 +1,223 @@
import glob
import json
import os
import re
import safetensors
import torch
from safetensors.torch import load_file as safetensors_load_file
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader,
)
from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import (
LatentUpsampler,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
UPSAMPLER_CONSTRUCTOR_KEYS = {
"in_channels",
"mid_channels",
"num_blocks_per_stage",
"dims",
"spatial_upsample",
"temporal_upsample",
"spatial_scale",
"rational_resampler",
}
_HF_BLOB_URL_RE = re.compile(
r"https?://huggingface\.co/([^/]+/[^/]+)/blob/([^/]+)/(.*)"
)
_HF_RESOLVE_URL_RE = re.compile(
r"https?://huggingface\.co/([^/]+/[^/]+)/resolve/([^/]+)/(.*)"
)
def _parse_hf_url(path: str):
m = _HF_BLOB_URL_RE.match(path) or _HF_RESOLVE_URL_RE.match(path)
if m:
return m.group(1), m.group(2), m.group(3)
return None
def _download_hf_file(repo_id: str, filename: str, revision: str = "main") -> str:
from huggingface_hub import hf_hub_download
logger.info("Downloading %s from %s (revision=%s)", filename, repo_id, revision)
return hf_hub_download(repo_id=repo_id, filename=filename, revision=revision)
def _find_safetensors_file(path: str) -> str:
"""Resolve path to a single safetensors file (local path, directory, HF URL, or HF repo id)."""
if os.path.isfile(path) and path.endswith(".safetensors"):
return path
if os.path.isdir(path):
files = sorted(glob.glob(os.path.join(path, "*.safetensors")))
if len(files) == 1:
return files[0]
elif len(files) > 1:
raise ValueError(
f"Found {len(files)} safetensors files in {path}, expected 1"
)
hf = _parse_hf_url(path)
if hf:
repo_id, revision, filename = hf
return _download_hf_file(repo_id, filename, revision)
try:
maybe_downloaded = maybe_download_model(path)
if os.path.isdir(maybe_downloaded):
files = sorted(glob.glob(os.path.join(maybe_downloaded, "*.safetensors")))
if len(files) == 1:
return files[0]
elif len(files) > 1:
raise ValueError(
f"Found {len(files)} safetensors files in {maybe_downloaded}, expected 1"
)
except Exception:
pass
raise FileNotFoundError(
f"No safetensors file found at {path}. "
"Provide a local .safetensors file, a directory containing one, "
"a HuggingFace URL (https://huggingface.co/<repo>/blob/main/<path>), "
"or a HuggingFace repo id."
)
def _normalize_config(raw: dict) -> dict:
"""Map diffusers / original-repo config fields to LatentUpsampler kwargs."""
config = {k: v for k, v in raw.items() if k in UPSAMPLER_CONSTRUCTOR_KEYS}
# diffusers uses rational_spatial_scale instead of rational_resampler + spatial_scale
if "rational_spatial_scale" in raw and "rational_resampler" not in config:
config["rational_resampler"] = True
config.setdefault("spatial_scale", raw["rational_spatial_scale"])
return config
def _infer_config_from_state_dict(state_dict: dict[str, torch.Tensor]) -> dict:
"""Infer LatentUpsampler kwargs from weight shapes and key names.
Works even when no config.json or safetensors metadata is available.
"""
config: dict = {}
w = state_dict.get("initial_conv.weight")
if w is not None:
config["mid_channels"] = w.shape[0]
config["in_channels"] = w.shape[1]
config["dims"] = 3 if w.ndim == 5 else 2
num_blocks = sum(
1
for k in state_dict
if k.startswith("res_blocks.") and k.endswith(".conv1.weight")
)
if num_blocks > 0:
config["num_blocks_per_stage"] = num_blocks
# Detect upsampler type from key patterns
has_rational = any(k.startswith("upsampler.blur_down.") for k in state_dict)
if has_rational:
config["rational_resampler"] = True
config["spatial_upsample"] = True
config["temporal_upsample"] = False
config["spatial_scale"] = 2.0
else:
up_w = state_dict.get("upsampler.0.weight")
if up_w is not None and up_w.ndim == 5:
ratio = up_w.shape[0] // up_w.shape[1]
if ratio == 8:
config["spatial_upsample"] = True
config["temporal_upsample"] = True
elif ratio == 2:
config["spatial_upsample"] = False
config["temporal_upsample"] = True
else:
config["spatial_upsample"] = True
config["temporal_upsample"] = False
else:
config["spatial_upsample"] = True
config["temporal_upsample"] = False
return config
def _load_config(
safetensors_path: str,
original_path: str,
state_dict: dict[str, torch.Tensor],
) -> dict:
"""Load upsampler config with fallback chain:
1. safetensors metadata ("config" key) - original LTX-2 repo format
2. sibling config.json - diffusers format
3. config.json from HF (if original_path was a URL)
4. infer from state dict shapes (always works)
"""
with safetensors.safe_open(safetensors_path, framework="pt") as f:
meta = f.metadata()
if meta and "config" in meta:
logger.info("Using config from safetensors metadata")
return _normalize_config(json.loads(meta["config"]))
config_json_path = os.path.join(os.path.dirname(safetensors_path), "config.json")
if os.path.isfile(config_json_path):
with open(config_json_path) as fp:
logger.info("Using config from sibling config.json")
return _normalize_config(json.load(fp))
hf = _parse_hf_url(original_path)
if hf:
repo_id, revision, filename = hf
config_filename = os.path.dirname(filename) + "/config.json"
try:
local = _download_hf_file(repo_id, config_filename, revision)
with open(local) as fp:
logger.info("Using config from HF config.json")
return _normalize_config(json.load(fp))
except Exception:
pass
logger.info("No explicit config found, inferring from state dict")
return _infer_config_from_state_dict(state_dict)
class UpsamplerLoader(ComponentLoader):
component_names = ["spatial_upsampler"]
expected_library = "diffusers"
def should_offload(self, server_args: ServerArgs, model_config=None):
return server_args.vae_cpu_offload
def load_customized(
self,
component_model_path: str,
server_args: ServerArgs,
component_name: str,
):
safetensors_path = _find_safetensors_file(component_model_path)
state_dict = safetensors_load_file(safetensors_path)
config = _load_config(safetensors_path, component_model_path, state_dict)
logger.info("Loading LatentUpsampler with config: %s", config)
should_offload = self.should_offload(server_args)
target_device = self.target_device(should_offload)
with torch.device("meta"):
model = LatentUpsampler(**config)
model.load_state_dict(state_dict, assign=True)
model = model.to(device=target_device, dtype=torch.bfloat16).eval()
logger.info("Loaded LatentUpsampler to %s", target_device)
return model
@@ -46,7 +46,7 @@ def apply_split_rotary_emb(
r = last // 2 r = last // 2
# (..., 2, r) # (..., 2, r)
split_x = x.reshape(*x.shape[:-1], 2, r) split_x = x.reshape(*x.shape[:-1], 2, r).float()
first_x = split_x[..., :1, :] # (..., 1, r) first_x = split_x[..., :1, :] # (..., 1, r)
second_x = split_x[..., 1:, :] # (..., 1, r) second_x = split_x[..., 1:, :] # (..., 1, r)
@@ -153,12 +153,6 @@ class LTX2Attention(torch.nn.Module):
query_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, query_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
key_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, key_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor: ) -> torch.Tensor:
batch_size, sequence_length, _ = (
hidden_states.shape
if encoder_hidden_states is None
else encoder_hidden_states.shape
)
if encoder_hidden_states is None: if encoder_hidden_states is None:
encoder_hidden_states = hidden_states encoder_hidden_states = hidden_states
@@ -183,16 +177,26 @@ class LTX2Attention(torch.nn.Module):
key_rotary_emb if key_rotary_emb is not None else query_rotary_emb, key_rotary_emb if key_rotary_emb is not None else query_rotary_emb,
) )
query = query.unflatten(2, (self.heads, -1)) query = query.unflatten(2, (self.heads, -1)).transpose(1, 2)
key = key.unflatten(2, (self.heads, -1)) key = key.unflatten(2, (self.heads, -1)).transpose(1, 2)
value = value.unflatten(2, (self.heads, -1)) value = value.unflatten(2, (self.heads, -1)).transpose(1, 2)
hidden_states = self.attn( if attention_mask is not None:
if attention_mask.ndim == 2:
attention_mask = attention_mask[:, None, None, :]
elif attention_mask.ndim == 3:
attention_mask = attention_mask[:, None, :, :]
attention_mask = attention_mask.to(dtype=query.dtype)
hidden_states = F.scaled_dot_product_attention(
query, query,
key, key,
value, value,
attn_mask=attention_mask,
dropout_p=0.0,
is_causal=False,
) )
hidden_states = hidden_states.flatten(2, 3) hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
hidden_states = hidden_states.to(query.dtype) hidden_states = hidden_states.to(query.dtype)
hidden_states = self.to_out[0](hidden_states) hidden_states = self.to_out[0](hidden_states)
@@ -464,9 +468,7 @@ class LTX2ConnectorTransformer1d(nn.Module):
attention_mask = torch.zeros_like(attention_mask) attention_mask = torch.zeros_like(attention_mask)
# 2. Calculate 1D RoPE positional embeddings # 2. Calculate 1D RoPE positional embeddings
rotary_emb = self.rope( rotary_emb = self.rope(batch_size, seq_len, device=hidden_states.device)
batch_size, seq_len, device=hidden_states.device, dtype=hidden_states.dtype
)
# 3. Run 1D transformer blocks # 3. Run 1D transformer blocks
for block in self.transformer_blocks: for block in self.transformer_blocks:
@@ -21,7 +21,7 @@ from sglang.multimodal_gen.runtime.distributed import (
from sglang.multimodal_gen.runtime.distributed.communication_op import ( 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 USPAttention from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
RowParallelLinear, RowParallelLinear,
@@ -66,7 +66,7 @@ def apply_split_rotary_emb(
) )
r = last // 2 r = last // 2
split_x = x.reshape(*x.shape[:-1], 2, r) split_x = x.reshape(*x.shape[:-1], 2, r).float()
first_x = split_x[..., :1, :] first_x = split_x[..., :1, :]
second_x = split_x[..., 1:, :] second_x = split_x[..., 1:, :]
@@ -137,7 +137,6 @@ class LTX2AudioVideoRotaryPosEmbed(nn.Module):
self.causal_offset = int(causal_offset) self.causal_offset = int(causal_offset)
self.modality = modality self.modality = modality
self.coords_dtype = torch.bfloat16 if modality == "video" else torch.float32
if self.modality not in ["video", "audio"]: if self.modality not in ["video", "audio"]:
raise ValueError( raise ValueError(
f"Modality {modality} is not supported. Supported modalities are `video` and `audio`." f"Modality {modality} is not supported. Supported modalities are `video` and `audio`."
@@ -244,7 +243,6 @@ class LTX2AudioVideoRotaryPosEmbed(nn.Module):
device = device or coords.device device = device or coords.device
num_pos_dims = coords.shape[1] num_pos_dims = coords.shape[1]
coords = coords.to(self.coords_dtype)
if coords.ndim == 4: if coords.ndim == 4:
coords_start, coords_end = coords.chunk(2, dim=-1) coords_start, coords_end = coords.chunk(2, dim=-1)
coords = (coords_start + coords_end) / 2.0 coords = (coords_start + coords_end) / 2.0
@@ -309,9 +307,7 @@ class LTX2AudioVideoRotaryPosEmbed(nn.Module):
cos_freqs = torch.swapaxes(cos_freq, 1, 2) cos_freqs = torch.swapaxes(cos_freq, 1, 2)
sin_freqs = torch.swapaxes(sin_freq, 1, 2) sin_freqs = torch.swapaxes(sin_freq, 1, 2)
# Cast to bf16 to match model weights dtype. coords_dtype controls return cos_freqs, sin_freqs
# intermediate coordinate precision (fp32 for audio) and differs.
return cos_freqs.to(torch.bfloat16), sin_freqs.to(torch.bfloat16)
def rms_norm(x: torch.Tensor, eps: float) -> torch.Tensor: def rms_norm(x: torch.Tensor, eps: float) -> torch.Tensor:
@@ -450,6 +446,7 @@ class LTX2Attention(nn.Module):
dim_head: int = 64, dim_head: int = 64,
norm_eps: float = 1e-6, norm_eps: float = 1e-6,
qk_norm: bool = True, qk_norm: bool = True,
use_local_attention: bool = False,
supported_attention_backends: set[AttentionBackendEnum] | None = None, supported_attention_backends: set[AttentionBackendEnum] | None = None,
prefix: str = "", prefix: str = "",
quant_config: QuantizationConfig | None = None, quant_config: QuantizationConfig | None = None,
@@ -463,6 +460,7 @@ class LTX2Attention(nn.Module):
self.inner_dim = self.heads * self.dim_head self.inner_dim = self.heads * self.dim_head
self.norm_eps = float(norm_eps) self.norm_eps = float(norm_eps)
self.qk_norm = bool(qk_norm) self.qk_norm = bool(qk_norm)
self.use_local_attention = bool(use_local_attention)
tp_size = get_tp_world_size() tp_size = get_tp_world_size()
if tp_size <= 0: if tp_size <= 0:
@@ -531,16 +529,27 @@ class LTX2Attention(nn.Module):
nn.Identity(), nn.Identity(),
) )
self.attn = USPAttention( if self.use_local_attention:
num_heads=self.local_heads, self.attn = LocalAttention(
head_size=self.dim_head, num_heads=self.local_heads,
num_kv_heads=self.local_heads, head_size=self.dim_head,
dropout_rate=0, num_kv_heads=self.local_heads,
softmax_scale=None, softmax_scale=None,
causal=False, causal=False,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn", prefix=f"{prefix}.attn",
) )
else:
self.attn = USPAttention(
num_heads=self.local_heads,
head_size=self.dim_head,
num_kv_heads=self.local_heads,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn",
)
def forward( def forward(
self, self,
@@ -549,66 +558,56 @@ class LTX2Attention(nn.Module):
mask: torch.Tensor | None = None, mask: torch.Tensor | None = None,
pe: tuple[torch.Tensor, torch.Tensor] | None = None, pe: tuple[torch.Tensor, torch.Tensor] | None = None,
k_pe: tuple[torch.Tensor, torch.Tensor] | None = None, k_pe: tuple[torch.Tensor, torch.Tensor] | None = None,
perturbation_mask: torch.Tensor | None = None,
all_perturbed: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
q, _ = self.to_q(x)
context_ = x if context is None else context context_ = x if context is None else context
k, _ = self.to_k(context_)
v, _ = self.to_v(context_) v, _ = self.to_v(context_)
use_attention = not all_perturbed
if self.qk_norm: if use_attention:
assert self.q_norm is not None and self.k_norm is not None q, _ = self.to_q(x)
q = self.q_norm(q) k, _ = self.to_k(context_)
k = self.k_norm(k)
if pe is not None: if self.qk_norm:
cos, sin = pe assert self.q_norm is not None and self.k_norm is not None
k_cos, k_sin = pe if k_pe is None else k_pe q = self.q_norm(q)
tp_size = get_tp_world_size() k = self.k_norm(k)
if tp_size > 1:
tp_rank = get_tp_rank() if pe is not None:
cos, sin = self._slice_rope_for_tp( cos, sin = pe
cos, sin, tp_rank=tp_rank, tp_size=tp_size k_cos, k_sin = pe if k_pe is None else k_pe
) tp_size = get_tp_world_size()
k_cos, k_sin = self._slice_rope_for_tp( if tp_size > 1:
k_cos, k_sin, tp_rank=tp_rank, tp_size=tp_size tp_rank = get_tp_rank()
) cos, sin = self._slice_rope_for_tp(
if cos.dim() == 3: cos, sin, tp_rank=tp_rank, tp_size=tp_size
q = apply_interleaved_rotary_emb(q, (cos, sin)) )
k = apply_interleaved_rotary_emb(k, (k_cos, k_sin)) k_cos, k_sin = self._slice_rope_for_tp(
else: k_cos, k_sin, tp_rank=tp_rank, tp_size=tp_size
q = apply_split_rotary_emb(q, (cos, sin)) )
k = apply_split_rotary_emb(k, (k_cos, k_sin)) if cos.dim() == 3:
q = apply_interleaved_rotary_emb(q, (cos, sin))
k = apply_interleaved_rotary_emb(k, (k_cos, k_sin))
else:
q = apply_split_rotary_emb(q, (cos, sin))
k = apply_split_rotary_emb(k, (k_cos, k_sin))
q = q.view(*q.shape[:-1], self.local_heads, self.dim_head)
k = k.view(*k.shape[:-1], self.local_heads, self.dim_head)
v = v.view(*v.shape[:-1], self.local_heads, self.dim_head) v = v.view(*v.shape[:-1], self.local_heads, self.dim_head)
if use_attention:
q = q.view(*q.shape[:-1], self.local_heads, self.dim_head)
k = k.view(*k.shape[:-1], self.local_heads, self.dim_head)
if mask is not None: if self.use_local_attention:
# Fallback to SDPA for masked attention out = self.attn(q, k, v, attn_mask=mask)
q_ = q.transpose(1, 2)
k_ = k.transpose(1, 2)
v_ = v.transpose(1, 2)
if torch.is_floating_point(mask):
m = mask
if m.dim() == 2:
m = m[:, None, None, :]
elif m.dim() == 3:
m = m[:, None, :, :]
sdpa_mask = m.to(dtype=q_.dtype, device=q_.device)
else: else:
m = mask.to(dtype=q_.dtype, device=q_.device) out = self.attn(q, k, v)
if m.dim() == 2:
m = m[:, None, None, :]
elif m.dim() == 3:
m = m[:, None, :, :]
sdpa_mask = (m - 1.0) * torch.finfo(q_.dtype).max
out = torch.nn.functional.scaled_dot_product_attention( if perturbation_mask is not None:
q_, k_, v_, attn_mask=sdpa_mask, dropout_p=0.0, is_causal=False out = out * perturbation_mask + v * (1 - perturbation_mask)
).transpose(1, 2)
else: if not use_attention:
out = self.attn(q, k, v) out = v
out = out.flatten(2) out = out.flatten(2)
out, _ = self.to_out[0](out) out, _ = self.to_out[0](out)
@@ -720,6 +719,8 @@ class LTX2TransformerBlock(nn.Module):
) )
# 2. Prompt Cross-Attention # 2. Prompt Cross-Attention
# Prompt KV is replicated across SP ranks, so prompt cross-attn should
# stay local and preserve the explicit KV mask semantics from official.
self.attn2 = LTX2Attention( self.attn2 = LTX2Attention(
query_dim=dim, query_dim=dim,
context_dim=cross_attention_dim, context_dim=cross_attention_dim,
@@ -727,6 +728,7 @@ class LTX2TransformerBlock(nn.Module):
dim_head=attention_head_dim, dim_head=attention_head_dim,
norm_eps=norm_eps, norm_eps=norm_eps,
qk_norm=qk_norm, qk_norm=qk_norm,
use_local_attention=True,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn2", prefix=f"{prefix}.attn2",
quant_config=quant_config, quant_config=quant_config,
@@ -738,6 +740,7 @@ class LTX2TransformerBlock(nn.Module):
dim_head=audio_attention_head_dim, dim_head=audio_attention_head_dim,
norm_eps=norm_eps, norm_eps=norm_eps,
qk_norm=qk_norm, qk_norm=qk_norm,
use_local_attention=True,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.audio_attn2", prefix=f"{prefix}.audio_attn2",
quant_config=quant_config, quant_config=quant_config,
@@ -822,6 +825,10 @@ class LTX2TransformerBlock(nn.Module):
audio_encoder_attention_mask: Optional[torch.Tensor] = None, audio_encoder_attention_mask: Optional[torch.Tensor] = None,
a2v_cross_attention_mask: Optional[torch.Tensor] = None, a2v_cross_attention_mask: Optional[torch.Tensor] = None,
v2a_cross_attention_mask: Optional[torch.Tensor] = None, v2a_cross_attention_mask: Optional[torch.Tensor] = None,
skip_video_self_attn: bool = False,
skip_audio_self_attn: bool = False,
skip_a2v_cross_attn: bool = False,
skip_v2a_cross_attn: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
batch_size = hidden_states.size(0) batch_size = hidden_states.size(0)
@@ -833,7 +840,11 @@ class LTX2TransformerBlock(nn.Module):
norm_hidden_states = ( norm_hidden_states = (
rms_norm(hidden_states, self.norm_eps) * (1 + vscale_msa) + vshift_msa rms_norm(hidden_states, self.norm_eps) * (1 + vscale_msa) + vshift_msa
) )
attn_hidden_states = self.attn1(norm_hidden_states, pe=video_rotary_emb) attn_hidden_states = self.attn1(
norm_hidden_states,
pe=video_rotary_emb,
all_perturbed=skip_video_self_attn,
)
hidden_states = hidden_states + attn_hidden_states * vgate_msa hidden_states = hidden_states + attn_hidden_states * vgate_msa
ashift_msa, ascale_msa, agate_msa = self.get_ada_values( ashift_msa, ascale_msa, agate_msa = self.get_ada_values(
@@ -843,10 +854,11 @@ class LTX2TransformerBlock(nn.Module):
rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_msa) + ashift_msa rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_msa) + ashift_msa
) )
attn_audio_hidden_states = self.audio_attn1( attn_audio_hidden_states = self.audio_attn1(
norm_audio_hidden_states, pe=audio_rotary_emb norm_audio_hidden_states,
pe=audio_rotary_emb,
all_perturbed=skip_audio_self_attn,
) )
audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa
# 2. Prompt Cross-Attention # 2. Prompt Cross-Attention
norm_hidden_states = rms_norm(hidden_states, self.norm_eps) norm_hidden_states = rms_norm(hidden_states, self.norm_eps)
attn_hidden_states = self.attn2( attn_hidden_states = self.attn2(
@@ -863,7 +875,6 @@ class LTX2TransformerBlock(nn.Module):
mask=audio_encoder_attention_mask, mask=audio_encoder_attention_mask,
) )
audio_hidden_states = audio_hidden_states + attn_audio_hidden_states audio_hidden_states = audio_hidden_states + attn_audio_hidden_states
# 3. Audio-to-Video and Video-to-Audio Cross-Attention # 3. Audio-to-Video and Video-to-Audio Cross-Attention
norm_hidden_states = rms_norm(hidden_states, self.norm_eps) norm_hidden_states = rms_norm(hidden_states, self.norm_eps)
norm_audio_hidden_states = rms_norm(audio_hidden_states, self.norm_eps) norm_audio_hidden_states = rms_norm(audio_hidden_states, self.norm_eps)
@@ -934,14 +945,15 @@ class LTX2TransformerBlock(nn.Module):
norm_audio_hidden_states * (1 + audio_a2v_ca_scale) + audio_a2v_ca_shift norm_audio_hidden_states * (1 + audio_a2v_ca_scale) + audio_a2v_ca_shift
) )
a2v_attn_hidden_states = self.audio_to_video_attn( if not skip_a2v_cross_attn:
mod_norm_hidden_states, a2v_attn_hidden_states = self.audio_to_video_attn(
context=mod_norm_audio_hidden_states, mod_norm_hidden_states,
pe=ca_video_rotary_emb, context=mod_norm_audio_hidden_states,
k_pe=ca_audio_rotary_emb, pe=ca_video_rotary_emb,
mask=a2v_cross_attention_mask, k_pe=ca_audio_rotary_emb,
) mask=a2v_cross_attention_mask,
hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states )
hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states
# V2A # V2A
mod_norm_hidden_states = ( mod_norm_hidden_states = (
@@ -951,15 +963,17 @@ class LTX2TransformerBlock(nn.Module):
norm_audio_hidden_states * (1 + audio_v2a_ca_scale) + audio_v2a_ca_shift norm_audio_hidden_states * (1 + audio_v2a_ca_scale) + audio_v2a_ca_shift
) )
v2a_attn_hidden_states = self.video_to_audio_attn( if not skip_v2a_cross_attn:
mod_norm_audio_hidden_states, v2a_attn_hidden_states = self.video_to_audio_attn(
context=mod_norm_hidden_states, mod_norm_audio_hidden_states,
pe=ca_audio_rotary_emb, context=mod_norm_hidden_states,
k_pe=ca_video_rotary_emb, pe=ca_audio_rotary_emb,
mask=v2a_cross_attention_mask, k_pe=ca_video_rotary_emb,
) mask=v2a_cross_attention_mask,
audio_hidden_states = audio_hidden_states + v2a_gate * v2a_attn_hidden_states )
audio_hidden_states = (
audio_hidden_states + v2a_gate * v2a_attn_hidden_states
)
# 4. Feedforward # 4. Feedforward
vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values( vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(3, None) self.scale_shift_table, batch_size, temb, slice(3, None)
@@ -978,7 +992,6 @@ class LTX2TransformerBlock(nn.Module):
) )
audio_ff_output = self.audio_ff(norm_audio_hidden_states) audio_ff_output = self.audio_ff(norm_audio_hidden_states)
audio_hidden_states = audio_hidden_states + audio_ff_output * agate_mlp audio_hidden_states = audio_hidden_states + audio_ff_output * agate_mlp
return hidden_states, audio_hidden_states return hidden_states, audio_hidden_states
@@ -1275,6 +1288,10 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_num_frames: Optional[int] = None, audio_num_frames: Optional[int] = None,
video_coords: Optional[torch.Tensor] = None, video_coords: Optional[torch.Tensor] = None,
audio_coords: Optional[torch.Tensor] = None, audio_coords: Optional[torch.Tensor] = None,
skip_video_self_attn_blocks: Optional[tuple[int, ...]] = None,
skip_audio_self_attn_blocks: Optional[tuple[int, ...]] = None,
disable_a2v_cross_attn: bool = False,
disable_v2a_cross_attn: bool = False,
**kwargs, **kwargs,
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ) -> tuple[torch.Tensor | None, torch.Tensor | None]:
@@ -1333,7 +1350,6 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
# 2. Patchify input projections # 2. Patchify input projections
hidden_states, _ = self.patchify_proj(hidden_states) hidden_states, _ = self.patchify_proj(hidden_states)
audio_hidden_states, _ = self.audio_patchify_proj(audio_hidden_states) audio_hidden_states, _ = self.audio_patchify_proj(audio_hidden_states)
# 3. Prepare timestep embeddings # 3. Prepare timestep embeddings
# 3.1. Prepare global modality (video and audio) timestep embedding and modulation parameters # 3.1. Prepare global modality (video and audio) timestep embedding and modulation parameters
temb, embedded_timestep = self.adaln_single( temb, embedded_timestep = self.adaln_single(
@@ -1391,8 +1407,9 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_encoder_hidden_states = self.audio_caption_projection( audio_encoder_hidden_states = self.audio_caption_projection(
audio_encoder_hidden_states audio_encoder_hidden_states
) )
# 5. Run blocks # 5. Run blocks
skip_video_self_attn_blocks = set(skip_video_self_attn_blocks or ())
skip_audio_self_attn_blocks = set(skip_audio_self_attn_blocks or ())
for block in self.transformer_blocks: for block in self.transformer_blocks:
hidden_states, audio_hidden_states = block( hidden_states, audio_hidden_states = block(
hidden_states, hidden_states,
@@ -1414,6 +1431,10 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
ca_audio_rotary_emb=ca_audio_rotary_emb, ca_audio_rotary_emb=ca_audio_rotary_emb,
encoder_attention_mask=encoder_attention_mask, encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=audio_encoder_attention_mask, audio_encoder_attention_mask=audio_encoder_attention_mask,
skip_video_self_attn=block.idx in skip_video_self_attn_blocks,
skip_audio_self_attn=block.idx in skip_audio_self_attn_blocks,
skip_a2v_cross_attn=disable_a2v_cross_attn,
skip_v2a_cross_attn=disable_v2a_cross_attn,
) )
# 6. Output layers # 6. Output layers
@@ -1439,7 +1460,6 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_hidden_states = self.audio_norm_out(audio_hidden_states) audio_hidden_states = self.audio_norm_out(audio_hidden_states)
audio_hidden_states = audio_hidden_states * (1 + audio_scale) + audio_shift audio_hidden_states = audio_hidden_states * (1 + audio_scale) + audio_shift
audio_hidden_states, _ = self.audio_proj_out(audio_hidden_states) audio_hidden_states, _ = self.audio_proj_out(audio_hidden_states)
# Unpatchify if requested (default True for pipeline compatibility) # Unpatchify if requested (default True for pipeline compatibility)
return_latents = kwargs.get("return_latents", True) return_latents = kwargs.get("return_latents", True)
@@ -146,23 +146,61 @@ class Gemma3Attention(nn.Module):
prefix=f"{prefix}.o_proj", prefix=f"{prefix}.o_proj",
) )
self.layer_type = (
config.text_config.layer_types[layer_id]
if hasattr(config.text_config, "layer_types")
else None
)
self.is_sliding = ( self.is_sliding = (
config.text_config.layer_types[layer_id] == "sliding_attention" config.text_config.layer_types[layer_id] == "sliding_attention"
) )
rope_parameters = getattr(config.text_config, "rope_parameters", None) or {}
layer_rope_params = {}
if self.layer_type is not None and isinstance(rope_parameters, dict):
layer_rope_params = dict(rope_parameters.get(self.layer_type) or {})
# Initialize the rotary embedding. # Initialize the rotary embedding.
if self.is_sliding: if self.is_sliding:
# Local attention. # Local attention.
self.rope_theta = config.text_config.rope_local_base_freq self.rope_theta = float(
rope_scaling = None # Default layer_rope_params.get(
"rope_theta",
getattr(
config.text_config,
"rope_local_base_freq",
getattr(
getattr(config.text_config, "default_theta", {}),
"get",
lambda *_: 10_000.0,
)("local", 10_000.0),
),
)
)
rope_scaling = layer_rope_params or None
# sliding window # sliding window
self.sliding_window = get_attention_sliding_window_size(config.text_config) self.sliding_window = get_attention_sliding_window_size(config.text_config)
# (left, right) = (window, 0) effectively for causal # (left, right) = (window, 0) effectively for causal
self.window_size = (self.sliding_window, 0) self.window_size = (self.sliding_window, 0)
else: else:
# Global attention. # Global attention.
self.rope_theta = config.text_config.rope_theta self.rope_theta = float(
rope_scaling = config.text_config.rope_scaling layer_rope_params.get(
"rope_theta",
getattr(
config.text_config,
"rope_theta",
getattr(
getattr(config.text_config, "default_theta", {}),
"get",
lambda *_: 1_000_000.0,
)("global", 1_000_000.0),
),
)
)
rope_scaling = layer_rope_params or getattr(
config.text_config, "rope_scaling", None
)
self.sliding_window = None self.sliding_window = None
self.window_size = (-1, -1) self.window_size = (-1, -1)
@@ -734,7 +772,9 @@ class Gemma3TextModel(nn.Module):
layer_id=i, layer_id=i,
config=config, config=config,
quant_config=self.quant_config, quant_config=self.quant_config,
prefix=f"{config.text_config.prefix}.layers.{i}", prefix=add_prefix(
f"layers.{i}", getattr(config.text_config, "prefix", "")
),
) )
for i in range(config.text_config.num_hidden_layers) for i in range(config.text_config.num_hidden_layers)
] ]
@@ -0,0 +1,5 @@
from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import (
LatentUpsampler,
)
__all__ = ["LatentUpsampler"]
@@ -0,0 +1,268 @@
# Ported from https://github.com/Lightricks/LTX-2
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Optional, Tuple
import torch
import torch.nn.functional as F
from einops import rearrange
class BlurDownsample(torch.nn.Module):
"""Anti-aliased spatial downsampling by integer stride using a fixed separable binomial kernel."""
def __init__(self, dims: int, stride: int, kernel_size: int = 5) -> None:
super().__init__()
assert dims in (2, 3)
assert isinstance(stride, int) and stride >= 1
assert kernel_size >= 3 and kernel_size % 2 == 1
self.dims = dims
self.stride = stride
self.kernel_size = kernel_size
k = torch.tensor([math.comb(kernel_size - 1, i) for i in range(kernel_size)])
k2d = k[:, None] @ k[None, :]
k2d = (k2d / k2d.sum()).float()
self.register_buffer("kernel", k2d[None, None, :, :])
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.stride == 1:
return x
if self.dims == 2:
return self._apply_2d(x)
b, _, f, _, _ = x.shape
x = rearrange(x, "b c f h w -> (b f) c h w")
x = self._apply_2d(x)
h2, w2 = x.shape[-2:]
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f, h=h2, w=w2)
return x
def _apply_2d(self, x2d: torch.Tensor) -> torch.Tensor:
c = x2d.shape[1]
weight = self.kernel.expand(c, 1, self.kernel_size, self.kernel_size)
x2d = F.conv2d(
x2d,
weight=weight,
bias=None,
stride=self.stride,
padding=self.kernel_size // 2,
groups=c,
)
return x2d
class PixelShuffleND(torch.nn.Module):
"""N-dimensional pixel shuffle for upsampling tensors."""
def __init__(self, dims: int, upscale_factors: Tuple[int, int, int] = (2, 2, 2)):
super().__init__()
assert dims in [1, 2, 3]
self.dims = dims
self.upscale_factors = upscale_factors
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.dims == 3:
return rearrange(
x,
"b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3)",
p1=self.upscale_factors[0],
p2=self.upscale_factors[1],
p3=self.upscale_factors[2],
)
elif self.dims == 2:
return rearrange(
x,
"b (c p1 p2) h w -> b c (h p1) (w p2)",
p1=self.upscale_factors[0],
p2=self.upscale_factors[1],
)
elif self.dims == 1:
return rearrange(
x,
"b (c p1) f h w -> b c (f p1) h w",
p1=self.upscale_factors[0],
)
else:
raise ValueError(f"Unsupported dims: {self.dims}")
class ResBlock(torch.nn.Module):
"""Residual block with two conv layers, group norm, and SiLU activation."""
def __init__(
self, channels: int, mid_channels: Optional[int] = None, dims: int = 3
):
super().__init__()
if mid_channels is None:
mid_channels = channels
conv = torch.nn.Conv2d if dims == 2 else torch.nn.Conv3d
self.conv1 = conv(channels, mid_channels, kernel_size=3, padding=1)
self.norm1 = torch.nn.GroupNorm(32, mid_channels)
self.conv2 = conv(mid_channels, channels, kernel_size=3, padding=1)
self.norm2 = torch.nn.GroupNorm(32, channels)
self.activation = torch.nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
residual = x
x = self.conv1(x)
x = self.norm1(x)
x = self.activation(x)
x = self.conv2(x)
x = self.norm2(x)
x = self.activation(x + residual)
return x
def _rational_for_scale(scale: float) -> Tuple[int, int]:
mapping = {0.75: (3, 4), 1.5: (3, 2), 2.0: (2, 1), 4.0: (4, 1)}
if float(scale) not in mapping:
raise ValueError(
f"Unsupported scale {scale}. Choose from {list(mapping.keys())}"
)
return mapping[float(scale)]
class SpatialRationalResampler(torch.nn.Module):
"""Fully-learned rational spatial scaling via PixelShuffle + anti-aliased downsample."""
def __init__(self, mid_channels: int, scale: float):
super().__init__()
self.scale = float(scale)
self.num, self.den = _rational_for_scale(self.scale)
self.conv = torch.nn.Conv2d(
mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1
)
self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num))
self.blur_down = BlurDownsample(dims=2, stride=self.den)
def forward(self, x: torch.Tensor) -> torch.Tensor:
b, _, f, _, _ = x.shape
x = rearrange(x, "b c f h w -> (b f) c h w")
x = self.conv(x)
x = self.pixel_shuffle(x)
x = self.blur_down(x)
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
return x
class LatentUpsampler(torch.nn.Module):
"""
Upsample VAE latents spatially and/or temporally.
Args:
in_channels: Number of channels in the input latent.
mid_channels: Number of channels in the middle layers.
num_blocks_per_stage: Number of ResBlocks per stage (pre/post upsampling).
dims: Dimensionality of convolutions (2 or 3).
spatial_upsample: Whether to spatially upsample.
temporal_upsample: Whether to temporally upsample.
spatial_scale: Scale factor for spatial upsampling.
rational_resampler: Whether to use rational resampler for spatial upsampling.
"""
def __init__(
self,
in_channels: int = 128,
mid_channels: int = 512,
num_blocks_per_stage: int = 4,
dims: int = 3,
spatial_upsample: bool = True,
temporal_upsample: bool = False,
spatial_scale: float = 2.0,
rational_resampler: bool = False,
):
super().__init__()
self.in_channels = in_channels
self.mid_channels = mid_channels
self.num_blocks_per_stage = num_blocks_per_stage
self.dims = dims
self.spatial_upsample = spatial_upsample
self.temporal_upsample = temporal_upsample
self.spatial_scale = float(spatial_scale)
self.rational_resampler = rational_resampler
conv = torch.nn.Conv2d if dims == 2 else torch.nn.Conv3d
self.initial_conv = conv(in_channels, mid_channels, kernel_size=3, padding=1)
self.initial_norm = torch.nn.GroupNorm(32, mid_channels)
self.initial_activation = torch.nn.SiLU()
self.res_blocks = torch.nn.ModuleList(
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
)
if spatial_upsample and temporal_upsample:
self.upsampler = torch.nn.Sequential(
torch.nn.Conv3d(
mid_channels, 8 * mid_channels, kernel_size=3, padding=1
),
PixelShuffleND(3),
)
elif spatial_upsample:
if rational_resampler:
self.upsampler = SpatialRationalResampler(
mid_channels=mid_channels, scale=self.spatial_scale
)
else:
self.upsampler = torch.nn.Sequential(
torch.nn.Conv2d(
mid_channels, 4 * mid_channels, kernel_size=3, padding=1
),
PixelShuffleND(2),
)
elif temporal_upsample:
self.upsampler = torch.nn.Sequential(
torch.nn.Conv3d(
mid_channels, 2 * mid_channels, kernel_size=3, padding=1
),
PixelShuffleND(1),
)
else:
raise ValueError(
"Either spatial_upsample or temporal_upsample must be True"
)
self.post_upsample_res_blocks = torch.nn.ModuleList(
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
)
self.final_conv = conv(mid_channels, in_channels, kernel_size=3, padding=1)
def forward(self, latent: torch.Tensor) -> torch.Tensor:
b, _, f, _, _ = latent.shape
if self.dims == 2:
x = rearrange(latent, "b c f h w -> (b f) c h w")
x = self.initial_conv(x)
x = self.initial_norm(x)
x = self.initial_activation(x)
for block in self.res_blocks:
x = block(x)
x = self.upsampler(x)
for block in self.post_upsample_res_blocks:
x = block(x)
x = self.final_conv(x)
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
else:
x = self.initial_conv(latent)
x = self.initial_norm(x)
x = self.initial_activation(x)
for block in self.res_blocks:
x = block(x)
if self.temporal_upsample:
x = self.upsampler(x)
x = x[:, :, 1:, :, :]
elif isinstance(self.upsampler, SpatialRationalResampler):
x = self.upsampler(x)
else:
x = rearrange(x, "b c f h w -> (b f) c h w")
x = self.upsampler(x)
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
for block in self.post_upsample_res_blocks:
x = block(x)
x = self.final_conv(x)
return x
@@ -843,8 +843,8 @@ class AutoencoderKLLTX2Audio(ParallelTiledVAE):
# Per-channel statistics for normalizing and denormalizing the latent representation. This statistics is computed over # Per-channel statistics for normalizing and denormalizing the latent representation. This statistics is computed over
# the entire dataset and stored in model's checkpoint under AudioVAE state_dict # the entire dataset and stored in model's checkpoint under AudioVAE state_dict
latents_std = torch.zeros((base_channels,)) latents_std = torch.ones((base_channels,))
latents_mean = torch.ones((base_channels,)) latents_mean = torch.zeros((base_channels,))
self.register_buffer("latents_mean", latents_mean, persistent=True) self.register_buffer("latents_mean", latents_mean, persistent=True)
self.register_buffer("latents_std", latents_std, persistent=True) self.register_buffer("latents_std", latents_std, persistent=True)
@@ -1135,6 +1135,28 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
config.arch_config.decoder_spatio_temporal_scaling config.arch_config.decoder_spatio_temporal_scaling
) )
decoder_layers_per_block = config.arch_config.decoder_layers_per_block decoder_layers_per_block = config.arch_config.decoder_layers_per_block
decoder_inject_noise = getattr(
config.arch_config, "decoder_inject_noise", (False, False, False, False)
)
if isinstance(decoder_inject_noise, bool):
decoder_inject_noise = (decoder_inject_noise,) * 4
else:
decoder_inject_noise = tuple(decoder_inject_noise)
upsample_residual = getattr(
config.arch_config, "upsample_residual", (True, True, True)
)
if isinstance(upsample_residual, bool):
upsample_residual = (upsample_residual,) * 3
else:
upsample_residual = tuple(upsample_residual)
upsample_factor = getattr(config.arch_config, "upsample_factor", (2, 2, 2))
if isinstance(upsample_factor, int):
upsample_factor = (upsample_factor,) * 3
else:
upsample_factor = tuple(upsample_factor)
timestep_conditioning = getattr(
config.arch_config, "timestep_conditioning", False
)
decoder_causal = config.arch_config.decoder_causal decoder_causal = config.arch_config.decoder_causal
decoder_spatial_padding_mode = config.arch_config.decoder_spatial_padding_mode decoder_spatial_padding_mode = config.arch_config.decoder_spatial_padding_mode
@@ -1154,16 +1176,20 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
) )
self.decoder = LTX2VideoDecoder3d( self.decoder = LTX2VideoDecoder3d(
latent_channels, in_channels=latent_channels,
out_channels, out_channels=out_channels,
decoder_block_out_channels, block_out_channels=decoder_block_out_channels,
decoder_spatio_temporal_scaling, spatio_temporal_scaling=decoder_spatio_temporal_scaling,
decoder_layers_per_block, layers_per_block=decoder_layers_per_block,
patch_size, patch_size=patch_size,
patch_size_t, patch_size_t=patch_size_t,
resnet_norm_eps, resnet_norm_eps=resnet_norm_eps,
decoder_causal, is_causal=decoder_causal,
decoder_spatial_padding_mode, inject_noise=decoder_inject_noise,
timestep_conditioning=timestep_conditioning,
upsample_residual=upsample_residual,
upsample_factor=upsample_factor,
spatial_padding_mode=decoder_spatial_padding_mode,
) )
latents_mean = torch.zeros((latent_channels,), requires_grad=False) latents_mean = torch.zeros((latent_channels,), requires_grad=False)
@@ -1,5 +1,3 @@
import inspect
import json
import math import math
import os import os
@@ -7,103 +5,183 @@ import numpy as np
import torch import torch
from diffusers import FlowMatchEulerDiscreteScheduler from diffusers import FlowMatchEulerDiscreteScheduler
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase, ComposedPipelineBase,
) )
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages import ( from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage, InputValidationStage,
LTX2AVDecodingStage, LTX2AVDecodingStage,
LTX2AVDenoisingStage, LTX2AVDenoisingStage,
LTX2AVLatentPreparationStage, LTX2AVLatentPreparationStage,
LTX2HalveResolutionStage,
LTX2LoRASwitchStage,
LTX2RefinementStage,
LTX2TextConnectorStage, LTX2TextConnectorStage,
LTX2UpsampleStage,
TextEncodingStage, TextEncodingStage,
) )
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
BASE_SHIFT_ANCHOR = 1024
def calculate_shift( MAX_SHIFT_ANCHOR = 4096
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
def prepare_mu(batch: Req, server_args: ServerArgs): def _resolve_ltx2_two_stage_component_paths(
height = batch.height model_path: str, component_paths: dict[str, str]
width = batch.width ) -> dict[str, str]:
num_frames = batch.num_frames resolved = dict(component_paths)
auto_resolved = []
vae_arch = getattr( if "spatial_upsampler" not in resolved:
getattr(server_args.pipeline_config, "vae_config", None), "arch_config", None spatial_candidates = [
os.path.join(model_path, "latent_upsampler"),
os.path.join(model_path, "ltx-2-spatial-upscaler-x2-1.0.safetensors"),
]
for candidate in spatial_candidates:
if os.path.exists(candidate):
resolved["spatial_upsampler"] = candidate
auto_resolved.append(f"spatial_upsampler={candidate}")
break
if "distilled_lora" not in resolved:
distilled_lora = os.path.join(
model_path, "ltx-2-19b-distilled-lora-384.safetensors"
)
if os.path.exists(distilled_lora):
resolved["distilled_lora"] = distilled_lora
auto_resolved.append(f"distilled_lora={distilled_lora}")
if auto_resolved:
logger.info(
"Auto-resolved LTX2 two-stage components: %s", ", ".join(auto_resolved)
)
return resolved
def calculate_ltx2_shift(
image_seq_len: int,
base_seq_len: int = BASE_SHIFT_ANCHOR,
max_seq_len: int = MAX_SHIFT_ANCHOR,
base_shift: float = 0.95,
max_shift: float = 2.05,
) -> float:
mm = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - mm * base_seq_len
return image_seq_len * mm + b
def prepare_ltx2_mu(batch: Req, server_args: ServerArgs):
latent_num_frames = (int(batch.num_frames) - 1) // int(
server_args.pipeline_config.vae_temporal_compression
) + 1
latent_height = int(batch.height) // int(
server_args.pipeline_config.vae_scale_factor
) )
vae_scale_factor = ( latent_width = int(batch.width) // int(server_args.pipeline_config.vae_scale_factor)
getattr(vae_arch, "spatial_compression_ratio", None) video_sequence_length = latent_num_frames * latent_height * latent_width
or getattr(vae_arch, "vae_scale_factor", None) return "mu", calculate_ltx2_shift(video_sequence_length)
or getattr(server_args.pipeline_config, "vae_scale_factor", None)
class LTX2SigmaPreparationStage(PipelineStage):
"""Prepare native LTX-2 sigma schedule before timestep setup."""
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.extra["ltx2_phase"] = "stage1"
batch.sigmas = np.linspace(
1.0,
1.0 / int(batch.num_inference_steps),
int(batch.num_inference_steps),
).tolist()
return batch
def _add_ltx2_front_stages(pipeline: ComposedPipelineBase):
pipeline.add_stages(
[
InputValidationStage(),
TextEncodingStage(
text_encoders=[pipeline.get_module("text_encoder")],
tokenizers=[pipeline.get_module("tokenizer")],
),
LTX2TextConnectorStage(connectors=pipeline.get_module("connectors")),
]
) )
vae_temporal_compression = getattr(
vae_arch, "temporal_compression_ratio", None
) or getattr(server_args.pipeline_config, "vae_temporal_compression", None)
# Values from LTX2Pipeline in diffusers
mu = calculate_shift( def _add_ltx2_stage1_generation_stages(pipeline: ComposedPipelineBase):
4096, pipeline.add_stage(LTX2SigmaPreparationStage())
base_seq_len=1024, pipeline.add_standard_timestep_preparation_stage(
max_seq_len=4096, prepare_extra_kwargs=[prepare_ltx2_mu]
base_shift=0.95, )
max_shift=2.05, pipeline.add_stages(
[
LTX2AVLatentPreparationStage(
scheduler=pipeline.get_module("scheduler"),
transformer=pipeline.get_module("transformer"),
audio_vae=pipeline.get_module("audio_vae"),
),
LTX2AVDenoisingStage(
transformer=pipeline.get_module("transformer"),
scheduler=pipeline.get_module("scheduler"),
vae=pipeline.get_module("vae"),
audio_vae=pipeline.get_module("audio_vae"),
pipeline=pipeline,
),
]
) )
return "mu", mu
def _load_component_config(model_path: str, component_name: str): def _add_ltx2_decoding_stage(pipeline: ComposedPipelineBase):
"""Helper to load component config from model_index.json or config.json""" pipeline.add_stage(
try: LTX2AVDecodingStage(
# Try loading model_index.json first vae=pipeline.get_module("vae"),
index_path = os.path.join(model_path, "model_index.json") audio_vae=pipeline.get_module("audio_vae"),
if os.path.exists(index_path): vocoder=pipeline.get_module("vocoder"),
with open(index_path, "r") as f: pipeline=pipeline,
index = json.load(f) )
)
if component_name in index:
# It's a subfolder
subfolder = index[component_name][1]
config_path = os.path.join(model_path, subfolder, "config.json")
if os.path.exists(config_path):
with open(config_path, "r") as f:
return json.load(f)
# Fallback to direct config.json in subfolder if standard structure
config_path = os.path.join(model_path, component_name, "config.json")
if os.path.exists(config_path):
with open(config_path, "r") as f:
return json.load(f)
except Exception as e:
logger.warning(f"Failed to load config for {component_name}: {e}")
return {}
def _filter_kwargs_for_cls(cls, kwargs):
"""Filter kwargs to only include those accepted by cls.__init__"""
sig = inspect.signature(cls.__init__)
return {k: v for k, v in kwargs.items() if k in sig.parameters}
class LTX2FlowMatchScheduler(FlowMatchEulerDiscreteScheduler): class LTX2FlowMatchScheduler(FlowMatchEulerDiscreteScheduler):
"""Override ``_time_shift_exponential`` to use torch f32 instead of numpy f64.""" """Override ``_time_shift_exponential`` to use torch f32 instead of numpy f64."""
def set_timesteps(
self,
num_inference_steps=None,
device=None,
sigmas=None,
mu=None,
timesteps=None,
):
if sigmas is not None and timesteps is None and mu is None:
sigmas = torch.tensor(sigmas, dtype=torch.float32, device=device)
timesteps = sigmas * self.config.num_train_timesteps
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
self.num_inference_steps = len(timesteps)
self.timesteps = timesteps
self.sigmas = sigmas
self._step_index = None
self._begin_index = None
return
return super().set_timesteps(
num_inference_steps=num_inference_steps,
device=device,
sigmas=sigmas,
mu=mu,
timesteps=timesteps,
)
def _time_shift_exponential(self, mu, sigma, t): def _time_shift_exponential(self, mu, sigma, t):
if isinstance(t, np.ndarray): if isinstance(t, np.ndarray):
t_torch = torch.from_numpy(t).to(torch.float32) t_torch = torch.from_numpy(t).to(torch.float32)
@@ -112,10 +190,7 @@ class LTX2FlowMatchScheduler(FlowMatchEulerDiscreteScheduler):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma) return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
class LTX2Pipeline(ComposedPipelineBase): class _BaseLTX2Pipeline(LoRAPipeline):
# NOTE: must match `model_index.json`'s `_class_name` for native dispatch.
pipeline_name = "LTX2Pipeline"
_required_config_modules = [ _required_config_modules = [
"transformer", "transformer",
"text_encoder", "text_encoder",
@@ -131,41 +206,127 @@ class LTX2Pipeline(ComposedPipelineBase):
orig = self.get_module("scheduler") orig = self.get_module("scheduler")
self.modules["scheduler"] = LTX2FlowMatchScheduler.from_config(orig.config) self.modules["scheduler"] = LTX2FlowMatchScheduler.from_config(orig.config)
class LTX2Pipeline(_BaseLTX2Pipeline):
# Must match model_index.json `_class_name`.
pipeline_name = "LTX2Pipeline"
def create_pipeline_stages(self, server_args: ServerArgs): def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stages( _add_ltx2_front_stages(self)
[ _add_ltx2_stage1_generation_stages(self)
InputValidationStage(), _add_ltx2_decoding_stage(self)
TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")], class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
), pipeline_name = "LTX2TwoStagePipeline"
LTX2TextConnectorStage(connectors=self.get_module("connectors")), STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
]
def initialize_pipeline(self, server_args: ServerArgs):
super().initialize_pipeline(server_args)
server_args.component_paths = _resolve_ltx2_two_stage_component_paths(
self.model_path, server_args.component_paths
) )
self.add_standard_timestep_preparation_stage(prepare_extra_kwargs=[prepare_mu]) upsampler_path = server_args.component_paths.get("spatial_upsampler")
if not upsampler_path:
raise ValueError(
"LTX2TwoStagePipeline requires --spatial-upsampler-path "
"(component_paths['spatial_upsampler'])."
)
module, memory_usage = PipelineComponentLoader.load_component(
component_name="spatial_upsampler",
component_model_path=upsampler_path,
transformers_or_diffusers="diffusers",
server_args=server_args,
)
self.modules["spatial_upsampler"] = module
self.memory_usages["spatial_upsampler"] = memory_usage
distilled_lora_path = server_args.component_paths.get("distilled_lora")
if not distilled_lora_path:
raise ValueError(
"LTX2TwoStagePipeline requires --distilled-lora-path "
"(component_paths['distilled_lora'])."
)
self._distilled_lora_path = distilled_lora_path
self._stage1_lora_path = server_args.lora_path
self._stage1_lora_scale = float(server_args.lora_scale)
self._active_lora_phase = None
def switch_lora_phase(self, phase: str) -> None:
if phase == self._active_lora_phase:
return
if phase == "stage1":
if self._stage1_lora_path:
self.set_lora(
lora_nickname="ltx2_stage1_base",
lora_path=self._stage1_lora_path,
target="transformer",
strength=self._stage1_lora_scale,
)
else:
# Stage 1 must run on the base transformer weights. If stage 2 left the
# distilled adapter active, stage 1 quality drifts away from the official
# two-stage pipeline immediately.
self.deactivate_lora_weights(target="transformer")
elif phase == "stage2":
lora_nicknames = []
lora_paths = []
lora_strengths = []
lora_targets = []
if self._stage1_lora_path:
lora_nicknames.append("ltx2_stage1_base")
lora_paths.append(self._stage1_lora_path)
lora_strengths.append(self._stage1_lora_scale)
lora_targets.append("transformer")
lora_nicknames.append("ltx2_stage2_distilled")
lora_paths.append(self._distilled_lora_path)
lora_strengths.append(1.0)
lora_targets.append("transformer")
self.set_lora(
lora_nickname=lora_nicknames,
lora_path=lora_paths,
target=lora_targets,
strength=lora_strengths,
# Keep the distilled adapter unmerged when it is the only active LoRA.
# Merging it into the base weights makes the subsequent switch back to
# stage 1 depend on unmerge bookkeeping instead of the original base.
merge_weights=self._stage1_lora_path is not None,
)
else:
raise ValueError(f"Unknown LTX2 two-stage LoRA phase: {phase}")
self._active_lora_phase = phase
def create_pipeline_stages(self, server_args: ServerArgs):
_add_ltx2_front_stages(self)
self.add_stage(LTX2HalveResolutionStage())
self.add_stage(
LTX2LoRASwitchStage(pipeline=self, phase="stage1"),
)
_add_ltx2_stage1_generation_stages(self)
self.add_stages( self.add_stages(
[ [
LTX2AVLatentPreparationStage( LTX2UpsampleStage(
scheduler=self.get_module("scheduler"), spatial_upsampler=self.get_module("spatial_upsampler"),
transformer=self.get_module("transformer"),
audio_vae=self.get_module("audio_vae"),
),
LTX2AVDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"), vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"), audio_vae=self.get_module("audio_vae"),
), ),
LTX2AVDecodingStage( (
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
"ltx2_lora_switch_stage2",
),
LTX2RefinementStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
distilled_sigmas=self.STAGE_2_DISTILLED_SIGMA_VALUES,
vae=self.get_module("vae"), vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"), audio_vae=self.get_module("audio_vae"),
vocoder=self.get_module("vocoder"),
pipeline=self,
), ),
] ]
) )
_add_ltx2_decoding_stage(self)
EntryClass = LTX2Pipeline EntryClass = [LTX2Pipeline, LTX2TwoStagePipeline]
@@ -8,7 +8,6 @@ This module defines the base class for pipelines that are composed of multiple s
""" """
import os import os
import re
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Callable, Literal, cast from typing import Any, Callable, Literal, cast
@@ -333,12 +332,7 @@ class ComposedPipelineBase(ABC):
@staticmethod @staticmethod
def _infer_stage_name(stage: PipelineStage) -> str: def _infer_stage_name(stage: PipelineStage) -> str:
class_name = stage.__class__.__name__ return stage.__class__.__name__
# snake_case
name = re.sub(r"(?<!^)(?=[A-Z])", "_", class_name).lower()
if not name.endswith("_stage"):
name += "_stage"
return name
def add_stage( def add_stage(
self, stage: PipelineStage, stage_name: str | None = None self, stage: PipelineStage, stage_name: str | None = None
@@ -397,6 +397,7 @@ class LoRAPipeline(ComposedPipelineBase):
rank: int, rank: int,
strengths: list[float], strengths: list[float],
clear_existing: bool = False, clear_existing: bool = False,
merge_weights: bool = True,
) -> int: ) -> int:
""" """
Apply LoRA weights to the given lora_layers. Supports multiple LoRA adapters. Apply LoRA weights to the given lora_layers. Supports multiple LoRA adapters.
@@ -435,38 +436,29 @@ class LoRAPipeline(ComposedPipelineBase):
lora_A_name in self.lora_adapters[nickname] lora_A_name in self.lora_adapters[nickname]
and lora_B_name in self.lora_adapters[nickname] and lora_B_name in self.lora_adapters[nickname]
): ):
# Some LoRA checkpoints (e.g. Lightning distill) store per-layer alpha as "<layer>.alpha". inferred_rank = int(
# If present, we must apply the standard LoRA scaling: scale = alpha / rank. self.lora_adapters[nickname][lora_A_name].shape[0]
try: )
inferred_rank = int(
self.lora_adapters[nickname][lora_A_name].shape[0]
)
except Exception:
inferred_rank = None
# Default to None for some checkpoints without "<layer>.alpha"
inferred_alpha: int | None = None
alpha_key = name + ".alpha" alpha_key = name + ".alpha"
if alpha_key in self.lora_adapters[nickname]: if alpha_key in self.lora_adapters[nickname]:
try: inferred_alpha = int(
inferred_alpha = int( self.lora_adapters[nickname][alpha_key].item()
self.lora_adapters[nickname][alpha_key].item()
)
except Exception:
inferred_alpha = None
if inferred_rank is not None:
layer.lora_rank = inferred_rank
layer.lora_alpha = (
inferred_alpha
if inferred_alpha is not None
else inferred_rank
) )
else:
# Some distilled LoRAs omit per-layer alpha and rely on the
# default LoRA scale of alpha == rank. Falling back to rank
# keeps the effective delta consistent with the official path.
inferred_alpha = inferred_rank
layer.lora_rank = inferred_rank
layer.lora_alpha = inferred_alpha
layer.set_lora_weights( layer.set_lora_weights(
self.lora_adapters[nickname][lora_A_name], self.lora_adapters[nickname][lora_A_name],
self.lora_adapters[nickname][lora_B_name], self.lora_adapters[nickname][lora_B_name],
lora_path=path, lora_path=path,
strength=lora_strength, strength=lora_strength,
merge_weights=merge_weights,
clear_existing=( clear_existing=(
clear_existing and idx == 0 clear_existing and idx == 0
), # Only clear on first LoRA ), # Only clear on first LoRA
@@ -589,6 +581,7 @@ class LoRAPipeline(ComposedPipelineBase):
lora_path: str | None | list[str | None] = None, lora_path: str | None | list[str | None] = None,
target: str | list[str] = "all", target: str | list[str] = "all",
strength: float | list[float] = 1.0, strength: float | list[float] = 1.0,
merge_weights: bool = True,
): # type: ignore ): # type: ignore
""" """
Load LoRA adapter(s) into the pipeline and apply them to the specified transformer(s). Load LoRA adapter(s) into the pipeline and apply them to the specified transformer(s).
@@ -682,6 +675,7 @@ class LoRAPipeline(ComposedPipelineBase):
rank, rank,
tgt_strengths, tgt_strengths,
clear_existing=True, clear_existing=True,
merge_weights=merge_weights,
) )
adapted_count += count adapted_count += count
self.cur_adapter_name[module_name] = merged_name self.cur_adapter_name[module_name] = merged_name
@@ -689,7 +683,7 @@ class LoRAPipeline(ComposedPipelineBase):
str(p or self.loaded_adapter_paths.get(n, "")) str(p or self.loaded_adapter_paths.get(n, ""))
for n, p in zip(tgt_nicknames, tgt_paths) for n, p in zip(tgt_nicknames, tgt_paths)
) )
self.is_lora_merged[module_name] = True self.is_lora_merged[module_name] = merge_weights
self.cur_adapter_strength[module_name] = tgt_strengths[0] self.cur_adapter_strength[module_name] = tgt_strengths[0]
# Store full configuration for multi-LoRA support (preserves order and all strengths) # Store full configuration for multi-LoRA support (preserves order and all strengths)
self.cur_adapter_config[module_name] = ( self.cur_adapter_config[module_name] = (
@@ -698,7 +692,7 @@ class LoRAPipeline(ComposedPipelineBase):
) )
logger.info( logger.info(
"Rank %d: LoRA adapter(s) %s applied to %d layers (targets: %s, strengths: %s)", "Rank %d: LoRA adapter(s) %s applied to %d layers (targets: %s, strengths: %s, merge_weights=%s)",
rank, rank,
", ".join(map(str, lora_paths)) if lora_paths else None, ", ".join(map(str, lora_paths)) if lora_paths else None,
adapted_count, adapted_count,
@@ -708,8 +702,42 @@ class LoRAPipeline(ComposedPipelineBase):
if len(strengths) > 1 if len(strengths) > 1
else f"{strengths[0]:.2f}" else f"{strengths[0]:.2f}"
), ),
merge_weights,
) )
def deactivate_lora_weights(self, target: str = "all") -> None:
"""
Disable LoRA for the specified target, regardless of whether weights were
merged into the base model or are still active in the wrapped LoRA path.
"""
target_modules, error = self._get_target_lora_layers(target)
if error:
logger.warning("deactivate_lora_weights: %s", error)
if not target_modules:
return
modules_requiring_unmerge = []
for module_name, lora_layers_dict in target_modules:
if self.is_lora_merged.get(module_name, False) or any(
layer.merged for layer in lora_layers_dict.values()
):
modules_requiring_unmerge.append((module_name, lora_layers_dict))
offload_context = self._temporarily_disable_offload(
target_modules=modules_requiring_unmerge
)
with offload_context:
for module_name, lora_layers_dict in target_modules:
for layer in lora_layers_dict.values():
if layer.merged:
layer.unmerge_lora_weights()
if not layer.disable_lora:
layer.disable_lora = True
self.is_lora_merged[module_name] = False
self.cur_adapter_strength.pop(module_name, None)
self.cur_adapter_config.pop(module_name, None)
logger.info("LoRA weights deactivated for %s", module_name)
def merge_lora_weights(self, target: str = "all", strength: float = 1.0) -> None: def merge_lora_weights(self, target: str = "all", strength: float = 1.0) -> None:
""" """
Merge LoRA weights into the base model for the specified target. Merge LoRA weights into the base model for the specified target.
@@ -22,6 +22,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding_av import (
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising_av import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising_av import (
LTX2AVDenoisingStage, LTX2AVDenoisingStage,
LTX2RefinementStage,
) )
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising_dmd import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising_dmd import (
DmdDenoisingStage, DmdDenoisingStage,
@@ -64,6 +65,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
from sglang.multimodal_gen.runtime.pipelines_core.stages.timestep_preparation import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.timestep_preparation import (
TimestepPreparationStage, TimestepPreparationStage,
) )
from sglang.multimodal_gen.runtime.pipelines_core.stages.upsampling import (
LTX2HalveResolutionStage,
LTX2LoRASwitchStage,
LTX2UpsampleStage,
)
__all__ = [ __all__ = [
"PipelineStage", "PipelineStage",
@@ -92,4 +98,9 @@ __all__ = [
"Hunyuan3DPaintPreprocessStage", "Hunyuan3DPaintPreprocessStage",
"Hunyuan3DPaintTexGenStage", "Hunyuan3DPaintTexGenStage",
"Hunyuan3DPaintPostprocessStage", "Hunyuan3DPaintPostprocessStage",
# LTX-2 two-stage
"LTX2RefinementStage",
"LTX2HalveResolutionStage",
"LTX2LoRASwitchStage",
"LTX2UpsampleStage",
] ]
@@ -106,6 +106,23 @@ class LTX2AVDecodingStage(DecodingStage):
logger.warning( logger.warning(
"audio_vae.latents_std is all zeros; audio denorm may be incorrect." "audio_vae.latents_std is all zeros; audio denorm may be incorrect."
) )
try:
latents_mean = self.audio_vae.latents_mean
except AttributeError:
latents_mean = None
if isinstance(latents_mean, torch.Tensor) and isinstance(
latents_std, torch.Tensor
):
latents_mean = latents_mean.to(device=device, dtype=dtype)
latents_std = latents_std.to(device=device, dtype=dtype)
if audio_latents.ndim == 4:
latents_mean = latents_mean.view(
1, audio_latents.shape[1], 1, audio_latents.shape[3]
)
latents_std = latents_std.view(
1, audio_latents.shape[1], 1, audio_latents.shape[3]
)
audio_latents = audio_latents * latents_std + latents_mean
with torch.no_grad(): with torch.no_grad():
# Decode latents to spectrogram # Decode latents to spectrogram
@@ -710,6 +710,8 @@ class DenoisingStage(PipelineStage):
trajectory_timesteps: list, trajectory_timesteps: list,
server_args: ServerArgs, server_args: ServerArgs,
is_warmup: bool = False, is_warmup: bool = False,
*args,
**kwargs,
): ):
# Gather results if using sequence parallelism # Gather results if using sequence parallelism
if trajectory_latents: if trajectory_latents:
@@ -9,6 +9,7 @@ import PIL.Image
import torch import torch
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
from diffusers.models.modeling_outputs import AutoencoderKLOutput from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.utils.torch_utils import randn_tensor
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.models.vision_utils import ( from sglang.multimodal_gen.runtime.models.vision_utils import (
@@ -70,7 +71,11 @@ class LTX2AVDenoisingStage(DenoisingStage):
return int(batch.sp_video_latent_num_frames) return int(batch.sp_video_latent_num_frames)
pc = server_args.pipeline_config pc = server_args.pipeline_config
return int((batch.num_frames - 1) // int(pc.vae_temporal_compression) + 1) return int(
(batch.num_frames - 1)
// int(pc.vae_config.arch_config.temporal_compression_ratio)
+ 1
)
@staticmethod @staticmethod
def _truncate_sp_padded_token_latents( def _truncate_sp_padded_token_latents(
@@ -106,6 +111,57 @@ class LTX2AVDenoisingStage(DenoisingStage):
return return
return super()._maybe_enable_cache_dit(num_inference_steps, batch) return super()._maybe_enable_cache_dit(num_inference_steps, batch)
def _get_ltx2_stage1_guider_params(
self, batch: Req, server_args: ServerArgs, stage: str
) -> dict[str, object] | None:
if stage != "stage1":
return None
pipeline_ref = getattr(self, "pipeline", None)
pipeline = pipeline_ref() if callable(pipeline_ref) else pipeline_ref
pipeline_name = getattr(pipeline, "pipeline_name", None)
if pipeline_name != "LTX2TwoStagePipeline":
return None
return batch.extra.get("ltx2_stage1_guider_params")
@staticmethod
def _ltx2_should_skip_step(step_index: int, skip_step: int) -> bool:
if skip_step == 0:
return False
return step_index % (skip_step + 1) != 0
@staticmethod
def _ltx2_apply_rescale(
cond: torch.Tensor, pred: torch.Tensor, rescale_scale: float
) -> torch.Tensor:
if rescale_scale == 0.0:
return pred
factor = cond.std() / pred.std()
factor = rescale_scale * factor + (1.0 - rescale_scale)
return pred * factor
@classmethod
def _ltx2_calculate_guided_x0(
cls,
*,
cond: torch.Tensor,
uncond_text: torch.Tensor | float,
uncond_perturbed: torch.Tensor | float,
uncond_modality: torch.Tensor | float,
cfg_scale: float,
stg_scale: float,
rescale_scale: float,
modality_scale: float,
) -> torch.Tensor:
pred = (
cond
+ (cfg_scale - 1.0) * (cond - uncond_text)
+ stg_scale * (cond - uncond_perturbed)
+ (modality_scale - 1.0) * (cond - uncond_modality)
)
return cls._ltx2_apply_rescale(cond, pred, rescale_scale)
@staticmethod @staticmethod
def _resize_center_crop( def _resize_center_crop(
img: PIL.Image.Image, *, width: int, height: int img: PIL.Image.Image, *, width: int, height: int
@@ -328,22 +384,16 @@ class LTX2AVDenoisingStage(DenoisingStage):
# Prepare variables for the denoising loop # Prepare variables for the denoising loop
prepared_vars = self._prepare_denoising_loop(batch, server_args) prepared_vars = self._prepare_denoising_loop(batch, server_args)
extra_step_kwargs = prepared_vars["extra_step_kwargs"]
target_dtype = prepared_vars["target_dtype"] target_dtype = prepared_vars["target_dtype"]
autocast_enabled = prepared_vars["autocast_enabled"] autocast_enabled = prepared_vars["autocast_enabled"]
timesteps = prepared_vars["timesteps"] timesteps = prepared_vars["timesteps"]
num_inference_steps = prepared_vars["num_inference_steps"] num_inference_steps = prepared_vars["num_inference_steps"]
num_warmup_steps = prepared_vars["num_warmup_steps"] num_warmup_steps = prepared_vars["num_warmup_steps"]
image_kwargs = prepared_vars["image_kwargs"]
pos_cond_kwargs = prepared_vars["pos_cond_kwargs"]
neg_cond_kwargs = prepared_vars["neg_cond_kwargs"]
latents = prepared_vars["latents"] latents = prepared_vars["latents"]
boundary_timestep = prepared_vars["boundary_timestep"] boundary_timestep = prepared_vars["boundary_timestep"]
z = prepared_vars["z"] z = prepared_vars["z"]
reserved_frames_mask = prepared_vars["reserved_frames_mask"] reserved_frames_mask = prepared_vars["reserved_frames_mask"]
seq_len = prepared_vars["seq_len"] stage = batch.extra.get("ltx2_phase", "stage1")
guidance = prepared_vars["guidance"]
audio_latents = batch.audio_latents audio_latents = batch.audio_latents
audio_scheduler = copy.deepcopy(self.scheduler) audio_scheduler = copy.deepcopy(self.scheduler)
@@ -356,8 +406,14 @@ class LTX2AVDenoisingStage(DenoisingStage):
latent_num_frames_for_model = self._get_video_latent_num_frames_for_model( latent_num_frames_for_model = self._get_video_latent_num_frames_for_model(
batch=batch, server_args=server_args, latents=latents batch=batch, server_args=server_args, latents=latents
) )
latent_height = batch.height // server_args.pipeline_config.vae_scale_factor latent_height = (
latent_width = batch.width // server_args.pipeline_config.vae_scale_factor batch.height
// server_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
)
latent_width = (
batch.width
// server_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
)
# Initialize lists for ODE trajectory # Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = [] trajectory_timesteps: list[torch.Tensor] = []
@@ -394,7 +450,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
clean_latent[:, :num_img_tokens, :] = batch.image_latent[ clean_latent[:, :num_img_tokens, :] = batch.image_latent[
:, :num_img_tokens, : :, :num_img_tokens, :
].to(device=latents.device, dtype=latents.dtype) ].to(device=latents.device, dtype=latents.dtype)
with torch.autocast( with torch.autocast(
device_type=current_platform.device_type, device_type=current_platform.device_type,
dtype=target_dtype, dtype=target_dtype,
@@ -438,7 +493,9 @@ class LTX2AVDenoisingStage(DenoisingStage):
latent_model_input = latents.to(target_dtype) latent_model_input = latents.to(target_dtype)
audio_latent_model_input = audio_latents.to(target_dtype) audio_latent_model_input = audio_latents.to(target_dtype)
stage1_guider_params = self._get_ltx2_stage1_guider_params(
batch, server_args, stage
)
latent_num_frames = latent_num_frames_for_model latent_num_frames = latent_num_frames_for_model
# Audio latent dims # Audio latent dims
@@ -468,56 +525,58 @@ class LTX2AVDenoisingStage(DenoisingStage):
timestep_video = timestep timestep_video = timestep
timestep_audio = timestep timestep_audio = timestep
# Conditions use_official_cfg_path = stage1_guider_params is None
encoder_hidden_states = batch.prompt_embeds[0] if use_official_cfg_path:
audio_encoder_hidden_states = batch.audio_prompt_embeds[0] encoder_hidden_states = batch.prompt_embeds[0]
encoder_attention_mask = batch.prompt_attention_mask audio_encoder_hidden_states = batch.audio_prompt_embeds[0]
encoder_attention_mask = batch.prompt_attention_mask
# Follow ltx-pipelines structure: separate pos/neg forward passes,
# then apply CFG on denoised (x0) predictions.
with set_forward_context(
current_timestep=i, attn_metadata=attn_metadata
):
v_pos, a_v_pos = current_model(
hidden_states=latent_model_input,
audio_hidden_states=audio_latent_model_input,
encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states,
timestep=timestep_video,
audio_timestep=timestep_audio,
encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=encoder_attention_mask,
num_frames=latent_num_frames,
height=latent_height,
width=latent_width,
fps=batch.fps,
audio_num_frames=audio_num_frames_latent,
video_coords=video_coords,
audio_coords=audio_coords,
return_latents=False,
return_dict=False,
)
if batch.do_classifier_free_guidance: if batch.do_classifier_free_guidance:
neg_encoder_hidden_states = ( latent_model_input = torch.cat(
batch.negative_prompt_embeds[0] [latent_model_input] * 2, dim=0
) )
neg_audio_encoder_hidden_states = ( audio_latent_model_input = torch.cat(
batch.negative_audio_prompt_embeds[0] [audio_latent_model_input] * 2, dim=0
) )
neg_encoder_attention_mask = ( encoder_hidden_states = torch.cat(
batch.negative_attention_mask [
batch.negative_prompt_embeds[0],
encoder_hidden_states,
],
dim=0,
)
audio_encoder_hidden_states = torch.cat(
[
batch.negative_audio_prompt_embeds[0],
audio_encoder_hidden_states,
],
dim=0,
)
encoder_attention_mask = torch.cat(
[
batch.negative_attention_mask,
encoder_attention_mask,
],
dim=0,
)
timestep_video = timestep_video.expand(
int(latent_model_input.shape[0])
)
timestep_audio = timestep_audio.expand(
int(latent_model_input.shape[0])
) )
v_neg, a_v_neg = current_model( with set_forward_context(
current_timestep=i, attn_metadata=attn_metadata
):
model_video, model_audio = current_model(
hidden_states=latent_model_input, hidden_states=latent_model_input,
audio_hidden_states=audio_latent_model_input, audio_hidden_states=audio_latent_model_input,
encoder_hidden_states=neg_encoder_hidden_states, encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=neg_audio_encoder_hidden_states, audio_encoder_hidden_states=audio_encoder_hidden_states,
timestep=timestep_video, timestep=timestep_video,
audio_timestep=timestep_audio, audio_timestep=timestep_audio,
encoder_attention_mask=neg_encoder_attention_mask, encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=neg_encoder_attention_mask, audio_encoder_attention_mask=encoder_attention_mask,
num_frames=latent_num_frames, num_frames=latent_num_frames,
height=latent_height, height=latent_height,
width=latent_width, width=latent_width,
@@ -528,16 +587,132 @@ class LTX2AVDenoisingStage(DenoisingStage):
return_latents=False, return_latents=False,
return_dict=False, return_dict=False,
) )
else:
v_neg = None
a_v_neg = None
v_pos = v_pos.float() model_video = model_video.float()
a_v_pos = a_v_pos.float() model_audio = model_audio.float()
if v_neg is not None: if batch.do_classifier_free_guidance:
v_neg = v_neg.float() model_video_uncond, model_video_text = (
if a_v_neg is not None: model_video.chunk(2)
a_v_neg = a_v_neg.float() )
model_audio_uncond, model_audio_text = (
model_audio.chunk(2)
)
model_video = model_video_uncond + (
batch.guidance_scale
* (model_video_text - model_video_uncond)
)
model_audio = model_audio_uncond + (
batch.guidance_scale
* (model_audio_text - model_audio_uncond)
)
v_pos = model_video
a_v_pos = model_audio
v_neg = None
a_v_neg = None
latents = self.scheduler.step(
v_pos, t_device, latents, return_dict=False
)[0]
audio_latents = audio_scheduler.step(
a_v_pos, t_device, audio_latents, return_dict=False
)[0]
if do_ti2v:
latents[:, :num_img_tokens, :] = batch.image_latent[
:, :num_img_tokens, :
].to(device=latents.device, dtype=latents.dtype)
latents = self.post_forward_for_ti2v_task(
batch, server_args, reserved_frames_mask, latents, z
)
if batch.return_trajectory_latents:
trajectory_timesteps.append(t_host)
trajectory_latents.append(latents)
if audio_latents is not None:
trajectory_audio_latents.append(audio_latents)
if i == num_timesteps - 1 or (
(i + 1) > num_warmup_steps
and (i + 1) % self.scheduler.order == 0
and progress_bar is not None
):
progress_bar.update()
if not is_warmup:
self.step_profile()
continue
else:
# Follow ltx-pipelines structure: separate pos/neg forward passes,
# then apply CFG on denoised (x0) predictions.
encoder_hidden_states = batch.prompt_embeds[0]
audio_encoder_hidden_states = batch.audio_prompt_embeds[0]
encoder_attention_mask = batch.prompt_attention_mask
with set_forward_context(
current_timestep=i, attn_metadata=attn_metadata
):
v_pos, a_v_pos = current_model(
hidden_states=latent_model_input,
audio_hidden_states=audio_latent_model_input,
encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states,
timestep=timestep_video,
audio_timestep=timestep_audio,
encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=encoder_attention_mask,
num_frames=latent_num_frames,
height=latent_height,
width=latent_width,
fps=batch.fps,
audio_num_frames=audio_num_frames_latent,
video_coords=video_coords,
audio_coords=audio_coords,
return_latents=False,
return_dict=False,
)
if (
stage1_guider_params is not None
or batch.do_classifier_free_guidance
):
neg_encoder_hidden_states = (
batch.negative_prompt_embeds[0]
)
neg_audio_encoder_hidden_states = (
batch.negative_audio_prompt_embeds[0]
)
neg_encoder_attention_mask = (
batch.negative_attention_mask
)
v_neg, a_v_neg = current_model(
hidden_states=latent_model_input,
audio_hidden_states=audio_latent_model_input,
encoder_hidden_states=neg_encoder_hidden_states,
audio_encoder_hidden_states=neg_audio_encoder_hidden_states,
timestep=timestep_video,
audio_timestep=timestep_audio,
encoder_attention_mask=neg_encoder_attention_mask,
audio_encoder_attention_mask=neg_encoder_attention_mask,
num_frames=latent_num_frames,
height=latent_height,
width=latent_width,
fps=batch.fps,
audio_num_frames=audio_num_frames_latent,
video_coords=video_coords,
audio_coords=audio_coords,
return_latents=False,
return_dict=False,
)
else:
v_neg = None
a_v_neg = None
v_pos = v_pos.float()
a_v_pos = a_v_pos.float()
if v_neg is not None:
v_neg = v_neg.float()
if a_v_neg is not None:
a_v_neg = a_v_neg.float()
# Velocity -> denoised (x0): x0 = x - sigma * v # Velocity -> denoised (x0): x0 = x - sigma * v
sigma_val = float(sigma.item()) sigma_val = float(sigma.item())
@@ -547,9 +722,18 @@ class LTX2AVDenoisingStage(DenoisingStage):
denoised_audio = ( denoised_audio = (
audio_latents.float() - sigma_val * a_v_pos audio_latents.float() - sigma_val * a_v_pos
).to(audio_latents.dtype) ).to(audio_latents.dtype)
denoised_video_neg = None
denoised_audio_neg = None
denoised_video_perturbed = None
denoised_audio_perturbed = None
denoised_video_modality = None
denoised_audio_modality = None
if ( if (
batch.do_classifier_free_guidance (
stage1_guider_params is not None
or batch.do_classifier_free_guidance
)
and v_neg is not None and v_neg is not None
and a_v_neg is not None and a_v_neg is not None
): ):
@@ -559,6 +743,159 @@ class LTX2AVDenoisingStage(DenoisingStage):
denoised_audio_neg = ( denoised_audio_neg = (
audio_latents.float() - sigma_val * a_v_neg audio_latents.float() - sigma_val * a_v_neg
).to(audio_latents.dtype) ).to(audio_latents.dtype)
if stage1_guider_params is not None:
video_skip = self._ltx2_should_skip_step(
i, int(stage1_guider_params["video_skip_step"])
)
audio_skip = self._ltx2_should_skip_step(
i, int(stage1_guider_params["audio_skip_step"])
)
need_perturbed = (
float(stage1_guider_params["video_stg_scale"]) != 0.0
or float(stage1_guider_params["audio_stg_scale"]) != 0.0
)
if need_perturbed:
with set_forward_context(
current_timestep=i, attn_metadata=attn_metadata
):
v_ptb, a_v_ptb = current_model(
hidden_states=latent_model_input,
audio_hidden_states=audio_latent_model_input,
encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states,
timestep=timestep_video,
audio_timestep=timestep_audio,
encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=encoder_attention_mask,
num_frames=latent_num_frames,
height=latent_height,
width=latent_width,
fps=batch.fps,
audio_num_frames=audio_num_frames_latent,
video_coords=video_coords,
audio_coords=audio_coords,
return_latents=False,
return_dict=False,
skip_video_self_attn_blocks=tuple(
stage1_guider_params["video_stg_blocks"]
),
skip_audio_self_attn_blocks=tuple(
stage1_guider_params["audio_stg_blocks"]
),
)
denoised_video_perturbed = (
latents.float() - sigma_val * v_ptb.float()
).to(latents.dtype)
denoised_audio_perturbed = (
audio_latents.float() - sigma_val * a_v_ptb.float()
).to(audio_latents.dtype)
need_modality = (
float(stage1_guider_params["video_modality_scale"])
!= 1.0
or float(stage1_guider_params["audio_modality_scale"])
!= 1.0
)
if need_modality:
with set_forward_context(
current_timestep=i, attn_metadata=attn_metadata
):
v_mod, a_v_mod = current_model(
hidden_states=latent_model_input,
audio_hidden_states=audio_latent_model_input,
encoder_hidden_states=encoder_hidden_states,
audio_encoder_hidden_states=audio_encoder_hidden_states,
timestep=timestep_video,
audio_timestep=timestep_audio,
encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=encoder_attention_mask,
num_frames=latent_num_frames,
height=latent_height,
width=latent_width,
fps=batch.fps,
audio_num_frames=audio_num_frames_latent,
video_coords=video_coords,
audio_coords=audio_coords,
return_latents=False,
return_dict=False,
disable_a2v_cross_attn=True,
disable_v2a_cross_attn=True,
)
denoised_video_modality = (
latents.float() - sigma_val * v_mod.float()
).to(latents.dtype)
denoised_audio_modality = (
audio_latents.float() - sigma_val * a_v_mod.float()
).to(audio_latents.dtype)
if not video_skip:
denoised_video = self._ltx2_calculate_guided_x0(
cond=denoised_video,
uncond_text=(
denoised_video_neg
if denoised_video_neg is not None
else denoised_video
),
uncond_perturbed=(
denoised_video_perturbed
if denoised_video_perturbed is not None
else 0.0
),
uncond_modality=(
denoised_video_modality
if denoised_video_modality is not None
else 0.0
),
cfg_scale=float(
stage1_guider_params["video_cfg_scale"]
),
stg_scale=float(
stage1_guider_params["video_stg_scale"]
),
rescale_scale=float(
stage1_guider_params["video_rescale_scale"]
),
modality_scale=float(
stage1_guider_params["video_modality_scale"]
),
)
if not audio_skip:
denoised_audio = self._ltx2_calculate_guided_x0(
cond=denoised_audio,
uncond_text=(
denoised_audio_neg
if denoised_audio_neg is not None
else denoised_audio
),
uncond_perturbed=(
denoised_audio_perturbed
if denoised_audio_perturbed is not None
else 0.0
),
uncond_modality=(
denoised_audio_modality
if denoised_audio_modality is not None
else 0.0
),
cfg_scale=float(
stage1_guider_params["audio_cfg_scale"]
),
stg_scale=float(
stage1_guider_params["audio_stg_scale"]
),
rescale_scale=float(
stage1_guider_params["audio_rescale_scale"]
),
modality_scale=float(
stage1_guider_params["audio_modality_scale"]
),
)
elif (
batch.do_classifier_free_guidance
and denoised_video_neg is not None
and denoised_audio_neg is not None
):
denoised_video = denoised_video + ( denoised_video = denoised_video + (
batch.guidance_scale - 1.0 batch.guidance_scale - 1.0
) * (denoised_video - denoised_video_neg) ) * (denoised_video - denoised_video_neg)
@@ -576,7 +913,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
denoised_video * denoise_mask denoised_video * denoise_mask
+ clean_latent.float() * (1.0 - denoise_mask) + clean_latent.float() * (1.0 - denoise_mask)
) )
# Euler step in sigma space: x_next = x + (sigma_next - sigma) * v, # Euler step in sigma space: x_next = x + (sigma_next - sigma) * v,
# where v = (x - x0) / sigma. # where v = (x - x0) / sigma.
if sigma_val == 0.0: if sigma_val == 0.0:
@@ -655,6 +991,8 @@ class LTX2AVDenoisingStage(DenoisingStage):
trajectory_audio_latents: list, trajectory_audio_latents: list,
server_args: ServerArgs, server_args: ServerArgs,
is_warmup: bool = False, is_warmup: bool = False,
*args,
**kwargs,
): ):
# 1. Handle Trajectory (Video) - Copy from base # 1. Handle Trajectory (Video) - Copy from base
if trajectory_latents: if trajectory_latents:
@@ -705,7 +1043,8 @@ class LTX2AVDenoisingStage(DenoisingStage):
batch.latents = latents batch.latents = latents
batch.audio_latents = audio_latents batch.audio_latents = audio_latents
# 4. Cleanup
# TODO: make this a general denoising-stage hook
if isinstance(self.transformer, OffloadableDiTMixin): if isinstance(self.transformer, OffloadableDiTMixin):
for manager in self.transformer.layerwise_offload_managers: for manager in self.transformer.layerwise_offload_managers:
manager.release_all() manager.release_all()
@@ -753,9 +1092,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
) )
return result return result
def do_classifier_free_guidance(self, batch: Req) -> bool:
return batch.guidance_scale > 1.0
class LTX2RefinementStage(LTX2AVDenoisingStage): class LTX2RefinementStage(LTX2AVDenoisingStage):
def __init__( def __init__(
@@ -764,34 +1100,109 @@ class LTX2RefinementStage(LTX2AVDenoisingStage):
super().__init__(transformer, scheduler, vae, audio_vae) super().__init__(transformer, scheduler, vae, audio_vae)
self.distilled_sigmas = torch.tensor(distilled_sigmas) self.distilled_sigmas = torch.tensor(distilled_sigmas)
@staticmethod
def _randn_like_with_batch_generators(
reference_tensor: torch.Tensor, batch: Req
) -> torch.Tensor:
generator = getattr(batch, "generator", None)
if isinstance(generator, list):
bsz = int(reference_tensor.shape[0])
valid_generators = [g for g in generator if isinstance(g, torch.Generator)]
if len(valid_generators) == 1:
generator = valid_generators[0]
elif len(valid_generators) >= bsz:
generator = valid_generators[:bsz]
else:
generator = None
elif not isinstance(generator, torch.Generator):
generator = None
return randn_tensor(
reference_tensor.shape,
generator=generator,
device=reference_tensor.device,
dtype=reference_tensor.dtype,
)
@staticmethod
def _reset_stage2_generators(batch: Req) -> None:
generator = getattr(batch, "generator", None)
if isinstance(generator, list) and generator:
generator_device = str(generator[0].device)
elif isinstance(generator, torch.Generator):
generator_device = str(generator.device)
else:
generator_device = "cpu"
seeds = getattr(batch, "seeds", None)
if not seeds:
seed = getattr(batch, "seed", None)
if seed is None:
return
seeds = [int(seed)]
batch.generator = [
torch.Generator(device=generator_device).manual_seed(int(seed))
for seed in seeds
]
def forward(self, batch: Req, server_args: ServerArgs) -> Req: def forward(self, batch: Req, server_args: ServerArgs) -> Req:
# 1. Add noise to latents batch.extra["ltx2_phase"] = "stage2"
self._reset_stage2_generators(batch)
noise_scale = self.distilled_sigmas[0].to(batch.latents.device) noise_scale = self.distilled_sigmas[0].to(batch.latents.device)
noise = torch.randn_like(batch.latents) video_noise = self._randn_like_with_batch_generators(batch.latents, batch)
batch.latents = batch.latents + noise * noise_scale batch.latents = video_noise * noise_scale + batch.latents * (1 - noise_scale)
# 2. Run denoising loop with distilled_sigmas if isinstance(batch.audio_latents, torch.Tensor):
# Save original sigmas audio_noise = self._randn_like_with_batch_generators(
original_sigmas = self.scheduler.sigmas batch.audio_latents, batch
original_timesteps = self.scheduler.timesteps )
original_num_inference_steps = self.scheduler.num_inference_steps audio_noise_scale = noise_scale.to(
batch.audio_latents.device, batch.audio_latents.dtype
)
batch.audio_latents = (
audio_noise * audio_noise_scale
+ batch.audio_latents * (1 - audio_noise_scale)
)
batch.latents = batch.latents.to(
device=batch.latents.device, dtype=torch.float32
)
if isinstance(batch.audio_latents, torch.Tensor):
batch.audio_latents = batch.audio_latents.to(
device=batch.audio_latents.device, dtype=torch.float32
)
# Set distilled sigmas # Stage 2 runs at full resolution, so Stage 1 TI2V conditioning is invalid.
self.scheduler.sigmas = self.distilled_sigmas.to(self.scheduler.sigmas.device) batch.image_latent = None
# Approximation for timesteps batch.ltx2_num_image_tokens = 0
self.scheduler.timesteps = self.scheduler.sigmas * 1000
self.scheduler.num_inference_steps = len(self.distilled_sigmas) - 1 # Use a private scheduler copy to avoid mutating shared state.
original_scheduler = self.scheduler
original_batch_timesteps = batch.timesteps
original_batch_num_inference_steps = batch.num_inference_steps
self.scheduler = copy.deepcopy(original_scheduler)
distilled_device = self.scheduler.sigmas.device
self.scheduler.sigmas = self.distilled_sigmas.to(distilled_device)
num_steps = len(self.distilled_sigmas) - 1
self.scheduler.num_inference_steps = num_steps
self.scheduler.timesteps = (self.distilled_sigmas[:num_steps] * 1000).to(
distilled_device
)
self.scheduler._step_index = None
self.scheduler._begin_index = None
batch.timesteps = self.scheduler.timesteps
batch.num_inference_steps = num_steps
original_do_cfg = batch.do_classifier_free_guidance
batch.do_classifier_free_guidance = False
# Call parent forward
try: try:
batch = super().forward(batch, server_args) batch = super().forward(batch, server_args)
finally: finally:
# Restore original sigmas self.scheduler = original_scheduler
self.scheduler.sigmas = original_sigmas batch.timesteps = original_batch_timesteps
self.scheduler.timesteps = original_timesteps batch.num_inference_steps = original_batch_num_inference_steps
self.scheduler.num_inference_steps = original_num_inference_steps batch.do_classifier_free_guidance = original_do_cfg
return batch return batch
def do_classifier_free_guidance(self, batch: Req) -> bool:
return False # Stage 2 uses simple denoising (no CFG)
@@ -79,6 +79,11 @@ class InputValidationStage(PipelineStage):
# Create generators based on generator_device parameter # Create generators based on generator_device parameter
# Note: This will overwrite any existing batch.generator # Note: This will overwrite any existing batch.generator
generator_device = batch.generator_device generator_device = batch.generator_device
if generator_device is None:
generator_device = (
getattr(server_args.pipeline_config, "generator_device", None)
or current_platform.device_type
)
if generator_device == "cpu": if generator_device == "cpu":
device_str = "cpu" device_str = "cpu"
@@ -37,6 +37,15 @@ class LatentPreparationStage(PipelineStage):
self.scheduler = scheduler self.scheduler = scheduler
self.transformer = transformer self.transformer = transformer
def _get_latent_dtype(
self,
batch: Req,
server_args: ServerArgs,
):
return server_args.pipeline_config.get_latent_dtype(
batch.prompt_embeds[0].dtype
)
def forward( def forward(
self, self,
batch: Req, batch: Req,
@@ -57,9 +66,7 @@ class LatentPreparationStage(PipelineStage):
batch_size = batch.batch_size batch_size = batch.batch_size
# Get required parameters # Get required parameters
dtype = server_args.pipeline_config.get_latent_dtype( dtype = self._get_latent_dtype(batch, server_args)
batch.prompt_embeds[0].dtype
)
device = get_local_torch_device() device = get_local_torch_device()
generator = batch.generator generator = batch.generator
latents = batch.latents latents = batch.latents
@@ -53,6 +53,13 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
result.add_check("latents", batch.latents, V.none_or_tensor) result.add_check("latents", batch.latents, V.none_or_tensor)
return result return result
def _get_latent_dtype(
self,
batch: Req,
server_args: ServerArgs,
):
return torch.float32
def forward(self, batch: Req, server_args: ServerArgs) -> Req: def forward(self, batch: Req, server_args: ServerArgs) -> Req:
# 1. Prepare Video Latents using base class logic # 1. Prepare Video Latents using base class logic
# This sets batch.latents and batch.raw_latent_shape # This sets batch.latents and batch.raw_latent_shape
@@ -70,12 +77,7 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
return batch return batch
device = get_local_torch_device() device = get_local_torch_device()
if isinstance(batch.prompt_embeds, list) and batch.prompt_embeds: dtype = self._get_latent_dtype(batch, server_args)
dtype = batch.prompt_embeds[0].dtype
elif isinstance(batch.prompt_embeds, torch.Tensor):
dtype = batch.prompt_embeds.dtype
else:
dtype = torch.float16
generator = batch.generator generator = batch.generator
audio_latents = batch.audio_latents audio_latents = batch.audio_latents
@@ -7,6 +7,8 @@ Prompt encoding stages for diffusion pipelines.
This module contains implementations of prompt encoding stages for diffusion pipelines. This module contains implementations of prompt encoding stages for diffusion pipelines.
""" """
import inspect
import torch import torch
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
@@ -263,14 +265,22 @@ class TextEncodingStage(PipelineStage):
attention_mask = torch.ones(input_ids.shape[:2], device=target_device) attention_mask = torch.ones(input_ids.shape[:2], device=target_device)
else: else:
attention_mask = text_inputs["attention_mask"] attention_mask = text_inputs["attention_mask"]
encoder_forward_kwargs = {
"input_ids": input_ids,
"attention_mask": attention_mask,
"output_hidden_states": True,
}
if "use_cache" in inspect.signature(text_encoder.forward).parameters:
encoder_forward_kwargs["use_cache"] = False
with set_forward_context(current_timestep=0, attn_metadata=None): with set_forward_context(current_timestep=0, attn_metadata=None):
outputs: BaseEncoderOutput = text_encoder( outputs: BaseEncoderOutput = text_encoder(**encoder_forward_kwargs)
input_ids=input_ids, postprocess_sig = inspect.signature(postprocess_func)
attention_mask=attention_mask,
output_hidden_states=True, postprocess_kwargs = {}
use_cache=False, if "pipeline_config" in postprocess_sig.parameters:
) # required by models like LTX
prompt_embeds = postprocess_func(outputs, text_inputs) postprocess_kwargs["pipeline_config"] = server_args.pipeline_config
prompt_embeds = postprocess_func(outputs, text_inputs, **postprocess_kwargs)
if dtype is not None: if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype) prompt_embeds = prompt_embeds.to(dtype=dtype)
@@ -0,0 +1,126 @@
import torch
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class LTX2HalveResolutionStage(PipelineStage):
"""Halve batch height/width for two-stage Stage 1 (low-res generation)."""
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
original_h, original_w = batch.height, batch.width
vae_scale_factor = getattr(server_args.pipeline_config, "vae_scale_factor", 32)
required_alignment = max(64, int(vae_scale_factor) * 2)
if original_h % required_alignment != 0 or original_w % required_alignment != 0:
raise ValueError(
"LTX-2 two-stage requires resolution divisible by "
f"{required_alignment}, got ({original_h}x{original_w})."
)
batch.height = batch.height // 2
batch.width = batch.width // 2
logger.info(
"Halved resolution: %dx%d -> %dx%d",
original_h,
original_w,
batch.height,
batch.width,
)
return batch
class LTX2LoRASwitchStage(PipelineStage):
"""Switch LoRA configuration for the requested two-stage phase."""
def __init__(self, pipeline, phase: str):
super().__init__()
self.pipeline = pipeline
self.phase = phase
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
switch_fn = getattr(self.pipeline, "switch_lora_phase", None)
if not callable(switch_fn):
raise ValueError(
"LTX2LoRASwitchStage requires pipeline.switch_lora_phase()"
)
switch_fn(self.phase)
batch.extra["ltx2_phase"] = self.phase
return batch
class LTX2UpsampleStage(PipelineStage):
"""Upsample Stage-1 video latents and prepare Stage-2 inputs."""
def __init__(self, spatial_upsampler, vae, audio_vae=None):
super().__init__()
self.spatial_upsampler = spatial_upsampler
self.vae = vae
self.audio_vae = audio_vae
def _upsample_video_latents(
self, latents: torch.Tensor, server_args: ServerArgs, device: torch.device
) -> torch.Tensor:
vae_mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(
device=device, dtype=latents.dtype
)
vae_std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(
device=device, dtype=latents.dtype
)
latents = latents * vae_std + vae_mean
self.spatial_upsampler = self.spatial_upsampler.to(
device=device, dtype=latents.dtype
)
latents = self.spatial_upsampler(latents)
if server_args.vae_cpu_offload:
self.spatial_upsampler = self.spatial_upsampler.to("cpu")
latents = (latents - vae_mean) / vae_std
return latents
@staticmethod
def _restore_full_resolution(batch: Req) -> None:
batch.height *= 2
batch.width *= 2
@staticmethod
def _pack_video_latents(
batch: Req, latents: torch.Tensor, server_args: ServerArgs
) -> None:
batch_size = latents.shape[0]
latents = server_args.pipeline_config.maybe_pack_latents(
latents, batch_size, batch
)
batch.latents = latents
batch.raw_latent_shape = latents.shape
def _repack_audio_latents(self, batch: Req, server_args: ServerArgs) -> None:
if batch.audio_latents is None or self.audio_vae is None:
return
audio_latents = server_args.pipeline_config.maybe_pack_audio_latents(
batch.audio_latents, batch.audio_latents.shape[0], batch
)
batch.audio_latents = audio_latents
batch.raw_audio_latent_shape = audio_latents.shape
logger.info(
"Re-packed audio latents for Stage 2: %s", list(audio_latents.shape)
)
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
device = get_local_torch_device()
latents = self._upsample_video_latents(batch.latents, server_args, device)
logger.info("Upsampled video latents: %s", list(latents.shape))
self._restore_full_resolution(batch)
self._pack_video_latents(batch, latents, server_args)
logger.info(
"Packed video latents for Stage 2: %s (resolution %dx%d)",
list(batch.latents.shape),
batch.height,
batch.width,
)
self._repack_audio_latents(batch, server_args)
return batch
@@ -48,7 +48,6 @@ from sglang.multimodal_gen.utils import (
FlexibleArgumentParser, FlexibleArgumentParser,
StoreBoolean, StoreBoolean,
expand_path_fields, expand_path_fields,
expand_path_kwargs,
) )
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -573,6 +572,15 @@ class ServerArgs:
"(e.g. 'Qwen-Image' for 'Qwen/Qwen-Image')." "(e.g. 'Qwen-Image' for 'Qwen/Qwen-Image')."
), ),
) )
parser.add_argument(
"--pipeline-class-name",
type=str,
default=ServerArgs.pipeline_class_name,
help=(
"Override pipeline class selection from model_index.json. "
"Must match a registered pipeline_name."
),
)
# attention # attention
parser.add_argument( parser.add_argument(
"--attention-backend", "--attention-backend",
@@ -936,7 +944,11 @@ class ServerArgs:
unknown_args: list[str], unknown_args: list[str],
) -> tuple[dict[str, str], list[str]]: ) -> tuple[dict[str, str], list[str]]:
""" """
Extract dynamic ``--<component>-path`` args from unrecognised CLI args. Extract dynamic component path args from unrecognised CLI args.
Supported forms:
- ``--<component>-path /path/to/component``
- ``--component-paths.<component> /path/to/component`` (expanded from config)
""" """
component_paths: dict[str, str] = {} component_paths: dict[str, str] = {}
remaining: list[str] = [] remaining: list[str] = []
@@ -944,8 +956,15 @@ class ServerArgs:
while i < len(unknown_args): while i < len(unknown_args):
arg = unknown_args[i] arg = unknown_args[i]
key_part = arg.split("=", 1)[0] if "=" in arg else arg key_part = arg.split("=", 1)[0] if "=" in arg else arg
if key_part.startswith("--") and key_part.endswith("-path"): component = None
if key_part.startswith("--component-paths."):
component = key_part[len("--component-paths.") :].replace("-", "_")
elif key_part.startswith("--component_paths."):
component = key_part[len("--component_paths.") :].replace("-", "_")
elif key_part.startswith("--") and key_part.endswith("-path"):
component = key_part[2:-5].replace("-", "_") component = key_part[2:-5].replace("-", "_")
if component is not None:
if "=" in arg: if "=" in arg:
component_paths[component] = arg.split("=", 1)[1] component_paths[component] = arg.split("=", 1)[1]
elif i + 1 < len(unknown_args) and not unknown_args[i + 1].startswith( elif i + 1 < len(unknown_args) and not unknown_args[i + 1].startswith(
@@ -997,7 +1016,6 @@ class ServerArgs:
@classmethod @classmethod
def from_dict(cls, kwargs: dict[str, Any]) -> "ServerArgs": def from_dict(cls, kwargs: dict[str, Any]) -> "ServerArgs":
"""Create a ServerArgs object from a dictionary.""" """Create a ServerArgs object from a dictionary."""
kwargs = expand_path_kwargs(dict(kwargs))
attrs = [attr.name for attr in dataclasses.fields(cls)] attrs = [attr.name for attr in dataclasses.fields(cls)]
server_args_kwargs: dict[str, Any] = {} server_args_kwargs: dict[str, Any] = {}
@@ -545,7 +545,7 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
) )
return config return config
except EntryNotFoundError: except EntryNotFoundError:
logger.warning( logger.debug(
"model_index.json not found for %s. Assuming it is a single model and downloading it.", "model_index.json not found for %s. Assuming it is a single model and downloading it.",
model_name_or_path, model_name_or_path,
) )
@@ -1089,6 +1089,70 @@
"expected_avg_denoise_ms": 319.61, "expected_avg_denoise_ms": 319.61,
"expected_median_denoise_ms": 127.39 "expected_median_denoise_ms": 127.39
}, },
"ltx_2_two_stage_t2v": {
"stages_ms": {
"InputValidationStage": 0.05,
"TextEncodingStage": 1827.03,
"LTX2TextConnectorStage": 11.54,
"LTX2HalveResolutionStage": 0.1,
"LTX2LoRASwitchStage": 13014.75,
"LTX2SigmaPreparationStage": 0.25,
"TimestepPreparationStage": 19.26,
"LTX2AVLatentPreparationStage": 0.37,
"LTX2AVDenoisingStage": 53324.98,
"LTX2UpsampleStage": 1894.17,
"LTX2RefinementStage": 4330.57,
"LTX2AVDecodingStage": 337.68
},
"denoise_step_ms": {
"0": 1206.17,
"1": 1335.08,
"2": 1336.0,
"3": 1337.77,
"4": 1335.76,
"5": 1334.02,
"6": 1332.78,
"7": 1333.1,
"8": 1334.86,
"9": 1335.26,
"10": 1335.62,
"11": 1334.21,
"12": 1336.52,
"13": 1335.84,
"14": 1338.81,
"15": 1330.74,
"16": 1334.75,
"17": 1337.27,
"18": 1334.05,
"19": 1335.62,
"20": 1349.22,
"21": 1338.2,
"22": 1341.3,
"23": 1358.03,
"24": 1341.14,
"25": 1339.98,
"26": 1332.75,
"27": 1333.31,
"28": 1333.9,
"29": 1333.36,
"30": 1333.96,
"31": 1336.93,
"32": 1335.04,
"33": 1334.48,
"34": 1334.99,
"35": 1333.58,
"36": 1334.49,
"37": 1333.83,
"38": 1332.2,
"39": 1333.38,
"40": 1652.67,
"41": 1330.67,
"42": 1340.97
},
"expected_e2e_ms": 75229.22,
"expected_avg_denoise_ms": 1340.53,
"expected_median_denoise_ms": 1334.99
},
"wan2_2_ti2v_5b": { "wan2_2_ti2v_5b": {
"stages_ms": { "stages_ms": {
"InputValidationStage": 96.27, "InputValidationStage": 96.27,
@@ -384,6 +384,10 @@ MULTI_FRAME_I2I_sampling_params = DiffusionSamplingParams(
T2V_PROMPT = "A curious raccoon" T2V_PROMPT = "A curious raccoon"
T2V_sampling_params = DiffusionSamplingParams(
prompt=T2V_PROMPT,
)
TI2V_sampling_params = DiffusionSamplingParams( TI2V_sampling_params = DiffusionSamplingParams(
prompt="The man in the picture slowly turns his head, his expression enigmatic and otherworldly. The camera performs a slow, cinematic dolly out, focusing on his face. Moody lighting, neon signs glowing in the background, shallow depth of field.", prompt="The man in the picture slowly turns his head, his expression enigmatic and otherworldly. The camera performs a slow, cinematic dolly out, focusing on his face. Moody lighting, neon signs glowing in the background, shallow depth of field.",
image_path="https://is1-ssl.mzstatic.com/image/thumb/Music114/v4/5f/fa/56/5ffa56c2-ea1f-7a17-6bad-192ff9b6476d/825646124206.jpg/600x600bb.jpg", image_path="https://is1-ssl.mzstatic.com/image/thumb/Music114/v4/5f/fa/56/5ffa56c2-ea1f-7a17-6bad-192ff9b6476d/825646124206.jpg/600x600bb.jpg",
@@ -399,7 +403,6 @@ TURBOWAN_I2V_sampling_params = DiffusionSamplingParams(
fps=4, fps=4,
) )
# All test cases with clean default values # All test cases with clean default values
# To test different models, simply add more DiffusionCase entries # To test different models, simply add more DiffusionCase entries
ONE_GPU_CASES_A: list[DiffusionTestCase] = [ ONE_GPU_CASES_A: list[DiffusionTestCase] = [
@@ -577,9 +580,7 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [
modality="video", modality="video",
custom_validator="video", custom_validator="video",
), ),
DiffusionSamplingParams( T2V_sampling_params,
prompt=T2V_PROMPT,
),
), ),
DiffusionTestCase( DiffusionTestCase(
"wan2_1_t2v_1.3b_text_encoder_cpu_offload", "wan2_1_t2v_1.3b_text_encoder_cpu_offload",
@@ -589,9 +590,7 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [
custom_validator="video", custom_validator="video",
text_encoder_cpu_offload=True, text_encoder_cpu_offload=True,
), ),
DiffusionSamplingParams( T2V_sampling_params,
prompt=T2V_PROMPT,
),
), ),
# TeaCache acceleration test for Wan video model # TeaCache acceleration test for Wan video model
DiffusionTestCase( DiffusionTestCase(
@@ -706,9 +705,7 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [
modality="video", modality="video",
custom_validator="video", custom_validator="video",
), ),
DiffusionSamplingParams( T2V_sampling_params,
prompt=T2V_PROMPT,
),
), ),
# === Text and Image to Video (TI2V) === # === Text and Image to Video (TI2V) ===
DiffusionTestCase( DiffusionTestCase(
@@ -792,9 +789,7 @@ if not current_platform.is_hip():
modality="video", modality="video",
custom_validator="video", custom_validator="video",
), ),
DiffusionSamplingParams( T2V_sampling_params,
prompt=T2V_PROMPT,
),
) )
) )
@@ -828,9 +823,7 @@ TWO_GPU_CASES_A = [
custom_validator="video", custom_validator="video",
num_gpus=2, num_gpus=2,
), ),
DiffusionSamplingParams( T2V_sampling_params,
prompt=T2V_PROMPT,
),
), ),
# TeaCache smoke test for Wan2.2 T2V A14B — verifies enable_teacache=True # TeaCache smoke test for Wan2.2 T2V A14B — verifies enable_teacache=True
# doesn't crash. Perf check disabled because Wan2.2-specific TeaCache # doesn't crash. Perf check disabled because Wan2.2-specific TeaCache
@@ -886,9 +879,7 @@ TWO_GPU_CASES_A = [
num_gpus=2, num_gpus=2,
cfg_parallel=True, cfg_parallel=True,
), ),
DiffusionSamplingParams( T2V_sampling_params,
prompt=T2V_PROMPT,
),
), ),
DiffusionTestCase( DiffusionTestCase(
"fsdp-inference", "fsdp-inference",
@@ -938,6 +929,17 @@ TWO_GPU_CASES_A = [
TI2V_sampling_params, TI2V_sampling_params,
run_perf_check=False, run_perf_check=False,
), ),
DiffusionTestCase(
"ltx_2_two_stage_t2v",
DiffusionServerArgs(
model_path="Lightricks/LTX-2",
modality="video",
num_gpus=2,
dit_layerwise_offload=True,
extras=["--pipeline-class-name LTX2TwoStagePipeline"],
),
T2V_sampling_params,
),
] ]
TWO_GPU_CASES_B = [ TWO_GPU_CASES_B = [
@@ -1054,7 +1056,6 @@ if not current_platform.is_hip():
) )
) )
# Load global configuration # Load global configuration
BASELINE_CONFIG = BaselineConfig.load( BASELINE_CONFIG = BaselineConfig.load(
Path(__file__).with_name("perf_baselines.json") Path(__file__).with_name("perf_baselines.json")
@@ -102,5 +102,25 @@ class TestPipelineResolutionCliOverride(unittest.TestCase):
self.assertEqual(server_args.pipeline_config.resolution, 768) self.assertEqual(server_args.pipeline_config.resolution, 768)
class TestComponentPathParsing(unittest.TestCase):
def test_extract_component_paths_accepts_config_expanded_keys(self):
component_paths, remaining = ServerArgs._extract_component_paths(
[
"--component-paths.spatial-upsampler",
"/tmp/latent_upsampler",
"--component_paths.distilled-lora=/tmp/distilled.safetensors",
]
)
self.assertEqual(
component_paths,
{
"spatial_upsampler": "/tmp/latent_upsampler",
"distilled_lora": "/tmp/distilled.safetensors",
},
)
self.assertEqual(remaining, [])
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+12 -16
View File
@@ -35,25 +35,21 @@ logger = init_logger(__name__)
T = TypeVar("T") T = TypeVar("T")
def _expand_path_value(field_name: str, value: Any) -> Any:
eu = os.path.expanduser
if field_name.endswith("_path") and isinstance(value, str):
return eu(value)
if field_name.endswith("_path") and isinstance(value, list):
return [eu(x) if isinstance(x, str) else x for x in value]
if field_name.endswith("_paths") and isinstance(value, dict):
return {k: eu(p) if isinstance(p, str) else p for k, p in value.items()}
return value
def expand_path_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]:
return {key: _expand_path_value(key, value) for key, value in kwargs.items()}
def expand_path_fields(obj) -> None: def expand_path_fields(obj) -> None:
"""In-place expanduser on all dataclass fields whose name ends with '_path' or '_paths'.""" """In-place expanduser on all dataclass fields whose name ends with '_path' or '_paths'."""
eu = os.path.expanduser
for f in fields(obj): for f in fields(obj):
setattr(obj, f.name, _expand_path_value(f.name, getattr(obj, f.name))) v = getattr(obj, f.name)
if f.name.endswith("_path") and isinstance(v, str):
setattr(obj, f.name, eu(v))
elif f.name.endswith("_path") and isinstance(v, list):
setattr(obj, f.name, [eu(x) if isinstance(x, str) else x for x in v])
elif f.name.endswith("_paths") and isinstance(v, dict):
setattr(
obj,
f.name,
{k: eu(p) if isinstance(p, str) else p for k, p in v.items()},
)
# TODO(will): used to convert server_args.precision to torch.dtype. Find a # TODO(will): used to convert server_args.precision to torch.dtype. Find a
+153 -153
View File
@@ -1,156 +1,156 @@
{ {
"_comment": "Per-model comparison config. Only frameworks listed under each case are tested. vLLM-Omni disabled until dep install issues resolved.", "_comment": "Per-model comparison config. Only frameworks listed under each case are tested. vLLM-Omni disabled until dep install issues resolved.",
"test_image_url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png", "test_image_url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png",
"cases": [ "cases": [
{ {
"id": "flux1_dev_t2i_1024", "id": "flux1_dev_t2i_1024",
"model": "black-forest-labs/FLUX.1-dev", "model": "black-forest-labs/FLUX.1-dev",
"task": "text-to-image", "task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets", "prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024, "width": 1024,
"height": 1024, "height": 1024,
"num_inference_steps": 50, "num_inference_steps": 50,
"guidance_scale": 4.0, "guidance_scale": 4.0,
"seed": 42, "seed": 42,
"num_gpus": 1, "num_gpus": 1,
"frameworks": { "frameworks": {
"sglang": { "sglang": {
"serve_args": "--enable-torch-compile --warmup --dit-layerwise-offload false", "serve_args": "--enable-torch-compile --warmup --dit-layerwise-offload false",
"extra_env": {} "extra_env": {}
}
}
},
{
"id": "flux2_dev_t2i_1024",
"model": "black-forest-labs/FLUX.2-dev",
"task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup --dit-layerwise-offload false",
"extra_env": {}
}
}
},
{
"id": "qwen_image_2512_t2i_1024",
"model": "Qwen/Qwen-Image-2512",
"task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "qwen_image_edit_2511",
"model": "Qwen/Qwen-Image-Edit-2511",
"task": "image-edit",
"prompt": "Make the cat wear a red hat",
"reference_image": true,
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "zimage_turbo_t2i_1024",
"model": "Tongyi-MAI/Z-Image-Turbo",
"task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024,
"height": 1024,
"num_inference_steps": 9,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "wan22_t2v_a14b_720p",
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
"task": "text-to-video",
"prompt": "A cat and a dog baking a cake together in a kitchen.",
"width": 1280,
"height": 720,
"num_frames": 81,
"num_inference_steps": 2,
"guidance_scale": 5.0,
"seed": 42,
"num_gpus": 4,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup --enable-cfg-parallel --ulysses-degree 2 --text-encoder-cpu-offload --pin-cpu-memory",
"extra_env": {}
}
}
},
{
"id": "wan22_ti2v_5b_720p",
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
"task": "text-image-to-video",
"prompt": "The cat starts walking slowly towards the camera.",
"reference_image": true,
"width": 1280,
"height": 720,
"num_frames": 81,
"num_inference_steps": 50,
"guidance_scale": 5.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "wan22_i2v_a14b_720p",
"model": "Wan-AI/Wan2.2-I2V-A14B-Diffusers",
"task": "image-to-video",
"prompt": "The cat starts walking slowly towards the camera.",
"reference_image": true,
"width": 1280,
"height": 720,
"num_frames": 81,
"num_inference_steps": 2,
"guidance_scale": 5.0,
"seed": 42,
"num_gpus": 4,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup --enable-cfg-parallel --ulysses-degree 2 --text-encoder-cpu-offload --pin-cpu-memory",
"extra_env": {}
}
}
} }
} ]
},
{
"id": "flux2_dev_t2i_1024",
"model": "black-forest-labs/FLUX.2-dev",
"task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup --dit-layerwise-offload false",
"extra_env": {}
}
}
},
{
"id": "qwen_image_2512_t2i_1024",
"model": "Qwen/Qwen-Image-2512",
"task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "qwen_image_edit_2511",
"model": "Qwen/Qwen-Image-Edit-2511",
"task": "image-edit",
"prompt": "Make the cat wear a red hat",
"reference_image": true,
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "zimage_turbo_t2i_1024",
"model": "Tongyi-MAI/Z-Image-Turbo",
"task": "text-to-image",
"prompt": "A futuristic cyberpunk city at night, neon lights reflecting on wet streets",
"width": 1024,
"height": 1024,
"num_inference_steps": 9,
"guidance_scale": 4.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "wan22_t2v_a14b_720p",
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
"task": "text-to-video",
"prompt": "A cat and a dog baking a cake together in a kitchen.",
"width": 1280,
"height": 720,
"num_frames": 81,
"num_inference_steps": 2,
"guidance_scale": 5.0,
"seed": 42,
"num_gpus": 4,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup --enable-cfg-parallel --ulysses-degree 2 --text-encoder-cpu-offload --pin-cpu-memory",
"extra_env": {}
}
}
},
{
"id": "wan22_ti2v_5b_720p",
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
"task": "text-image-to-video",
"prompt": "The cat starts walking slowly towards the camera.",
"reference_image": true,
"width": 1280,
"height": 720,
"num_frames": 81,
"num_inference_steps": 50,
"guidance_scale": 5.0,
"seed": 42,
"num_gpus": 1,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup",
"extra_env": {}
}
}
},
{
"id": "wan22_i2v_a14b_720p",
"model": "Wan-AI/Wan2.2-I2V-A14B-Diffusers",
"task": "image-to-video",
"prompt": "The cat starts walking slowly towards the camera.",
"reference_image": true,
"width": 1280,
"height": 720,
"num_frames": 81,
"num_inference_steps": 2,
"guidance_scale": 5.0,
"seed": 42,
"num_gpus": 4,
"frameworks": {
"sglang": {
"serve_args": "--enable-torch-compile --warmup --enable-cfg-parallel --ulysses-degree 2 --text-encoder-cpu-offload --pin-cpu-memory",
"extra_env": {}
}
}
}
]
} }