[diffusion] model: support LTX2.3 (#22111)

This commit is contained in:
Mick
2026-04-06 12:26:30 +08:00
committed by GitHub
parent f407461ec8
commit 82c41a2d9e
28 changed files with 2686 additions and 206 deletions
+4 -1
View File
@@ -33,12 +33,15 @@ 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 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | LTX-2 | `Lightricks/LTX-2` | 768×512<br>1536×1024 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| LTX-2.3 | `Lightricks/LTX-2.3` | 768×512<br>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 is 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`. 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`.
4. `Lightricks/LTX-2.3` is supported through the bundled native overlay materialization path. One-stage generation uses the default `LTX2Pipeline`; two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`.
5. For LTX models, the `Resolutions` column uses output video `width×height` semantics, matching `sglang generate --width ... --height ...`. One-stage generation is validated at `768×512`; two-stage generation is validated at `1536×1024`.
### Image Generation Models ### Image Generation Models
-5
View File
@@ -76,11 +76,6 @@ 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`:
@@ -12,13 +12,17 @@ class LTX2ConnectorArchConfig(AdapterArchConfig):
audio_connector_num_attention_heads: int = 30 audio_connector_num_attention_heads: int = 30
audio_connector_num_layers: int = 2 audio_connector_num_layers: int = 2
audio_connector_num_learnable_registers: int = 128 audio_connector_num_learnable_registers: int = 128
audio_feature_extractor_out_features: int = 0
caption_channels: int = 3840 caption_channels: int = 3840
causal_temporal_positioning: bool = False causal_temporal_positioning: bool = False
connector_rope_base_seq_len: int = 4096 connector_rope_base_seq_len: int = 4096
connector_apply_gated_attention: bool = False
feature_extractor_in_features: int = 0
rope_double_precision: bool = True rope_double_precision: bool = True
rope_theta: float = 10000.0 rope_theta: float = 10000.0
rope_type: str = "split" rope_type: str = "split"
text_proj_in_factor: int = 49 text_proj_in_factor: int = 49
video_feature_extractor_out_features: int = 0
video_connector_attention_head_dim: int = 128 video_connector_attention_head_dim: int = 128
video_connector_num_attention_heads: int = 30 video_connector_num_attention_heads: int = 30
video_connector_num_layers: int = 2 video_connector_num_layers: int = 2
@@ -63,6 +63,7 @@ class LTX2ArchConfig(DiTArchConfig):
# We use upstream variable names (patchify_proj, adaln_single) but HF uses different keys. # We use upstream variable names (patchify_proj, adaln_single) but HF uses different keys.
# #
# HF key -> SGLang key (upstream naming) # HF key -> SGLang key (upstream naming)
r"^model\.diffusion_model\.(.*)$": r"\1",
r"^proj_in\.(.*)$": r"patchify_proj.\1", r"^proj_in\.(.*)$": r"patchify_proj.\1",
r"^time_embed\.(.*)$": r"adaln_single.\1", r"^time_embed\.(.*)$": r"adaln_single.\1",
r"^audio_proj_in\.(.*)$": r"audio_patchify_proj.\1", r"^audio_proj_in\.(.*)$": r"audio_patchify_proj.\1",
@@ -123,6 +124,10 @@ class LTX2ArchConfig(DiTArchConfig):
attention_type: LTX2AttentionFunction = LTX2AttentionFunction.DEFAULT attention_type: LTX2AttentionFunction = LTX2AttentionFunction.DEFAULT
rope_type: LTX2RopeType = LTX2RopeType.INTERLEAVED rope_type: LTX2RopeType = LTX2RopeType.INTERLEAVED
double_precision_rope: bool = False double_precision_rope: bool = False
quantize_video_rope_coords_to_hidden_dtype: bool = False
apply_gated_attention: bool = False
cross_attention_adaln: bool = False
caption_proj_before_connector: bool = False
# Video parameters # Video parameters
num_attention_heads: int = 32 num_attention_heads: int = 32
@@ -147,6 +152,14 @@ class LTX2ArchConfig(DiTArchConfig):
audio_positional_embedding_max_pos: list[int] | None = None audio_positional_embedding_max_pos: list[int] | None = None
av_ca_timestep_scale_multiplier: int = 1 av_ca_timestep_scale_multiplier: int = 1
# 2.3 connector-related fields may show up in transformer/config.json.
connector_attention_head_dim: int = 128
connector_num_attention_heads: int = 30
connector_num_layers: int = 2
audio_connector_attention_head_dim: int = 128
audio_connector_num_attention_heads: int = 30
audio_connector_num_layers: int = 2
# SGLang-specific parameters # SGLang-specific parameters
patch_size: tuple[int, int, int] = (1, 2, 2) patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512 text_len: int = 512
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import List from typing import Any, List
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
@@ -52,6 +52,12 @@ class LTXVideoVAEArchConfig(VAEArchConfig):
decoder_causal: bool = False decoder_causal: bool = False
decoder_spatial_padding_mode: str = "reflect" decoder_spatial_padding_mode: str = "reflect"
# Native LTX variant metadata.
ltx_variant: str = "ltx_2"
condition_encoder_subdir: str = ""
video_decoder_variant: str = "ltx_2"
video_decoder_config: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
class LTXVideoVAEConfig(VAEConfig): class LTXVideoVAEConfig(VAEConfig):
@@ -93,20 +93,48 @@ def pack_text_embeds(
return normalized_hidden_states return normalized_hidden_states
def pack_text_embeds_v2(
text_hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
eps: float = 1e-6,
) -> torch.Tensor:
"""
LTX-2.3 feature extractor pre-processing.
Upstream `FeatureExtractorV2` applies per-token RMS normalization on each
Gemma layer and then flattens `[hidden_dim, num_layers]` into the channel
dimension, zeroing out padded positions afterwards.
"""
variance = torch.mean(text_hidden_states**2, dim=2, keepdim=True)
normalized_hidden_states = text_hidden_states * torch.rsqrt(variance + eps)
normalized_hidden_states = normalized_hidden_states.flatten(2)
mask = attention_mask.bool().unsqueeze(-1)
return torch.where(
mask, normalized_hidden_states, torch.zeros_like(normalized_hidden_states)
)
def is_ltx23_native_variant(arch_config: object) -> bool:
return str(getattr(arch_config, "ltx_variant", "ltx_2")) == "ltx_2_3"
def _gemma_postprocess_func( def _gemma_postprocess_func(
outputs: BaseEncoderOutput, outputs: BaseEncoderOutput,
text_inputs: dict, text_inputs: dict,
pipeline_config: Optional["LTX2PipelineConfig"] = None, 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
# We need to stack them along the last dimension and pack them
hidden_states = torch.stack(outputs.hidden_states, dim=-1) hidden_states = torch.stack(outputs.hidden_states, dim=-1)
attention_mask = text_inputs["attention_mask"] attention_mask = text_inputs["attention_mask"]
if (
pipeline_config is not None
and pipeline_config.dit_config.arch_config.caption_proj_before_connector
):
return pack_text_embeds_v2(hidden_states, attention_mask)
sequence_lengths = attention_mask.sum(dim=-1) sequence_lengths = attention_mask.sum(dim=-1)
# Assuming left padding for Gemma as per Diffusers
return pack_text_embeds(hidden_states, sequence_lengths, padding_side="left") return pack_text_embeds(hidden_states, sequence_lengths, padding_side="left")
else: else:
raise AttributeError( raise AttributeError(
@@ -1,4 +1,6 @@
import dataclasses import dataclasses
from dataclasses import field
from typing import Any
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@@ -39,3 +41,44 @@ class LTX2SamplingParams(SamplingParams):
"pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, flat lighting, " "pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, flat lighting, "
"inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts." "inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts."
) )
@dataclasses.dataclass
class LTX23SamplingParams(LTX2SamplingParams):
"""Sampling parameters matching official LTX-2.3 one-stage defaults."""
generator_device: str = "cuda"
guidance_scale: float = 3.0
num_inference_steps: int = 30
video_cfg_scale: float = 3.0
video_stg_scale: float = 1.0
video_rescale_scale: float = 0.7
video_modality_scale: float = 3.0
video_skip_step: int = 0
video_stg_blocks: list[int] = field(default_factory=lambda: [28])
audio_cfg_scale: float = 7.0
audio_stg_scale: float = 1.0
audio_rescale_scale: float = 0.7
audio_modality_scale: float = 3.0
audio_skip_step: int = 0
audio_stg_blocks: list[int] = field(default_factory=lambda: [28])
def build_request_extra(self) -> dict[str, Any]:
extra = super().build_request_extra()
extra["ltx2_stage1_guider_params"] = {
"video_cfg_scale": self.video_cfg_scale,
"video_stg_scale": self.video_stg_scale,
"video_rescale_scale": self.video_rescale_scale,
"video_modality_scale": self.video_modality_scale,
"video_skip_step": self.video_skip_step,
"video_stg_blocks": self.video_stg_blocks,
"audio_cfg_scale": self.audio_cfg_scale,
"audio_stg_scale": self.audio_stg_scale,
"audio_rescale_scale": self.audio_rescale_scale,
"audio_modality_scale": self.audio_modality_scale,
"audio_skip_step": self.audio_skip_step,
"audio_stg_blocks": self.audio_stg_blocks,
}
return extra
@@ -253,6 +253,18 @@ class SamplingParams:
if env_steps is not None and self.num_inference_steps is not None: if env_steps is not None and self.num_inference_steps is not None:
self.num_inference_steps = int(env_steps) self.num_inference_steps = int(env_steps)
def build_request_extra(self) -> dict[str, Any]:
"""Return optional request-scoped extras for downstream pipeline stages."""
extra = {}
diffusers_kwargs = getattr(self, "diffusers_kwargs", None)
if diffusers_kwargs:
extra["diffusers_kwargs"] = diffusers_kwargs
return extra
def apply_request_extra(self, req: Any) -> None:
"""Merge request extras (model specific, e.g., LTX2.3) into an already-created pipeline request."""
req.extra.update(self.build_request_extra())
def _adjust_output_quality(self, output_quality: str, data_type: DataType) -> int: def _adjust_output_quality(self, output_quality: str, data_type: DataType) -> int:
"""Convert output_quality string to compression level.""" """Convert output_quality string to compression level."""
output_quality_mapper = {"maximum": 100, "high": 90, "medium": 55, "low": 35} output_quality_mapper = {"maximum": 100, "high": 90, "medium": 55, "low": 35}
@@ -0,0 +1,302 @@
import json
import os
from huggingface_hub import snapshot_download
from safetensors import safe_open
from safetensors.torch import save_file
from sglang.multimodal_gen.runtime.utils.model_overlay import (
_copytree_link_or_copy,
_ensure_dir,
_link_or_copy_file,
)
AUXILIARY_MODEL_ID = "Lightricks/LTX-2"
CONFIG_DONOR_MODEL_ID = "FastVideo/LTX-2.3-Distilled-Diffusers"
AUXILIARY_PATTERNS = [
"audio_vae/**",
"scheduler/**",
"text_encoder/**",
"tokenizer/**",
"vae/config.json",
"vae/diffusion_pytorch_model.safetensors",
]
CONFIG_DONOR_PATTERNS = [
"transformer/config.json",
"text_encoder/config.json",
"vae/**",
"vocoder/**",
]
MONOLITH_PREFIX = "model.diffusion_model."
VIDEO_CONNECTOR_PREFIX = f"{MONOLITH_PREFIX}video_embeddings_connector."
AUDIO_CONNECTOR_PREFIX = f"{MONOLITH_PREFIX}audio_embeddings_connector."
TEXT_PROJ_IN_PREFIX = f"{MONOLITH_PREFIX}text_proj_in."
VIDEO_AGGREGATE_PREFIX = "text_embedding_projection.video_aggregate_embed."
AUDIO_AGGREGATE_PREFIX = "text_embedding_projection.audio_aggregate_embed."
def _load_json(path: str) -> dict:
with open(path) as f:
return json.load(f)
def _write_json(path: str, payload: dict) -> None:
with open(path, "w") as f:
json.dump(payload, f, indent=2)
f.write("\n")
def _rename_connector_key(key: str) -> str | None:
if key.startswith(VIDEO_CONNECTOR_PREFIX):
suffix = key[len(VIDEO_CONNECTOR_PREFIX) :]
suffix = suffix.replace("transformer_1d_blocks", "transformer_blocks")
suffix = suffix.replace(".attn1.q_norm.", ".attn1.norm_q.")
suffix = suffix.replace(".attn1.k_norm.", ".attn1.norm_k.")
return f"video_connector.{suffix}"
if key.startswith(AUDIO_CONNECTOR_PREFIX):
suffix = key[len(AUDIO_CONNECTOR_PREFIX) :]
suffix = suffix.replace("transformer_1d_blocks", "transformer_blocks")
suffix = suffix.replace(".attn1.q_norm.", ".attn1.norm_q.")
suffix = suffix.replace(".attn1.k_norm.", ".attn1.norm_k.")
return f"audio_connector.{suffix}"
if key.startswith(TEXT_PROJ_IN_PREFIX):
return key[len(MONOLITH_PREFIX) :]
if key.startswith(VIDEO_AGGREGATE_PREFIX):
return f"video_aggregate_embed.{key[len(VIDEO_AGGREGATE_PREFIX):]}"
if key.startswith(AUDIO_AGGREGATE_PREFIX):
return f"audio_aggregate_embed.{key[len(AUDIO_AGGREGATE_PREFIX):]}"
return None
def _repack_transformer_weights(source_path: str, output_path: str) -> None:
tensors = {}
with safe_open(source_path, framework="pt") as f:
for key in f.keys():
if not key.startswith(MONOLITH_PREFIX):
continue
if key.startswith(VIDEO_CONNECTOR_PREFIX):
continue
if key.startswith(AUDIO_CONNECTOR_PREFIX):
continue
if key.startswith(TEXT_PROJ_IN_PREFIX):
continue
tensors[key[len(MONOLITH_PREFIX) :]] = f.get_tensor(key)
if not tensors:
raise ValueError("No transformer tensors found in LTX-2.3 source checkpoint.")
save_file(tensors, output_path)
def _repack_connectors_weights(source_path: str, output_path: str) -> None:
tensors = {}
with safe_open(source_path, framework="pt") as f:
for key in f.keys():
renamed = _rename_connector_key(key)
if renamed is None:
continue
tensors[renamed] = f.get_tensor(key)
if not tensors:
raise ValueError("No connector tensors found in LTX-2.3 source checkpoint.")
save_file(tensors, output_path)
def _build_transformer_config(config_donor_dir: str) -> dict:
config = _load_json(os.path.join(config_donor_dir, "transformer", "config.json"))
config["_class_name"] = "LTX2VideoTransformer3DModel"
config["force_sdpa_v2a_cross_attention"] = True
config["quantize_video_rope_coords_to_hidden_dtype"] = True
return config
def _build_connectors_config(config_donor_dir: str) -> dict:
text_encoder_config = _load_json(
os.path.join(config_donor_dir, "text_encoder", "config.json")
)
return {
"_class_name": "LTX2TextConnectors",
"_diffusers_version": "0.37.0.dev0",
"audio_connector_attention_head_dim": text_encoder_config[
"audio_connector_attention_head_dim"
],
"audio_connector_num_attention_heads": text_encoder_config[
"audio_connector_num_attention_heads"
],
"audio_connector_num_layers": text_encoder_config["audio_connector_num_layers"],
"audio_connector_num_learnable_registers": text_encoder_config[
"connector_num_learnable_registers"
],
"audio_feature_extractor_out_features": text_encoder_config[
"audio_feature_extractor_out_features"
],
"caption_channels": text_encoder_config["hidden_size"],
"causal_temporal_positioning": False,
"connector_apply_gated_attention": text_encoder_config[
"connector_apply_gated_attention"
],
"feature_extractor_in_features": text_encoder_config[
"feature_extractor_in_features"
],
"connector_rope_base_seq_len": text_encoder_config[
"connector_positional_embedding_max_pos"
][0],
"rope_double_precision": text_encoder_config["connector_double_precision_rope"],
"rope_theta": text_encoder_config["connector_positional_embedding_theta"],
"rope_type": text_encoder_config["connector_rope_type"],
"text_proj_in_factor": text_encoder_config["feature_extractor_in_features"]
// text_encoder_config["hidden_size"],
"video_feature_extractor_out_features": text_encoder_config[
"video_feature_extractor_out_features"
],
"video_connector_attention_head_dim": text_encoder_config[
"connector_attention_head_dim"
],
"video_connector_num_attention_heads": text_encoder_config[
"connector_num_attention_heads"
],
"video_connector_num_layers": text_encoder_config["connector_num_layers"],
"video_connector_num_learnable_registers": text_encoder_config[
"connector_num_learnable_registers"
],
}
def _build_vae_config(auxiliary_dir: str, config_donor_dir: str) -> dict:
config = _load_json(os.path.join(auxiliary_dir, "vae", "config.json"))
config["ltx_variant"] = "ltx_2_3"
config["condition_encoder_subdir"] = "ltx23_image_encoder"
config["video_decoder_variant"] = "ltx_2_3"
config["video_decoder_config"] = _load_json(
os.path.join(config_donor_dir, "vae", "config.json")
)["vae"]
return config
def _repack_ltx23_image_encoder_weights(source_path: str, output_path: str) -> None:
tensors = {}
with safe_open(source_path, framework="pt") as f:
for key in f.keys():
if key.startswith("encoder."):
tensors[key[len("encoder.") :]] = f.get_tensor(key)
continue
if key.startswith("per_channel_statistics."):
tensors[key] = f.get_tensor(key)
if not tensors:
raise ValueError("No LTX-2.3 image-encoder tensors found in donor checkpoint.")
save_file(tensors, output_path)
def _repack_ltx23_video_decoder_weights(
auxiliary_encoder_path: str,
donor_decoder_path: str,
output_path: str,
) -> None:
tensors = {}
with safe_open(auxiliary_encoder_path, framework="pt") as f:
for key in f.keys():
if key.startswith("encoder."):
tensors[key] = f.get_tensor(key)
with safe_open(donor_decoder_path, framework="pt") as f:
for key in f.keys():
if key.startswith("decoder."):
tensors[key] = f.get_tensor(key)
continue
if key == "per_channel_statistics.mean-of-means":
tensor = f.get_tensor(key)
tensors["decoder.per_channel_statistics.mean_of_means"] = tensor
tensors["latents_mean"] = tensor.clone()
continue
if key == "per_channel_statistics.std-of-means":
tensor = f.get_tensor(key)
tensors["decoder.per_channel_statistics.std_of_means"] = tensor
tensors["latents_std"] = tensor.clone()
continue
if not tensors:
raise ValueError("No LTX-2.3 decoder tensors found in donor checkpoint.")
save_file(tensors, output_path)
def materialize(
*,
overlay_dir: str,
source_dir: str,
output_dir: str,
manifest: dict,
) -> None:
_ = overlay_dir, manifest
auxiliary_dir = snapshot_download(
repo_id=AUXILIARY_MODEL_ID,
allow_patterns=AUXILIARY_PATTERNS,
max_workers=8,
)
config_donor_dir = snapshot_download(
repo_id=CONFIG_DONOR_MODEL_ID,
allow_patterns=CONFIG_DONOR_PATTERNS,
max_workers=8,
)
for component_name in ("audio_vae", "scheduler", "text_encoder", "tokenizer"):
_copytree_link_or_copy(
os.path.join(auxiliary_dir, component_name),
os.path.join(output_dir, component_name),
)
_copytree_link_or_copy(
os.path.join(config_donor_dir, "vocoder"),
os.path.join(output_dir, "vocoder"),
)
source_checkpoint = os.path.join(source_dir, "ltx-2.3-22b-dev.safetensors")
transformer_dir = os.path.join(output_dir, "transformer")
_ensure_dir(transformer_dir)
_write_json(
os.path.join(transformer_dir, "config.json"),
_build_transformer_config(config_donor_dir),
)
_repack_transformer_weights(
source_checkpoint, os.path.join(transformer_dir, "model.safetensors")
)
connectors_dir = os.path.join(output_dir, "connectors")
_ensure_dir(connectors_dir)
_write_json(
os.path.join(connectors_dir, "config.json"),
_build_connectors_config(config_donor_dir),
)
_repack_connectors_weights(
source_checkpoint, os.path.join(connectors_dir, "model.safetensors")
)
vae_dir = os.path.join(output_dir, "vae")
_ensure_dir(vae_dir)
_write_json(
os.path.join(vae_dir, "config.json"),
_build_vae_config(auxiliary_dir, config_donor_dir),
)
_repack_ltx23_video_decoder_weights(
os.path.join(auxiliary_dir, "vae", "diffusion_pytorch_model.safetensors"),
os.path.join(config_donor_dir, "vae", "model.safetensors"),
os.path.join(vae_dir, "model.safetensors"),
)
image_encoder_dir = os.path.join(vae_dir, "ltx23_image_encoder")
_ensure_dir(image_encoder_dir)
_link_or_copy_file(
os.path.join(config_donor_dir, "vae", "config.json"),
os.path.join(image_encoder_dir, "config.json"),
)
_repack_ltx23_image_encoder_weights(
os.path.join(config_donor_dir, "vae", "model.safetensors"),
os.path.join(image_encoder_dir, "model.safetensors"),
)
_link_or_copy_file(
os.path.join(source_dir, "ltx-2.3-22b-distilled-lora-384.safetensors"),
os.path.join(output_dir, "ltx-2.3-22b-distilled-lora-384.safetensors"),
)
_link_or_copy_file(
os.path.join(source_dir, "ltx-2.3-spatial-upscaler-x2-1.1.safetensors"),
os.path.join(output_dir, "ltx-2.3-spatial-upscaler-x2-1.1.safetensors"),
)
+26 -6
View File
@@ -90,7 +90,10 @@ from sglang.multimodal_gen.configs.sample.hunyuan import (
HunyuanSamplingParams, HunyuanSamplingParams,
) )
from sglang.multimodal_gen.configs.sample.hunyuan3d import Hunyuan3DSamplingParams from sglang.multimodal_gen.configs.sample.hunyuan3d import Hunyuan3DSamplingParams
from sglang.multimodal_gen.configs.sample.ltx_2 import LTX2SamplingParams from sglang.multimodal_gen.configs.sample.ltx_2 import (
LTX2SamplingParams,
LTX23SamplingParams,
)
from sglang.multimodal_gen.configs.sample.mova import ( from sglang.multimodal_gen.configs.sample.mova import (
MOVA_360P_SamplingParams, MOVA_360P_SamplingParams,
MOVA_720P_SamplingParams, MOVA_720P_SamplingParams,
@@ -155,7 +158,18 @@ def _discover_and_register_pipelines():
package.__path__, package.__name__ + "." package.__path__, package.__name__ + "."
): ):
if not ispkg: if not ispkg:
pipeline_module = importlib.import_module(module_name) try:
pipeline_module = importlib.import_module(module_name)
except Exception as exc:
logger.warning(
"Skipping pipeline module %s during discovery due to import failure: %s",
module_name,
exc,
)
logger.debug(
"Pipeline import failure details for %s", module_name, exc_info=True
)
continue
if hasattr(pipeline_module, "EntryClass"): if hasattr(pipeline_module, "EntryClass"):
entry_cls = pipeline_module.EntryClass entry_cls = pipeline_module.EntryClass
entry_cls_list = ( entry_cls_list = (
@@ -594,12 +608,18 @@ 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=[ hf_model_paths=["Lightricks/LTX-2"],
"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() and "ltx-2.3" not in path.lower(),
],
)
register_configs(
sampling_param_cls=LTX23SamplingParams,
pipeline_config_cls=LTX2PipelineConfig,
hf_model_paths=["Lightricks/LTX-2.3"],
model_detectors=[
lambda path: "ltx-2.3" in path.lower(),
], ],
) )
@@ -552,13 +552,15 @@ class DiffGenerator:
self.shutdown() self.shutdown()
def __del__(self): def __del__(self):
if self.owns_scheduler_client: owns_scheduler_client = bool(getattr(self, "owns_scheduler_client", False))
local_scheduler_process = getattr(self, "local_scheduler_process", None)
if owns_scheduler_client:
logger.warning( logger.warning(
"Generator was garbage collected without being shut down. " "Generator was garbage collected without being shut down. "
"Attempting to shut down the local server and client." "Attempting to shut down the local server and client."
) )
self.shutdown() self.shutdown()
elif self.local_scheduler_process: elif local_scheduler_process:
logger.warning( logger.warning(
"Generator was garbage collected without being shut down. " "Generator was garbage collected without being shut down. "
"Attempting to shut down the local server." "Attempting to shut down the local server."
@@ -288,12 +288,7 @@ def prepare_request(
sampling_params=sampling_params, sampling_params=sampling_params,
VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, VSA_sparsity=server_args.attention_backend_config.VSA_sparsity,
) )
try: sampling_params.apply_request_extra(req)
diffusers_kwargs = sampling_params.diffusers_kwargs
except AttributeError:
diffusers_kwargs = None
if diffusers_kwargs:
req.extra["diffusers_kwargs"] = diffusers_kwargs
req.adjust_size(server_args) req.adjust_size(server_args)
@@ -1,5 +1,8 @@
from safetensors.torch import load_file as safetensors_load_file from safetensors.torch import load_file as safetensors_load_file
from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import (
LTX2ConnectorConfig,
)
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.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader, ComponentLoader,
@@ -50,10 +53,9 @@ class AdapterLoader(ComponentLoader):
target_device = get_local_torch_device() target_device = get_local_torch_device()
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision] default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
from types import SimpleNamespace
with set_default_torch_dtype(default_dtype), skip_init_modules(): with set_default_torch_dtype(default_dtype), skip_init_modules():
connector_cfg = SimpleNamespace(**config) connector_cfg = LTX2ConnectorConfig()
connector_cfg.update_model_arch(config)
model = model_cls(connector_cfg).to( model = model_cls(connector_cfg).to(
device=target_device, dtype=default_dtype device=target_device, dtype=default_dtype
) )
@@ -1,3 +1,4 @@
import math
from typing import Optional, Tuple, Union from typing import Optional, Tuple, Union
import torch import torch
@@ -89,6 +90,7 @@ class LTX2Attention(torch.nn.Module):
norm_eps: float = 1e-6, norm_eps: float = 1e-6,
norm_elementwise_affine: bool = True, norm_elementwise_affine: bool = True,
rope_type: str = "interleaved", rope_type: str = "interleaved",
apply_gated_attention: bool = False,
processor=None, processor=None,
): ):
super().__init__() super().__init__()
@@ -125,6 +127,9 @@ class LTX2Attention(torch.nn.Module):
self.to_v = torch.nn.Linear( self.to_v = torch.nn.Linear(
self.cross_attention_dim, self.inner_kv_dim, bias=bias self.cross_attention_dim, self.inner_kv_dim, bias=bias
) )
self.to_gate_logits = None
if apply_gated_attention:
self.to_gate_logits = torch.nn.Linear(query_dim, heads, bias=True)
self.to_out = torch.nn.ModuleList([]) self.to_out = torch.nn.ModuleList([])
self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) self.to_out.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(torch.nn.Dropout(dropout)) self.to_out.append(torch.nn.Dropout(dropout))
@@ -153,6 +158,7 @@ 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:
gate_input = hidden_states
if encoder_hidden_states is None: if encoder_hidden_states is None:
encoder_hidden_states = hidden_states encoder_hidden_states = hidden_states
@@ -199,6 +205,15 @@ class LTX2Attention(torch.nn.Module):
hidden_states = hidden_states.transpose(1, 2).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)
if self.to_gate_logits is not None:
gate_logits = self.to_gate_logits(gate_input)
b, t, _ = hidden_states.shape
hidden_states = hidden_states.view(b, t, self.heads, self.head_dim)
hidden_states = hidden_states * (
2.0 * torch.sigmoid(gate_logits).unsqueeze(-1)
)
hidden_states = hidden_states.view(b, t, self.heads * self.head_dim)
hidden_states = self.to_out[0](hidden_states) hidden_states = self.to_out[0](hidden_states)
hidden_states = self.to_out[1](hidden_states) hidden_states = self.to_out[1](hidden_states)
return hidden_states return hidden_states
@@ -317,6 +332,7 @@ class LTX2TransformerBlock1d(nn.Module):
activation_fn: str = "gelu-approximate", activation_fn: str = "gelu-approximate",
eps: float = 1e-6, eps: float = 1e-6,
rope_type: str = "interleaved", rope_type: str = "interleaved",
apply_gated_attention: bool = False,
): ):
super().__init__() super().__init__()
@@ -327,6 +343,7 @@ class LTX2TransformerBlock1d(nn.Module):
kv_heads=num_attention_heads, kv_heads=num_attention_heads,
dim_head=attention_head_dim, dim_head=attention_head_dim,
rope_type=rope_type, rope_type=rope_type,
apply_gated_attention=apply_gated_attention,
) )
self.norm2 = torch.nn.RMSNorm(dim, eps=eps, elementwise_affine=False) self.norm2 = torch.nn.RMSNorm(dim, eps=eps, elementwise_affine=False)
@@ -373,6 +390,7 @@ class LTX2ConnectorTransformer1d(nn.Module):
eps: float = 1e-6, eps: float = 1e-6,
causal_temporal_positioning: bool = False, causal_temporal_positioning: bool = False,
rope_type: str = "interleaved", rope_type: str = "interleaved",
apply_gated_attention: bool = False,
): ):
super().__init__() super().__init__()
self.num_attention_heads = num_attention_heads self.num_attention_heads = num_attention_heads
@@ -403,6 +421,7 @@ class LTX2ConnectorTransformer1d(nn.Module):
num_attention_heads=num_attention_heads, num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim, attention_head_dim=attention_head_dim,
rope_type=rope_type, rope_type=rope_type,
apply_gated_attention=apply_gated_attention,
) )
for _ in range(num_layers) for _ in range(num_layers)
] ]
@@ -516,10 +535,37 @@ class LTX2TextConnectors(nn.Module):
rope_double_precision = config.rope_double_precision rope_double_precision = config.rope_double_precision
causal_temporal_positioning = config.causal_temporal_positioning causal_temporal_positioning = config.causal_temporal_positioning
rope_type = config.rope_type rope_type = config.rope_type
connector_apply_gated_attention = config.connector_apply_gated_attention
self.text_proj_in = nn.Linear( feature_extractor_in_features = config.feature_extractor_in_features
caption_channels * text_proj_in_factor, caption_channels, bias=False video_feature_extractor_out_features = (
config.video_feature_extractor_out_features
) )
audio_feature_extractor_out_features = (
config.audio_feature_extractor_out_features
)
self.text_proj_in: nn.Linear | None = None
self.video_aggregate_embed: nn.Linear | None = None
self.audio_aggregate_embed: nn.Linear | None = None
if (
feature_extractor_in_features > 0
and video_feature_extractor_out_features > 0
and audio_feature_extractor_out_features > 0
):
self.video_aggregate_embed = nn.Linear(
feature_extractor_in_features,
video_feature_extractor_out_features,
bias=True,
)
self.audio_aggregate_embed = nn.Linear(
feature_extractor_in_features,
audio_feature_extractor_out_features,
bias=True,
)
else:
self.text_proj_in = nn.Linear(
caption_channels * text_proj_in_factor, caption_channels, bias=False
)
self.video_connector = LTX2ConnectorTransformer1d( self.video_connector = LTX2ConnectorTransformer1d(
num_attention_heads=video_connector_num_attention_heads, num_attention_heads=video_connector_num_attention_heads,
attention_head_dim=video_connector_attention_head_dim, attention_head_dim=video_connector_attention_head_dim,
@@ -530,6 +576,7 @@ class LTX2TextConnectors(nn.Module):
rope_double_precision=rope_double_precision, rope_double_precision=rope_double_precision,
causal_temporal_positioning=causal_temporal_positioning, causal_temporal_positioning=causal_temporal_positioning,
rope_type=rope_type, rope_type=rope_type,
apply_gated_attention=connector_apply_gated_attention,
) )
self.audio_connector = LTX2ConnectorTransformer1d( self.audio_connector = LTX2ConnectorTransformer1d(
num_attention_heads=audio_connector_num_attention_heads, num_attention_heads=audio_connector_num_attention_heads,
@@ -541,8 +588,15 @@ class LTX2TextConnectors(nn.Module):
rope_double_precision=rope_double_precision, rope_double_precision=rope_double_precision,
causal_temporal_positioning=causal_temporal_positioning, causal_temporal_positioning=causal_temporal_positioning,
rope_type=rope_type, rope_type=rope_type,
apply_gated_attention=connector_apply_gated_attention,
) )
@staticmethod
def _rescale_v2_features(
x: torch.Tensor, target_dim: int, source_dim: int
) -> torch.Tensor:
return x * math.sqrt(target_dim / source_dim)
def forward( def forward(
self, self,
text_encoder_hidden_states: torch.Tensor, text_encoder_hidden_states: torch.Tensor,
@@ -557,12 +611,6 @@ class LTX2TextConnectors(nn.Module):
) )
attention_mask = attention_mask.to(text_dtype) * torch.finfo(text_dtype).max attention_mask = attention_mask.to(text_dtype) * torch.finfo(text_dtype).max
# Ensure input dtype matches the layer's weight dtype
if text_encoder_hidden_states.dtype != self.text_proj_in.weight.dtype:
text_encoder_hidden_states = text_encoder_hidden_states.to(
self.text_proj_in.weight.dtype
)
# Ensure sequence length is divisible by num_learnable_registers (128) # Ensure sequence length is divisible by num_learnable_registers (128)
seq_len = text_encoder_hidden_states.shape[1] seq_len = text_encoder_hidden_states.shape[1]
num_learnable_registers = self.video_connector.num_learnable_registers num_learnable_registers = self.video_connector.num_learnable_registers
@@ -579,10 +627,44 @@ class LTX2TextConnectors(nn.Module):
# Pad with a large negative value to mask out the new tokens # Pad with a large negative value to mask out the new tokens
attention_mask = F.pad(attention_mask, (0, pad_len), value=-1000000.0) attention_mask = F.pad(attention_mask, (0, pad_len), value=-1000000.0)
text_encoder_hidden_states = self.text_proj_in(text_encoder_hidden_states) if (
self.video_aggregate_embed is not None
and self.audio_aggregate_embed is not None
):
video_hidden_states = text_encoder_hidden_states
audio_hidden_states = text_encoder_hidden_states
if video_hidden_states.dtype != self.video_aggregate_embed.weight.dtype:
video_hidden_states = video_hidden_states.to(
self.video_aggregate_embed.weight.dtype
)
if audio_hidden_states.dtype != self.audio_aggregate_embed.weight.dtype:
audio_hidden_states = audio_hidden_states.to(
self.audio_aggregate_embed.weight.dtype
)
source_dim = self.video_aggregate_embed.out_features
video_hidden_states = self._rescale_v2_features(
video_hidden_states,
self.video_aggregate_embed.out_features,
source_dim,
)
audio_hidden_states = self._rescale_v2_features(
audio_hidden_states,
self.audio_aggregate_embed.out_features,
source_dim,
)
video_hidden_states = self.video_aggregate_embed(video_hidden_states)
audio_hidden_states = self.audio_aggregate_embed(audio_hidden_states)
else:
assert self.text_proj_in is not None
if text_encoder_hidden_states.dtype != self.text_proj_in.weight.dtype:
text_encoder_hidden_states = text_encoder_hidden_states.to(
self.text_proj_in.weight.dtype
)
video_hidden_states = self.text_proj_in(text_encoder_hidden_states)
audio_hidden_states = video_hidden_states
video_text_embedding, new_attn_mask = self.video_connector( video_text_embedding, new_attn_mask = self.video_connector(
text_encoder_hidden_states, attention_mask video_hidden_states, attention_mask
) )
attn_mask = (new_attn_mask < 1e-6).to(torch.int64) attn_mask = (new_attn_mask < 1e-6).to(torch.int64)
@@ -593,7 +675,7 @@ class LTX2TextConnectors(nn.Module):
new_attn_mask = attn_mask.squeeze(-1) new_attn_mask = attn_mask.squeeze(-1)
audio_text_embedding, _ = self.audio_connector( audio_text_embedding, _ = self.audio_connector(
text_encoder_hidden_states, attention_mask audio_hidden_states, attention_mask
) )
return video_text_embedding, audio_text_embedding, new_attn_mask return video_text_embedding, audio_text_embedding, new_attn_mask
@@ -37,6 +37,15 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
ADALN_NUM_BASE_PARAMS = 6
ADALN_NUM_CROSS_ATTN_PARAMS = 3
def adaln_embedding_coefficient(cross_attention_adaln: bool) -> int:
return ADALN_NUM_BASE_PARAMS + (
ADALN_NUM_CROSS_ATTN_PARAMS if cross_attention_adaln else 0
)
def apply_interleaved_rotary_emb( def apply_interleaved_rotary_emb(
x: torch.Tensor, freqs: Tuple[torch.Tensor, torch.Tensor] x: torch.Tensor, freqs: Tuple[torch.Tensor, torch.Tensor]
@@ -447,6 +456,7 @@ class LTX2Attention(nn.Module):
norm_eps: float = 1e-6, norm_eps: float = 1e-6,
qk_norm: bool = True, qk_norm: bool = True,
use_local_attention: bool = False, use_local_attention: bool = False,
apply_gated_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,
@@ -461,6 +471,8 @@ class LTX2Attention(nn.Module):
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) self.use_local_attention = bool(use_local_attention)
self.apply_gated_attention = bool(apply_gated_attention)
self.prefix = prefix
tp_size = get_tp_world_size() tp_size = get_tp_world_size()
if tp_size <= 0: if tp_size <= 0:
@@ -499,6 +511,15 @@ class LTX2Attention(nn.Module):
gather_output=False, gather_output=False,
quant_config=quant_config, quant_config=quant_config,
) )
self.to_gate_logits: ColumnParallelLinear | None = None
if self.apply_gated_attention:
self.to_gate_logits = ColumnParallelLinear(
self.query_dim,
self.heads,
bias=True,
gather_output=False,
quant_config=quant_config,
)
self.q_norm: nn.Module | None = None self.q_norm: nn.Module | None = None
self.k_norm: nn.Module | None = None self.k_norm: nn.Module | None = None
@@ -561,6 +582,7 @@ class LTX2Attention(nn.Module):
perturbation_mask: torch.Tensor | None = None, perturbation_mask: torch.Tensor | None = None,
all_perturbed: bool = False, all_perturbed: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
gate_input = x
context_ = x if context is None else context context_ = x if context is None else context
v, _ = self.to_v(context_) v, _ = self.to_v(context_)
use_attention = not all_perturbed use_attention = not all_perturbed
@@ -609,9 +631,17 @@ class LTX2Attention(nn.Module):
if not use_attention: if not use_attention:
out = v out = v
out = out.flatten(2) if self.to_gate_logits is not None:
out, _ = self.to_out[0](out) gate_logits, _ = self.to_gate_logits(gate_input)
return out b, t = out.shape[:2]
out = out.view(b, t, self.local_heads, self.dim_head)
out = out * (2.0 * torch.sigmoid(gate_logits).unsqueeze(-1))
out = out.view(b, t, self.local_heads * self.dim_head)
out_flat = out.flatten(2)
out_proj, _ = self.to_out[0](out_flat)
return out_proj
def _slice_rope_for_tp( def _slice_rope_for_tp(
self, self,
@@ -688,6 +718,10 @@ class LTX2TransformerBlock(nn.Module):
audio_cross_attention_dim: int, audio_cross_attention_dim: int,
qk_norm: bool = True, qk_norm: bool = True,
norm_eps: float = 1e-6, norm_eps: float = 1e-6,
apply_gated_attention: bool = False,
cross_attention_adaln: bool = False,
use_local_av_cross_attention: bool = False,
force_sdpa_v2a_cross_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,
@@ -695,6 +729,9 @@ class LTX2TransformerBlock(nn.Module):
super().__init__() super().__init__()
self.idx = idx self.idx = idx
self.norm_eps = norm_eps self.norm_eps = norm_eps
# LTX2.3
self.cross_attention_adaln = cross_attention_adaln
self.use_local_av_cross_attention = use_local_av_cross_attention
# 1. Self-Attention (video and audio) # 1. Self-Attention (video and audio)
self.attn1 = LTX2Attention( self.attn1 = LTX2Attention(
@@ -703,6 +740,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,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1", prefix=f"{prefix}.attn1",
quant_config=quant_config, quant_config=quant_config,
@@ -713,6 +751,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,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.audio_attn1", prefix=f"{prefix}.audio_attn1",
quant_config=quant_config, quant_config=quant_config,
@@ -729,6 +768,7 @@ class LTX2TransformerBlock(nn.Module):
norm_eps=norm_eps, norm_eps=norm_eps,
qk_norm=qk_norm, qk_norm=qk_norm,
use_local_attention=True, use_local_attention=True,
apply_gated_attention=apply_gated_attention,
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,
@@ -741,6 +781,7 @@ class LTX2TransformerBlock(nn.Module):
norm_eps=norm_eps, norm_eps=norm_eps,
qk_norm=qk_norm, qk_norm=qk_norm,
use_local_attention=True, use_local_attention=True,
apply_gated_attention=apply_gated_attention,
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,
@@ -754,6 +795,8 @@ 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=use_local_av_cross_attention,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.audio_to_video_attn", prefix=f"{prefix}.audio_to_video_attn",
quant_config=quant_config, quant_config=quant_config,
@@ -765,7 +808,13 @@ 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,
supported_attention_backends=supported_attention_backends, use_local_attention=use_local_av_cross_attention,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=(
{AttentionBackendEnum.TORCH_SDPA}
if force_sdpa_v2a_cross_attention
else supported_attention_backends
),
prefix=f"{prefix}.video_to_audio_attn", prefix=f"{prefix}.video_to_audio_attn",
quant_config=quant_config, quant_config=quant_config,
) )
@@ -777,14 +826,23 @@ class LTX2TransformerBlock(nn.Module):
) )
# 5. Modulation Parameters # 5. Modulation Parameters
self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5) num_ada_params = adaln_embedding_coefficient(cross_attention_adaln)
self.scale_shift_table = nn.Parameter(
torch.randn(num_ada_params, dim) / dim**0.5
)
self.audio_scale_shift_table = nn.Parameter( self.audio_scale_shift_table = nn.Parameter(
torch.randn(6, audio_dim) / audio_dim**0.5 torch.randn(num_ada_params, audio_dim) / audio_dim**0.5
) )
self.video_a2v_cross_attn_scale_shift_table = nn.Parameter(torch.randn(5, dim)) self.video_a2v_cross_attn_scale_shift_table = nn.Parameter(torch.randn(5, dim))
self.audio_a2v_cross_attn_scale_shift_table = nn.Parameter( self.audio_a2v_cross_attn_scale_shift_table = nn.Parameter(
torch.randn(5, audio_dim) torch.randn(5, audio_dim)
) )
if self.cross_attention_adaln:
# LTX2.3
self.prompt_scale_shift_table = nn.Parameter(torch.randn(2, dim))
self.audio_prompt_scale_shift_table = nn.Parameter(
torch.randn(2, audio_dim)
)
def get_ada_values( def get_ada_values(
self, self,
@@ -813,6 +871,8 @@ class LTX2TransformerBlock(nn.Module):
audio_encoder_hidden_states: torch.Tensor, audio_encoder_hidden_states: torch.Tensor,
temb: torch.Tensor, temb: torch.Tensor,
temb_audio: torch.Tensor, temb_audio: torch.Tensor,
temb_prompt: torch.Tensor | None,
temb_audio_prompt: torch.Tensor | None,
temb_ca_scale_shift: torch.Tensor, temb_ca_scale_shift: torch.Tensor,
temb_ca_audio_scale_shift: torch.Tensor, temb_ca_audio_scale_shift: torch.Tensor,
temb_ca_gate: torch.Tensor, temb_ca_gate: torch.Tensor,
@@ -860,21 +920,70 @@ class LTX2TransformerBlock(nn.Module):
) )
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) if self.cross_attention_adaln:
attn_hidden_states = self.attn2( # LTX2.3
norm_hidden_states, if temb_prompt is None or temb_audio_prompt is None:
context=encoder_hidden_states, raise ValueError(
mask=encoder_attention_mask, "cross_attention_adaln requires prompt modulation tensors."
) )
hidden_states = hidden_states + attn_hidden_states vshift_q, vscale_q, vgate_q = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(6, 9)
)
v_prompt_shift, v_prompt_scale = self.get_ada_values(
self.prompt_scale_shift_table, batch_size, temb_prompt, slice(None)
)
norm_hidden_states = (
rms_norm(hidden_states, self.norm_eps) * (1 + vscale_q) + vshift_q
)
mod_encoder_hidden_states = (
encoder_hidden_states * (1 + v_prompt_scale) + v_prompt_shift
)
attn_hidden_states = self.attn2(
norm_hidden_states,
context=mod_encoder_hidden_states,
mask=encoder_attention_mask,
)
hidden_states = hidden_states + attn_hidden_states * vgate_q
norm_audio_hidden_states = rms_norm(audio_hidden_states, self.norm_eps) ashift_q, ascale_q, agate_q = self.get_ada_values(
attn_audio_hidden_states = self.audio_attn2( self.audio_scale_shift_table, batch_size, temb_audio, slice(6, 9)
norm_audio_hidden_states, )
context=audio_encoder_hidden_states, a_prompt_shift, a_prompt_scale = self.get_ada_values(
mask=audio_encoder_attention_mask, self.audio_prompt_scale_shift_table,
) batch_size,
audio_hidden_states = audio_hidden_states + attn_audio_hidden_states temb_audio_prompt,
slice(None),
)
norm_audio_hidden_states = (
rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_q) + ashift_q
)
mod_audio_encoder_hidden_states = (
audio_encoder_hidden_states * (1 + a_prompt_scale) + a_prompt_shift
)
attn_audio_hidden_states = self.audio_attn2(
norm_audio_hidden_states,
context=mod_audio_encoder_hidden_states,
mask=audio_encoder_attention_mask,
)
audio_hidden_states = (
audio_hidden_states + attn_audio_hidden_states * agate_q
)
else:
norm_hidden_states = rms_norm(hidden_states, self.norm_eps)
attn_hidden_states = self.attn2(
norm_hidden_states,
context=encoder_hidden_states,
mask=encoder_attention_mask,
)
hidden_states = hidden_states + attn_hidden_states
norm_audio_hidden_states = rms_norm(audio_hidden_states, self.norm_eps)
attn_audio_hidden_states = self.audio_attn2(
norm_audio_hidden_states,
context=audio_encoder_hidden_states,
mask=audio_encoder_attention_mask,
)
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)
@@ -976,7 +1085,7 @@ class LTX2TransformerBlock(nn.Module):
) )
# 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, 6)
) )
norm_hidden_states = ( norm_hidden_states = (
rms_norm(hidden_states, self.norm_eps) * (1 + vscale_mlp) + vshift_mlp rms_norm(hidden_states, self.norm_eps) * (1 + vscale_mlp) + vshift_mlp
@@ -985,7 +1094,7 @@ class LTX2TransformerBlock(nn.Module):
hidden_states = hidden_states + ff_output * vgate_mlp hidden_states = hidden_states + ff_output * vgate_mlp
ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values( ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(3, None) self.audio_scale_shift_table, batch_size, temb_audio, slice(3, 6)
) )
norm_audio_hidden_states = ( norm_audio_hidden_states = (
rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_mlp) + ashift_mlp rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_mlp) + ashift_mlp
@@ -1003,6 +1112,12 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
reverse_param_names_mapping = LTX2ArchConfig().reverse_param_names_mapping reverse_param_names_mapping = LTX2ArchConfig().reverse_param_names_mapping
lora_param_names_mapping = LTX2ArchConfig().lora_param_names_mapping lora_param_names_mapping = LTX2ArchConfig().lora_param_names_mapping
@staticmethod
def _collapse_prompt_timestep(timestep: torch.Tensor) -> torch.Tensor:
if timestep.ndim <= 1:
return timestep
return timestep.amax(dim=tuple(range(1, timestep.ndim)))
def _validate_tp_config(self, *, arch: LTX2ArchConfig, tp_size: int) -> None: def _validate_tp_config(self, *, arch: LTX2ArchConfig, tp_size: int) -> None:
"""Validate TP-related dimension constraints (fail-fast).""" """Validate TP-related dimension constraints (fail-fast)."""
if tp_size < 1: if tp_size < 1:
@@ -1089,20 +1204,38 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
) )
# 2. Prompt embeddings # 2. Prompt embeddings
self.caption_projection = LTX2TextProjection( self.caption_projection: LTX2TextProjection | None = None
in_features=arch.caption_channels, hidden_size=self.hidden_size self.audio_caption_projection: LTX2TextProjection | None = None
) if not arch.caption_proj_before_connector:
self.audio_caption_projection = LTX2TextProjection( self.caption_projection = LTX2TextProjection(
in_features=arch.caption_channels, hidden_size=self.audio_hidden_size in_features=arch.caption_channels, hidden_size=self.hidden_size
) )
self.audio_caption_projection = LTX2TextProjection(
in_features=arch.caption_channels, hidden_size=self.audio_hidden_size
)
# 3. Timestep Modulation Params and Embedding # 3. Timestep Modulation Params and Embedding
self.adaln_single = LTX2AdaLayerNormSingle( self.adaln_single = LTX2AdaLayerNormSingle(
self.hidden_size, embedding_coefficient=6 self.hidden_size,
embedding_coefficient=adaln_embedding_coefficient(
arch.cross_attention_adaln
),
) )
self.audio_adaln_single = LTX2AdaLayerNormSingle( self.audio_adaln_single = LTX2AdaLayerNormSingle(
self.audio_hidden_size, embedding_coefficient=6 self.audio_hidden_size,
embedding_coefficient=adaln_embedding_coefficient(
arch.cross_attention_adaln
),
) )
self.prompt_adaln_single: LTX2AdaLayerNormSingle | None = None
self.audio_prompt_adaln_single: LTX2AdaLayerNormSingle | None = None
if arch.cross_attention_adaln:
self.prompt_adaln_single = LTX2AdaLayerNormSingle(
self.hidden_size, embedding_coefficient=2
)
self.audio_prompt_adaln_single = LTX2AdaLayerNormSingle(
self.audio_hidden_size, embedding_coefficient=2
)
# Global Cross Attention Modulation Parameters # Global Cross Attention Modulation Parameters
self.av_ca_video_scale_shift_adaln_single = LTX2AdaLayerNormSingle( self.av_ca_video_scale_shift_adaln_single = LTX2AdaLayerNormSingle(
@@ -1141,6 +1274,9 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
rope_double_precision = bool( rope_double_precision = bool(
hf_config.get("rope_double_precision", arch.double_precision_rope) hf_config.get("rope_double_precision", arch.double_precision_rope)
) )
self.quantize_video_rope_coords_to_hidden_dtype = bool(
hf_config.get("quantize_video_rope_coords_to_hidden_dtype", False)
)
causal_offset = int(hf_config.get("causal_offset", 1)) causal_offset = int(hf_config.get("causal_offset", 1))
pos_embed_max_pos = int(arch.positional_embedding_max_pos[0]) pos_embed_max_pos = int(arch.positional_embedding_max_pos[0])
@@ -1231,6 +1367,14 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_cross_attention_dim=arch.audio_cross_attention_dim, audio_cross_attention_dim=arch.audio_cross_attention_dim,
norm_eps=self.norm_eps, norm_eps=self.norm_eps,
qk_norm=True, # Always True in LTX2 qk_norm=True, # Always True in LTX2
apply_gated_attention=arch.apply_gated_attention,
cross_attention_adaln=arch.cross_attention_adaln,
use_local_av_cross_attention=bool(
getattr(arch, "use_local_av_cross_attention", False)
),
force_sdpa_v2a_cross_attention=bool(
getattr(arch, "force_sdpa_v2a_cross_attention", False)
),
supported_attention_backends=self._supported_attention_backends, supported_attention_backends=self._supported_attention_backends,
prefix=config.prefix, prefix=config.prefix,
quant_config=quant_config, quant_config=quant_config,
@@ -1336,6 +1480,14 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
device=audio_hidden_states.device, device=audio_hidden_states.device,
) )
if self.quantize_video_rope_coords_to_hidden_dtype:
video_coords = video_coords.to(
device=hidden_states.device, dtype=hidden_states.dtype
)
else:
video_coords = video_coords.to(device=hidden_states.device)
audio_coords = audio_coords.to(device=audio_hidden_states.device)
video_rotary_emb = self.rope(video_coords, device=hidden_states.device) video_rotary_emb = self.rope(video_coords, device=hidden_states.device)
audio_rotary_emb = self.audio_rope( audio_rotary_emb = self.audio_rope(
audio_coords, device=audio_hidden_states.device audio_coords, device=audio_hidden_states.device
@@ -1367,12 +1519,25 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_embedded_timestep = audio_embedded_timestep.view( audio_embedded_timestep = audio_embedded_timestep.view(
batch_size, -1, audio_embedded_timestep.size(-1) batch_size, -1, audio_embedded_timestep.size(-1)
) )
temb_prompt = None
temb_audio_prompt = None
if self.prompt_adaln_single is not None:
prompt_timestep = self._collapse_prompt_timestep(timestep)
temb_prompt, _ = self.prompt_adaln_single(
prompt_timestep.flatten(), hidden_dtype=hidden_states.dtype
)
temb_prompt = temb_prompt.view(batch_size, -1, temb_prompt.size(-1))
if self.audio_prompt_adaln_single is not None:
audio_prompt_timestep = self._collapse_prompt_timestep(audio_timestep)
temb_audio_prompt, _ = self.audio_prompt_adaln_single(
audio_prompt_timestep.flatten(),
hidden_dtype=audio_hidden_states.dtype,
)
temb_audio_prompt = temb_audio_prompt.view(
batch_size, -1, temb_audio_prompt.size(-1)
)
# 3.2. Prepare global modality cross attention modulation parameters # 3.2. Prepare global modality cross attention modulation parameters
ts_ca_mult = (
self.av_ca_timestep_scale_multiplier / self.timestep_scale_multiplier
)
hidden_dtype = hidden_states.dtype hidden_dtype = hidden_states.dtype
temb_ca_scale_shift, _ = self.av_ca_video_scale_shift_adaln_single( temb_ca_scale_shift, _ = self.av_ca_video_scale_shift_adaln_single(
timestep.flatten(), hidden_dtype=hidden_dtype timestep.flatten(), hidden_dtype=hidden_dtype
@@ -1403,10 +1568,12 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
) )
# 4. Prepare prompt embeddings # 4. Prepare prompt embeddings
encoder_hidden_states = self.caption_projection(encoder_hidden_states) if self.caption_projection is not None:
audio_encoder_hidden_states = self.audio_caption_projection( encoder_hidden_states = self.caption_projection(encoder_hidden_states)
audio_encoder_hidden_states if self.audio_caption_projection is not None:
) audio_encoder_hidden_states = self.audio_caption_projection(
audio_encoder_hidden_states
)
# 5. Run blocks # 5. Run blocks
skip_video_self_attn_blocks = set(skip_video_self_attn_blocks or ()) skip_video_self_attn_blocks = set(skip_video_self_attn_blocks or ())
skip_audio_self_attn_blocks = set(skip_audio_self_attn_blocks or ()) skip_audio_self_attn_blocks = set(skip_audio_self_attn_blocks or ())
@@ -1421,6 +1588,8 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
# under ForwardPattern.Pattern_0. # under ForwardPattern.Pattern_0.
temb=temb, temb=temb,
temb_audio=temb_audio, temb_audio=temb_audio,
temb_prompt=temb_prompt,
temb_audio_prompt=temb_audio_prompt,
temb_ca_scale_shift=temb_ca_scale_shift, temb_ca_scale_shift=temb_ca_scale_shift,
temb_ca_audio_scale_shift=temb_ca_audio_scale_shift, temb_ca_audio_scale_shift=temb_ca_audio_scale_shift,
temb_ca_gate=temb_ca_gate, temb_ca_gate=temb_ca_gate,
@@ -0,0 +1,204 @@
from typing import Any
import torch
import torch.nn as nn
from sglang.multimodal_gen.runtime.models.vaes.ltx_2_vae import (
LTX2VideoCausalConv3d,
LTX2VideoResnetBlock3d,
LTXVideoDownsampler3d,
)
def _patchify_video(sample: torch.Tensor, patch_size: int) -> torch.Tensor:
if patch_size == 1:
return sample
batch_size, channels, num_frames, height, width = sample.shape
sample = sample.reshape(
batch_size,
channels,
num_frames,
1,
height // patch_size,
patch_size,
width // patch_size,
patch_size,
)
return sample.permute(0, 1, 3, 7, 5, 2, 4, 6).flatten(1, 4)
class LTX23VideoPixelNorm(nn.Module):
def __init__(self, dim: int = 1, eps: float = 1e-8) -> None:
super().__init__()
self.dim = dim
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
mean_sq = torch.mean(x**2, dim=self.dim, keepdim=True)
rms = torch.sqrt(mean_sq + self.eps)
return x / rms
class LTX23PerChannelStatistics(nn.Module):
def __init__(self, latent_channels: int) -> None:
super().__init__()
self.register_buffer("std-of-means", torch.empty(latent_channels))
self.register_buffer("mean-of-means", torch.empty(latent_channels))
def normalize(self, x: torch.Tensor) -> torch.Tensor:
mean = self.get_buffer("mean-of-means").view(1, -1, 1, 1, 1).to(x)
std = self.get_buffer("std-of-means").view(1, -1, 1, 1, 1).to(x)
return (x - mean) / std
class LTX23VideoResBlockStack(nn.Module):
def __init__(
self, channels: int, num_layers: int, spatial_padding_mode: str
) -> None:
super().__init__()
self.res_blocks = nn.ModuleList(
[
LTX2VideoResnetBlock3d(
in_channels=channels,
out_channels=channels,
spatial_padding_mode=spatial_padding_mode,
)
for _ in range(num_layers)
]
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for res_block in self.res_blocks:
hidden_states = res_block(hidden_states, causal=True)
return hidden_states
def _make_ltx23_encoder_block(
block_name: str,
block_config: dict[str, Any],
in_channels: int,
spatial_padding_mode: str,
) -> tuple[nn.Module, int]:
if block_name == "res_x":
return (
LTX23VideoResBlockStack(
channels=in_channels,
num_layers=int(block_config["num_layers"]),
spatial_padding_mode=spatial_padding_mode,
),
in_channels,
)
multiplier = int(block_config.get("multiplier", 2))
stride_map = {
"compress_space_res": (1, 2, 2),
"compress_time_res": (2, 1, 1),
"compress_all_res": (2, 2, 2),
}
stride = stride_map.get(block_name)
if stride is None:
raise ValueError(f"Unsupported LTX-2.3 encoder block: {block_name}")
out_channels = in_channels * multiplier
return (
LTXVideoDownsampler3d(
in_channels=in_channels,
out_channels=out_channels,
stride=stride,
spatial_padding_mode=spatial_padding_mode,
),
out_channels,
)
class LTX23VideoConditionEncoder(nn.Module):
def __init__(self, config: dict[str, Any]) -> None:
super().__init__()
vae_config = config.get("vae", config)
latent_channels = int(vae_config["latent_channels"])
patch_size = int(vae_config.get("patch_size", 4))
spatial_padding_mode = str(vae_config.get("spatial_padding_mode", "zeros"))
encoder_blocks = list(vae_config["encoder_blocks"])
latent_log_var = str(vae_config.get("latent_log_var", "uniform"))
self.patch_size = patch_size
self.latency_channels = latent_channels
self.latent_log_var = latent_log_var
self.per_channel_statistics = LTX23PerChannelStatistics(latent_channels)
feature_channels = latent_channels
self.conv_in = LTX2VideoCausalConv3d(
in_channels=int(vae_config.get("in_channels", 3)) * patch_size**2,
out_channels=feature_channels,
kernel_size=3,
stride=1,
spatial_padding_mode=spatial_padding_mode,
)
self.down_blocks = nn.ModuleList()
for block_name, block_params in encoder_blocks:
block_config = (
{"num_layers": block_params}
if isinstance(block_params, int)
else dict(block_params)
)
block, feature_channels = _make_ltx23_encoder_block(
block_name=block_name,
block_config=block_config,
in_channels=feature_channels,
spatial_padding_mode=spatial_padding_mode,
)
self.down_blocks.append(block)
self.conv_norm_out = LTX23VideoPixelNorm(dim=1, eps=1e-8)
self.conv_act = nn.SiLU()
conv_out_channels = latent_channels
if latent_log_var == "per_channel":
conv_out_channels *= 2
elif latent_log_var in {"uniform", "constant"}:
conv_out_channels += 1
elif latent_log_var != "none":
raise ValueError(f"Unsupported latent_log_var: {latent_log_var}")
self.conv_out = LTX2VideoCausalConv3d(
in_channels=feature_channels,
out_channels=conv_out_channels,
kernel_size=3,
stride=1,
spatial_padding_mode=spatial_padding_mode,
)
def forward(self, sample: torch.Tensor) -> torch.Tensor:
frames_count = int(sample.shape[2])
if (frames_count - 1) % 8 != 0:
frames_to_crop = (frames_count - 1) % 8
sample = sample[:, :, :-frames_to_crop, ...]
hidden_states = _patchify_video(sample, self.patch_size)
hidden_states = self.conv_in(hidden_states, causal=True)
for block in self.down_blocks:
hidden_states = block(hidden_states)
hidden_states = self.conv_norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states, causal=True)
if self.latent_log_var == "uniform":
means = hidden_states[:, :-1, ...]
logvar = hidden_states[:, -1:, ...]
hidden_states = torch.cat(
[
means,
logvar.repeat(1, means.shape[1], *([1] * (means.ndim - 2))),
],
dim=1,
)
elif self.latent_log_var == "constant":
means = hidden_states[:, :-1, ...]
logvar = torch.full_like(means, -30.0)
hidden_states = torch.cat([means, logvar], dim=1)
means, _ = torch.chunk(hidden_states, 2, dim=1)
return self.per_channel_statistics.normalize(means)
@@ -394,6 +394,18 @@ class LTXVideoUpsampler3d(nn.Module):
return hidden_states return hidden_states
class LTX23PerChannelStatistics(nn.Module):
def __init__(self, latent_channels: int) -> None:
super().__init__()
self.register_buffer("mean_of_means", torch.empty(latent_channels))
self.register_buffer("std_of_means", torch.empty(latent_channels))
def un_normalize(self, x: torch.Tensor) -> torch.Tensor:
mean = self.mean_of_means.view(1, -1, 1, 1, 1).to(x)
std = self.std_of_means.view(1, -1, 1, 1, 1).to(x)
return x * std + mean
# Like LTX 1.0 LTXVideo095DownBlock3D, but with the updated LTX2VideoResnetBlock3d # Like LTX 1.0 LTXVideo095DownBlock3D, but with the updated LTX2VideoResnetBlock3d
class LTX2VideoDownBlock3D(nn.Module): class LTX2VideoDownBlock3D(nn.Module):
r""" r"""
@@ -609,6 +621,64 @@ class LTX2VideoMidBlock3d(nn.Module):
return hidden_states return hidden_states
class LTX23VideoMidBlock3d(nn.Module):
def __init__(
self,
in_channels: int,
num_layers: int = 1,
dropout: float = 0.0,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "swish",
inject_noise: bool = False,
timestep_conditioning: bool = False,
spatial_padding_mode: str = "zeros",
) -> None:
super().__init__()
self.time_embedder = None
if timestep_conditioning:
self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(
in_channels * 4, 0
)
self.res_blocks = nn.ModuleList(
[
LTX2VideoResnetBlock3d(
in_channels=in_channels,
out_channels=in_channels,
dropout=dropout,
eps=resnet_eps,
non_linearity=resnet_act_fn,
inject_noise=inject_noise,
timestep_conditioning=timestep_conditioning,
spatial_padding_mode=spatial_padding_mode,
)
for _ in range(num_layers)
]
)
def forward(
self,
hidden_states: torch.Tensor,
temb: Optional[torch.Tensor] = None,
causal: bool = True,
) -> torch.Tensor:
if self.time_embedder is not None:
temb = self.time_embedder(
timestep=temb.flatten(),
resolution=None,
aspect_ratio=None,
batch_size=hidden_states.size(0),
hidden_dtype=hidden_states.dtype,
)
temb = temb.view(hidden_states.size(0), -1, 1, 1, 1)
for res_block in self.res_blocks:
hidden_states = res_block(hidden_states, temb, causal=causal)
return hidden_states
# Like LTXVideoUpBlock3d but with no conv_in and the updated LTX2VideoResnetBlock3d # Like LTXVideoUpBlock3d but with no conv_in and the updated LTX2VideoResnetBlock3d
class LTX2VideoUpBlock3d(nn.Module): class LTX2VideoUpBlock3d(nn.Module):
r""" r"""
@@ -1104,6 +1174,192 @@ class LTX2VideoDecoder3d(nn.Module):
return hidden_states return hidden_states
def _make_ltx23_decoder_block(
block_name: str,
block_config: dict,
in_channels: int,
resnet_norm_eps: float,
timestep_conditioning: bool,
spatial_padding_mode: str,
) -> tuple[nn.Module, int]:
out_channels = in_channels
if block_name == "res_x":
block = LTX23VideoMidBlock3d(
in_channels=in_channels,
num_layers=int(block_config["num_layers"]),
resnet_eps=resnet_norm_eps,
timestep_conditioning=timestep_conditioning,
spatial_padding_mode=spatial_padding_mode,
)
elif block_name == "res_x_y":
out_channels = in_channels // int(block_config.get("multiplier", 2))
block = LTX2VideoResnetBlock3d(
in_channels=in_channels,
out_channels=out_channels,
eps=resnet_norm_eps,
timestep_conditioning=False,
spatial_padding_mode=spatial_padding_mode,
)
elif block_name == "compress_time":
out_channels = in_channels // int(block_config.get("multiplier", 1))
block = LTXVideoUpsampler3d(
in_channels=in_channels,
stride=(2, 1, 1),
residual=False,
upscale_factor=int(block_config.get("multiplier", 1)),
spatial_padding_mode=spatial_padding_mode,
)
elif block_name == "compress_space":
out_channels = in_channels // int(block_config.get("multiplier", 1))
block = LTXVideoUpsampler3d(
in_channels=in_channels,
stride=(1, 2, 2),
residual=False,
upscale_factor=int(block_config.get("multiplier", 1)),
spatial_padding_mode=spatial_padding_mode,
)
elif block_name == "compress_all":
out_channels = in_channels // int(block_config.get("multiplier", 1))
block = LTXVideoUpsampler3d(
in_channels=in_channels,
stride=(2, 2, 2),
residual=bool(block_config.get("residual", False)),
upscale_factor=int(block_config.get("multiplier", 1)),
spatial_padding_mode=spatial_padding_mode,
)
else:
raise ValueError(f"Unsupported LTX-2.3 decoder block: {block_name}")
return block, out_channels
class LTX23VideoDecoder3d(nn.Module):
def __init__(
self,
in_channels: int = 128,
out_channels: int = 3,
decoder_blocks: tuple[tuple[str, dict], ...] = (),
patch_size: int = 4,
patch_size_t: int = 1,
resnet_norm_eps: float = 1e-6,
is_causal: bool = False,
timestep_conditioning: bool = False,
base_channels: int = 128,
spatial_padding_mode: str = "zeros",
) -> None:
super().__init__()
self.patch_size = patch_size
self.patch_size_t = patch_size_t
self.out_channels = out_channels * patch_size**2
self.is_causal = is_causal
self.per_channel_statistics = LTX23PerChannelStatistics(in_channels)
feature_channels = base_channels * 8
self.conv_in = LTX2VideoCausalConv3d(
in_channels=in_channels,
out_channels=feature_channels,
kernel_size=3,
stride=1,
spatial_padding_mode=spatial_padding_mode,
)
self.up_blocks = nn.ModuleList([])
for block_name, block_params in reversed(tuple(decoder_blocks)):
block_config = (
{"num_layers": block_params}
if isinstance(block_params, int)
else dict(block_params)
)
block, feature_channels = _make_ltx23_decoder_block(
block_name=block_name,
block_config=block_config,
in_channels=feature_channels,
resnet_norm_eps=resnet_norm_eps,
timestep_conditioning=timestep_conditioning,
spatial_padding_mode=spatial_padding_mode,
)
self.up_blocks.append(block)
self.norm_out = PerChannelRMSNorm()
self.conv_act = nn.SiLU()
self.conv_out = LTX2VideoCausalConv3d(
in_channels=feature_channels,
out_channels=self.out_channels,
kernel_size=3,
stride=1,
spatial_padding_mode=spatial_padding_mode,
)
self.time_embedder = None
self.scale_shift_table = None
self.timestep_scale_multiplier = None
if timestep_conditioning:
self.timestep_scale_multiplier = nn.Parameter(
torch.tensor(1000.0, dtype=torch.float32)
)
self.time_embedder = PixArtAlphaCombinedTimestepSizeEmbeddings(
feature_channels * 2, 0
)
self.scale_shift_table = nn.Parameter(
torch.randn(2, feature_channels) / feature_channels**0.5
)
def forward(
self,
hidden_states: torch.Tensor,
temb: Optional[torch.Tensor] = None,
causal: Optional[bool] = None,
) -> torch.Tensor:
causal = self.is_causal if causal is None else causal
hidden_states = self.per_channel_statistics.un_normalize(hidden_states)
hidden_states = self.conv_in(hidden_states, causal=causal)
if self.timestep_scale_multiplier is not None and temb is not None:
temb = temb * self.timestep_scale_multiplier
for up_block in self.up_blocks:
if isinstance(up_block, LTX23VideoMidBlock3d):
hidden_states = up_block(hidden_states, temb, causal=causal)
elif isinstance(up_block, LTX2VideoResnetBlock3d):
hidden_states = up_block(hidden_states, None, causal=causal)
else:
hidden_states = up_block(hidden_states, causal=causal)
hidden_states = self.norm_out(hidden_states)
if self.time_embedder is not None and temb is not None:
temb = self.time_embedder(
timestep=temb.flatten(),
resolution=None,
aspect_ratio=None,
batch_size=hidden_states.size(0),
hidden_dtype=hidden_states.dtype,
)
temb = temb.view(hidden_states.size(0), -1, 1, 1, 1).unflatten(1, (2, -1))
temb = temb + self.scale_shift_table[None, ..., None, None, None]
shift, scale = temb.unbind(dim=1)
hidden_states = hidden_states * (1 + scale) + shift
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states, causal=causal)
p = self.patch_size
p_t = self.patch_size_t
batch_size, _, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.reshape(
batch_size, -1, p_t, p, p, num_frames, height, width
)
hidden_states = (
hidden_states.permute(0, 1, 5, 2, 6, 4, 7, 3)
.flatten(6, 7)
.flatten(4, 5)
.flatten(2, 3)
)
return hidden_states
class AutoencoderKLLTX2Video(ParallelTiledVAE): class AutoencoderKLLTX2Video(ParallelTiledVAE):
r""" r"""
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos.
@@ -1157,6 +1413,10 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
timestep_conditioning = getattr( timestep_conditioning = getattr(
config.arch_config, "timestep_conditioning", False config.arch_config, "timestep_conditioning", False
) )
use_ltx23_video_decoder = (
str(getattr(config.arch_config, "video_decoder_variant", "ltx_2"))
== "ltx_2_3"
)
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
@@ -1175,22 +1435,53 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
encoder_spatial_padding_mode, encoder_spatial_padding_mode,
) )
self.decoder = LTX2VideoDecoder3d( if use_ltx23_video_decoder:
in_channels=latent_channels, video_decoder_config = dict(config.arch_config.video_decoder_config)
out_channels=out_channels, if not video_decoder_config:
block_out_channels=decoder_block_out_channels, raise ValueError(
spatio_temporal_scaling=decoder_spatio_temporal_scaling, "LTX-2.3 native video decoder requires video_decoder_config."
layers_per_block=decoder_layers_per_block, )
patch_size=patch_size, self.decoder = LTX23VideoDecoder3d(
patch_size_t=patch_size_t, in_channels=latent_channels,
resnet_norm_eps=resnet_norm_eps, out_channels=out_channels,
is_causal=decoder_causal, decoder_blocks=tuple(video_decoder_config["decoder_blocks"]),
inject_noise=decoder_inject_noise, patch_size=int(video_decoder_config.get("patch_size", patch_size)),
timestep_conditioning=timestep_conditioning, patch_size_t=patch_size_t,
upsample_residual=upsample_residual, resnet_norm_eps=resnet_norm_eps,
upsample_factor=upsample_factor, is_causal=bool(
spatial_padding_mode=decoder_spatial_padding_mode, video_decoder_config.get("causal_decoder", decoder_causal)
) ),
timestep_conditioning=bool(
video_decoder_config.get(
"timestep_conditioning", timestep_conditioning
)
),
base_channels=int(
video_decoder_config.get("decoder_base_channels", 128)
),
spatial_padding_mode=str(
video_decoder_config.get(
"spatial_padding_mode", decoder_spatial_padding_mode
)
),
)
else:
self.decoder = LTX2VideoDecoder3d(
in_channels=latent_channels,
out_channels=out_channels,
block_out_channels=decoder_block_out_channels,
spatio_temporal_scaling=decoder_spatio_temporal_scaling,
layers_per_block=decoder_layers_per_block,
patch_size=patch_size,
patch_size_t=patch_size_t,
resnet_norm_eps=resnet_norm_eps,
is_causal=decoder_causal,
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)
latents_std = torch.ones((latent_channels,), requires_grad=False) latents_std = torch.ones((latent_channels,), requires_grad=False)
@@ -1,13 +1,237 @@
import math import math
from abc import ABC from abc import ABC
from contextlib import nullcontext
from typing import Tuple from typing import Tuple
import einops
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from sglang.multimodal_gen.configs.models.vocoder.ltx_vocoder import LTXVocoderConfig from sglang.multimodal_gen.configs.models.vocoder.ltx_vocoder import LTXVocoderConfig
LRELU_SLOPE = 0.1
def get_padding(kernel_size: int, dilation: int = 1) -> int:
return int((kernel_size * dilation - dilation) / 2)
def _sinc(x: torch.Tensor) -> torch.Tensor:
return torch.where(
x == 0,
torch.tensor(1.0, device=x.device, dtype=x.dtype),
torch.sin(math.pi * x) / math.pi / x,
)
def kaiser_sinc_filter1d(
cutoff: float, half_width: float, kernel_size: int
) -> torch.Tensor:
even = kernel_size % 2 == 0
half_size = kernel_size // 2
delta_f = 4 * half_width
amplitude = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
if amplitude > 50.0:
beta = 0.1102 * (amplitude - 8.7)
elif amplitude >= 21.0:
beta = 0.5842 * (amplitude - 21) ** 0.4 + 0.07886 * (amplitude - 21.0)
else:
beta = 0.0
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
time = (
torch.arange(-half_size, half_size) + 0.5
if even
else torch.arange(kernel_size) - half_size
)
if cutoff == 0:
filter_ = torch.zeros_like(time)
else:
filter_ = 2 * cutoff * window * _sinc(2 * cutoff * time)
filter_ /= filter_.sum()
return filter_.view(1, 1, kernel_size)
class LowPassFilter1d(nn.Module):
def __init__(
self,
cutoff: float = 0.5,
half_width: float = 0.6,
stride: int = 1,
padding: bool = True,
padding_mode: str = "replicate",
kernel_size: int = 12,
):
super().__init__()
self.kernel_size = kernel_size
self.even = kernel_size % 2 == 0
self.pad_left = kernel_size // 2 - int(self.even)
self.pad_right = kernel_size // 2
self.stride = stride
self.padding = padding
self.padding_mode = padding_mode
self.register_buffer(
"filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size)
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
_, channels, _ = x.shape
if self.padding:
x = F.pad(x, (self.pad_left, self.pad_right), mode=self.padding_mode)
return F.conv1d(
x,
self.filter.expand(channels, -1, -1),
stride=self.stride,
groups=channels,
)
class UpSample1d(nn.Module):
def __init__(
self,
ratio: int = 2,
kernel_size: int | None = None,
persistent: bool = True,
window_type: str = "kaiser",
):
super().__init__()
self.ratio = ratio
self.stride = ratio
if window_type == "hann":
rolloff = 0.99
lowpass_filter_width = 6
width = math.ceil(lowpass_filter_width / rolloff)
self.kernel_size = 2 * width * ratio + 1
self.pad = width
self.pad_left = 2 * width * ratio
self.pad_right = self.kernel_size - ratio
time_axis = (torch.arange(self.kernel_size) / ratio - width) * rolloff
time_clamped = time_axis.clamp(-lowpass_filter_width, lowpass_filter_width)
window = torch.cos(time_clamped * math.pi / lowpass_filter_width / 2) ** 2
sinc_filter = (torch.sinc(time_axis) * window * rolloff / ratio).view(
1, 1, -1
)
else:
self.kernel_size = (
int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
)
self.pad = self.kernel_size // ratio - 1
self.pad_left = (
self.pad * self.stride + (self.kernel_size - self.stride) // 2
)
self.pad_right = (
self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
)
sinc_filter = kaiser_sinc_filter1d(
cutoff=0.5 / ratio,
half_width=0.6 / ratio,
kernel_size=self.kernel_size,
)
self.register_buffer("filter", sinc_filter, persistent=persistent)
def forward(self, x: torch.Tensor) -> torch.Tensor:
_, channels, _ = x.shape
x = F.pad(x, (self.pad, self.pad), mode="replicate")
filt = self.filter.to(dtype=x.dtype, device=x.device).expand(channels, -1, -1)
x = self.ratio * F.conv_transpose1d(
x, filt, stride=self.stride, groups=channels
)
return x[..., self.pad_left : -self.pad_right]
class DownSample1d(nn.Module):
def __init__(self, ratio: int = 2, kernel_size: int | None = None):
super().__init__()
self.lowpass = LowPassFilter1d(
cutoff=0.5 / ratio,
half_width=0.6 / ratio,
stride=ratio,
kernel_size=int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.lowpass(x)
class Activation1d(nn.Module):
def __init__(
self,
activation: nn.Module,
up_ratio: int = 2,
down_ratio: int = 2,
up_kernel_size: int = 12,
down_kernel_size: int = 12,
):
super().__init__()
self.act = activation
self.upsample = UpSample1d(up_ratio, up_kernel_size)
self.downsample = DownSample1d(down_ratio, down_kernel_size)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.upsample(x)
x = self.act(x)
return self.downsample(x)
class Snake(nn.Module):
def __init__(
self,
in_features: int,
alpha: float = 1.0,
alpha_trainable: bool = True,
alpha_logscale: bool = True,
):
super().__init__()
self.alpha_logscale = alpha_logscale
self.alpha = nn.Parameter(
torch.zeros(in_features)
if alpha_logscale
else torch.ones(in_features) * alpha
)
self.alpha.requires_grad = alpha_trainable
self.eps = 1e-9
def forward(self, x: torch.Tensor) -> torch.Tensor:
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
return x + (1.0 / (alpha + self.eps)) * torch.sin(x * alpha).pow(2)
class SnakeBeta(nn.Module):
def __init__(
self,
in_features: int,
alpha: float = 1.0,
alpha_trainable: bool = True,
alpha_logscale: bool = True,
):
super().__init__()
self.alpha_logscale = alpha_logscale
self.alpha = nn.Parameter(
torch.zeros(in_features)
if alpha_logscale
else torch.ones(in_features) * alpha
)
self.alpha.requires_grad = alpha_trainable
self.beta = nn.Parameter(
torch.zeros(in_features)
if alpha_logscale
else torch.ones(in_features) * alpha
)
self.beta.requires_grad = alpha_trainable
self.eps = 1e-9
def forward(self, x: torch.Tensor) -> torch.Tensor:
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
beta = self.beta.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
beta = torch.exp(beta)
return x + (1.0 / (beta + self.eps)) * torch.sin(x * alpha).pow(2)
class ResBlock(nn.Module): class ResBlock(nn.Module):
def __init__( def __init__(
@@ -61,6 +285,252 @@ class ResBlock(nn.Module):
return x return x
class AMPBlock1(nn.Module):
def __init__(
self,
channels: int,
kernel_size: int = 3,
dilation: tuple[int, int, int] = (1, 3, 5),
activation: str = "snake",
):
super().__init__()
act_cls = SnakeBeta if activation == "snakebeta" else Snake
self.convs1 = nn.ModuleList(
[
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]),
),
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1]),
),
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[2],
padding=get_padding(kernel_size, dilation[2]),
),
]
)
self.convs2 = nn.ModuleList(
[
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1),
),
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1),
),
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1),
),
]
)
self.acts1 = nn.ModuleList(
[Activation1d(act_cls(channels)) for _ in range(len(self.convs1))]
)
self.acts2 = nn.ModuleList(
[Activation1d(act_cls(channels)) for _ in range(len(self.convs2))]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for conv1, conv2, act1, act2 in zip(
self.convs1, self.convs2, self.acts1, self.acts2
):
xt = act1(x)
xt = conv1(xt)
xt = act2(xt)
xt = conv2(xt)
x = x + xt
return x
class LTX23MelSTFT(nn.Module):
class STFTFn(nn.Module):
def __init__(self, filter_length: int, hop_length: int, win_length: int):
super().__init__()
self.hop_length = hop_length
self.win_length = win_length
n_freqs = filter_length // 2 + 1
self.register_buffer(
"forward_basis", torch.zeros(n_freqs * 2, 1, filter_length)
)
self.register_buffer(
"inverse_basis", torch.zeros(n_freqs * 2, 1, filter_length)
)
def forward(self, y: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if y.dim() == 2:
y = y.unsqueeze(1)
left_pad = max(0, self.win_length - self.hop_length)
y = F.pad(y, (left_pad, 0))
spec = F.conv1d(y, self.forward_basis, stride=self.hop_length, padding=0)
n_freqs = spec.shape[1] // 2
real, imag = spec[:, :n_freqs], spec[:, n_freqs:]
magnitude = torch.sqrt(real**2 + imag**2)
phase = torch.atan2(imag.float(), real.float()).to(real.dtype)
return magnitude, phase
def __init__(
self, filter_length: int, hop_length: int, win_length: int, n_mel_channels: int
):
super().__init__()
self.stft_fn = self.STFTFn(filter_length, hop_length, win_length)
n_freqs = filter_length // 2 + 1
self.register_buffer("mel_basis", torch.zeros(n_mel_channels, n_freqs))
def mel_spectrogram(
self, y: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
magnitude, phase = self.stft_fn(y)
energy = torch.norm(magnitude, dim=1)
mel = torch.matmul(self.mel_basis.to(magnitude.dtype), magnitude)
log_mel = torch.log(torch.clamp(mel, min=1e-5))
return log_mel, magnitude, phase, energy
class LTX23VocoderCore(nn.Module):
def __init__( # noqa: PLR0913
self,
resblock_kernel_sizes: list[int] | None = None,
upsample_rates: list[int] | None = None,
upsample_kernel_sizes: list[int] | None = None,
resblock_dilation_sizes: list[list[int]] | None = None,
upsample_initial_channel: int = 1024,
resblock: str = "1",
output_sampling_rate: int = 24000,
activation: str = "snake",
use_tanh_at_final: bool = True,
apply_final_activation: bool = True,
use_bias_at_final: bool = True,
):
super().__init__()
if resblock_kernel_sizes is None:
resblock_kernel_sizes = [3, 7, 11]
if upsample_rates is None:
upsample_rates = [6, 5, 2, 2, 2]
if upsample_kernel_sizes is None:
upsample_kernel_sizes = [16, 15, 8, 4, 4]
if resblock_dilation_sizes is None:
resblock_dilation_sizes = [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
self.output_sampling_rate = output_sampling_rate
self.num_kernels = len(resblock_kernel_sizes)
self.num_upsamples = len(upsample_rates)
self.use_tanh_at_final = use_tanh_at_final
self.apply_final_activation = apply_final_activation
self.is_amp = resblock == "AMP1"
self.conv_pre = nn.Conv1d(
in_channels=128,
out_channels=upsample_initial_channel,
kernel_size=7,
stride=1,
padding=3,
)
self.ups = nn.ModuleList(
nn.ConvTranspose1d(
upsample_initial_channel // (2**i),
upsample_initial_channel // (2 ** (i + 1)),
kernel_size,
stride,
padding=(kernel_size - stride) // 2,
)
for i, (stride, kernel_size) in enumerate(
zip(upsample_rates, upsample_kernel_sizes, strict=True)
)
)
final_channels = upsample_initial_channel // (2 ** len(upsample_rates))
self.resblocks = nn.ModuleList()
for i in range(len(upsample_rates)):
channels = upsample_initial_channel // (2 ** (i + 1))
for kernel_size, dilations in zip(
resblock_kernel_sizes, resblock_dilation_sizes, strict=True
):
if self.is_amp:
self.resblocks.append(
AMPBlock1(
channels,
kernel_size,
tuple(dilations),
activation=activation,
)
)
else:
self.resblocks.append(
ResBlock(
channels,
kernel_size=kernel_size,
dilations=tuple(dilations),
leaky_relu_negative_slope=LRELU_SLOPE,
padding_mode=get_padding(kernel_size, 1),
)
)
self.act_post = (
Activation1d(SnakeBeta(final_channels)) if self.is_amp else nn.LeakyReLU()
)
self.conv_post = nn.Conv1d(
in_channels=final_channels,
out_channels=2,
kernel_size=7,
stride=1,
padding=3,
bias=use_bias_at_final,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x.transpose(2, 3)
if x.dim() == 4:
assert x.shape[1] == 2, "Input must have 2 channels for stereo"
x = einops.rearrange(x, "b s c t -> b (s c) t")
x = self.conv_pre(x)
for i in range(self.num_upsamples):
if not self.is_amp:
x = F.leaky_relu(x, LRELU_SLOPE)
x = self.ups[i](x)
start = i * self.num_kernels
end = start + self.num_kernels
block_outputs = torch.stack(
[self.resblocks[idx](x) for idx in range(start, end)],
dim=0,
)
x = block_outputs.mean(dim=0)
x = self.act_post(x)
x = self.conv_post(x)
if self.apply_final_activation:
x = torch.tanh(x) if self.use_tanh_at_final else torch.clamp(x, -1, 1)
return x
class LTX2Vocoder(ABC, nn.Module): class LTX2Vocoder(ABC, nn.Module):
r""" r"""
LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms. LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms.
@@ -72,10 +542,61 @@ class LTX2Vocoder(ABC, nn.Module):
): ):
super().__init__() super().__init__()
self.config = config self.config = config
nested_vocoder_cfg = getattr(config.arch_config, "vocoder", None)
if isinstance(nested_vocoder_cfg, dict) and "bwe" in nested_vocoder_cfg:
vocoder_cfg = nested_vocoder_cfg.get("vocoder", {})
bwe_cfg = nested_vocoder_cfg["bwe"]
self.vocoder = LTX23VocoderCore(
resblock_kernel_sizes=vocoder_cfg.get("resblock_kernel_sizes"),
upsample_rates=vocoder_cfg.get("upsample_rates"),
upsample_kernel_sizes=vocoder_cfg.get("upsample_kernel_sizes"),
resblock_dilation_sizes=vocoder_cfg.get("resblock_dilation_sizes"),
upsample_initial_channel=vocoder_cfg.get(
"upsample_initial_channel", 1024
),
resblock=vocoder_cfg.get("resblock", "1"),
output_sampling_rate=bwe_cfg["input_sampling_rate"],
activation=vocoder_cfg.get("activation", "snake"),
use_tanh_at_final=vocoder_cfg.get("use_tanh_at_final", True),
apply_final_activation=vocoder_cfg.get("apply_final_activation", True),
use_bias_at_final=vocoder_cfg.get("use_bias_at_final", True),
)
self.bwe_generator = LTX23VocoderCore(
resblock_kernel_sizes=bwe_cfg.get("resblock_kernel_sizes"),
upsample_rates=bwe_cfg.get("upsample_rates"),
upsample_kernel_sizes=bwe_cfg.get("upsample_kernel_sizes"),
resblock_dilation_sizes=bwe_cfg.get("resblock_dilation_sizes"),
upsample_initial_channel=bwe_cfg.get("upsample_initial_channel", 1024),
resblock=bwe_cfg.get("resblock", "1"),
output_sampling_rate=bwe_cfg["output_sampling_rate"],
activation=bwe_cfg.get("activation", "snake"),
use_tanh_at_final=bwe_cfg.get("use_tanh_at_final", True),
apply_final_activation=bwe_cfg.get("apply_final_activation", True),
use_bias_at_final=bwe_cfg.get("use_bias_at_final", True),
)
self.mel_stft = LTX23MelSTFT(
filter_length=bwe_cfg["n_fft"],
hop_length=bwe_cfg["hop_length"],
win_length=bwe_cfg.get("win_size", bwe_cfg["n_fft"]),
n_mel_channels=bwe_cfg["num_mels"],
)
self.input_sampling_rate = bwe_cfg["input_sampling_rate"]
self.output_sampling_rate = bwe_cfg["output_sampling_rate"]
self.hop_length = bwe_cfg["hop_length"]
with torch.device("cpu"):
self.resampler = UpSample1d(
ratio=self.output_sampling_rate // self.input_sampling_rate,
persistent=False,
window_type="hann",
)
self.sample_rate = self.output_sampling_rate
return
self.sample_rate = ( self.sample_rate = (
getattr(config.arch_config, "sample_rate", None) getattr(config.arch_config, "sample_rate", None)
or getattr(config.arch_config, "sampling_rate", None) or getattr(config.arch_config, "sampling_rate", None)
or getattr(config.arch_config, "audio_sample_rate", None) or getattr(config.arch_config, "audio_sample_rate", None)
or getattr(config.arch_config, "output_sampling_rate", None)
) )
in_channels = config.arch_config.in_channels in_channels = config.arch_config.in_channels
@@ -139,6 +660,12 @@ class LTX2Vocoder(ABC, nn.Module):
self.conv_out = nn.Conv1d(output_channels, out_channels, 7, stride=1, padding=3) self.conv_out = nn.Conv1d(output_channels, out_channels, 7, stride=1, padding=3)
def _compute_ltx23_mel(self, audio: torch.Tensor) -> torch.Tensor:
batch, channels, _ = audio.shape
flat = audio.reshape(batch * channels, -1)
mel, _, _, _ = self.mel_stft.mel_spectrogram(flat)
return mel.reshape(batch, channels, mel.shape[1], mel.shape[2])
def forward( def forward(
self, hidden_states: torch.Tensor, time_last: bool = False self, hidden_states: torch.Tensor, time_last: bool = False
) -> torch.Tensor: ) -> torch.Tensor:
@@ -157,6 +684,32 @@ class LTX2Vocoder(ABC, nn.Module):
`torch.Tensor`: `torch.Tensor`:
Audio waveform tensor of shape (batch_size, out_channels, audio_length) Audio waveform tensor of shape (batch_size, out_channels, audio_length)
""" """
if hasattr(self, "bwe_generator"):
input_dtype = hidden_states.dtype
autocast_ctx = (
torch.autocast(
device_type=hidden_states.device.type, dtype=torch.float32
)
if hidden_states.device.type != "cpu"
else nullcontext()
)
with autocast_ctx:
waveform = self.vocoder(hidden_states.float())
length_low_rate = waveform.shape[-1]
output_length = (
length_low_rate
* self.output_sampling_rate
// self.input_sampling_rate
)
remainder = length_low_rate % self.hop_length
if remainder != 0:
waveform = F.pad(waveform, (0, self.hop_length - remainder))
mel = self._compute_ltx23_mel(waveform)
residual = self.bwe_generator(mel.transpose(2, 3))
skip = self.resampler(waveform)
assert residual.shape == skip.shape
waveform = torch.clamp(residual + skip, -1, 1)[..., :output_length]
return waveform.to(input_dtype)
# Ensure that the time/frame dimension is last # Ensure that the time/frame dimension is last
if not time_last: if not time_last:
@@ -5,6 +5,9 @@ import numpy as np
import torch import torch
from diffusers import FlowMatchEulerDiscreteScheduler from diffusers import FlowMatchEulerDiscreteScheduler
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader, PipelineComponentLoader,
) )
@@ -44,6 +47,8 @@ def _resolve_ltx2_two_stage_component_paths(
if "spatial_upsampler" not in resolved: if "spatial_upsampler" not in resolved:
spatial_candidates = [ spatial_candidates = [
os.path.join(model_path, "latent_upsampler"), os.path.join(model_path, "latent_upsampler"),
os.path.join(model_path, "ltx-2.3-spatial-upscaler-x2-1.1.safetensors"),
os.path.join(model_path, "ltx-2.3-spatial-upscaler-x2-1.0.safetensors"),
os.path.join(model_path, "ltx-2-spatial-upscaler-x2-1.0.safetensors"), os.path.join(model_path, "ltx-2-spatial-upscaler-x2-1.0.safetensors"),
] ]
for candidate in spatial_candidates: for candidate in spatial_candidates:
@@ -53,12 +58,15 @@ def _resolve_ltx2_two_stage_component_paths(
break break
if "distilled_lora" not in resolved: if "distilled_lora" not in resolved:
distilled_lora = os.path.join( distilled_lora_candidates = [
model_path, "ltx-2-19b-distilled-lora-384.safetensors" os.path.join(model_path, "ltx-2.3-22b-distilled-lora-384.safetensors"),
) os.path.join(model_path, "ltx-2-19b-distilled-lora-384.safetensors"),
if os.path.exists(distilled_lora): ]
resolved["distilled_lora"] = distilled_lora for distilled_lora in distilled_lora_candidates:
auto_resolved.append(f"distilled_lora={distilled_lora}") if os.path.exists(distilled_lora):
resolved["distilled_lora"] = distilled_lora
auto_resolved.append(f"distilled_lora={distilled_lora}")
break
if auto_resolved: if auto_resolved:
logger.info( logger.info(
@@ -81,6 +89,8 @@ def calculate_ltx2_shift(
def prepare_ltx2_mu(batch: Req, server_args: ServerArgs): def prepare_ltx2_mu(batch: Req, server_args: ServerArgs):
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
return "mu", None
latent_num_frames = (int(batch.num_frames) - 1) // int( latent_num_frames = (int(batch.num_frames) - 1) // int(
server_args.pipeline_config.vae_temporal_compression server_args.pipeline_config.vae_temporal_compression
) + 1 ) + 1
@@ -92,16 +102,49 @@ def prepare_ltx2_mu(batch: Req, server_args: ServerArgs):
return "mu", calculate_ltx2_shift(video_sequence_length) return "mu", calculate_ltx2_shift(video_sequence_length)
def build_official_ltx2_sigmas(
steps: int,
*,
max_shift: float = 2.05,
base_shift: float = 0.95,
stretch: bool = True,
terminal: float = 0.1,
default_number_of_tokens: int = MAX_SHIFT_ANCHOR,
) -> list[float]:
sigmas = torch.linspace(1.0, 0.0, steps + 1, dtype=torch.float32)
mm = (max_shift - base_shift) / (MAX_SHIFT_ANCHOR - BASE_SHIFT_ANCHOR)
b = base_shift - mm * BASE_SHIFT_ANCHOR
sigma_shift = float(default_number_of_tokens) * mm + b
non_zero_mask = sigmas != 0
shifted = torch.where(
non_zero_mask,
math.exp(sigma_shift) / (math.exp(sigma_shift) + (1.0 / sigmas - 1.0)),
torch.zeros_like(sigmas),
)
if stretch:
one_minus_z = 1.0 - shifted[non_zero_mask]
scale_factor = one_minus_z[-1] / (1.0 - terminal)
shifted[non_zero_mask] = 1.0 - (one_minus_z / scale_factor)
return shifted[:-1].tolist()
class LTX2SigmaPreparationStage(PipelineStage): class LTX2SigmaPreparationStage(PipelineStage):
"""Prepare native LTX-2 sigma schedule before timestep setup.""" """Prepare native LTX-2 sigma schedule before timestep setup."""
def forward(self, batch: Req, server_args: ServerArgs) -> Req: def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.extra["ltx2_phase"] = "stage1" batch.extra["ltx2_phase"] = "stage1"
batch.sigmas = np.linspace( if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
1.0, batch.sigmas = build_official_ltx2_sigmas(int(batch.num_inference_steps))
1.0 / int(batch.num_inference_steps), else:
int(batch.num_inference_steps), batch.sigmas = np.linspace(
).tolist() 1.0,
1.0 / int(batch.num_inference_steps),
int(batch.num_inference_steps),
).tolist()
return batch return batch
@@ -66,6 +66,7 @@ def build_pipeline(
) )
else: else:
logger.info("No pipeline_class_name specified, using model_index.json") logger.info("No pipeline_class_name specified, using model_index.json")
model_info = get_model_info( model_info = get_model_info(
model_path, model_path,
backend=server_args.backend, backend=server_args.backend,
@@ -135,7 +135,7 @@ class Req:
trajectory_latents: torch.Tensor | None = None trajectory_latents: torch.Tensor | None = None
trajectory_audio_latents: torch.Tensor | None = None trajectory_audio_latents: torch.Tensor | None = None
# Extra parameters that might be needed by specific pipeline implementations # Extra parameters that might be needed by specific pipeline implementations (e.g., LTX2.3 DenoisingAVStage)
extra: dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
is_warmup: bool = False is_warmup: bool = False
@@ -25,6 +25,11 @@ class LTX2AVDecodingStage(DecodingStage):
self.video_processor = VideoProcessor(vae_scale_factor=32) self.video_processor = VideoProcessor(vae_scale_factor=32)
@staticmethod
def _ltx2_should_externally_denorm_video_latents(server_args: ServerArgs) -> bool:
arch_config = server_args.pipeline_config.vae_config.arch_config
return str(getattr(arch_config, "video_decoder_variant", "ltx_2")) != "ltx_2_3"
def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch: def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch:
self.load_model() self.load_model()
@@ -40,9 +45,10 @@ class LTX2AVDecodingStage(DecodingStage):
original_dtype = vae_dtype original_dtype = vae_dtype
self.vae.to(torch.bfloat16) self.vae.to(torch.bfloat16)
latents = latents.to(torch.bfloat16) latents = latents.to(torch.bfloat16)
std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latents) if self._ltx2_should_externally_denorm_video_latents(server_args):
mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latents) std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latents)
latents = latents * std + mean mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latents)
latents = latents * std + mean
latents = server_args.pipeline_config.preprocess_decoding( latents = server_args.pipeline_config.preprocess_decoding(
latents, server_args, vae=self.vae latents, server_args, vae=self.vae
) )
@@ -1,5 +1,7 @@
import copy import copy
import json
import math import math
import os
import time import time
from io import BytesIO from io import BytesIO
@@ -10,8 +12,15 @@ 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 diffusers.utils.torch_utils import randn_tensor
from safetensors.torch import load_file as safetensors_load_file
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
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.vaes.ltx_2_3_condition_encoder import (
LTX23VideoConditionEncoder,
)
from sglang.multimodal_gen.runtime.models.vision_utils import ( from sglang.multimodal_gen.runtime.models.vision_utils import (
load_image, load_image,
normalize, normalize,
@@ -46,6 +55,8 @@ class LTX2AVDenoisingStage(DenoisingStage):
transformer=transformer, scheduler=scheduler, vae=vae, **kwargs transformer=transformer, scheduler=scheduler, vae=vae, **kwargs
) )
self.audio_vae = audio_vae self.audio_vae = audio_vae
self._condition_image_encoder = None
self._condition_image_encoder_dir = None
@staticmethod @staticmethod
def _get_video_latent_num_frames_for_model( def _get_video_latent_num_frames_for_model(
@@ -116,13 +127,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
) -> dict[str, object] | None: ) -> dict[str, object] | None:
if stage != "stage1": if stage != "stage1":
return None 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") return batch.extra.get("ltx2_stage1_guider_params")
@staticmethod @staticmethod
@@ -141,6 +145,56 @@ class LTX2AVDenoisingStage(DenoisingStage):
factor = rescale_scale * factor + (1.0 - rescale_scale) factor = rescale_scale * factor + (1.0 - rescale_scale)
return pred * factor return pred * factor
@staticmethod
def _prepare_ltx2_ti2v_clean_state(
latents: torch.Tensor,
image_latent: torch.Tensor,
num_img_tokens: int,
zero_clean_latent: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
latents = latents.clone()
conditioned = image_latent[:, :num_img_tokens, :].to(
device=latents.device, dtype=latents.dtype
)
latents[:, :num_img_tokens, :] = conditioned
denoise_mask = torch.ones(
(latents.shape[0], latents.shape[1], 1),
device=latents.device,
dtype=torch.float32,
)
denoise_mask[:, :num_img_tokens, :] = 0.0
if zero_clean_latent:
clean_latent = torch.zeros_like(latents)
else:
clean_latent = latents.detach().clone()
clean_latent[:, :num_img_tokens, :] = conditioned
return latents, denoise_mask, clean_latent
@staticmethod
def _ltx2_velocity_to_x0(
sample: torch.Tensor,
velocity: torch.Tensor,
sigma: float | torch.Tensor,
) -> torch.Tensor:
if isinstance(sigma, torch.Tensor):
sigma = sigma.to(device=sample.device, dtype=torch.float32)
while sigma.ndim < sample.ndim:
sigma = sigma.unsqueeze(-1)
return (sample.float() - sigma * velocity.float()).to(sample.dtype)
return (sample.float() - float(sigma) * velocity.float()).to(sample.dtype)
@staticmethod
def _repeat_batch_dim(tensor: torch.Tensor, target_batch_size: int) -> torch.Tensor:
"""Repeat along batch dim while preserving any tokenwise timestep layout."""
if tensor.shape[0] == int(target_batch_size):
return tensor
if tensor.shape[0] <= 0 or int(target_batch_size) % int(tensor.shape[0]) != 0:
raise ValueError(
f"Cannot repeat tensor with batch={tensor.shape[0]} to target_batch_size={target_batch_size}"
)
repeat_factor = int(target_batch_size) // int(tensor.shape[0])
return tensor.repeat(repeat_factor, *([1] * (tensor.ndim - 1)))
@classmethod @classmethod
def _ltx2_calculate_guided_x0( def _ltx2_calculate_guided_x0(
cls, cls,
@@ -252,6 +306,40 @@ class LTX2AVDenoisingStage(DenoisingStage):
return True return True
return int(getattr(batch, "sp_video_start_frame", 0)) == 0 return int(getattr(batch, "sp_video_start_frame", 0)) == 0
def _get_condition_image_encoder(
self,
server_args: ServerArgs,
*,
device: torch.device,
dtype: torch.dtype,
) -> LTX23VideoConditionEncoder | None:
arch_config = server_args.pipeline_config.vae_config.arch_config
encoder_subdir = str(getattr(arch_config, "condition_encoder_subdir", ""))
if not encoder_subdir:
return None
vae_model_path = server_args.model_paths["vae"]
encoder_dir = os.path.join(vae_model_path, encoder_subdir)
config_path = os.path.join(encoder_dir, "config.json")
weights_path = os.path.join(encoder_dir, "model.safetensors")
if not os.path.exists(config_path) or not os.path.exists(weights_path):
raise ValueError(
f"LTX-2 condition encoder files not found under {encoder_dir}"
)
cached_dir = self._condition_image_encoder_dir
encoder = self._condition_image_encoder
if encoder is None or cached_dir != encoder_dir:
with open(config_path, encoding="utf-8") as f:
config = json.load(f)
encoder = LTX23VideoConditionEncoder(config)
encoder.load_state_dict(safetensors_load_file(weights_path), strict=True)
self._condition_image_encoder = encoder
self._condition_image_encoder_dir = encoder_dir
encoder = encoder.to(device=device, dtype=dtype)
return encoder
def _prepare_ltx2_image_latent(self, batch: Req, server_args: ServerArgs) -> None: def _prepare_ltx2_image_latent(self, batch: Req, server_args: ServerArgs) -> None:
"""Encode `batch.image_path` into packed token latents for LTX-2 TI2V.""" """Encode `batch.image_path` into packed token latents for LTX-2 TI2V."""
if ( if (
@@ -276,8 +364,11 @@ class LTX2AVDenoisingStage(DenoisingStage):
) )
img = load_image(image_path) img = load_image(image_path)
img_array = np.array(img).astype(np.uint8)[..., :3]
img_array = self._apply_video_codec_compression(img_array, crf=33)
conditioned_img = PIL.Image.fromarray(img_array)
batch.condition_image = self._resize_center_crop( batch.condition_image = self._resize_center_crop(
img, width=int(batch.width), height=int(batch.height) conditioned_img, width=int(batch.width), height=int(batch.height)
) )
latents_device = ( latents_device = (
@@ -287,17 +378,22 @@ class LTX2AVDenoisingStage(DenoisingStage):
) )
encode_dtype = batch.latents.dtype encode_dtype = batch.latents.dtype
original_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] original_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
self.vae = self.vae.to(device=latents_device, dtype=encode_dtype)
vae_autocast_enabled = ( vae_autocast_enabled = (
original_dtype != torch.float32 original_dtype != torch.float32
) and not server_args.disable_autocast ) and not server_args.disable_autocast
condition_image_encoder = self._get_condition_image_encoder(
server_args, device=latents_device, dtype=encode_dtype
)
if condition_image_encoder is None:
self.vae = self.vae.to(device=latents_device, dtype=encode_dtype)
video_condition = self._resize_center_crop_tensor( video_condition = self._resize_center_crop_tensor(
img, conditioned_img,
width=int(batch.width), width=int(batch.width),
height=int(batch.height), height=int(batch.height),
device=latents_device, device=latents_device,
dtype=encode_dtype, dtype=encode_dtype,
apply_codec_compression=False,
) )
with torch.autocast( with torch.autocast(
@@ -306,31 +402,42 @@ class LTX2AVDenoisingStage(DenoisingStage):
enabled=vae_autocast_enabled, enabled=vae_autocast_enabled,
): ):
try: try:
if server_args.pipeline_config.vae_tiling: if (
condition_image_encoder is None
and server_args.pipeline_config.vae_tiling
):
self.vae.enable_tiling() self.vae.enable_tiling()
except Exception: except Exception:
pass pass
if not vae_autocast_enabled: if not vae_autocast_enabled:
video_condition = video_condition.to(encode_dtype) video_condition = video_condition.to(encode_dtype)
latent_dist: DiagonalGaussianDistribution = self.vae.encode(video_condition) if condition_image_encoder is not None:
if isinstance(latent_dist, AutoencoderKLOutput): latent = condition_image_encoder(video_condition)
latent_dist = latent_dist.latent_dist else:
latent_dist: DiagonalGaussianDistribution = self.vae.encode(
video_condition
)
if isinstance(latent_dist, AutoencoderKLOutput):
latent_dist = latent_dist.latent_dist
mode = server_args.pipeline_config.vae_config.encode_sample_mode() if condition_image_encoder is None:
if mode == "argmax": mode = server_args.pipeline_config.vae_config.encode_sample_mode()
latent = latent_dist.mode() if mode == "argmax":
elif mode == "sample": latent = latent_dist.mode()
if batch.generator is None: elif mode == "sample":
raise ValueError("Generator must be provided for VAE sampling.") if batch.generator is None:
latent = latent_dist.sample(batch.generator) raise ValueError("Generator must be provided for VAE sampling.")
latent = latent_dist.sample(batch.generator)
else:
raise ValueError(f"Unsupported encode_sample_mode: {mode}")
# Per-channel normalization: normalized = (x - mean) / std
mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latent)
std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latent)
latent = (latent - mean) / std
else: else:
raise ValueError(f"Unsupported encode_sample_mode: {mode}") latent = latent.to(dtype=encode_dtype)
# Per-channel normalization: normalized = (x - mean) / std
mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latent)
std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latent)
latent = (latent - mean) / std
packed = server_args.pipeline_config.maybe_pack_latents( packed = server_args.pipeline_config.maybe_pack_latents(
latent, latent.shape[0], batch latent, latent.shape[0], batch
@@ -362,9 +469,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
batch.height, batch.height,
) )
self.vae.to(original_dtype) if condition_image_encoder is None:
self.vae.to(original_dtype)
if server_args.vae_cpu_offload: if server_args.vae_cpu_offload:
self.vae = self.vae.to("cpu") self.vae = self.vae.to("cpu")
if condition_image_encoder is not None:
self._condition_image_encoder = condition_image_encoder.to("cpu")
@torch.no_grad() @torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> Req: def forward(self, batch: Req, server_args: ServerArgs) -> Req:
@@ -437,19 +547,15 @@ class LTX2AVDenoisingStage(DenoisingStage):
if do_ti2v: if do_ti2v:
if not (isinstance(latents, torch.Tensor) and latents.ndim == 3): if not (isinstance(latents, torch.Tensor) and latents.ndim == 3):
raise ValueError("LTX-2 TI2V expects packed token latents [B, S, D].") raise ValueError("LTX-2 TI2V expects packed token latents [B, S, D].")
latents[:, :num_img_tokens, :] = batch.image_latent[ use_zero_clean_latent = is_ltx23_native_variant(
:, :num_img_tokens, : server_args.pipeline_config.vae_config.arch_config
].to(device=latents.device, dtype=latents.dtype) )
denoise_mask = torch.ones( latents, denoise_mask, clean_latent = self._prepare_ltx2_ti2v_clean_state(
(latents.shape[0], latents.shape[1], 1), latents=latents,
device=latents.device, image_latent=batch.image_latent,
dtype=torch.float32, num_img_tokens=num_img_tokens,
zero_clean_latent=use_zero_clean_latent,
) )
denoise_mask[:, :num_img_tokens, :] = 0.0
clean_latent = latents.detach().clone()
clean_latent[:, :num_img_tokens, :] = batch.image_latent[
:, :num_img_tokens, :
].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,
@@ -558,11 +664,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
], ],
dim=0, dim=0,
) )
timestep_video = timestep_video.expand( cfg_batch_size = int(latent_model_input.shape[0])
int(latent_model_input.shape[0]) timestep_video = self._repeat_batch_dim(
timestep_video, cfg_batch_size
) )
timestep_audio = timestep_audio.expand( timestep_audio = self._repeat_batch_dim(
int(latent_model_input.shape[0]) timestep_audio, cfg_batch_size
) )
with set_forward_context( with set_forward_context(
@@ -616,11 +723,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
audio_latents = audio_scheduler.step( audio_latents = audio_scheduler.step(
a_v_pos, t_device, audio_latents, return_dict=False a_v_pos, t_device, audio_latents, return_dict=False
)[0] )[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( latents = self.post_forward_for_ti2v_task(
batch, server_args, reserved_frames_mask, latents, z batch, server_args, reserved_frames_mask, latents, z
) )
@@ -714,14 +816,20 @@ class LTX2AVDenoisingStage(DenoisingStage):
if a_v_neg is not None: if a_v_neg is not None:
a_v_neg = a_v_neg.float() a_v_neg = a_v_neg.float()
# Velocity -> denoised (x0): x0 = x - sigma * v
sigma_val = float(sigma.item()) sigma_val = float(sigma.item())
denoised_video = (latents.float() - sigma_val * v_pos).to( video_sigma_for_x0: float | torch.Tensor = sigma_val
latents.dtype if do_ti2v and denoise_mask is not None:
video_sigma_for_x0 = sigma.to(
device=latents.device, dtype=torch.float32
) * denoise_mask.squeeze(-1)
denoised_video = self._ltx2_velocity_to_x0(
latents, v_pos, video_sigma_for_x0
) )
denoised_audio = ( denoised_audio = self._ltx2_velocity_to_x0(
audio_latents.float() - sigma_val * a_v_pos audio_latents, a_v_pos, sigma_val
).to(audio_latents.dtype) )
denoised_video_cond = denoised_video
denoised_audio_cond = denoised_audio
denoised_video_neg = None denoised_video_neg = None
denoised_audio_neg = None denoised_audio_neg = None
denoised_video_perturbed = None denoised_video_perturbed = None
@@ -737,12 +845,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
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
): ):
denoised_video_neg = ( denoised_video_neg = self._ltx2_velocity_to_x0(
latents.float() - sigma_val * v_neg latents, v_neg, video_sigma_for_x0
).to(latents.dtype) )
denoised_audio_neg = ( denoised_audio_neg = self._ltx2_velocity_to_x0(
audio_latents.float() - sigma_val * a_v_neg audio_latents, a_v_neg, sigma_val
).to(audio_latents.dtype) )
if stage1_guider_params is not None: if stage1_guider_params is not None:
video_skip = self._ltx2_should_skip_step( video_skip = self._ltx2_should_skip_step(
i, int(stage1_guider_params["video_skip_step"]) i, int(stage1_guider_params["video_skip_step"])
@@ -784,12 +892,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
stage1_guider_params["audio_stg_blocks"] stage1_guider_params["audio_stg_blocks"]
), ),
) )
denoised_video_perturbed = ( denoised_video_perturbed = self._ltx2_velocity_to_x0(
latents.float() - sigma_val * v_ptb.float() latents, v_ptb.float(), video_sigma_for_x0
).to(latents.dtype) )
denoised_audio_perturbed = ( denoised_audio_perturbed = self._ltx2_velocity_to_x0(
audio_latents.float() - sigma_val * a_v_ptb.float() audio_latents, a_v_ptb.float(), sigma_val
).to(audio_latents.dtype) )
need_modality = ( need_modality = (
float(stage1_guider_params["video_modality_scale"]) float(stage1_guider_params["video_modality_scale"])
@@ -822,12 +930,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
disable_a2v_cross_attn=True, disable_a2v_cross_attn=True,
disable_v2a_cross_attn=True, disable_v2a_cross_attn=True,
) )
denoised_video_modality = ( denoised_video_modality = self._ltx2_velocity_to_x0(
latents.float() - sigma_val * v_mod.float() latents, v_mod.float(), video_sigma_for_x0
).to(latents.dtype) )
denoised_audio_modality = ( denoised_audio_modality = self._ltx2_velocity_to_x0(
audio_latents.float() - sigma_val * a_v_mod.float() audio_latents, a_v_mod.float(), sigma_val
).to(audio_latents.dtype) )
if not video_skip: if not video_skip:
denoised_video = self._ltx2_calculate_guided_x0( denoised_video = self._ltx2_calculate_guided_x0(
@@ -934,11 +1042,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
audio_latents.float() + v_audio.float() * dt audio_latents.float() + v_audio.float() * dt
).to(dtype=audio_latents.dtype) ).to(dtype=audio_latents.dtype)
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( latents = self.post_forward_for_ti2v_task(
batch, server_args, reserved_frames_mask, latents, z batch, server_args, reserved_frames_mask, latents, z
) )
@@ -1,6 +1,9 @@
import torch import torch
from diffusers.utils.torch_utils import randn_tensor from diffusers.utils.torch_utils import randn_tensor
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
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.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.latent_preparation import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import (
@@ -60,10 +63,112 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
): ):
return torch.float32 return torch.float32
@staticmethod
def _packed_video_latent_shape(
latent_shape: tuple[int, int, int, int, int],
pipeline_config,
) -> tuple[int, int, int]:
batch_size, channels, num_frames, height, width = latent_shape
patch_size_t = int(pipeline_config.patch_size_t)
patch_size = int(pipeline_config.patch_size)
return (
batch_size,
(num_frames // patch_size_t)
* (height // patch_size)
* (width // patch_size),
channels * patch_size_t * patch_size * patch_size,
)
@staticmethod
def _packed_audio_latent_shape(
latent_shape: tuple[int, int, int, int],
) -> tuple[int, int, int]:
batch_size, channels, latent_length, mel_bins = latent_shape
return (batch_size, latent_length, channels * mel_bins)
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 if not is_ltx23_native_variant(
# This sets batch.latents and batch.raw_latent_shape server_args.pipeline_config.vae_config.arch_config
batch = super().forward(batch, server_args) ):
batch = super().forward(batch, server_args)
try:
generate_audio = batch.generate_audio
except AttributeError:
generate_audio = True
if not generate_audio:
batch.audio_latents = None
batch.raw_audio_latent_shape = None
return batch
device = get_local_torch_device()
dtype = self._get_latent_dtype(batch, server_args)
generator = batch.generator
audio_latents = batch.audio_latents
batch_size = batch.batch_size
num_frames = batch.num_frames
if audio_latents is None:
shape = server_args.pipeline_config.prepare_audio_latent_shape(
batch, batch_size, num_frames
)
audio_latents = randn_tensor(
shape, generator=generator, device=device, dtype=dtype
)
else:
audio_latents = audio_latents.to(device)
audio_latents = server_args.pipeline_config.maybe_pack_audio_latents(
audio_latents, batch_size, batch
)
batch.audio_latents = audio_latents
batch.raw_audio_latent_shape = audio_latents.shape
return batch
# 1. Prepare video latents directly in packed token space.
# Official LTX-2.3 pipelines sample noise after patchify; generating unpacked
# [B, C, F, H, W] noise and packing afterwards changes token ordering.
latent_num_frames = self.adjust_video_length(batch, server_args)
batch_size = batch.batch_size
dtype = self._get_latent_dtype(batch, server_args)
device = get_local_torch_device()
generator = batch.generator
latents = batch.latents
num_frames = (
latent_num_frames if latent_num_frames is not None else batch.num_frames
)
if latents is None:
latent_shape = server_args.pipeline_config.prepare_latent_shape(
batch, batch_size, num_frames
)
latents = randn_tensor(
self._packed_video_latent_shape(
latent_shape, server_args.pipeline_config
),
generator=generator,
device=device,
dtype=dtype,
)
latent_ids = server_args.pipeline_config.maybe_prepare_latent_ids(latents)
if latent_ids is not None:
batch.latent_ids = latent_ids.to(device=device)
else:
latents = latents.to(device)
latents = server_args.pipeline_config.maybe_pack_latents(
latents, batch_size, batch
)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
batch.latents = latents
batch.raw_latent_shape = latents.shape
# 2. Prepare Audio Latents (optional) # 2. Prepare Audio Latents (optional)
# Default to True if not specified # Default to True if not specified
@@ -76,28 +181,24 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
batch.raw_audio_latent_shape = None batch.raw_audio_latent_shape = None
return batch return batch
device = get_local_torch_device()
dtype = self._get_latent_dtype(batch, server_args)
generator = batch.generator
audio_latents = batch.audio_latents audio_latents = batch.audio_latents
batch_size = batch.batch_size
num_frames = batch.num_frames
if audio_latents is None: if audio_latents is None:
shape = server_args.pipeline_config.prepare_audio_latent_shape( latent_shape = server_args.pipeline_config.prepare_audio_latent_shape(
batch, batch_size, num_frames batch, batch_size, batch.num_frames
) )
audio_latents = randn_tensor( audio_latents = randn_tensor(
shape, generator=generator, device=device, dtype=dtype self._packed_audio_latent_shape(latent_shape),
generator=generator,
device=device,
dtype=dtype,
) )
else: else:
audio_latents = audio_latents.to(device) audio_latents = audio_latents.to(device)
audio_latents = server_args.pipeline_config.maybe_pack_audio_latents(
audio_latents = server_args.pipeline_config.maybe_pack_audio_latents( audio_latents, batch_size, batch
audio_latents, batch_size, batch )
)
# Store in batch # Store in batch
batch.audio_latents = audio_latents batch.audio_latents = audio_latents
@@ -24,8 +24,13 @@ from sglang.utils import load_diffusion_overlay_registry_from_env
logger = init_logger(__name__) logger = init_logger(__name__)
# Built-in diffusion model overlay registry. # Built-in diffusion model overlay registry.
# Keep this empty until concrete overlay repos are ready to ship. BUILTIN_MODEL_OVERLAY_REGISTRY: dict[str, dict[str, Any]] = {
BUILTIN_MODEL_OVERLAY_REGISTRY: dict[str, dict[str, Any]] = {} "Lightricks/LTX-2.3": {
"overlay_repo_id": "MickJ/LTX-2.3-overlay",
"overlay_revision": "main",
"bundled_overlay_subdir": "ltx_2_3",
},
}
MODEL_OVERLAY_METADATA_PATTERNS = [ MODEL_OVERLAY_METADATA_PATTERNS = [
@@ -42,6 +47,43 @@ MODEL_OVERLAY_METADATA_PATTERNS = [
_MODEL_OVERLAY_REGISTRY_CACHE: dict[str, dict[str, Any]] | None = None _MODEL_OVERLAY_REGISTRY_CACHE: dict[str, dict[str, Any]] | None = None
def _compute_overlay_fingerprint(overlay_dir: str) -> str:
hasher = hashlib.sha256()
for root, dir_names, file_names in os.walk(overlay_dir):
dir_names[:] = sorted(
d for d in dir_names if d != "__pycache__" and not d.endswith(".egg-info")
)
for file_name in sorted(file_names):
if file_name.endswith((".safetensors", ".bin", ".pth", ".pt")):
continue
file_path = os.path.join(root, file_name)
rel_path = os.path.relpath(file_path, overlay_dir).replace(os.sep, "/")
hasher.update(rel_path.encode("utf-8"))
with open(file_path, "rb") as f:
hasher.update(hashlib.sha256(f.read()).digest())
return hasher.hexdigest()
def _resolve_bundled_overlay_dir(overlay_spec: dict[str, Any]) -> str | None:
bundled_overlay_subdir = overlay_spec.get("bundled_overlay_subdir")
if not bundled_overlay_subdir:
return None
bundled_overlay_dir = os.path.abspath(
os.path.join(
os.path.dirname(__file__),
"..",
"..",
"model_overlays",
str(bundled_overlay_subdir),
)
)
if not os.path.isdir(bundled_overlay_dir):
return None
if load_overlay_manifest_if_present(bundled_overlay_dir) is None:
return None
return bundled_overlay_dir
def get_diffusion_cache_root() -> str: def get_diffusion_cache_root() -> str:
return os.path.expanduser( return os.path.expanduser(
os.getenv("SGLANG_DIFFUSION_CACHE_ROOT", "~/.cache/sgl_diffusion") os.getenv("SGLANG_DIFFUSION_CACHE_ROOT", "~/.cache/sgl_diffusion")
@@ -300,6 +342,15 @@ def download_overlay_metadata(
*, *,
snapshot_download_fn: Callable[..., str], snapshot_download_fn: Callable[..., str],
) -> str: ) -> str:
bundled_overlay_dir = _resolve_bundled_overlay_dir(overlay_spec)
if bundled_overlay_dir is not None:
logger.info(
"Using bundled overlay metadata for %s from %s",
source_model_id,
bundled_overlay_dir,
)
return bundled_overlay_dir
overlay_repo_id = str(overlay_spec["overlay_repo_id"]) overlay_repo_id = str(overlay_spec["overlay_repo_id"])
if os.path.exists(overlay_repo_id): if os.path.exists(overlay_repo_id):
logger.info( logger.info(
@@ -419,6 +470,7 @@ def materialize_overlay_model(
materializer_version = str(manifest.get("materializer_version", "v1")) materializer_version = str(manifest.get("materializer_version", "v1"))
overlay_repo_id = str(overlay_spec["overlay_repo_id"]) overlay_repo_id = str(overlay_spec["overlay_repo_id"])
overlay_revision = str(overlay_spec.get("overlay_revision", "main")) overlay_revision = str(overlay_spec.get("overlay_revision", "main"))
overlay_fingerprint = _compute_overlay_fingerprint(overlay_dir)
cache_key = hashlib.sha256( cache_key = hashlib.sha256(
json.dumps( json.dumps(
{ {
@@ -426,6 +478,7 @@ def materialize_overlay_model(
"overlay_repo_id": overlay_repo_id, "overlay_repo_id": overlay_repo_id,
"overlay_revision": overlay_revision, "overlay_revision": overlay_revision,
"materializer_version": materializer_version, "materializer_version": materializer_version,
"overlay_fingerprint": overlay_fingerprint,
}, },
sort_keys=True, sort_keys=True,
).encode("utf-8") ).encode("utf-8")
@@ -502,6 +555,7 @@ def materialize_overlay_model(
"overlay_repo_id": overlay_repo_id, "overlay_repo_id": overlay_repo_id,
"overlay_revision": overlay_revision, "overlay_revision": overlay_revision,
"materializer_version": materializer_version, "materializer_version": materializer_version,
"overlay_fingerprint": overlay_fingerprint,
}, },
f, f,
indent=2, indent=2,
@@ -2504,6 +2504,54 @@
"expected_e2e_ms": 8091.46, "expected_e2e_ms": 8091.46,
"expected_avg_denoise_ms": 141.37, "expected_avg_denoise_ms": 141.37,
"expected_median_denoise_ms": 142.63 "expected_median_denoise_ms": 142.63
},
"ltx_2.3_one_stage_ti2v": {
"stages_ms": {
"InputValidationStage": 3.27,
"TextEncodingStage": 1766.71,
"LTX2TextConnectorStage": 27.34,
"LTX2SigmaPreparationStage": 0.14,
"TimestepPreparationStage": 15.78,
"LTX2AVLatentPreparationStage": 0.25,
"LTX2AVDenoisingStage": 23757.65,
"LTX2AVDecodingStage": 1171.65,
"per_frame_generation": null
},
"denoise_step_ms": {
"0": 709.91,
"1": 704.93,
"2": 707.02,
"3": 702.54,
"4": 702.07,
"5": 746.58,
"6": 774.04,
"7": 766.37,
"8": 751.37,
"9": 734.67,
"10": 710.79,
"11": 689.84,
"12": 688.12,
"13": 689.74,
"14": 687.5,
"15": 699.65,
"16": 697.09,
"17": 686.87,
"18": 691.43,
"19": 712.52,
"20": 705.43,
"21": 727.21,
"22": 706.73,
"23": 705.63,
"24": 714.95,
"25": 707.69,
"26": 748.35,
"27": 738.02,
"28": 738.05,
"29": 726.96
},
"expected_e2e_ms": 26916.58,
"expected_avg_denoise_ms": 715.73,
"expected_median_denoise_ms": 707.35
} }
} }
} }
@@ -935,7 +935,6 @@ TWO_GPU_CASES_A = [
model_path="Lightricks/LTX-2", model_path="Lightricks/LTX-2",
modality="video", modality="video",
num_gpus=2, num_gpus=2,
dit_layerwise_offload=True,
extras=["--pipeline-class-name LTX2TwoStagePipeline"], extras=["--pipeline-class-name LTX2TwoStagePipeline"],
), ),
T2V_sampling_params, T2V_sampling_params,
@@ -1040,6 +1039,15 @@ TWO_GPU_CASES_B = [
), ),
TI2I_sampling_params, TI2I_sampling_params,
), ),
DiffusionTestCase(
"ltx_2.3_one_stage_ti2v",
DiffusionServerArgs(
model_path="Lightricks/LTX-2.3",
modality="video",
num_gpus=2,
),
TI2V_sampling_params,
),
] ]
if not current_platform.is_hip(): if not current_platform.is_hip():
@@ -0,0 +1,392 @@
import json
import os
import tempfile
from types import SimpleNamespace
import pytest
import torch
from safetensors import safe_open
from safetensors.torch import save_file
pytest.importorskip("triton.compiler")
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.model_overlays.ltx_2_3._overlay.materialize import (
_build_transformer_config,
_build_vae_config,
_rename_connector_key,
_repack_ltx23_image_encoder_weights,
_repack_ltx23_video_decoder_weights,
)
from sglang.multimodal_gen.registry import get_model_info
from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import (
_resolve_ltx2_two_stage_component_paths,
build_official_ltx2_sigmas,
prepare_ltx2_mu,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding_av import (
LTX2AVDecodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising_av import (
LTX2AVDenoisingStage,
)
from sglang.multimodal_gen.runtime.utils.model_overlay import (
resolve_model_overlay_target,
)
def _make_req(**sampling_kwargs) -> Req:
return Req(
sampling_params=SamplingParams(**sampling_kwargs),
prompt="prompt",
prompt_embeds=[torch.zeros(1, 1, 1)],
)
def test_ltx23_builtin_overlay_target_is_hf_repo():
target = resolve_model_overlay_target("Lightricks/LTX-2.3")
assert target is not None
source_model_id, overlay_spec = target
assert source_model_id == "Lightricks/LTX-2.3"
assert str(overlay_spec["overlay_repo_id"]) == "MickJ/LTX-2.3-overlay"
assert str(overlay_spec["overlay_revision"]) == "main"
assert str(overlay_spec["bundled_overlay_subdir"]) == "ltx_2_3"
def test_ltx23_model_info_resolves_to_native_pipeline_and_sampling_params():
model_info = get_model_info("Lightricks/LTX-2.3", backend="sglang")
assert model_info is not None
assert model_info.pipeline_cls.__name__ == "LTX2Pipeline"
assert model_info.sampling_param_cls.__name__ == "LTX23SamplingParams"
def test_ltx23_sampling_defaults_use_cuda_generator():
sampling_params = SamplingParams.from_pretrained(
"Lightricks/LTX-2.3",
backend="sglang",
)
assert sampling_params.generator_device == "cuda"
assert sampling_params.guidance_scale == 3.0
assert sampling_params.num_inference_steps == 30
def test_ltx2_sampling_defaults_keep_cpu_generator():
sampling_params = SamplingParams.from_pretrained(
"Lightricks/LTX-2",
backend="sglang",
)
assert sampling_params.generator_device == "cpu"
def test_ltx23_build_request_extra_sets_stage1_guider_defaults():
sampling_params = SamplingParams.from_pretrained(
"Lightricks/LTX-2.3",
backend="sglang",
)
assert sampling_params.build_request_extra()["ltx2_stage1_guider_params"] == {
"video_cfg_scale": 3.0,
"video_stg_scale": 1.0,
"video_rescale_scale": 0.7,
"video_modality_scale": 3.0,
"video_skip_step": 0,
"video_stg_blocks": [28],
"audio_cfg_scale": 7.0,
"audio_stg_scale": 1.0,
"audio_rescale_scale": 0.7,
"audio_modality_scale": 3.0,
"audio_skip_step": 0,
"audio_stg_blocks": [28],
}
def test_sampling_params_apply_request_extra_populates_req_extra():
sampling_params = SamplingParams.from_pretrained(
"Lightricks/LTX-2.3",
backend="sglang",
)
req = Req(sampling_params=sampling_params, prompt="prompt")
sampling_params.apply_request_extra(req)
assert req.extra["ltx2_stage1_guider_params"]["video_cfg_scale"] == 3.0
assert req.extra["ltx2_stage1_guider_params"]["audio_cfg_scale"] == 7.0
def test_ltx23_uses_official_sigma_schedule():
sigmas = build_official_ltx2_sigmas(30)
assert len(sigmas) == 30
assert sigmas[0] == pytest.approx(1.0)
assert sigmas[1] == pytest.approx(0.99495703, abs=1e-6)
assert sigmas[-1] == pytest.approx(0.1, abs=1e-6)
def test_ltx23_native_variant_uses_explicit_marker_only():
assert is_ltx23_native_variant(SimpleNamespace(ltx_variant="ltx_2_3")) is True
assert is_ltx23_native_variant(SimpleNamespace(ltx_variant="ltx_2")) is False
def test_prepare_ltx2_mu_respects_variant_marker():
ltx23_server_args = SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_config=SimpleNamespace(
arch_config=SimpleNamespace(ltx_variant="ltx_2_3")
)
)
)
legacy_server_args = SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_config=SimpleNamespace(
arch_config=SimpleNamespace(ltx_variant="ltx_2")
),
vae_temporal_compression=8,
vae_scale_factor=32,
)
)
assert prepare_ltx2_mu(
_make_req(num_frames=121, height=512, width=768),
ltx23_server_args,
) == ("mu", None)
key, mu = prepare_ltx2_mu(
_make_req(num_frames=121, height=512, width=768),
legacy_server_args,
)
assert key == "mu"
assert isinstance(mu, float)
assert mu > 0.0
def test_ltx23_ti2v_clean_latent_uses_zero_background():
latents = torch.arange(24, dtype=torch.float32).view(1, 6, 4)
image_latent = torch.full((1, 2, 4), 99.0)
conditioned, denoise_mask, clean_latent = (
LTX2AVDenoisingStage._prepare_ltx2_ti2v_clean_state(
latents=latents,
image_latent=image_latent,
num_img_tokens=2,
zero_clean_latent=True,
)
)
assert torch.equal(conditioned[:, :2], image_latent)
assert torch.equal(clean_latent[:, :2], image_latent)
assert torch.equal(clean_latent[:, 2:], torch.zeros_like(clean_latent[:, 2:]))
assert torch.equal(denoise_mask[:, :2], torch.zeros_like(denoise_mask[:, :2]))
assert torch.equal(denoise_mask[:, 2:], torch.ones_like(denoise_mask[:, 2:]))
def test_ltx2_ti2v_clean_latent_keeps_legacy_background_when_requested():
latents = torch.arange(24, dtype=torch.float32).view(1, 6, 4)
image_latent = torch.full((1, 2, 4), 99.0)
conditioned, _, clean_latent = LTX2AVDenoisingStage._prepare_ltx2_ti2v_clean_state(
latents=latents,
image_latent=image_latent,
num_img_tokens=2,
zero_clean_latent=False,
)
assert torch.equal(conditioned[:, :2], image_latent)
assert torch.equal(clean_latent[:, :2], image_latent)
assert torch.equal(clean_latent[:, 2:], latents[:, 2:])
def test_ltx23_velocity_to_x0_supports_tokenwise_sigma():
sample = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]], dtype=torch.float32)
velocity = torch.tensor([[[0.5, 0.5], [1.0, 1.0]]], dtype=torch.float32)
sigma = torch.tensor([[0.0, 0.5]], dtype=torch.float32)
denoised = LTX2AVDenoisingStage._ltx2_velocity_to_x0(sample, velocity, sigma)
expected = torch.tensor([[[1.0, 2.0], [2.5, 3.5]]], dtype=torch.float32)
assert torch.allclose(denoised, expected)
def test_ltx23_connector_repack_renames_qk_norm_keys():
assert (
_rename_connector_key(
"model.diffusion_model.video_embeddings_connector.transformer_1d_blocks.0.attn1.q_norm.weight"
)
== "video_connector.transformer_blocks.0.attn1.norm_q.weight"
)
assert (
_rename_connector_key(
"model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.1.attn1.k_norm.weight"
)
== "audio_connector.transformer_blocks.1.attn1.norm_k.weight"
)
def test_ltx23_transformer_config_forces_sdpa_for_v2a_cross_attention():
with tempfile.TemporaryDirectory() as tmpdir:
donor_dir = os.path.join(tmpdir, "donor")
os.makedirs(os.path.join(donor_dir, "transformer"), exist_ok=True)
with open(os.path.join(donor_dir, "transformer", "config.json"), "w") as f:
json.dump({"_class_name": "OldClass", "num_layers": 1}, f)
config = _build_transformer_config(donor_dir)
assert config["_class_name"] == "LTX2VideoTransformer3DModel"
assert config["force_sdpa_v2a_cross_attention"] is True
def test_ltx23_vae_config_adds_required_markers():
with tempfile.TemporaryDirectory() as tmpdir:
auxiliary_dir = os.path.join(tmpdir, "aux")
config_donor_dir = os.path.join(tmpdir, "donor")
os.makedirs(os.path.join(auxiliary_dir, "vae"), exist_ok=True)
os.makedirs(os.path.join(config_donor_dir, "vae"), exist_ok=True)
with open(os.path.join(auxiliary_dir, "vae", "config.json"), "w") as f:
json.dump(
{
"_class_name": "AutoencoderKLLTX2Video",
"scaling_factor": 1.0,
"patch_size": 4,
"decoder_causal": False,
"timestep_conditioning": False,
"encoder_spatial_padding_mode": "zeros",
"decoder_spatial_padding_mode": "reflect",
},
f,
)
with open(os.path.join(config_donor_dir, "vae", "config.json"), "w") as f:
json.dump(
{
"vae": {
"decoder_blocks": [["res_x", {"num_layers": 2}]],
"decoder_base_channels": 128,
"patch_size": 4,
"spatial_padding_mode": "zeros",
}
},
f,
)
config = _build_vae_config(auxiliary_dir, config_donor_dir)
assert config["ltx_variant"] == "ltx_2_3"
assert config["condition_encoder_subdir"] == "ltx23_image_encoder"
assert config["video_decoder_variant"] == "ltx_2_3"
assert config["video_decoder_config"]["decoder_base_channels"] == 128
def test_ltx23_repack_image_encoder_keeps_only_encoder_tensors():
with tempfile.TemporaryDirectory() as tmpdir:
source_path = os.path.join(tmpdir, "source.safetensors")
output_path = os.path.join(tmpdir, "output.safetensors")
save_file(
{
"encoder.conv_in.conv.weight": torch.ones(1),
"decoder.conv_in.conv.weight": torch.full((1,), 2.0),
"per_channel_statistics.mean-of-means": torch.full((2,), 3.0),
},
source_path,
)
_repack_ltx23_image_encoder_weights(source_path, output_path)
with safe_open(output_path, framework="pt") as f:
assert sorted(f.keys()) == [
"conv_in.conv.weight",
"per_channel_statistics.mean-of-means",
]
def test_ltx23_repack_video_decoder_keeps_decoder_and_stats():
with tempfile.TemporaryDirectory() as tmpdir:
auxiliary_path = os.path.join(tmpdir, "aux.safetensors")
donor_path = os.path.join(tmpdir, "donor.safetensors")
output_path = os.path.join(tmpdir, "output.safetensors")
save_file(
{
"encoder.conv_in.conv.weight": torch.full((1,), 5.0),
},
auxiliary_path,
)
save_file(
{
"decoder.conv_in.conv.weight": torch.ones(1),
"per_channel_statistics.mean-of-means": torch.full((2,), 3.0),
"per_channel_statistics.std-of-means": torch.full((2,), 4.0),
},
donor_path,
)
_repack_ltx23_video_decoder_weights(auxiliary_path, donor_path, output_path)
with safe_open(output_path, framework="pt") as f:
assert sorted(f.keys()) == [
"decoder.conv_in.conv.weight",
"decoder.per_channel_statistics.mean_of_means",
"decoder.per_channel_statistics.std_of_means",
"encoder.conv_in.conv.weight",
"latents_mean",
"latents_std",
]
def test_ltx23_decode_skips_external_denorm():
ltx23_server_args = SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_config=SimpleNamespace(
arch_config=SimpleNamespace(video_decoder_variant="ltx_2_3")
)
)
)
legacy_server_args = SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_config=SimpleNamespace(
arch_config=SimpleNamespace(video_decoder_variant="ltx_2")
)
)
)
assert (
LTX2AVDecodingStage._ltx2_should_externally_denorm_video_latents(
ltx23_server_args
)
is False
)
assert (
LTX2AVDecodingStage._ltx2_should_externally_denorm_video_latents(
legacy_server_args
)
is True
)
def test_ltx2_two_stage_component_auto_resolution_preserves_legacy_candidates(tmp_path):
legacy_spatial = tmp_path / "ltx-2-spatial-upscaler-x2-1.0.safetensors"
legacy_lora = tmp_path / "ltx-2-19b-distilled-lora-384.safetensors"
legacy_spatial.touch()
legacy_lora.touch()
resolved = _resolve_ltx2_two_stage_component_paths(str(tmp_path), {})
assert resolved["spatial_upsampler"] == str(legacy_spatial)
assert resolved["distilled_lora"] == str(legacy_lora)
def test_ltx23_two_stage_component_auto_resolution_prefers_23_assets(tmp_path):
spatial = tmp_path / "ltx-2.3-spatial-upscaler-x2-1.1.safetensors"
lora = tmp_path / "ltx-2.3-22b-distilled-lora-384.safetensors"
spatial.touch()
lora.touch()
resolved = _resolve_ltx2_two_stage_component_paths(str(tmp_path), {})
assert resolved["spatial_upsampler"] == str(spatial)
assert resolved["distilled_lora"] == str(lora)