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