[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 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-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**:
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`
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
-5
View File
@@ -76,11 +76,6 @@ sglang generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--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
Apply LoRA adapters via `--lora-path`:
@@ -12,13 +12,17 @@ class LTX2ConnectorArchConfig(AdapterArchConfig):
audio_connector_num_attention_heads: int = 30
audio_connector_num_layers: int = 2
audio_connector_num_learnable_registers: int = 128
audio_feature_extractor_out_features: int = 0
caption_channels: int = 3840
causal_temporal_positioning: bool = False
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_theta: float = 10000.0
rope_type: str = "split"
text_proj_in_factor: int = 49
video_feature_extractor_out_features: int = 0
video_connector_attention_head_dim: int = 128
video_connector_num_attention_heads: int = 30
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.
#
# HF key -> SGLang key (upstream naming)
r"^model\.diffusion_model\.(.*)$": r"\1",
r"^proj_in\.(.*)$": r"patchify_proj.\1",
r"^time_embed\.(.*)$": r"adaln_single.\1",
r"^audio_proj_in\.(.*)$": r"audio_patchify_proj.\1",
@@ -123,6 +124,10 @@ class LTX2ArchConfig(DiTArchConfig):
attention_type: LTX2AttentionFunction = LTX2AttentionFunction.DEFAULT
rope_type: LTX2RopeType = LTX2RopeType.INTERLEAVED
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
num_attention_heads: int = 32
@@ -147,6 +152,14 @@ class LTX2ArchConfig(DiTArchConfig):
audio_positional_embedding_max_pos: list[int] | None = None
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
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
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
@@ -52,6 +52,12 @@ class LTXVideoVAEArchConfig(VAEArchConfig):
decoder_causal: bool = False
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
class LTXVideoVAEConfig(VAEConfig):
@@ -93,20 +93,48 @@ def pack_text_embeds(
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(
outputs: BaseEncoderOutput,
text_inputs: dict,
pipeline_config: Optional["LTX2PipelineConfig"] = None,
) -> torch.Tensor:
_ = pipeline_config
# LTX-2 requires all hidden states concatenated for the connector
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)
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)
# Assuming left padding for Gemma as per Diffusers
return pack_text_embeds(hidden_states, sequence_lengths, padding_side="left")
else:
raise AttributeError(
@@ -1,4 +1,6 @@
import dataclasses
from dataclasses import field
from typing import Any
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, "
"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:
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:
"""Convert output_quality string to compression level."""
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"),
)
+25 -5
View File
@@ -90,7 +90,10 @@ from sglang.multimodal_gen.configs.sample.hunyuan import (
HunyuanSamplingParams,
)
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 (
MOVA_360P_SamplingParams,
MOVA_720P_SamplingParams,
@@ -155,7 +158,18 @@ def _discover_and_register_pipelines():
package.__path__, package.__name__ + "."
):
if not ispkg:
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"):
entry_cls = pipeline_module.EntryClass
entry_cls_list = (
@@ -594,12 +608,18 @@ def _register_configs():
register_configs(
sampling_param_cls=LTX2SamplingParams,
pipeline_config_cls=LTX2PipelineConfig,
hf_model_paths=[
"Lightricks/LTX-2",
],
hf_model_paths=["Lightricks/LTX-2"],
model_detectors=[
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()
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(
"Generator was garbage collected without being shut down. "
"Attempting to shut down the local server and client."
)
self.shutdown()
elif self.local_scheduler_process:
elif local_scheduler_process:
logger.warning(
"Generator was garbage collected without being shut down. "
"Attempting to shut down the local server."
@@ -288,12 +288,7 @@ def prepare_request(
sampling_params=sampling_params,
VSA_sparsity=server_args.attention_backend_config.VSA_sparsity,
)
try:
diffusers_kwargs = sampling_params.diffusers_kwargs
except AttributeError:
diffusers_kwargs = None
if diffusers_kwargs:
req.extra["diffusers_kwargs"] = diffusers_kwargs
sampling_params.apply_request_extra(req)
req.adjust_size(server_args)
@@ -1,5 +1,8 @@
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.loader.component_loaders.component_loader import (
ComponentLoader,
@@ -50,10 +53,9 @@ class AdapterLoader(ComponentLoader):
target_device = get_local_torch_device()
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():
connector_cfg = SimpleNamespace(**config)
connector_cfg = LTX2ConnectorConfig()
connector_cfg.update_model_arch(config)
model = model_cls(connector_cfg).to(
device=target_device, dtype=default_dtype
)
@@ -1,3 +1,4 @@
import math
from typing import Optional, Tuple, Union
import torch
@@ -89,6 +90,7 @@ class LTX2Attention(torch.nn.Module):
norm_eps: float = 1e-6,
norm_elementwise_affine: bool = True,
rope_type: str = "interleaved",
apply_gated_attention: bool = False,
processor=None,
):
super().__init__()
@@ -125,6 +127,9 @@ class LTX2Attention(torch.nn.Module):
self.to_v = torch.nn.Linear(
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.append(torch.nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
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,
key_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
gate_input = hidden_states
if encoder_hidden_states is None:
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.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[1](hidden_states)
return hidden_states
@@ -317,6 +332,7 @@ class LTX2TransformerBlock1d(nn.Module):
activation_fn: str = "gelu-approximate",
eps: float = 1e-6,
rope_type: str = "interleaved",
apply_gated_attention: bool = False,
):
super().__init__()
@@ -327,6 +343,7 @@ class LTX2TransformerBlock1d(nn.Module):
kv_heads=num_attention_heads,
dim_head=attention_head_dim,
rope_type=rope_type,
apply_gated_attention=apply_gated_attention,
)
self.norm2 = torch.nn.RMSNorm(dim, eps=eps, elementwise_affine=False)
@@ -373,6 +390,7 @@ class LTX2ConnectorTransformer1d(nn.Module):
eps: float = 1e-6,
causal_temporal_positioning: bool = False,
rope_type: str = "interleaved",
apply_gated_attention: bool = False,
):
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -403,6 +421,7 @@ class LTX2ConnectorTransformer1d(nn.Module):
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
rope_type=rope_type,
apply_gated_attention=apply_gated_attention,
)
for _ in range(num_layers)
]
@@ -516,7 +535,34 @@ class LTX2TextConnectors(nn.Module):
rope_double_precision = config.rope_double_precision
causal_temporal_positioning = config.causal_temporal_positioning
rope_type = config.rope_type
connector_apply_gated_attention = config.connector_apply_gated_attention
feature_extractor_in_features = config.feature_extractor_in_features
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
)
@@ -530,6 +576,7 @@ class LTX2TextConnectors(nn.Module):
rope_double_precision=rope_double_precision,
causal_temporal_positioning=causal_temporal_positioning,
rope_type=rope_type,
apply_gated_attention=connector_apply_gated_attention,
)
self.audio_connector = LTX2ConnectorTransformer1d(
num_attention_heads=audio_connector_num_attention_heads,
@@ -541,8 +588,15 @@ class LTX2TextConnectors(nn.Module):
rope_double_precision=rope_double_precision,
causal_temporal_positioning=causal_temporal_positioning,
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(
self,
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
# 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)
seq_len = text_encoder_hidden_states.shape[1]
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
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(
text_encoder_hidden_states, attention_mask
video_hidden_states, attention_mask
)
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)
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
@@ -37,6 +37,15 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
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(
x: torch.Tensor, freqs: Tuple[torch.Tensor, torch.Tensor]
@@ -447,6 +456,7 @@ class LTX2Attention(nn.Module):
norm_eps: float = 1e-6,
qk_norm: bool = True,
use_local_attention: bool = False,
apply_gated_attention: bool = False,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
prefix: str = "",
quant_config: QuantizationConfig | None = None,
@@ -461,6 +471,8 @@ class LTX2Attention(nn.Module):
self.norm_eps = float(norm_eps)
self.qk_norm = bool(qk_norm)
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()
if tp_size <= 0:
@@ -499,6 +511,15 @@ class LTX2Attention(nn.Module):
gather_output=False,
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.k_norm: nn.Module | None = None
@@ -561,6 +582,7 @@ class LTX2Attention(nn.Module):
perturbation_mask: torch.Tensor | None = None,
all_perturbed: bool = False,
) -> torch.Tensor:
gate_input = x
context_ = x if context is None else context
v, _ = self.to_v(context_)
use_attention = not all_perturbed
@@ -609,9 +631,17 @@ class LTX2Attention(nn.Module):
if not use_attention:
out = v
out = out.flatten(2)
out, _ = self.to_out[0](out)
return out
if self.to_gate_logits is not None:
gate_logits, _ = self.to_gate_logits(gate_input)
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(
self,
@@ -688,6 +718,10 @@ class LTX2TransformerBlock(nn.Module):
audio_cross_attention_dim: int,
qk_norm: bool = True,
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,
prefix: str = "",
quant_config: QuantizationConfig | None = None,
@@ -695,6 +729,9 @@ class LTX2TransformerBlock(nn.Module):
super().__init__()
self.idx = idx
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)
self.attn1 = LTX2Attention(
@@ -703,6 +740,7 @@ class LTX2TransformerBlock(nn.Module):
dim_head=attention_head_dim,
norm_eps=norm_eps,
qk_norm=qk_norm,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1",
quant_config=quant_config,
@@ -713,6 +751,7 @@ class LTX2TransformerBlock(nn.Module):
dim_head=audio_attention_head_dim,
norm_eps=norm_eps,
qk_norm=qk_norm,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.audio_attn1",
quant_config=quant_config,
@@ -729,6 +768,7 @@ class LTX2TransformerBlock(nn.Module):
norm_eps=norm_eps,
qk_norm=qk_norm,
use_local_attention=True,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn2",
quant_config=quant_config,
@@ -741,6 +781,7 @@ class LTX2TransformerBlock(nn.Module):
norm_eps=norm_eps,
qk_norm=qk_norm,
use_local_attention=True,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.audio_attn2",
quant_config=quant_config,
@@ -754,6 +795,8 @@ class LTX2TransformerBlock(nn.Module):
dim_head=audio_attention_head_dim,
norm_eps=norm_eps,
qk_norm=qk_norm,
use_local_attention=use_local_av_cross_attention,
apply_gated_attention=apply_gated_attention,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.audio_to_video_attn",
quant_config=quant_config,
@@ -765,7 +808,13 @@ class LTX2TransformerBlock(nn.Module):
dim_head=audio_attention_head_dim,
norm_eps=norm_eps,
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",
quant_config=quant_config,
)
@@ -777,14 +826,23 @@ class LTX2TransformerBlock(nn.Module):
)
# 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(
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.audio_a2v_cross_attn_scale_shift_table = nn.Parameter(
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(
self,
@@ -813,6 +871,8 @@ class LTX2TransformerBlock(nn.Module):
audio_encoder_hidden_states: torch.Tensor,
temb: 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_audio_scale_shift: torch.Tensor,
temb_ca_gate: torch.Tensor,
@@ -860,6 +920,55 @@ class LTX2TransformerBlock(nn.Module):
)
audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa
# 2. Prompt Cross-Attention
if self.cross_attention_adaln:
# LTX2.3
if temb_prompt is None or temb_audio_prompt is None:
raise ValueError(
"cross_attention_adaln requires prompt modulation tensors."
)
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
ashift_q, ascale_q, agate_q = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(6, 9)
)
a_prompt_shift, a_prompt_scale = self.get_ada_values(
self.audio_prompt_scale_shift_table,
batch_size,
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,
@@ -976,7 +1085,7 @@ class LTX2TransformerBlock(nn.Module):
)
# 4. Feedforward
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 = (
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
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 = (
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
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:
"""Validate TP-related dimension constraints (fail-fast)."""
if tp_size < 1:
@@ -1089,6 +1204,9 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
)
# 2. Prompt embeddings
self.caption_projection: LTX2TextProjection | None = None
self.audio_caption_projection: LTX2TextProjection | None = None
if not arch.caption_proj_before_connector:
self.caption_projection = LTX2TextProjection(
in_features=arch.caption_channels, hidden_size=self.hidden_size
)
@@ -1098,10 +1216,25 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
# 3. Timestep Modulation Params and Embedding
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_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
@@ -1141,6 +1274,9 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
rope_double_precision = bool(
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))
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,
norm_eps=self.norm_eps,
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,
prefix=config.prefix,
quant_config=quant_config,
@@ -1336,6 +1480,14 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
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)
audio_rotary_emb = self.audio_rope(
audio_coords, device=audio_hidden_states.device
@@ -1367,12 +1519,25 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_embedded_timestep = audio_embedded_timestep.view(
batch_size, -1, audio_embedded_timestep.size(-1)
)
# 3.2. Prepare global modality cross attention modulation parameters
ts_ca_mult = (
self.av_ca_timestep_scale_multiplier / self.timestep_scale_multiplier
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
hidden_dtype = hidden_states.dtype
temb_ca_scale_shift, _ = self.av_ca_video_scale_shift_adaln_single(
timestep.flatten(), hidden_dtype=hidden_dtype
@@ -1403,7 +1568,9 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
)
# 4. Prepare prompt embeddings
if self.caption_projection is not None:
encoder_hidden_states = self.caption_projection(encoder_hidden_states)
if self.audio_caption_projection is not None:
audio_encoder_hidden_states = self.audio_caption_projection(
audio_encoder_hidden_states
)
@@ -1421,6 +1588,8 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
# under ForwardPattern.Pattern_0.
temb=temb,
temb_audio=temb_audio,
temb_prompt=temb_prompt,
temb_audio_prompt=temb_audio_prompt,
temb_ca_scale_shift=temb_ca_scale_shift,
temb_ca_audio_scale_shift=temb_ca_audio_scale_shift,
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
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
class LTX2VideoDownBlock3D(nn.Module):
r"""
@@ -609,6 +621,64 @@ class LTX2VideoMidBlock3d(nn.Module):
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
class LTX2VideoUpBlock3d(nn.Module):
r"""
@@ -1104,6 +1174,192 @@ class LTX2VideoDecoder3d(nn.Module):
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):
r"""
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(
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_spatial_padding_mode = config.arch_config.decoder_spatial_padding_mode
@@ -1175,6 +1435,37 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
encoder_spatial_padding_mode,
)
if use_ltx23_video_decoder:
video_decoder_config = dict(config.arch_config.video_decoder_config)
if not video_decoder_config:
raise ValueError(
"LTX-2.3 native video decoder requires video_decoder_config."
)
self.decoder = LTX23VideoDecoder3d(
in_channels=latent_channels,
out_channels=out_channels,
decoder_blocks=tuple(video_decoder_config["decoder_blocks"]),
patch_size=int(video_decoder_config.get("patch_size", patch_size)),
patch_size_t=patch_size_t,
resnet_norm_eps=resnet_norm_eps,
is_causal=bool(
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,
@@ -1,13 +1,237 @@
import math
from abc import ABC
from contextlib import nullcontext
from typing import Tuple
import einops
import torch
import torch.nn as nn
import torch.nn.functional as F
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):
def __init__(
@@ -61,6 +285,252 @@ class ResBlock(nn.Module):
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):
r"""
LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms.
@@ -72,10 +542,61 @@ class LTX2Vocoder(ABC, nn.Module):
):
super().__init__()
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 = (
getattr(config.arch_config, "sample_rate", None)
or getattr(config.arch_config, "sampling_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
@@ -139,6 +660,12 @@ class LTX2Vocoder(ABC, nn.Module):
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(
self, hidden_states: torch.Tensor, time_last: bool = False
) -> torch.Tensor:
@@ -157,6 +684,32 @@ class LTX2Vocoder(ABC, nn.Module):
`torch.Tensor`:
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
if not time_last:
@@ -5,6 +5,9 @@ import numpy as np
import torch
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 (
PipelineComponentLoader,
)
@@ -44,6 +47,8 @@ def _resolve_ltx2_two_stage_component_paths(
if "spatial_upsampler" not in resolved:
spatial_candidates = [
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"),
]
for candidate in spatial_candidates:
@@ -53,12 +58,15 @@ def _resolve_ltx2_two_stage_component_paths(
break
if "distilled_lora" not in resolved:
distilled_lora = os.path.join(
model_path, "ltx-2-19b-distilled-lora-384.safetensors"
)
distilled_lora_candidates = [
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"),
]
for distilled_lora in distilled_lora_candidates:
if os.path.exists(distilled_lora):
resolved["distilled_lora"] = distilled_lora
auto_resolved.append(f"distilled_lora={distilled_lora}")
break
if auto_resolved:
logger.info(
@@ -81,6 +89,8 @@ def calculate_ltx2_shift(
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(
server_args.pipeline_config.vae_temporal_compression
) + 1
@@ -92,11 +102,44 @@ def prepare_ltx2_mu(batch: Req, server_args: ServerArgs):
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):
"""Prepare native LTX-2 sigma schedule before timestep setup."""
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.extra["ltx2_phase"] = "stage1"
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
batch.sigmas = build_official_ltx2_sigmas(int(batch.num_inference_steps))
else:
batch.sigmas = np.linspace(
1.0,
1.0 / int(batch.num_inference_steps),
@@ -66,6 +66,7 @@ def build_pipeline(
)
else:
logger.info("No pipeline_class_name specified, using model_index.json")
model_info = get_model_info(
model_path,
backend=server_args.backend,
@@ -135,7 +135,7 @@ class Req:
trajectory_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)
is_warmup: bool = False
@@ -25,6 +25,11 @@ class LTX2AVDecodingStage(DecodingStage):
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:
self.load_model()
@@ -40,6 +45,7 @@ class LTX2AVDecodingStage(DecodingStage):
original_dtype = vae_dtype
self.vae.to(torch.bfloat16)
latents = latents.to(torch.bfloat16)
if self._ltx2_should_externally_denorm_video_latents(server_args):
std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latents)
mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latents)
latents = latents * std + mean
@@ -1,5 +1,7 @@
import copy
import json
import math
import os
import time
from io import BytesIO
@@ -10,8 +12,15 @@ import torch
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
from diffusers.models.modeling_outputs import AutoencoderKLOutput
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.models.vaes.ltx_2_3_condition_encoder import (
LTX23VideoConditionEncoder,
)
from sglang.multimodal_gen.runtime.models.vision_utils import (
load_image,
normalize,
@@ -46,6 +55,8 @@ class LTX2AVDenoisingStage(DenoisingStage):
transformer=transformer, scheduler=scheduler, vae=vae, **kwargs
)
self.audio_vae = audio_vae
self._condition_image_encoder = None
self._condition_image_encoder_dir = None
@staticmethod
def _get_video_latent_num_frames_for_model(
@@ -116,13 +127,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
) -> dict[str, object] | None:
if stage != "stage1":
return None
pipeline_ref = getattr(self, "pipeline", None)
pipeline = pipeline_ref() if callable(pipeline_ref) else pipeline_ref
pipeline_name = getattr(pipeline, "pipeline_name", None)
if pipeline_name != "LTX2TwoStagePipeline":
return None
return batch.extra.get("ltx2_stage1_guider_params")
@staticmethod
@@ -141,6 +145,56 @@ class LTX2AVDenoisingStage(DenoisingStage):
factor = rescale_scale * factor + (1.0 - rescale_scale)
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
def _ltx2_calculate_guided_x0(
cls,
@@ -252,6 +306,40 @@ class LTX2AVDenoisingStage(DenoisingStage):
return True
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:
"""Encode `batch.image_path` into packed token latents for LTX-2 TI2V."""
if (
@@ -276,8 +364,11 @@ class LTX2AVDenoisingStage(DenoisingStage):
)
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(
img, width=int(batch.width), height=int(batch.height)
conditioned_img, width=int(batch.width), height=int(batch.height)
)
latents_device = (
@@ -287,17 +378,22 @@ class LTX2AVDenoisingStage(DenoisingStage):
)
encode_dtype = batch.latents.dtype
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 = (
original_dtype != torch.float32
) 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(
img,
conditioned_img,
width=int(batch.width),
height=int(batch.height),
device=latents_device,
dtype=encode_dtype,
apply_codec_compression=False,
)
with torch.autocast(
@@ -306,17 +402,26 @@ class LTX2AVDenoisingStage(DenoisingStage):
enabled=vae_autocast_enabled,
):
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()
except Exception:
pass
if not vae_autocast_enabled:
video_condition = video_condition.to(encode_dtype)
latent_dist: DiagonalGaussianDistribution = self.vae.encode(video_condition)
if condition_image_encoder is not None:
latent = condition_image_encoder(video_condition)
else:
latent_dist: DiagonalGaussianDistribution = self.vae.encode(
video_condition
)
if isinstance(latent_dist, AutoencoderKLOutput):
latent_dist = latent_dist.latent_dist
if condition_image_encoder is None:
mode = server_args.pipeline_config.vae_config.encode_sample_mode()
if mode == "argmax":
latent = latent_dist.mode()
@@ -331,6 +436,8 @@ class LTX2AVDenoisingStage(DenoisingStage):
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:
latent = latent.to(dtype=encode_dtype)
packed = server_args.pipeline_config.maybe_pack_latents(
latent, latent.shape[0], batch
@@ -362,9 +469,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
batch.height,
)
if condition_image_encoder is None:
self.vae.to(original_dtype)
if server_args.vae_cpu_offload:
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()
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
@@ -437,19 +547,15 @@ class LTX2AVDenoisingStage(DenoisingStage):
if do_ti2v:
if not (isinstance(latents, torch.Tensor) and latents.ndim == 3):
raise ValueError("LTX-2 TI2V expects packed token latents [B, S, D].")
latents[:, :num_img_tokens, :] = batch.image_latent[
:, :num_img_tokens, :
].to(device=latents.device, dtype=latents.dtype)
denoise_mask = torch.ones(
(latents.shape[0], latents.shape[1], 1),
device=latents.device,
dtype=torch.float32,
use_zero_clean_latent = is_ltx23_native_variant(
server_args.pipeline_config.vae_config.arch_config
)
latents, denoise_mask, clean_latent = self._prepare_ltx2_ti2v_clean_state(
latents=latents,
image_latent=batch.image_latent,
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(
device_type=current_platform.device_type,
dtype=target_dtype,
@@ -558,11 +664,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
],
dim=0,
)
timestep_video = timestep_video.expand(
int(latent_model_input.shape[0])
cfg_batch_size = int(latent_model_input.shape[0])
timestep_video = self._repeat_batch_dim(
timestep_video, cfg_batch_size
)
timestep_audio = timestep_audio.expand(
int(latent_model_input.shape[0])
timestep_audio = self._repeat_batch_dim(
timestep_audio, cfg_batch_size
)
with set_forward_context(
@@ -616,11 +723,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
audio_latents = audio_scheduler.step(
a_v_pos, t_device, audio_latents, return_dict=False
)[0]
if do_ti2v:
latents[:, :num_img_tokens, :] = batch.image_latent[
:, :num_img_tokens, :
].to(device=latents.device, dtype=latents.dtype)
latents = self.post_forward_for_ti2v_task(
batch, server_args, reserved_frames_mask, latents, z
)
@@ -714,14 +816,20 @@ class LTX2AVDenoisingStage(DenoisingStage):
if a_v_neg is not None:
a_v_neg = a_v_neg.float()
# Velocity -> denoised (x0): x0 = x - sigma * v
sigma_val = float(sigma.item())
denoised_video = (latents.float() - sigma_val * v_pos).to(
latents.dtype
video_sigma_for_x0: float | torch.Tensor = sigma_val
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 = (
audio_latents.float() - sigma_val * a_v_pos
).to(audio_latents.dtype)
denoised_audio = self._ltx2_velocity_to_x0(
audio_latents, a_v_pos, sigma_val
)
denoised_video_cond = denoised_video
denoised_audio_cond = denoised_audio
denoised_video_neg = None
denoised_audio_neg = None
denoised_video_perturbed = None
@@ -737,12 +845,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
and v_neg is not None
and a_v_neg is not None
):
denoised_video_neg = (
latents.float() - sigma_val * v_neg
).to(latents.dtype)
denoised_audio_neg = (
audio_latents.float() - sigma_val * a_v_neg
).to(audio_latents.dtype)
denoised_video_neg = self._ltx2_velocity_to_x0(
latents, v_neg, video_sigma_for_x0
)
denoised_audio_neg = self._ltx2_velocity_to_x0(
audio_latents, a_v_neg, sigma_val
)
if stage1_guider_params is not None:
video_skip = self._ltx2_should_skip_step(
i, int(stage1_guider_params["video_skip_step"])
@@ -784,12 +892,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
stage1_guider_params["audio_stg_blocks"]
),
)
denoised_video_perturbed = (
latents.float() - sigma_val * v_ptb.float()
).to(latents.dtype)
denoised_audio_perturbed = (
audio_latents.float() - sigma_val * a_v_ptb.float()
).to(audio_latents.dtype)
denoised_video_perturbed = self._ltx2_velocity_to_x0(
latents, v_ptb.float(), video_sigma_for_x0
)
denoised_audio_perturbed = self._ltx2_velocity_to_x0(
audio_latents, a_v_ptb.float(), sigma_val
)
need_modality = (
float(stage1_guider_params["video_modality_scale"])
@@ -822,12 +930,12 @@ class LTX2AVDenoisingStage(DenoisingStage):
disable_a2v_cross_attn=True,
disable_v2a_cross_attn=True,
)
denoised_video_modality = (
latents.float() - sigma_val * v_mod.float()
).to(latents.dtype)
denoised_audio_modality = (
audio_latents.float() - sigma_val * a_v_mod.float()
).to(audio_latents.dtype)
denoised_video_modality = self._ltx2_velocity_to_x0(
latents, v_mod.float(), video_sigma_for_x0
)
denoised_audio_modality = self._ltx2_velocity_to_x0(
audio_latents, a_v_mod.float(), sigma_val
)
if not video_skip:
denoised_video = self._ltx2_calculate_guided_x0(
@@ -934,11 +1042,6 @@ class LTX2AVDenoisingStage(DenoisingStage):
audio_latents.float() + v_audio.float() * dt
).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(
batch, server_args, reserved_frames_mask, latents, z
)
@@ -1,6 +1,9 @@
import torch
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.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import (
@@ -60,13 +63,35 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
):
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:
# 1. Prepare Video Latents using base class logic
# This sets batch.latents and batch.raw_latent_shape
if not is_ltx23_native_variant(
server_args.pipeline_config.vae_config.arch_config
):
batch = super().forward(batch, server_args)
# 2. Prepare Audio Latents (optional)
# Default to True if not specified
try:
generate_audio = batch.generate_audio
except AttributeError:
@@ -99,6 +124,82 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
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)
# Default to True if not specified
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
audio_latents = batch.audio_latents
if audio_latents is None:
latent_shape = server_args.pipeline_config.prepare_audio_latent_shape(
batch, batch_size, batch.num_frames
)
audio_latents = randn_tensor(
self._packed_audio_latent_shape(latent_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
)
# Store in batch
batch.audio_latents = audio_latents
batch.raw_audio_latent_shape = audio_latents.shape
@@ -24,8 +24,13 @@ from sglang.utils import load_diffusion_overlay_registry_from_env
logger = init_logger(__name__)
# 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 = [
@@ -42,6 +47,43 @@ MODEL_OVERLAY_METADATA_PATTERNS = [
_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:
return os.path.expanduser(
os.getenv("SGLANG_DIFFUSION_CACHE_ROOT", "~/.cache/sgl_diffusion")
@@ -300,6 +342,15 @@ def download_overlay_metadata(
*,
snapshot_download_fn: Callable[..., 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"])
if os.path.exists(overlay_repo_id):
logger.info(
@@ -419,6 +470,7 @@ def materialize_overlay_model(
materializer_version = str(manifest.get("materializer_version", "v1"))
overlay_repo_id = str(overlay_spec["overlay_repo_id"])
overlay_revision = str(overlay_spec.get("overlay_revision", "main"))
overlay_fingerprint = _compute_overlay_fingerprint(overlay_dir)
cache_key = hashlib.sha256(
json.dumps(
{
@@ -426,6 +478,7 @@ def materialize_overlay_model(
"overlay_repo_id": overlay_repo_id,
"overlay_revision": overlay_revision,
"materializer_version": materializer_version,
"overlay_fingerprint": overlay_fingerprint,
},
sort_keys=True,
).encode("utf-8")
@@ -502,6 +555,7 @@ def materialize_overlay_model(
"overlay_repo_id": overlay_repo_id,
"overlay_revision": overlay_revision,
"materializer_version": materializer_version,
"overlay_fingerprint": overlay_fingerprint,
},
f,
indent=2,
@@ -2504,6 +2504,54 @@
"expected_e2e_ms": 8091.46,
"expected_avg_denoise_ms": 141.37,
"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",
modality="video",
num_gpus=2,
dit_layerwise_offload=True,
extras=["--pipeline-class-name LTX2TwoStagePipeline"],
),
T2V_sampling_params,
@@ -1040,6 +1039,15 @@ TWO_GPU_CASES_B = [
),
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():
@@ -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)