[diffusion] model: support LTX2.3 two stage (#22182)

This commit is contained in:
Mick
2026-04-12 22:15:57 +08:00
committed by GitHub
parent 31453bb76a
commit 495ef8ec64
21 changed files with 1541 additions and 859 deletions
+4 -4
View File
@@ -33,16 +33,16 @@ default parameters when initializing and generating videos.
| TurboWan2.1 T2V 14B | `IPostYellow/TurboWan2.1-T2V-14B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ | | TurboWan2.1 T2V 14B | `IPostYellow/TurboWan2.1-T2V-14B-Diffusers` | 480p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
| TurboWan2.1 T2V 14B 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ | | TurboWan2.1 T2V 14B 720P | `IPostYellow/TurboWan2.1-T2V-14B-720P-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
| TurboWan2.2 I2V A14B | `IPostYellow/TurboWan2.2-I2V-A14B-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ | | TurboWan2.2 I2V A14B | `IPostYellow/TurboWan2.2-I2V-A14B-Diffusers` | 720p | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ⭕ |
| LTX-2 | `Lightricks/LTX-2` | 768×512<br>1536×1024 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | LTX-2 (one and two stages) | `Lightricks/LTX-2` | 768×512<br>1536×1024 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| LTX-2.3 | `Lightricks/LTX-2.3` | 768×512<br>1536×1024 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | LTX-2.3 (one and two stages) | `Lightricks/LTX-2.3` | 768×512<br>1536×1024 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
**Note**: **Note**:
1. Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue. 1. Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
2. SageSLA is based on SpargeAttn. Install it first with `pip install git+https://github.com/thu-ml/SpargeAttn.git --no-build-isolation` 2. SageSLA is based on SpargeAttn. Install it first with `pip install git+https://github.com/thu-ml/SpargeAttn.git --no-build-isolation`
3. LTX-2 two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`. The spatial upsampler and distilled LoRA are auto-resolved from the model snapshot by default, and can still be overridden with `--spatial-upsampler-path` and `--distilled-lora-path`. 3. LTX-2 two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`. The spatial upsampler and distilled LoRA are auto-resolved from the model snapshot by default, and can still be overridden with `--spatial-upsampler-path` and `--distilled-lora-path`.
4. `Lightricks/LTX-2.3` is supported through the bundled native overlay materialization path. One-stage generation uses the default `LTX2Pipeline`; two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`. 3. LTX-2 and LTX-2.3 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`.
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`. - For LTX models, the `Resolutions` column uses output video `width×height` semantics, matching `sglang generate --width ... --height ...`.
### Image Generation Models ### Image Generation Models
@@ -365,6 +365,42 @@ class PipelineConfig:
latents = sequence_model_parallel_all_gather(latents, dim=2) latents = sequence_model_parallel_all_gather(latents, dim=2)
return latents return latents
def can_shard_audio_latents_for_sp(self, audio_latents) -> bool:
"""Return whether this pipeline uses packed audio latents that can be SP-sharded."""
return False
def shard_audio_latents_for_sp(self, batch, audio_latents):
"""Shard packed audio latents for SP. Pipelines without packed audio latents should return the input unchanged."""
return audio_latents, False
def gather_audio_latents_for_sp(self, audio_latents, batch):
"""Gather SP-sharded audio latents back to full sequence length."""
return audio_latents
def prepare_video_rope_coords_for_sp(
self,
model,
batch,
latent_model_input,
*,
num_frames,
height,
width,
):
"""Prepare model-side video RoPE coordinates for the local SP shard when the pipeline requires them."""
return None
def prepare_audio_rope_coords_for_sp(
self,
model,
batch,
audio_latent_model_input,
*,
num_frames,
):
"""Prepare model-side audio RoPE coordinates for the local SP shard when the pipeline requires them."""
return None
def gather_noise_pred_for_sp(self, batch, noise_pred): def gather_noise_pred_for_sp(self, batch, noise_pred):
noise_pred = self.gather_latents_for_sp(noise_pred) noise_pred = self.gather_latents_for_sp(noise_pred)
raw_latent_shape = getattr(batch, "raw_latent_shape", None) raw_latent_shape = getattr(batch, "raw_latent_shape", None)
@@ -345,6 +345,7 @@ class LTX2PipelineConfig(PipelineConfig):
latent_frames, tokens_per_frame = ( latent_frames, tokens_per_frame = (
self._infer_video_latent_frames_and_tokens_per_frame(batch, seq_len) self._infer_video_latent_frames_and_tokens_per_frame(batch, seq_len)
) )
orig_latent_frames = int(latent_frames)
# Pad whole frames so `latent_frames` is divisible by `sp_world_size`. # Pad whole frames so `latent_frames` is divisible by `sp_world_size`.
pad_frames = (sp_world_size - (latent_frames % sp_world_size)) % sp_world_size pad_frames = (sp_world_size - (latent_frames % sp_world_size)) % sp_world_size
@@ -360,6 +361,9 @@ class LTX2PipelineConfig(PipelineConfig):
local_frames = int(latent_frames) // int(sp_world_size) local_frames = int(latent_frames) // int(sp_world_size)
start_frame = int(sp_rank) * int(local_frames) start_frame = int(sp_rank) * int(local_frames)
valid_local_frames = max(
min(int(orig_latent_frames) - int(start_frame), int(local_frames)), 0
)
start = int(start_frame) * int(tokens_per_frame) start = int(start_frame) * int(tokens_per_frame)
end = int(start) + int(local_frames) * int(tokens_per_frame) end = int(start) + int(local_frames) * int(tokens_per_frame)
latents = latents[:, start:end, :] latents = latents[:, start:end, :]
@@ -368,6 +372,9 @@ class LTX2PipelineConfig(PipelineConfig):
batch.sp_video_latent_num_frames = int(local_frames) batch.sp_video_latent_num_frames = int(local_frames)
batch.sp_video_start_frame = int(start_frame) batch.sp_video_start_frame = int(start_frame)
batch.sp_video_tokens_per_frame = int(tokens_per_frame) batch.sp_video_tokens_per_frame = int(tokens_per_frame)
batch.sp_video_valid_token_count = int(valid_local_frames) * int(
tokens_per_frame
)
return latents, True return latents, True
@@ -379,6 +386,104 @@ class LTX2PipelineConfig(PipelineConfig):
return sequence_model_parallel_all_gather(latents.contiguous(), dim=1) return sequence_model_parallel_all_gather(latents.contiguous(), dim=1)
return super().gather_latents_for_sp(latents, batch=batch) return super().gather_latents_for_sp(latents, batch=batch)
def shard_audio_latents_for_sp(self, batch, audio_latents):
sp_world_size = get_sp_world_size()
if sp_world_size <= 1:
return audio_latents, False
if not (isinstance(audio_latents, torch.Tensor) and audio_latents.ndim == 3):
return audio_latents, False
sp_rank = get_sp_parallel_rank()
seq_len = int(audio_latents.shape[1])
batch.sp_audio_orig_num_frames = int(seq_len)
pad_frames = (sp_world_size - (seq_len % sp_world_size)) % sp_world_size
if pad_frames:
pad = torch.zeros(
(audio_latents.shape[0], pad_frames, audio_latents.shape[2]),
device=audio_latents.device,
dtype=audio_latents.dtype,
)
audio_latents = torch.cat([audio_latents, pad], dim=1)
seq_len += int(pad_frames)
local_frames = seq_len // sp_world_size
start_frame = sp_rank * local_frames
end_frame = start_frame + local_frames
valid_local_frames = max(
min(
int(batch.sp_audio_orig_num_frames) - int(start_frame),
int(local_frames),
),
0,
)
audio_latents = audio_latents[:, start_frame:end_frame, :]
batch.sp_audio_latent_num_frames = int(local_frames)
batch.sp_audio_start_frame = int(start_frame)
batch.sp_audio_valid_token_count = int(valid_local_frames)
return audio_latents, True
def can_shard_audio_latents_for_sp(self, audio_latents) -> bool:
return (
get_sp_world_size() > 1
and isinstance(audio_latents, torch.Tensor)
and audio_latents.ndim == 3
)
def gather_audio_latents_for_sp(self, audio_latents, batch):
if get_sp_world_size() <= 1:
return audio_latents
if not (isinstance(audio_latents, torch.Tensor) and audio_latents.ndim == 3):
return audio_latents
audio_latents = sequence_model_parallel_all_gather(
audio_latents.contiguous(), dim=1
)
orig_num_frames = int(batch.sp_audio_orig_num_frames)
if orig_num_frames > 0:
audio_latents = audio_latents[:, :orig_num_frames, :]
return audio_latents
def prepare_video_rope_coords_for_sp(
self,
model,
batch,
latent_model_input,
*,
num_frames,
height,
width,
):
if not batch.did_sp_shard_latents:
return None
return model.rope.prepare_video_coords(
batch_size=int(latent_model_input.shape[0]),
num_frames=num_frames,
height=height,
width=width,
device=latent_model_input.device,
fps=batch.fps,
start_frame=int(batch.sp_video_start_frame),
)
def prepare_audio_rope_coords_for_sp(
self,
model,
batch,
audio_latent_model_input,
*,
num_frames,
):
if not batch.did_sp_shard_audio_latents:
return None
return model.audio_rope.prepare_audio_coords(
batch_size=int(audio_latent_model_input.shape[0]),
num_frames=num_frames,
device=audio_latent_model_input.device,
start_frame=int(batch.sp_audio_start_frame),
)
def maybe_pack_audio_latents(self, latents, batch_size, batch): def maybe_pack_audio_latents(self, latents, batch_size, batch):
# If already packed (3D shape [B, T, C*F]), skip packing # If already packed (3D shape [B, T, C*F]), skip packing
if latents.dim() == 3: if latents.dim() == 3:
@@ -1,302 +0,0 @@
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"),
)
@@ -398,8 +398,10 @@ class USPAttention(nn.Module):
q: torch.Tensor, q: torch.Tensor,
k: torch.Tensor, k: torch.Tensor,
v: torch.Tensor, v: torch.Tensor,
attn_mask: torch.Tensor | None = None,
num_replicated_prefix: int = 0, num_replicated_prefix: int = 0,
num_replicated_suffix: int = 0, num_replicated_suffix: int = 0,
skip_sequence_parallel_override: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Forward pass for USPAttention. Forward pass for USPAttention.
@@ -421,7 +423,82 @@ class USPAttention(nn.Module):
""" """
forward_context: ForwardContext = get_forward_context() forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata ctx_attn_metadata = forward_context.attn_metadata
if self.skip_sequence_parallel or get_sequence_parallel_world_size() == 1: effective_skip_sp = (
self.skip_sequence_parallel or skip_sequence_parallel_override
)
if attn_mask is not None:
def _prepare_sdpa_mask(
mask: torch.Tensor, *, dtype: torch.dtype, device: torch.device
) -> torch.Tensor:
mask = mask.to(device=device)
if torch.is_floating_point(mask):
mask = mask.to(dtype=dtype)
if mask.dim() == 2:
mask = mask[:, None, None, :]
elif mask.dim() == 3:
mask = mask[:, None, :, :]
return mask
mask = mask.to(dtype=dtype)
if mask.dim() == 2:
mask = mask[:, None, None, :]
elif mask.dim() == 3:
mask = mask[:, None, :, :]
return (mask - 1.0) * torch.finfo(dtype).max
sp_world_size = get_sequence_parallel_world_size()
if effective_skip_sp or sp_world_size == 1:
q_ = q.transpose(1, 2)
k_ = k.transpose(1, 2)
v_ = v.transpose(1, 2)
mask = _prepare_sdpa_mask(attn_mask, dtype=q_.dtype, device=q_.device)
return torch.nn.functional.scaled_dot_product_attention(
q_,
k_,
v_,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
scale=self.softmax_scale,
).transpose(1, 2)
if get_ring_parallel_world_size() > 1:
raise NotImplementedError(
"USPAttention masked path does not support ring parallelism yet."
)
if attn_mask.dim() != 2:
raise NotImplementedError(
"USPAttention masked SP path currently expects a [B, S_local] key mask."
)
sp_size = get_ulysses_parallel_world_size()
if sp_size > 1:
q = _usp_input_all_to_all(q, head_dim=2)
k = _usp_input_all_to_all(k, head_dim=2)
v = _usp_input_all_to_all(v, head_dim=2)
gathered_mask = sequence_model_parallel_all_gather(
attn_mask.contiguous(), dim=1
)
q_ = q.transpose(1, 2)
k_ = k.transpose(1, 2)
v_ = v.transpose(1, 2)
mask = _prepare_sdpa_mask(gathered_mask, dtype=q_.dtype, device=q_.device)
out = torch.nn.functional.scaled_dot_product_attention(
q_,
k_,
v_,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
scale=self.softmax_scale,
).transpose(1, 2)
if sp_size > 1:
out = _usp_output_all_to_all(out, head_dim=2)
return out
if effective_skip_sp or get_sequence_parallel_world_size() == 1:
# No sequence parallelism, just run local attention. # No sequence parallelism, just run local attention.
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata) out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return out return out
@@ -19,6 +19,7 @@ from sglang.multimodal_gen.runtime.distributed import (
model_parallel_is_initialized, model_parallel_is_initialized,
) )
from sglang.multimodal_gen.runtime.distributed.communication_op import ( from sglang.multimodal_gen.runtime.distributed.communication_op import (
sequence_model_parallel_all_gather,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention
@@ -581,6 +582,8 @@ class LTX2Attention(nn.Module):
k_pe: tuple[torch.Tensor, torch.Tensor] | None = None, k_pe: tuple[torch.Tensor, torch.Tensor] | None = None,
perturbation_mask: torch.Tensor | None = None, perturbation_mask: torch.Tensor | None = None,
all_perturbed: bool = False, all_perturbed: bool = False,
skip_sequence_parallel_override: bool = False,
gather_context_kv_for_sp: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
gate_input = x gate_input = x
context_ = x if context is None else context context_ = x if context is None else context
@@ -620,10 +623,34 @@ class LTX2Attention(nn.Module):
q = q.view(*q.shape[:-1], self.local_heads, self.dim_head) q = q.view(*q.shape[:-1], self.local_heads, self.dim_head)
k = k.view(*k.shape[:-1], self.local_heads, self.dim_head) k = k.view(*k.shape[:-1], self.local_heads, self.dim_head)
if gather_context_kv_for_sp:
k_full = sequence_model_parallel_all_gather(k.contiguous(), dim=1)
v_full = sequence_model_parallel_all_gather(v.contiguous(), dim=1)
gathered_mask = None
if mask is not None:
gathered_mask = sequence_model_parallel_all_gather(
mask.contiguous(), dim=1
)
if self.use_local_attention: if self.use_local_attention:
out = self.attn(q, k_full, v_full, attn_mask=gathered_mask)
else:
out = self.attn(
q,
k_full,
v_full,
attn_mask=gathered_mask,
skip_sequence_parallel_override=True,
)
elif self.use_local_attention:
out = self.attn(q, k, v, attn_mask=mask) out = self.attn(q, k, v, attn_mask=mask)
else: else:
out = self.attn(q, k, v) out = self.attn(
q,
k,
v,
attn_mask=mask,
skip_sequence_parallel_override=skip_sequence_parallel_override,
)
if perturbation_mask is not None: if perturbation_mask is not None:
out = out * perturbation_mask + v * (1 - perturbation_mask) out = out * perturbation_mask + v * (1 - perturbation_mask)
@@ -883,12 +910,15 @@ class LTX2TransformerBlock(nn.Module):
ca_audio_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, ca_audio_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
encoder_attention_mask: Optional[torch.Tensor] = None, encoder_attention_mask: Optional[torch.Tensor] = None,
audio_encoder_attention_mask: Optional[torch.Tensor] = None, audio_encoder_attention_mask: Optional[torch.Tensor] = None,
video_self_attention_mask: Optional[torch.Tensor] = None,
audio_self_attention_mask: Optional[torch.Tensor] = None,
a2v_cross_attention_mask: Optional[torch.Tensor] = None, a2v_cross_attention_mask: Optional[torch.Tensor] = None,
v2a_cross_attention_mask: Optional[torch.Tensor] = None, v2a_cross_attention_mask: Optional[torch.Tensor] = None,
skip_video_self_attn: bool = False, skip_video_self_attn: bool = False,
skip_audio_self_attn: bool = False, skip_audio_self_attn: bool = False,
skip_a2v_cross_attn: bool = False, skip_a2v_cross_attn: bool = False,
skip_v2a_cross_attn: bool = False, skip_v2a_cross_attn: bool = False,
audio_replicated_for_sp: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
batch_size = hidden_states.size(0) batch_size = hidden_states.size(0)
@@ -902,8 +932,10 @@ class LTX2TransformerBlock(nn.Module):
) )
attn_hidden_states = self.attn1( attn_hidden_states = self.attn1(
norm_hidden_states, norm_hidden_states,
mask=video_self_attention_mask,
pe=video_rotary_emb, pe=video_rotary_emb,
all_perturbed=skip_video_self_attn, all_perturbed=skip_video_self_attn,
gather_context_kv_for_sp=audio_replicated_for_sp,
) )
hidden_states = hidden_states + attn_hidden_states * vgate_msa hidden_states = hidden_states + attn_hidden_states * vgate_msa
@@ -915,8 +947,10 @@ class LTX2TransformerBlock(nn.Module):
) )
attn_audio_hidden_states = self.audio_attn1( attn_audio_hidden_states = self.audio_attn1(
norm_audio_hidden_states, norm_audio_hidden_states,
mask=audio_self_attention_mask,
pe=audio_rotary_emb, pe=audio_rotary_emb,
all_perturbed=skip_audio_self_attn, all_perturbed=skip_audio_self_attn,
skip_sequence_parallel_override=audio_replicated_for_sp,
) )
audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa
# 2. Prompt Cross-Attention # 2. Prompt Cross-Attention
@@ -1061,6 +1095,7 @@ class LTX2TransformerBlock(nn.Module):
pe=ca_video_rotary_emb, pe=ca_video_rotary_emb,
k_pe=ca_audio_rotary_emb, k_pe=ca_audio_rotary_emb,
mask=a2v_cross_attention_mask, mask=a2v_cross_attention_mask,
skip_sequence_parallel_override=audio_replicated_for_sp,
) )
hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states
@@ -1079,6 +1114,7 @@ class LTX2TransformerBlock(nn.Module):
pe=ca_audio_rotary_emb, pe=ca_audio_rotary_emb,
k_pe=ca_video_rotary_emb, k_pe=ca_video_rotary_emb,
mask=v2a_cross_attention_mask, mask=v2a_cross_attention_mask,
gather_context_kv_for_sp=audio_replicated_for_sp,
) )
audio_hidden_states = ( audio_hidden_states = (
audio_hidden_states + v2a_gate * v2a_attn_hidden_states audio_hidden_states + v2a_gate * v2a_attn_hidden_states
@@ -1415,6 +1451,45 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
self.layer_names = ["transformer_blocks"] self.layer_names = ["transformer_blocks"]
def _maybe_quantize_video_rope_coords(
self,
video_coords: torch.Tensor,
hidden_device: torch.device,
hidden_dtype: torch.dtype,
) -> torch.Tensor:
if self.quantize_video_rope_coords_to_hidden_dtype:
return video_coords.to(device=hidden_device, dtype=hidden_dtype)
return video_coords.to(device=hidden_device)
def _get_av_ca_gate_timestep_factor(self) -> float:
ltx_variant = str(getattr(self.config.arch_config, "ltx_variant", "ltx_2"))
if ltx_variant == "ltx_2_3":
return self.av_ca_timestep_scale_multiplier / self.timestep_scale_multiplier
return float(self.av_ca_timestep_scale_multiplier)
def _get_av_ca_timesteps(
self,
timestep: torch.Tensor,
audio_timestep: torch.Tensor,
prompt_timestep: torch.Tensor | None,
audio_prompt_timestep: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
ltx_variant = str(getattr(self.config.arch_config, "ltx_variant", "ltx_2"))
if ltx_variant != "ltx_2_3":
return timestep, audio_timestep
video_timestep = (
self._collapse_prompt_timestep(timestep)
if prompt_timestep is None
else prompt_timestep
)
audio_timestep_for_ca = (
self._collapse_prompt_timestep(audio_timestep)
if audio_prompt_timestep is None
else audio_prompt_timestep
)
return video_timestep, audio_timestep_for_ca
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -1423,6 +1498,8 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_encoder_hidden_states: torch.Tensor, audio_encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor, timestep: torch.LongTensor,
audio_timestep: Optional[torch.LongTensor] = None, audio_timestep: Optional[torch.LongTensor] = None,
prompt_timestep: Optional[torch.Tensor] = None,
audio_prompt_timestep: Optional[torch.Tensor] = None,
encoder_attention_mask: Optional[torch.Tensor] = None, encoder_attention_mask: Optional[torch.Tensor] = None,
audio_encoder_attention_mask: Optional[torch.Tensor] = None, audio_encoder_attention_mask: Optional[torch.Tensor] = None,
num_frames: Optional[int] = None, num_frames: Optional[int] = None,
@@ -1432,10 +1509,15 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
audio_num_frames: Optional[int] = None, audio_num_frames: Optional[int] = None,
video_coords: Optional[torch.Tensor] = None, video_coords: Optional[torch.Tensor] = None,
audio_coords: Optional[torch.Tensor] = None, audio_coords: Optional[torch.Tensor] = None,
video_self_attention_mask: Optional[torch.Tensor] = None,
audio_self_attention_mask: Optional[torch.Tensor] = None,
a2v_cross_attention_mask: Optional[torch.Tensor] = None,
v2a_cross_attention_mask: Optional[torch.Tensor] = None,
skip_video_self_attn_blocks: Optional[tuple[int, ...]] = None, skip_video_self_attn_blocks: Optional[tuple[int, ...]] = None,
skip_audio_self_attn_blocks: Optional[tuple[int, ...]] = None, skip_audio_self_attn_blocks: Optional[tuple[int, ...]] = None,
disable_a2v_cross_attn: bool = False, disable_a2v_cross_attn: bool = False,
disable_v2a_cross_attn: bool = False, disable_v2a_cross_attn: bool = False,
audio_replicated_for_sp: bool = False,
**kwargs, **kwargs,
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ) -> tuple[torch.Tensor | None, torch.Tensor | None]:
@@ -1480,14 +1562,10 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
device=audio_hidden_states.device, device=audio_hidden_states.device,
) )
if self.quantize_video_rope_coords_to_hidden_dtype: video_coords = self._maybe_quantize_video_rope_coords(
video_coords = video_coords.to( video_coords, hidden_states.device, hidden_states.dtype
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) audio_coords = audio_coords.to(device=audio_hidden_states.device)
video_rotary_emb = self.rope(video_coords, device=hidden_states.device) video_rotary_emb = self.rope(video_coords, device=hidden_states.device)
audio_rotary_emb = self.audio_rope( audio_rotary_emb = self.audio_rope(
audio_coords, device=audio_hidden_states.device audio_coords, device=audio_hidden_states.device
@@ -1506,6 +1584,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
# 3.1. Prepare global modality (video and audio) timestep embedding and modulation parameters # 3.1. Prepare global modality (video and audio) timestep embedding and modulation parameters
temb, embedded_timestep = self.adaln_single( temb, embedded_timestep = self.adaln_single(
timestep.flatten(), timestep.flatten(),
hidden_dtype=hidden_states.dtype,
) )
temb = temb.view(batch_size, -1, temb.size(-1)) temb = temb.view(batch_size, -1, temb.size(-1))
embedded_timestep = embedded_timestep.view( embedded_timestep = embedded_timestep.view(
@@ -1513,7 +1592,8 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
) )
temb_audio, audio_embedded_timestep = self.audio_adaln_single( temb_audio, audio_embedded_timestep = self.audio_adaln_single(
audio_timestep.flatten() audio_timestep.flatten(),
hidden_dtype=audio_hidden_states.dtype,
) )
temb_audio = temb_audio.view(batch_size, -1, temb_audio.size(-1)) temb_audio = temb_audio.view(batch_size, -1, temb_audio.size(-1))
audio_embedded_timestep = audio_embedded_timestep.view( audio_embedded_timestep = audio_embedded_timestep.view(
@@ -1522,13 +1602,21 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
temb_prompt = None temb_prompt = None
temb_audio_prompt = None temb_audio_prompt = None
if self.prompt_adaln_single is not None: if self.prompt_adaln_single is not None:
prompt_timestep = self._collapse_prompt_timestep(timestep) prompt_timestep = (
self._collapse_prompt_timestep(timestep)
if prompt_timestep is None
else prompt_timestep
)
temb_prompt, _ = self.prompt_adaln_single( temb_prompt, _ = self.prompt_adaln_single(
prompt_timestep.flatten(), hidden_dtype=hidden_states.dtype prompt_timestep.flatten(), hidden_dtype=hidden_states.dtype
) )
temb_prompt = temb_prompt.view(batch_size, -1, temb_prompt.size(-1)) temb_prompt = temb_prompt.view(batch_size, -1, temb_prompt.size(-1))
if self.audio_prompt_adaln_single is not None: if self.audio_prompt_adaln_single is not None:
audio_prompt_timestep = self._collapse_prompt_timestep(audio_timestep) audio_prompt_timestep = (
self._collapse_prompt_timestep(audio_timestep)
if audio_prompt_timestep is None
else audio_prompt_timestep
)
temb_audio_prompt, _ = self.audio_prompt_adaln_single( temb_audio_prompt, _ = self.audio_prompt_adaln_single(
audio_prompt_timestep.flatten(), audio_prompt_timestep.flatten(),
hidden_dtype=audio_hidden_states.dtype, hidden_dtype=audio_hidden_states.dtype,
@@ -1539,28 +1627,35 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
# 3.2. Prepare global modality cross attention modulation parameters # 3.2. Prepare global modality cross attention modulation parameters
hidden_dtype = hidden_states.dtype hidden_dtype = hidden_states.dtype
av_ca_video_timestep, av_ca_audio_timestep = self._get_av_ca_timesteps(
timestep,
audio_timestep,
prompt_timestep,
audio_prompt_timestep,
)
temb_ca_scale_shift, _ = self.av_ca_video_scale_shift_adaln_single( temb_ca_scale_shift, _ = self.av_ca_video_scale_shift_adaln_single(
timestep.flatten(), hidden_dtype=hidden_dtype av_ca_video_timestep.flatten(), hidden_dtype=hidden_dtype
) )
temb_ca_scale_shift = temb_ca_scale_shift.view( temb_ca_scale_shift = temb_ca_scale_shift.view(
batch_size, -1, temb_ca_scale_shift.shape[-1] batch_size, -1, temb_ca_scale_shift.shape[-1]
) )
av_ca_gate_factor = self._get_av_ca_gate_timestep_factor()
temb_ca_gate, _ = self.av_ca_a2v_gate_adaln_single( temb_ca_gate, _ = self.av_ca_a2v_gate_adaln_single(
timestep.flatten() * self.av_ca_timestep_scale_multiplier, av_ca_video_timestep.flatten() * av_ca_gate_factor,
hidden_dtype=hidden_dtype, hidden_dtype=hidden_dtype,
) )
temb_ca_gate = temb_ca_gate.view(batch_size, -1, temb_ca_gate.shape[-1]) temb_ca_gate = temb_ca_gate.view(batch_size, -1, temb_ca_gate.shape[-1])
temb_ca_audio_scale_shift, _ = self.av_ca_audio_scale_shift_adaln_single( temb_ca_audio_scale_shift, _ = self.av_ca_audio_scale_shift_adaln_single(
audio_timestep.flatten(), hidden_dtype=audio_hidden_states.dtype av_ca_audio_timestep.flatten(), hidden_dtype=audio_hidden_states.dtype
) )
temb_ca_audio_scale_shift = temb_ca_audio_scale_shift.view( temb_ca_audio_scale_shift = temb_ca_audio_scale_shift.view(
batch_size, -1, temb_ca_audio_scale_shift.shape[-1] batch_size, -1, temb_ca_audio_scale_shift.shape[-1]
) )
temb_ca_audio_gate, _ = self.av_ca_v2a_gate_adaln_single( temb_ca_audio_gate, _ = self.av_ca_v2a_gate_adaln_single(
audio_timestep.flatten() * self.av_ca_timestep_scale_multiplier, av_ca_audio_timestep.flatten() * av_ca_gate_factor,
hidden_dtype=audio_hidden_states.dtype, hidden_dtype=audio_hidden_states.dtype,
) )
temb_ca_audio_gate = temb_ca_audio_gate.view( temb_ca_audio_gate = temb_ca_audio_gate.view(
@@ -1600,10 +1695,15 @@ class LTX2VideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
ca_audio_rotary_emb=ca_audio_rotary_emb, ca_audio_rotary_emb=ca_audio_rotary_emb,
encoder_attention_mask=encoder_attention_mask, encoder_attention_mask=encoder_attention_mask,
audio_encoder_attention_mask=audio_encoder_attention_mask, audio_encoder_attention_mask=audio_encoder_attention_mask,
video_self_attention_mask=video_self_attention_mask,
audio_self_attention_mask=audio_self_attention_mask,
a2v_cross_attention_mask=a2v_cross_attention_mask,
v2a_cross_attention_mask=v2a_cross_attention_mask,
skip_video_self_attn=block.idx in skip_video_self_attn_blocks, skip_video_self_attn=block.idx in skip_video_self_attn_blocks,
skip_audio_self_attn=block.idx in skip_audio_self_attn_blocks, skip_audio_self_attn=block.idx in skip_audio_self_attn_blocks,
skip_a2v_cross_attn=disable_a2v_cross_attn, skip_a2v_cross_attn=disable_a2v_cross_attn,
skip_v2a_cross_attn=disable_v2a_cross_attn, skip_v2a_cross_attn=disable_v2a_cross_attn,
audio_replicated_for_sp=audio_replicated_for_sp,
) )
# 6. Output layers # 6. Output layers
@@ -46,9 +46,9 @@ def _resolve_ltx2_two_stage_component_paths(
if "spatial_upsampler" not in resolved: if "spatial_upsampler" not in resolved:
spatial_candidates = [ spatial_candidates = [
os.path.join(model_path, "latent_upsampler"),
os.path.join(model_path, "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.3-spatial-upscaler-x2-1.0.safetensors"),
os.path.join(model_path, "ltx-2.3-spatial-upscaler-x2-1.1.safetensors"),
os.path.join(model_path, "latent_upsampler"),
os.path.join(model_path, "ltx-2-spatial-upscaler-x2-1.0.safetensors"), os.path.join(model_path, "ltx-2-spatial-upscaler-x2-1.0.safetensors"),
] ]
for candidate in spatial_candidates: for candidate in spatial_candidates:
@@ -59,6 +59,7 @@ def _resolve_ltx2_two_stage_component_paths(
if "distilled_lora" not in resolved: if "distilled_lora" not in resolved:
distilled_lora_candidates = [ distilled_lora_candidates = [
os.path.join(model_path, "ltx-2.3-20b-distilled-lora-384.safetensors"),
os.path.join(model_path, "ltx-2.3-22b-distilled-lora-384.safetensors"), os.path.join(model_path, "ltx-2.3-22b-distilled-lora-384.safetensors"),
os.path.join(model_path, "ltx-2-19b-distilled-lora-384.safetensors"), os.path.join(model_path, "ltx-2-19b-distilled-lora-384.safetensors"),
] ]
@@ -264,6 +265,12 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
pipeline_name = "LTX2TwoStagePipeline" pipeline_name = "LTX2TwoStagePipeline"
STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0] STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
@staticmethod
def _should_merge_stage2_distilled_lora(server_args: ServerArgs) -> bool:
return is_ltx23_native_variant(
server_args.pipeline_config.vae_config.arch_config
)
def initialize_pipeline(self, server_args: ServerArgs): def initialize_pipeline(self, server_args: ServerArgs):
super().initialize_pipeline(server_args) super().initialize_pipeline(server_args)
server_args.component_paths = _resolve_ltx2_two_stage_component_paths( server_args.component_paths = _resolve_ltx2_two_stage_component_paths(
@@ -332,10 +339,12 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
lora_path=lora_paths, lora_path=lora_paths,
target=lora_targets, target=lora_targets,
strength=lora_strengths, strength=lora_strengths,
# Keep the distilled adapter unmerged when it is the only active LoRA. # Official LTX-2.3 two-stage builds stage 2 with distilled LoRA fused
# Merging it into the base weights makes the subsequent switch back to # into the transformer weights. Legacy LTX-2 should keep the
# stage 1 depend on unmerge bookkeeping instead of the original base. # preexisting unmerged behavior to avoid regressing stage 2 quality.
merge_weights=self._stage1_lora_path is not None, merge_weights=self._should_merge_stage2_distilled_lora(
self.server_args
),
) )
else: else:
raise ValueError(f"Unknown LTX2 two-stage LoRA phase: {phase}") raise ValueError(f"Unknown LTX2 two-stage LoRA phase: {phase}")
@@ -102,11 +102,16 @@ class Req:
audio_latents: torch.Tensor | None = None audio_latents: torch.Tensor | None = None
audio_noise: torch.Tensor | None = None audio_noise: torch.Tensor | None = None
raw_audio_latent_shape: tuple[int, ...] | None = None raw_audio_latent_shape: tuple[int, ...] | None = None
did_sp_shard_audio_latents: bool = False
sp_audio_start_frame: int = 0
sp_audio_orig_num_frames: int = 0
# Audio Parameters # Audio Parameters
generate_audio: bool = True generate_audio: bool = True
raw_latent_shape: torch.Tensor | None = None raw_latent_shape: torch.Tensor | None = None
did_sp_shard_latents: bool = False
sp_video_start_frame: int = 0
noise_pred: torch.Tensor | None = None noise_pred: torch.Tensor | None = None
# vae-encoded condition image # vae-encoded condition image
image_latent: torch.Tensor | list[torch.Tensor] | None = None image_latent: torch.Tensor | list[torch.Tensor] | None = None
@@ -855,9 +855,21 @@ class DenoisingStage(PipelineStage):
# image_latent must be sharded consistently with latents when it is # image_latent must be sharded consistently with latents when it is
# concatenated along the sequence dimension in the denoising loop. # concatenated along the sequence dimension in the denoising loop.
if batch.image_latent is not None: if batch.image_latent is not None:
sp_video_metadata = {
name: getattr(batch, name)
for name in (
"sp_video_latent_num_frames",
"sp_video_start_frame",
"sp_video_tokens_per_frame",
"sp_video_valid_token_count",
)
if hasattr(batch, name)
}
batch.image_latent, _ = server_args.pipeline_config.shard_latents_for_sp( batch.image_latent, _ = server_args.pipeline_config.shard_latents_for_sp(
batch, batch.image_latent batch, batch.image_latent
) )
for name, value in sp_video_metadata.items():
setattr(batch, name, value)
def _postprocess_sp_latents( def _postprocess_sp_latents(
self, self,
File diff suppressed because it is too large Load Diff
@@ -61,6 +61,12 @@ class LTX2AVLatentPreparationStage(LatentPreparationStage):
batch: Req, batch: Req,
server_args: ServerArgs, server_args: ServerArgs,
): ):
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
if server_args.pipeline_class_name == "LTX2TwoStagePipeline":
return server_args.pipeline_config.get_latent_dtype(
batch.prompt_embeds[0].dtype
)
return torch.float32
return torch.float32 return torch.float32
@staticmethod @staticmethod
@@ -375,6 +375,21 @@ class ServerArgs:
) )
if self.attention_backend is None and self.backend != Backend.DIFFUSERS: if self.attention_backend is None and self.backend != Backend.DIFFUSERS:
if (
current_platform.is_cuda()
and self.pipeline_class_name is None
and self.num_gpus == 1
and self.tp_size == 1
and self.sp_degree == 1
and self.ulysses_degree == 1
and self.ring_degree == 1
and self._is_ltx23_model_path(self.model_path)
):
self.attention_backend = "fa"
logger.info(
"Automatically set attention_backend=fa for LTX-2.3 one-stage on 1 GPU to preserve precision"
)
return
self._set_default_attention_backend() self._set_default_attention_backend()
def _adjust_warmup(self): def _adjust_warmup(self):
@@ -409,12 +424,17 @@ class ServerArgs:
self.master_port = self.settle_port(self.master_port, 37) self.master_port = self.settle_port(self.master_port, 37)
def _adjust_parallelism(self): def _adjust_parallelism(self):
if self.tp_size is None: tp_unspecified = self.tp_size is None
self.tp_size = 1 sp_unspecified = self.sp_degree is None
ulysses_unspecified = self.ulysses_degree is None
ring_unspecified = self.ring_degree is None
if self.hsdp_shard_dim is None: if self.hsdp_shard_dim is None:
self.hsdp_shard_dim = self.num_gpus self.hsdp_shard_dim = self.num_gpus
if self.tp_size is None:
self.tp_size = 1
# adjust sp_degree: allocate all remaining GPUs after TP and DP # adjust sp_degree: allocate all remaining GPUs after TP and DP
if self.sp_degree is None: if self.sp_degree is None:
num_gpus_per_group = self.dp_size * self.tp_size num_gpus_per_group = self.dp_size * self.tp_size
@@ -446,6 +466,20 @@ class ServerArgs:
self.ring_degree = 1 self.ring_degree = 1
logger.debug(f"Ring degree not set, using default value {self.ring_degree}") logger.debug(f"Ring degree not set, using default value {self.ring_degree}")
@staticmethod
def _is_ltx23_model_path(model_path: str | None) -> bool:
if not model_path:
return False
normalized = model_path.lower()
return any(
token in normalized
for token in (
"lightricks/ltx-2.3",
"models--lightricks--ltx-2.3",
"lightricks__ltx-2.3",
)
)
def _adjust_platform_specific(self): def _adjust_platform_specific(self):
if current_platform.is_mps(): if current_platform.is_mps():
self.use_fsdp_inference = False self.use_fsdp_inference = False
@@ -26,6 +26,7 @@ logger = init_logger(__name__)
# Built-in diffusion model overlay registry. # Built-in diffusion model overlay registry.
BUILTIN_MODEL_OVERLAY_REGISTRY: dict[str, dict[str, Any]] = { BUILTIN_MODEL_OVERLAY_REGISTRY: dict[str, dict[str, Any]] = {
"Lightricks/LTX-2.3": { "Lightricks/LTX-2.3": {
# TODO: consider move to lmsys hf repo
"overlay_repo_id": "MickJ/LTX-2.3-overlay", "overlay_repo_id": "MickJ/LTX-2.3-overlay",
"overlay_revision": "main", "overlay_revision": "main",
"bundled_overlay_subdir": "ltx_2_3", "bundled_overlay_subdir": "ltx_2_3",
@@ -73,14 +73,6 @@ SKIP_COMPONENTS: Dict[str, Dict[ComponentType, ComponentSkip]] = {
"HF reference transformer cannot be materialized from the video_dit repo layout" "HF reference transformer cannot be materialized from the video_dit repo layout"
) )
}, },
"ltx_2.3_one_stage_ti2v": {
ComponentType.VAE: ComponentSkip(
"LTX-2.3 VAE component diverges from the HF reference after local overlay materialization; weight transfer matched 96/176 (54.55%), below the minimum threshold for trustworthy comparison"
),
ComponentType.TRANSFORMER: ComponentSkip(
"LTX-2.3 transformer component does not match the HF reference architecture after local overlay materialization; scale_shift_table parameters load as [9, ...] in the checkpoint but [6, ...] in the reference model"
),
},
"qwen_image_t2i_cache_dit_enabled": { "qwen_image_t2i_cache_dit_enabled": {
ComponentType.VAE: ComponentSkip( ComponentType.VAE: ComponentSkip(
"Representative VAE accuracy is already covered by qwen_image_t2i for the same source component and topology" "Representative VAE accuracy is already covered by qwen_image_t2i for the same source component and topology"
@@ -2562,6 +2562,61 @@
"expected_e2e_ms": 26916.58, "expected_e2e_ms": 26916.58,
"expected_avg_denoise_ms": 715.73, "expected_avg_denoise_ms": 715.73,
"expected_median_denoise_ms": 707.35 "expected_median_denoise_ms": 707.35
},
"ltx_2.3_two_stage_t2v_2gpus": {
"stages_ms": {
"InputValidationStage": 0.05,
"TextEncodingStage": 2020.14,
"LTX2TextConnectorStage": 26.56,
"LTX2HalveResolutionStage": 0.06,
"LTX2LoRASwitchStage": 104.32,
"LTX2SigmaPreparationStage": 0.37,
"TimestepPreparationStage": 26.33,
"LTX2AVLatentPreparationStage": 0.13,
"LTX2AVDenoisingStage": 25176.87,
"LTX2UpsampleStage": 549.01,
"LTX2RefinementStage": 663.05,
"LTX2AVDecodingStage": 391.25,
"per_frame_generation": null
},
"denoise_step_ms": {
"0": 1744.42,
"1": 817.18,
"2": 854.8,
"3": 836.55,
"4": 808.83,
"5": 809.81,
"6": 796.47,
"7": 767.67,
"8": 802.66,
"9": 805.93,
"10": 808.6,
"11": 820.85,
"12": 846.88,
"13": 852.86,
"14": 844.04,
"15": 833.44,
"16": 803.25,
"17": 807.18,
"18": 815.48,
"19": 811.04,
"20": 804.2,
"21": 781.2,
"22": 767.35,
"23": 772.59,
"24": 785.54,
"25": 770.46,
"26": 779.59,
"27": 817.82,
"28": 806.95,
"29": 798.13,
"30": 222.32,
"31": 214.83,
"32": 222.11
},
"expected_e2e_ms": 34384.39,
"expected_avg_denoise_ms": 782.76,
"expected_median_denoise_ms": 806.95
} }
} }
} }
@@ -10,10 +10,10 @@ from sglang.multimodal_gen.test.server.accuracy_utils import (
run_text_encoder_accuracy_case, run_text_encoder_accuracy_case,
) )
from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine
from sglang.multimodal_gen.test.server.testcase_configs import ONE_GPU_CASES_A from sglang.multimodal_gen.test.server.testcase_configs import ACCURACY_ONE_GPU_CASES_A
@pytest.mark.parametrize("case", ONE_GPU_CASES_A, ids=lambda x: x.id) @pytest.mark.parametrize("case", ACCURACY_ONE_GPU_CASES_A, ids=lambda x: x.id)
class TestAccuracy1GPU_A: class TestAccuracy1GPU_A:
"""1-GPU Component Accuracy Suite (Set A).""" """1-GPU Component Accuracy Suite (Set A)."""
@@ -10,10 +10,10 @@ from sglang.multimodal_gen.test.server.accuracy_utils import (
run_text_encoder_accuracy_case, run_text_encoder_accuracy_case,
) )
from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine
from sglang.multimodal_gen.test.server.testcase_configs import ONE_GPU_CASES_B from sglang.multimodal_gen.test.server.testcase_configs import ACCURACY_ONE_GPU_CASES_B
@pytest.mark.parametrize("case", ONE_GPU_CASES_B, ids=lambda x: x.id) @pytest.mark.parametrize("case", ACCURACY_ONE_GPU_CASES_B, ids=lambda x: x.id)
class TestAccuracy1GPU_B: class TestAccuracy1GPU_B:
"""1-GPU Component Accuracy Suite (Set B).""" """1-GPU Component Accuracy Suite (Set B)."""
@@ -10,10 +10,10 @@ from sglang.multimodal_gen.test.server.accuracy_utils import (
run_text_encoder_accuracy_case, run_text_encoder_accuracy_case,
) )
from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine
from sglang.multimodal_gen.test.server.testcase_configs import TWO_GPU_CASES_A from sglang.multimodal_gen.test.server.testcase_configs import ACCURACY_TWO_GPU_CASES_A
@pytest.mark.parametrize("case", TWO_GPU_CASES_A, ids=lambda x: x.id) @pytest.mark.parametrize("case", ACCURACY_TWO_GPU_CASES_A, ids=lambda x: x.id)
class TestAccuracy2GPU_A: class TestAccuracy2GPU_A:
"""2-GPU Component Accuracy Suite (Set A).""" """2-GPU Component Accuracy Suite (Set A)."""
@@ -10,10 +10,10 @@ from sglang.multimodal_gen.test.server.accuracy_utils import (
run_text_encoder_accuracy_case, run_text_encoder_accuracy_case,
) )
from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine
from sglang.multimodal_gen.test.server.testcase_configs import TWO_GPU_CASES_B from sglang.multimodal_gen.test.server.testcase_configs import ACCURACY_TWO_GPU_CASES_B
@pytest.mark.parametrize("case", TWO_GPU_CASES_B, ids=lambda x: x.id) @pytest.mark.parametrize("case", ACCURACY_TWO_GPU_CASES_B, ids=lambda x: x.id)
class TestAccuracy2GPU_B: class TestAccuracy2GPU_B:
"""2-GPU Component Accuracy Suite (Set B).""" """2-GPU Component Accuracy Suite (Set B)."""
@@ -953,6 +953,16 @@ TWO_GPU_CASES_B = [
), ),
TI2V_sampling_params, TI2V_sampling_params,
), ),
DiffusionTestCase(
"ltx_2.3_two_stage_t2v_2gpus",
DiffusionServerArgs(
model_path="Lightricks/LTX-2.3",
modality="video",
num_gpus=2,
extras=["--pipeline-class-name LTX2TwoStagePipeline"],
),
T2V_sampling_params,
),
# I2V LoRA test case # I2V LoRA test case
DiffusionTestCase( DiffusionTestCase(
"wan2_1_i2v_14b_lora_2gpu", "wan2_1_i2v_14b_lora_2gpu",
@@ -1065,6 +1075,90 @@ if not current_platform.is_hip():
) )
) )
def _select_accuracy_cases(
cases: list[DiffusionTestCase], enabled_ids: tuple[str, ...]
) -> list[DiffusionTestCase]:
enabled = set(enabled_ids)
return [case for case in cases if case.id in enabled]
ACCURACY_ONE_GPU_CASES_A_IDS = (
"qwen_image_t2i",
"qwen_image_t2i_cache_dit_enabled",
"flux_image_t2i",
"flux_2_image_t2i",
"flux_2_klein_image_t2i",
"layerwise_offload",
"zimage_image_t2i",
"zimage_image_t2i_fp8",
"zimage_image_t2i_multi_lora",
"qwen_image_edit_ti2i",
"qwen_image_edit_2509_ti2i",
"qwen_image_edit_2511_ti2i",
"qwen_image_layered_i2i",
"flux_2_image_t2i_upscaling_4x",
"mova_360p_1gpu",
)
ACCURACY_ONE_GPU_CASES_B_IDS = (
"wan2_1_t2v_1.3b",
"wan2_1_t2v_1.3b_text_encoder_cpu_offload",
"wan2_1_t2v_1.3b_teacache_enabled",
"wan2_1_t2v_1.3b_frame_interp_2x",
"wan2_1_t2v_1.3b_upscaling_4x",
"wan2_1_t2v_1.3b_frame_interp_2x_upscaling_4x",
"wan2_1_t2v_1_3b_lora_1gpu",
"flux_2_ti2i",
"flux_2_t2i_customized_vae_path",
"fast_hunyuan_video",
"wan2_2_ti2v_5b",
"fastwan2_2_ti2v_5b",
"hunyuan3d_shape_gen",
"turbo_wan2_1_t2v_1.3b",
"flux_2_nvfp4_t2i",
"flux_2_ti2i_multi_image_cache_dit",
)
ACCURACY_TWO_GPU_CASES_A_IDS = (
"wan2_2_i2v_a14b_2gpu",
"wan2_2_t2v_a14b_2gpu",
"wan2_2_t2v_a14b_teacache_2gpu",
"wan2_2_t2v_a14b_lora_2gpu",
"wan2_1_t2v_14b_2gpu",
"wan2_1_t2v_1.3b_cfg_parallel",
"fsdp-inference",
"mova_360p_tp2",
"mova_360p_ring1_uly2",
"mova_360p_ring2_uly1",
"ltx_2_two_stage_t2v",
)
ACCURACY_TWO_GPU_CASES_B_IDS = (
"wan2_1_i2v_14b_480P_2gpu",
"wan2_1_i2v_14b_lora_2gpu",
"wan2_1_i2v_14b_720P_2gpu",
"qwen_image_t2i_2_gpus",
"zimage_image_t2i_2_gpus",
"zimage_image_t2i_2_gpus_non_square",
"flux_image_t2i_2_gpus",
"flux_2_image_t2i_2_gpus",
"flux_2_klein_ti2i_2_gpus",
)
ACCURACY_ONE_GPU_CASES_A = _select_accuracy_cases(
ONE_GPU_CASES_A, ACCURACY_ONE_GPU_CASES_A_IDS
)
ACCURACY_ONE_GPU_CASES_B = _select_accuracy_cases(
ONE_GPU_CASES_B, ACCURACY_ONE_GPU_CASES_B_IDS
)
ACCURACY_TWO_GPU_CASES_A = _select_accuracy_cases(
TWO_GPU_CASES_A, ACCURACY_TWO_GPU_CASES_A_IDS
)
ACCURACY_TWO_GPU_CASES_B = _select_accuracy_cases(
TWO_GPU_CASES_B, ACCURACY_TWO_GPU_CASES_B_IDS
)
# Load global configuration # Load global configuration
BASELINE_CONFIG = BaselineConfig.load( BASELINE_CONFIG = BaselineConfig.load(
Path(__file__).with_name("perf_baselines.json") Path(__file__).with_name("perf_baselines.json")
@@ -1,392 +0,0 @@
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)