Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3572c6821e | ||
|
|
deb901f6fc | ||
|
|
beae2943cf | ||
|
|
fadb71bb64 | ||
|
|
65c6fbaa46 | ||
|
|
e32a8a9504 | ||
|
|
6907d87871 | ||
|
|
9775fed31a | ||
|
|
5c6f635d73 | ||
|
|
c19708fd58 | ||
|
|
b2e4fb0743 | ||
|
|
a3ad4852b0 | ||
|
|
add2be21b5 | ||
|
|
41203d92b8 | ||
|
|
5fa8415c0b | ||
|
|
3a182925f3 | ||
|
|
c1e4787775 | ||
|
|
2c6bf47b9f | ||
|
|
548cc08817 | ||
|
|
7521b06693 | ||
|
|
fb9ad77086 |
@@ -0,0 +1,34 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -2,5 +2,16 @@ from fastvideo.configs.models.base import ModelConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig)
|
||||
|
||||
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
__all__ = [
|
||||
"ModelConfig",
|
||||
"VAEConfig",
|
||||
"DiTConfig",
|
||||
"EncoderConfig",
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.configs.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 audio VAE and vocoder configuration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioArchConfig(ArchConfig):
|
||||
architectures: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioEncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioEncoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioDecoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioDecoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VocoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2Vocoder"]))
|
||||
@@ -3,11 +3,12 @@ from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
|
||||
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
|
||||
"LongCatVideoConfig"
|
||||
"LongCatVideoConfig", "LTX2VideoConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 Transformer configuration for native FastVideo integration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_ltx2_blocks(name: str, _module) -> bool:
|
||||
"""FSDP shard condition for LTX-2 transformer blocks."""
|
||||
return "transformer_blocks" in name
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VideoArchConfig(DiTArchConfig):
|
||||
"""Architecture configuration for LTX-2 video transformer."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_ltx2_blocks])
|
||||
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
|
||||
|
||||
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo)
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^model\.diffusion_model\.(.*)$": r"model.\1",
|
||||
r"^diffusion_model\.(.*)$": r"model.\1",
|
||||
r"^model\.(.*)$": r"model.\1",
|
||||
r"^(.*)$": r"model.\1",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
lora_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Core transformer settings (defaults from LTX-2 metadata)
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
num_layers: int = 48
|
||||
cross_attention_dim: int = 4096
|
||||
caption_channels: int = 3840
|
||||
norm_eps: float = 1e-6
|
||||
attention_type: str = "default"
|
||||
rope_type: str = "split"
|
||||
double_precision_rope: bool = True
|
||||
|
||||
positional_embedding_theta: float = 10000.0
|
||||
positional_embedding_max_pos: list[int] = field(
|
||||
default_factory=lambda: [20, 2048, 2048])
|
||||
timestep_scale_multiplier: int = 1000
|
||||
use_middle_indices_grid: bool = True
|
||||
|
||||
# Patchification (video-only path)
|
||||
patch_size: tuple[int, int, int] = (1, 1, 1)
|
||||
num_channels_latents: int = 128
|
||||
in_channels: int | None = None
|
||||
out_channels: int | None = None
|
||||
|
||||
# Audio defaults (reserved for joint AV ports)
|
||||
audio_num_attention_heads: int = 32
|
||||
audio_attention_head_dim: int = 64
|
||||
audio_in_channels: int = 128
|
||||
audio_out_channels: int = 128
|
||||
audio_cross_attention_dim: int = 2048
|
||||
audio_positional_embedding_max_pos: list[int] = field(
|
||||
default_factory=lambda: [20])
|
||||
av_ca_timestep_scale_multiplier: int = 1
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
patch_volume = self.patch_size[0] * self.patch_size[
|
||||
1] * self.patch_size[2]
|
||||
if self.in_channels is None:
|
||||
self.in_channels = self.num_channels_latents * patch_volume
|
||||
if self.out_channels is None:
|
||||
self.out_channels = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VideoConfig(DiTConfig):
|
||||
"""Main configuration for LTX-2 transformer."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
|
||||
prefix: str = "ltx2"
|
||||
@@ -8,10 +8,11 @@ from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
|
||||
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config"
|
||||
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2GemmaArchConfig(TextEncoderArchConfig):
|
||||
architectures: list[str] = field(
|
||||
default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
|
||||
hidden_size: int = 3840
|
||||
num_hidden_layers: int = 48
|
||||
num_attention_heads: int = 30
|
||||
text_len: int = 1024
|
||||
pad_token_id: int = 0
|
||||
eos_token_id: int = 2
|
||||
|
||||
gemma_model_path: str = ""
|
||||
gemma_dtype: str = "bfloat16"
|
||||
padding_side: str = "left"
|
||||
|
||||
feature_extractor_in_features: int = 3840 * 49
|
||||
feature_extractor_out_features: int = 3840
|
||||
|
||||
connector_num_attention_heads: int = 30
|
||||
connector_attention_head_dim: int = 128
|
||||
connector_num_layers: int = 2
|
||||
connector_positional_embedding_theta: float = 10000.0
|
||||
connector_positional_embedding_max_pos: list[int] = field(
|
||||
default_factory=lambda: [4096])
|
||||
connector_rope_type: str = "split"
|
||||
connector_double_precision_rope: bool = False
|
||||
connector_num_learnable_registers: int | None = 128
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.tokenizer_kwargs["padding"] = "max_length"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2GemmaConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=LTX2GemmaArchConfig)
|
||||
|
||||
prefix: str = "ltx2_gemma"
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
@@ -12,4 +13,5 @@ __all__ = [
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 VAE configuration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VAEArchConfig(VAEArchConfig):
|
||||
# Mirrors LTX-2 safetensors metadata config under "vae"
|
||||
_class_name: str = "CausalVideoAutoencoder"
|
||||
dims: int = 3
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 128
|
||||
encoder_blocks: list = field(default_factory=list)
|
||||
decoder_blocks: list = field(default_factory=list)
|
||||
patch_size: int = 4
|
||||
norm_layer: str = "pixel_norm"
|
||||
latent_log_var: str = "uniform"
|
||||
encoder_spatial_padding_mode: str = "zeros"
|
||||
decoder_spatial_padding_mode: str = "reflect"
|
||||
causal_decoder: bool = False
|
||||
timestep_conditioning: bool = True
|
||||
use_quant_conv: bool = False
|
||||
scaling_factor: float = 1.0
|
||||
normalize_latent_channels: bool = False
|
||||
|
||||
# Match FastVideo naming for compression ratios (LTX-2 default)
|
||||
temporal_compression_ratio: int = 8
|
||||
spatial_compression_ratio: int = 32
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
|
||||
|
||||
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
|
||||
ltx2_spatial_tile_size_in_pixels: int = 512
|
||||
ltx2_spatial_tile_overlap_in_pixels: int = 64
|
||||
ltx2_temporal_tile_size_in_frames: int = 64
|
||||
ltx2_temporal_tile_overlap_in_frames: int = 24
|
||||
@@ -4,6 +4,7 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
@@ -16,5 +17,6 @@ __all__ = [
|
||||
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
|
||||
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
LTX2AudioDecoderConfig, LTX2VocoderConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.configs.models.dits import LTX2VideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, LTX2GemmaConfig
|
||||
from fastvideo.configs.models.vaes import LTX2VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def ltx2_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2T2VConfig(PipelineConfig):
|
||||
"""Configuration for LTX-2 T2V pipeline."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=LTX2VideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=LTX2VAEConfig)
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LTX2GemmaConfig(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(ltx2_postprocess_text, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", ))
|
||||
|
||||
audio_decoder_config: ModelConfig = field(
|
||||
default_factory=LTX2AudioDecoderConfig)
|
||||
vocoder_config: ModelConfig = field(default_factory=LTX2VocoderConfig)
|
||||
audio_decoder_precision: str = "bf16"
|
||||
vocoder_precision: str = "bf16"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -9,6 +9,7 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.configs.pipelines.turbodiffusion import (
|
||||
@@ -64,6 +65,9 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
|
||||
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
|
||||
# LTX-2 models
|
||||
"Lightricks/LTX-2": LTX2T2VConfig,
|
||||
"converted/ltx2_diffusers": LTX2T2VConfig,
|
||||
# TurboDiffusion models
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers": TurboDiffusionT2V_1_3B_Config,
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers": TurboDiffusionT2V_14B_Config,
|
||||
@@ -102,6 +106,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
lambda id: "cosmos25" in id.lower(),
|
||||
"turbodiffusion":
|
||||
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
|
||||
"ltx2":
|
||||
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -123,6 +129,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig,
|
||||
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
|
||||
"ltx2": LTX2T2VConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2SamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled T2V.
|
||||
"""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 1024
|
||||
width: int = 1536
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 8
|
||||
guidance_scale: float = 1.0
|
||||
# No default negative_prompt for distilled models
|
||||
negative_prompt: str = ""
|
||||
@@ -10,6 +10,7 @@ from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
|
||||
from fastvideo.configs.sample.cosmos2_5 import Cosmos_Predict2_5_2B_Diffusers_SamplingParam
|
||||
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
@@ -40,48 +41,36 @@ from fastvideo.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers":
|
||||
FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo":
|
||||
HunyuanSamplingParam,
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
|
||||
Hunyuan15_480P_SamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
|
||||
Hunyuan15_720P_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers":
|
||||
StepVideoT2VSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
|
||||
# Wan2.1
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
|
||||
WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
|
||||
# Wan2.2
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
|
||||
|
||||
# FastWan2.1
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
FastWanT2V480P_SamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480P_SamplingParam,
|
||||
|
||||
# FastWan2.2
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
|
||||
@@ -102,12 +91,9 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
|
||||
|
||||
# MatrixGame2.0 models
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers":
|
||||
MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers":
|
||||
MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers":
|
||||
MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
|
||||
|
||||
# TurboDiffusion models
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers":
|
||||
@@ -117,6 +103,10 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers":
|
||||
TurboDiffusionI2V_A14B_SamplingParam,
|
||||
|
||||
# LTX-2 models
|
||||
"Lightricks/LTX-2": LTX2SamplingParam,
|
||||
"FastVideo/LTX2-Distilled-Diffusers": LTX2SamplingParam,
|
||||
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -144,6 +134,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
lambda id: "cosmos2_5" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower() and "2_5" not in id.lower(),
|
||||
"ltx2":
|
||||
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -164,6 +156,7 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
TurboDiffusionT2V_1_3B_SamplingParam, # Default to T2V for fallback
|
||||
"cosmos25": Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
|
||||
"cosmos": Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
"ltx2": LTX2SamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,8 @@ import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -389,6 +391,11 @@ class VideoGenerator:
|
||||
if batch.save_video:
|
||||
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
audio = output_batch.extra.get("audio")
|
||||
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
|
||||
if (audio is not None and audio_sample_rate is not None and
|
||||
not self._mux_audio(output_path, audio, audio_sample_rate)):
|
||||
logger.warning("Audio mux failed; saved video without audio.")
|
||||
|
||||
if batch.return_frames:
|
||||
return frames
|
||||
@@ -396,6 +403,7 @@ class VideoGenerator:
|
||||
return {
|
||||
"samples": samples,
|
||||
"frames": frames,
|
||||
"audio": output_batch.extra.get("audio"),
|
||||
"prompts": prompt,
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
@@ -405,6 +413,98 @@ class VideoGenerator:
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mux_audio(
|
||||
video_path: str,
|
||||
audio: torch.Tensor | np.ndarray,
|
||||
sample_rate: int,
|
||||
) -> bool:
|
||||
"""Mux audio into video using PyAV."""
|
||||
try:
|
||||
import av
|
||||
except ImportError:
|
||||
logger.warning("PyAV not installed; cannot mux audio. "
|
||||
"Install with: pip install av")
|
||||
return False
|
||||
|
||||
if torch.is_tensor(audio):
|
||||
audio_np = audio.detach().cpu().float().numpy()
|
||||
else:
|
||||
audio_np = np.asarray(audio, dtype=np.float32)
|
||||
|
||||
if audio_np.ndim == 1:
|
||||
audio_np = audio_np[:, None]
|
||||
elif audio_np.ndim == 2:
|
||||
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
|
||||
audio_np = audio_np.T
|
||||
else:
|
||||
logger.warning("Unexpected audio shape %s; skipping mux.",
|
||||
audio_np.shape)
|
||||
return False
|
||||
|
||||
audio_np = np.clip(audio_np, -1.0, 1.0)
|
||||
audio_int16 = (audio_np * 32767.0).astype(np.int16)
|
||||
num_channels = audio_int16.shape[1]
|
||||
layout = "stereo" if num_channels == 2 else "mono"
|
||||
|
||||
try:
|
||||
import wave
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
out_path = os.path.join(tmpdir, "muxed.mp4")
|
||||
wav_path = os.path.join(tmpdir, "audio.wav")
|
||||
|
||||
# Write audio to WAV file
|
||||
with wave.open(wav_path, "wb") as wav_file:
|
||||
wav_file.setnchannels(num_channels)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(sample_rate)
|
||||
wav_file.writeframes(audio_int16.tobytes())
|
||||
|
||||
# Open input video and audio
|
||||
input_video = av.open(video_path)
|
||||
input_audio = av.open(wav_path)
|
||||
|
||||
# Create output with both streams
|
||||
output = av.open(out_path, mode="w")
|
||||
|
||||
# Add video stream (copy codec from input)
|
||||
in_video_stream = input_video.streams.video[0]
|
||||
out_video_stream = output.add_stream(
|
||||
codec_name=in_video_stream.codec_context.name,
|
||||
rate=in_video_stream.average_rate,
|
||||
)
|
||||
out_video_stream.width = in_video_stream.width
|
||||
out_video_stream.height = in_video_stream.height
|
||||
out_video_stream.pix_fmt = in_video_stream.pix_fmt
|
||||
|
||||
# Add audio stream (AAC)
|
||||
out_audio_stream = output.add_stream("aac", rate=sample_rate)
|
||||
out_audio_stream.layout = layout
|
||||
|
||||
# Remux video (decode and re-encode to be safe)
|
||||
for frame in input_video.decode(video=0):
|
||||
for packet in out_video_stream.encode(frame):
|
||||
output.mux(packet)
|
||||
for packet in out_video_stream.encode():
|
||||
output.mux(packet)
|
||||
|
||||
# Encode audio
|
||||
for frame in input_audio.decode(audio=0):
|
||||
frame.pts = None # Let encoder assign PTS
|
||||
for packet in out_audio_stream.encode(frame):
|
||||
output.mux(packet)
|
||||
for packet in out_audio_stream.encode():
|
||||
output.mux(packet)
|
||||
|
||||
input_video.close()
|
||||
input_audio.close()
|
||||
output.close()
|
||||
shutil.move(out_path, video_path)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning("Audio mux failed: %s", e)
|
||||
return False
|
||||
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None) -> None:
|
||||
|
||||
@@ -166,6 +166,14 @@ class FastVideoArgs:
|
||||
# Prompt text file for batch processing
|
||||
prompt_txt: str | None = None
|
||||
|
||||
# LTX-2 VAE tiling overrides
|
||||
ltx2_vae_tiling: bool | None = None
|
||||
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
|
||||
ltx2_vae_temporal_tile_size_in_frames: int | None = None
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
|
||||
ltx2_initial_latent_path: str | None = None
|
||||
|
||||
# model paths for correct deallocation
|
||||
model_paths: dict[str, str] = field(default_factory=dict)
|
||||
model_loaded: dict[str, bool] = field(default_factory=lambda: {
|
||||
@@ -203,8 +211,44 @@ class FastVideoArgs:
|
||||
logger.error("Failed to load V-MoBA config from %s: %s",
|
||||
self.moba_config_path, e)
|
||||
raise
|
||||
self._apply_ltx2_vae_overrides()
|
||||
self.check_fastvideo_args()
|
||||
|
||||
def _apply_ltx2_vae_overrides(self) -> None:
|
||||
if self.pipeline_config is None:
|
||||
return
|
||||
vae_config = self.pipeline_config.vae_config
|
||||
has_any = any(value is not None for value in (
|
||||
self.ltx2_vae_spatial_tile_size_in_pixels,
|
||||
self.ltx2_vae_spatial_tile_overlap_in_pixels,
|
||||
self.ltx2_vae_temporal_tile_size_in_frames,
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames,
|
||||
))
|
||||
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
|
||||
"vae_tiling"):
|
||||
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
|
||||
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
|
||||
self.pipeline_config.vae_tiling = True
|
||||
|
||||
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
|
||||
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
|
||||
vae_config.ltx2_spatial_tile_size_in_pixels = (
|
||||
self.ltx2_vae_spatial_tile_size_in_pixels)
|
||||
if hasattr(
|
||||
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
|
||||
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
|
||||
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
|
||||
self.ltx2_vae_spatial_tile_overlap_in_pixels)
|
||||
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
|
||||
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
|
||||
vae_config.ltx2_temporal_tile_size_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_size_in_frames)
|
||||
if hasattr(
|
||||
vae_config, "ltx2_temporal_tile_overlap_in_frames"
|
||||
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
|
||||
vae_config.ltx2_temporal_tile_overlap_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames)
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
# Model and path configuration
|
||||
@@ -325,6 +369,44 @@ class FastVideoArgs:
|
||||
"Path to a text file containing prompts (one per line) for batch processing",
|
||||
)
|
||||
|
||||
# LTX-2 VAE tiling overrides
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-tiling",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.ltx2_vae_tiling,
|
||||
help="Enable LTX-2 VAE tiling overrides.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-spatial-tile-size-in-pixels",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
|
||||
help="LTX-2 VAE spatial tile size in pixels.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-spatial-tile-overlap-in-pixels",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
|
||||
help="LTX-2 VAE spatial tile overlap in pixels.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-temporal-tile-size-in-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
|
||||
help="LTX-2 VAE temporal tile size in frames.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-temporal-tile-overlap-in-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
|
||||
help="LTX-2 VAE temporal tile overlap in frames.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-initial-latent-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_initial_latent_path,
|
||||
help="Path to load/save a precomputed LTX-2 initial latent.",
|
||||
)
|
||||
|
||||
# LoRA parameters (inference-time adapter loading)
|
||||
parser.add_argument(
|
||||
"--lora-path",
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoder,
|
||||
LTX2AudioEncoder,
|
||||
LTX2Vocoder,
|
||||
)
|
||||
|
||||
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,563 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from typing import Iterable
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import Gemma3ForConditionalGeneration
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.dits.ltx2 import (
|
||||
FeedForward,
|
||||
LTXRopeType,
|
||||
apply_ltx_rotary_emb,
|
||||
generate_ltx_freq_grid_np,
|
||||
generate_ltx_freq_grid_pytorch,
|
||||
precompute_ltx_freqs_cis,
|
||||
)
|
||||
from fastvideo.models.loader.weight_utils import default_weight_loader
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _debug_log_line(message: str) -> None:
|
||||
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") != "1":
|
||||
return
|
||||
log_path = os.getenv("LTX2_PIPELINE_DEBUG_PATH", "")
|
||||
if not log_path:
|
||||
return
|
||||
log_dir = os.path.dirname(log_path)
|
||||
if log_dir:
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
with open(log_path, "a", encoding="utf-8") as f:
|
||||
f.write(message + "\n")
|
||||
|
||||
|
||||
def _debug_gemma_log_line(message: str) -> None:
|
||||
log_path = os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", "")
|
||||
if not log_path:
|
||||
return
|
||||
log_dir = os.path.dirname(log_path)
|
||||
if log_dir:
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
with open(log_path, "a", encoding="utf-8") as f:
|
||||
f.write(message + "\n")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GemmaConnectorConfig:
|
||||
num_attention_heads: int
|
||||
attention_head_dim: int
|
||||
num_layers: int
|
||||
positional_embedding_theta: float
|
||||
positional_embedding_max_pos: list[int]
|
||||
rope_type: LTXRopeType
|
||||
double_precision_rope: bool
|
||||
num_learnable_registers: int | None
|
||||
|
||||
|
||||
class GemmaFeaturesExtractorProjLinear(nn.Module):
|
||||
"""Linear projection that aggregates stacked Gemma hidden states."""
|
||||
|
||||
def __init__(self, in_features: int, out_features: int) -> None:
|
||||
super().__init__()
|
||||
self.aggregate_embed = nn.Linear(in_features, out_features, bias=False)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.aggregate_embed(x)
|
||||
|
||||
|
||||
class _BasicTransformerBlock1D(nn.Module):
|
||||
"""1D transformer block for connector processing."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
heads: int,
|
||||
dim_head: int,
|
||||
rope_type: LTXRopeType,
|
||||
norm_eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.attn1 = _GemmaAttention(
|
||||
query_dim=dim,
|
||||
context_dim=None,
|
||||
heads=heads,
|
||||
dim_head=dim_head,
|
||||
norm_eps=norm_eps,
|
||||
rope_type=rope_type,
|
||||
)
|
||||
self.ff = FeedForward(dim, dim_out=dim)
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
norm_hidden_states = torch.nn.functional.rms_norm(
|
||||
hidden_states, (hidden_states.shape[-1],), eps=self.norm_eps
|
||||
)
|
||||
if norm_hidden_states.ndim == 4:
|
||||
norm_hidden_states = norm_hidden_states.squeeze(1)
|
||||
|
||||
attn_output = self.attn1(
|
||||
norm_hidden_states,
|
||||
mask=attention_mask,
|
||||
pe=pe,
|
||||
)
|
||||
hidden_states = attn_output + hidden_states
|
||||
if hidden_states.ndim == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
|
||||
norm_hidden_states = torch.nn.functional.rms_norm(
|
||||
hidden_states, (hidden_states.shape[-1],), eps=self.norm_eps
|
||||
)
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = ff_output + hidden_states
|
||||
if hidden_states.ndim == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class _GemmaAttention(nn.Module):
|
||||
"""Attention implementation aligned with LTX-2 text encoder."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
context_dim: int | None,
|
||||
heads: int,
|
||||
dim_head: int,
|
||||
norm_eps: float,
|
||||
rope_type: LTXRopeType,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = query_dim if context_dim is None else context_dim
|
||||
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.rope_type = rope_type
|
||||
|
||||
self.q_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
|
||||
self.k_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=True)
|
||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=True)
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=True)
|
||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim, bias=True), nn.Identity())
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
context: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
k_pe: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
q = self.to_q(x)
|
||||
context = x if context is None else context
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
if pe is not None:
|
||||
q = apply_ltx_rotary_emb(q, pe, self.rope_type)
|
||||
k = apply_ltx_rotary_emb(k, pe if k_pe is None else k_pe, self.rope_type)
|
||||
|
||||
b, q_len, _ = q.shape
|
||||
k_len = k.shape[1]
|
||||
q = q.view(b, q_len, self.heads, self.dim_head).transpose(1, 2)
|
||||
k = k.view(b, k_len, self.heads, self.dim_head).transpose(1, 2)
|
||||
v = v.view(b, k_len, self.heads, self.dim_head).transpose(1, 2)
|
||||
|
||||
if mask is not None:
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
if mask.ndim == 3:
|
||||
mask = mask.unsqueeze(1)
|
||||
|
||||
out = torch.nn.functional.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
)
|
||||
out = out.transpose(1, 2).reshape(b, q_len, -1)
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Embeddings1DConnector(nn.Module):
|
||||
"""Transformer connector that refines Gemma embeddings for LTX-2."""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
def __init__(self, config: GemmaConnectorConfig) -> None:
|
||||
super().__init__()
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.positional_embedding_theta = config.positional_embedding_theta
|
||||
self.positional_embedding_max_pos = config.positional_embedding_max_pos
|
||||
self.rope_type = config.rope_type
|
||||
self.double_precision_rope = config.double_precision_rope
|
||||
self.transformer_1d_blocks = nn.ModuleList(
|
||||
[
|
||||
_BasicTransformerBlock1D(
|
||||
dim=self.inner_dim,
|
||||
heads=config.num_attention_heads,
|
||||
dim_head=config.attention_head_dim,
|
||||
rope_type=config.rope_type,
|
||||
)
|
||||
for _ in range(config.num_layers)
|
||||
]
|
||||
)
|
||||
self.num_learnable_registers = config.num_learnable_registers
|
||||
if self.num_learnable_registers:
|
||||
self.learnable_registers = nn.Parameter(
|
||||
torch.rand(
|
||||
self.num_learnable_registers,
|
||||
self.inner_dim,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
* 2.0
|
||||
- 1.0
|
||||
)
|
||||
|
||||
def _replace_padded_with_learnable_registers(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
assert hidden_states.shape[1] % self.num_learnable_registers == 0, (
|
||||
f"Hidden states sequence length {hidden_states.shape[1]} must be divisible by "
|
||||
f"num_learnable_registers {self.num_learnable_registers}."
|
||||
)
|
||||
|
||||
num_registers_duplications = (
|
||||
hidden_states.shape[1] // self.num_learnable_registers
|
||||
)
|
||||
learnable_registers = torch.tile(
|
||||
self.learnable_registers, (num_registers_duplications, 1)
|
||||
)
|
||||
attention_mask_binary = (
|
||||
attention_mask.squeeze(1).squeeze(1).unsqueeze(-1) >= -9000.0
|
||||
).int()
|
||||
|
||||
non_zero_hidden_states = hidden_states[
|
||||
:, attention_mask_binary.squeeze().bool(), :
|
||||
]
|
||||
non_zero_nums = non_zero_hidden_states.shape[1]
|
||||
pad_length = hidden_states.shape[1] - non_zero_nums
|
||||
adjusted_hidden_states = torch.nn.functional.pad(
|
||||
non_zero_hidden_states, pad=(0, 0, 0, pad_length), value=0
|
||||
)
|
||||
flipped_mask = torch.flip(attention_mask_binary, dims=[1])
|
||||
hidden_states = flipped_mask * adjusted_hidden_states + (
|
||||
1 - flipped_mask
|
||||
) * learnable_registers
|
||||
|
||||
attention_mask = torch.full_like(
|
||||
attention_mask,
|
||||
0.0,
|
||||
dtype=attention_mask.dtype,
|
||||
device=attention_mask.device,
|
||||
)
|
||||
|
||||
return hidden_states, attention_mask
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if self.num_learnable_registers:
|
||||
hidden_states, attention_mask = (
|
||||
self._replace_padded_with_learnable_registers(
|
||||
hidden_states, attention_mask
|
||||
)
|
||||
)
|
||||
|
||||
indices_grid = torch.arange(
|
||||
hidden_states.shape[1],
|
||||
dtype=torch.float32,
|
||||
device=hidden_states.device,
|
||||
)
|
||||
indices_grid = indices_grid[None, None, :]
|
||||
freq_grid_generator = (
|
||||
generate_ltx_freq_grid_np
|
||||
if self.double_precision_rope
|
||||
else generate_ltx_freq_grid_pytorch
|
||||
)
|
||||
freqs_cis = precompute_ltx_freqs_cis(
|
||||
indices_grid=indices_grid,
|
||||
dim=self.inner_dim,
|
||||
out_dtype=hidden_states.dtype,
|
||||
theta=self.positional_embedding_theta,
|
||||
max_pos=self.positional_embedding_max_pos,
|
||||
num_attention_heads=self.num_attention_heads,
|
||||
rope_type=self.rope_type,
|
||||
freq_grid_generator=freq_grid_generator,
|
||||
)
|
||||
|
||||
for block in self.transformer_1d_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states, attention_mask=attention_mask, pe=freqs_cis
|
||||
)
|
||||
|
||||
hidden_states = torch.nn.functional.rms_norm(
|
||||
hidden_states, (hidden_states.shape[-1],), eps=1e-6
|
||||
)
|
||||
|
||||
return hidden_states, attention_mask
|
||||
|
||||
|
||||
class LTX2GemmaTextEncoderModel(TextEncoder):
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__(config)
|
||||
arch = config.arch_config
|
||||
|
||||
self.feature_extractor_linear = GemmaFeaturesExtractorProjLinear(
|
||||
in_features=arch.feature_extractor_in_features,
|
||||
out_features=arch.feature_extractor_out_features,
|
||||
)
|
||||
|
||||
connector_config = GemmaConnectorConfig(
|
||||
num_attention_heads=arch.connector_num_attention_heads,
|
||||
attention_head_dim=arch.connector_attention_head_dim,
|
||||
num_layers=arch.connector_num_layers,
|
||||
positional_embedding_theta=arch.connector_positional_embedding_theta,
|
||||
positional_embedding_max_pos=arch.connector_positional_embedding_max_pos,
|
||||
rope_type=LTXRopeType(arch.connector_rope_type),
|
||||
double_precision_rope=arch.connector_double_precision_rope,
|
||||
num_learnable_registers=arch.connector_num_learnable_registers,
|
||||
)
|
||||
self.embeddings_connector = Embeddings1DConnector(connector_config)
|
||||
self.audio_embeddings_connector = Embeddings1DConnector(connector_config)
|
||||
|
||||
self.gemma_model_path = arch.gemma_model_path
|
||||
self.gemma_dtype = arch.gemma_dtype
|
||||
self.padding_side = arch.padding_side
|
||||
self._gemma_model: Gemma3ForConditionalGeneration | None = None
|
||||
|
||||
def named_parameters(self, prefix: str = "", recurse: bool = True):
|
||||
for name, param in super().named_parameters(
|
||||
prefix=prefix, recurse=recurse
|
||||
):
|
||||
if name.startswith("gemma_model."):
|
||||
continue
|
||||
yield name, param
|
||||
|
||||
@property
|
||||
def gemma_model(self) -> Gemma3ForConditionalGeneration:
|
||||
if self._gemma_model is None:
|
||||
gemma_path = self.gemma_model_path
|
||||
if not gemma_path:
|
||||
raise ValueError(
|
||||
"gemma_model_path must be set (expected text_encoder/gemma)."
|
||||
)
|
||||
dtype = getattr(torch, self.gemma_dtype, torch.bfloat16)
|
||||
self._gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
|
||||
gemma_path,
|
||||
local_files_only=True,
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
# Configure model-level attention implementation when using TORCH_SDPA.
|
||||
# Note: torch.backends.cuda.enable_*_sdp() settings should be configured
|
||||
# at application/pipeline initialization level, not here, to avoid
|
||||
# unexpected side effects across the application.
|
||||
if os.getenv("FASTVIDEO_ATTENTION_BACKEND") == "TORCH_SDPA":
|
||||
if hasattr(self._gemma_model.config, "attn_implementation"):
|
||||
self._gemma_model.config.attn_implementation = "sdpa"
|
||||
if hasattr(self._gemma_model.config, "_attn_implementation"):
|
||||
self._gemma_model.config._attn_implementation = "sdpa"
|
||||
device = next(self.feature_extractor_linear.parameters()).device
|
||||
self._gemma_model.to(device=device)
|
||||
self._gemma_model.eval()
|
||||
return self._gemma_model
|
||||
|
||||
def _run_feature_extractor(
|
||||
self,
|
||||
hidden_states: tuple[torch.Tensor, ...],
|
||||
attention_mask: torch.Tensor,
|
||||
padding_side: str,
|
||||
) -> torch.Tensor:
|
||||
encoded_text_features = torch.stack(hidden_states, dim=-1)
|
||||
if os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", ""):
|
||||
for idx, layer in enumerate(hidden_states):
|
||||
_debug_gemma_log_line(
|
||||
f"fastvideo:gemma_hidden_state_{idx}"
|
||||
f":sum={layer.float().sum().item():.6f}"
|
||||
)
|
||||
_debug_gemma_log_line(
|
||||
"fastvideo:gemma_hidden_states_stack"
|
||||
f":sum={encoded_text_features.float().sum().item():.6f}"
|
||||
)
|
||||
encoded_text_features_dtype = encoded_text_features.dtype
|
||||
sequence_lengths = attention_mask.sum(dim=-1)
|
||||
normed_text_features = _norm_and_concat_padded_batch(
|
||||
encoded_text_features, sequence_lengths, padding_side=padding_side
|
||||
)
|
||||
return self.feature_extractor_linear(
|
||||
normed_text_features.to(encoded_text_features_dtype)
|
||||
)
|
||||
|
||||
def _convert_to_additive_mask(
|
||||
self, attention_mask: torch.Tensor, dtype: torch.dtype
|
||||
) -> torch.Tensor:
|
||||
return (attention_mask - 1).to(dtype).reshape(
|
||||
(attention_mask.shape[0], 1, -1, attention_mask.shape[-1])
|
||||
) * torch.finfo(dtype).max
|
||||
|
||||
def _run_connectors(
|
||||
self,
|
||||
encoded_input: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
connector_attention_mask = self._convert_to_additive_mask(
|
||||
attention_mask, encoded_input.dtype
|
||||
)
|
||||
encoded, encoded_connector_attention_mask = self.embeddings_connector(
|
||||
encoded_input, connector_attention_mask
|
||||
)
|
||||
|
||||
attention_mask = (encoded_connector_attention_mask < 0.000001).to(
|
||||
torch.int64
|
||||
)
|
||||
attention_mask = attention_mask.reshape(
|
||||
[encoded.shape[0], encoded.shape[1], 1]
|
||||
)
|
||||
encoded = encoded * attention_mask
|
||||
|
||||
encoded_for_audio, _ = self.audio_embeddings_connector(
|
||||
encoded_input, connector_attention_mask
|
||||
)
|
||||
|
||||
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
if input_ids is None:
|
||||
raise ValueError("input_ids is required for Gemma text encoding.")
|
||||
if attention_mask is None:
|
||||
attention_mask = torch.ones_like(input_ids)
|
||||
|
||||
model = self.gemma_model
|
||||
input_ids = input_ids.to(device=model.device)
|
||||
attention_mask = attention_mask.to(device=model.device)
|
||||
outputs = model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
|
||||
encoded_inputs = self._run_feature_extractor(
|
||||
outputs.hidden_states,
|
||||
attention_mask,
|
||||
padding_side=self.padding_side,
|
||||
)
|
||||
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
|
||||
_debug_log_line(
|
||||
"fastvideo:gemma_feature"
|
||||
f":sum={encoded_inputs.float().sum().item():.6f} "
|
||||
f"shape={tuple(encoded_inputs.shape)}"
|
||||
)
|
||||
video_encoding, audio_encoding, attention_mask = self._run_connectors(
|
||||
encoded_inputs, attention_mask
|
||||
)
|
||||
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
|
||||
_debug_log_line(
|
||||
"fastvideo:gemma_video_encoding"
|
||||
f":sum={video_encoding.float().sum().item():.6f} "
|
||||
f"shape={tuple(video_encoding.shape)}"
|
||||
)
|
||||
_debug_log_line(
|
||||
"fastvideo:gemma_audio_encoding"
|
||||
f":sum={audio_encoding.float().sum().item():.6f} "
|
||||
f"shape={tuple(audio_encoding.shape)}"
|
||||
)
|
||||
|
||||
hidden_states = (audio_encoding, ) if output_hidden_states else None
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=video_encoding,
|
||||
hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
def load_weights(
|
||||
self, weights: Iterable[tuple[str, torch.Tensor]]
|
||||
) -> set[str]:
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if name == "aggregate_embed.weight":
|
||||
name = "feature_extractor_linear.aggregate_embed.weight"
|
||||
if name not in params_dict:
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
|
||||
|
||||
def _norm_and_concat_padded_batch(
|
||||
encoded_text: torch.Tensor,
|
||||
sequence_lengths: torch.Tensor,
|
||||
padding_side: str = "right",
|
||||
) -> torch.Tensor:
|
||||
b, t, d, l = encoded_text.shape
|
||||
device = encoded_text.device
|
||||
|
||||
token_indices = torch.arange(t, device=device)[None, :]
|
||||
if padding_side == "right":
|
||||
mask = token_indices < sequence_lengths[:, None]
|
||||
elif padding_side == "left":
|
||||
start_indices = t - sequence_lengths[:, None]
|
||||
mask = token_indices >= start_indices
|
||||
else:
|
||||
raise ValueError(
|
||||
f"padding_side must be 'left' or 'right', got {padding_side}"
|
||||
)
|
||||
|
||||
mask = mask.reshape(b, t, 1, 1)
|
||||
eps = 1e-6
|
||||
|
||||
masked = encoded_text.masked_fill(~mask, 0.0)
|
||||
denom = (sequence_lengths * d).view(b, 1, 1, 1)
|
||||
mean = masked.sum(dim=(1, 2), keepdim=True) / (denom + eps)
|
||||
|
||||
x_min = encoded_text.masked_fill(~mask, float("inf")).amin(
|
||||
dim=(1, 2), keepdim=True
|
||||
)
|
||||
x_max = encoded_text.masked_fill(~mask, float("-inf")).amax(
|
||||
dim=(1, 2), keepdim=True
|
||||
)
|
||||
range_ = x_max - x_min
|
||||
|
||||
normed = 8 * (encoded_text - mean) / (range_ + eps)
|
||||
normed = normed.reshape(b, t, -1)
|
||||
|
||||
mask_flattened = mask.reshape(b, t, 1).expand(-1, -1, d * l)
|
||||
normed = normed.masked_fill(~mask_flattened, 0.0)
|
||||
return normed
|
||||
@@ -80,6 +80,9 @@ class ComponentLoader(ABC):
|
||||
"transformer": (TransformerLoader, "diffusers"),
|
||||
"transformer_2": (TransformerLoader, "diffusers"),
|
||||
"vae": (VAELoader, "diffusers"),
|
||||
"audio_vae": (AudioDecoderLoader, "diffusers"),
|
||||
"audio_decoder": (AudioDecoderLoader, "diffusers"),
|
||||
"vocoder": (VocoderLoader, "diffusers"),
|
||||
"text_encoder": (TextEncoderLoader, "transformers"),
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
@@ -242,6 +245,47 @@ class TextEncoderLoader(ComponentLoader):
|
||||
model_config.pop("model_type", None)
|
||||
model_config.pop("tokenizer_class", None)
|
||||
model_config.pop("torch_dtype", None)
|
||||
repo_root = os.path.dirname(model_path)
|
||||
index_path = os.path.join(repo_root, "model_index.json")
|
||||
gemma_path = ""
|
||||
gemma_path_from_candidate = False
|
||||
if os.path.isfile(index_path):
|
||||
try:
|
||||
with open(index_path, encoding="utf-8") as f:
|
||||
model_index = json.load(f)
|
||||
gemma_path = model_index.get("gemma_model_path", "")
|
||||
except json.JSONDecodeError:
|
||||
gemma_path = ""
|
||||
if not gemma_path:
|
||||
candidate = os.path.normpath(os.path.join(model_path, "gemma"))
|
||||
if os.path.isdir(candidate):
|
||||
gemma_path = candidate
|
||||
gemma_path_from_candidate = True
|
||||
model_config["gemma_model_path"] = gemma_path
|
||||
if gemma_path and not gemma_path_from_candidate:
|
||||
if not os.path.isabs(gemma_path):
|
||||
model_config["gemma_model_path"] = os.path.normpath(
|
||||
os.path.join(repo_root, gemma_path)
|
||||
)
|
||||
transformer_config_path = os.path.join(
|
||||
repo_root, "transformer", "config.json"
|
||||
)
|
||||
if os.path.isfile(transformer_config_path):
|
||||
try:
|
||||
with open(transformer_config_path, encoding="utf-8") as f:
|
||||
transformer_config = json.load(f)
|
||||
if (
|
||||
"connector_double_precision_rope" not in model_config
|
||||
or not model_config["connector_double_precision_rope"]
|
||||
):
|
||||
if transformer_config.get("double_precision_rope") is True:
|
||||
model_config["connector_double_precision_rope"] = True
|
||||
if "connector_rope_type" not in model_config:
|
||||
rope_type = transformer_config.get("rope_type")
|
||||
if rope_type is not None:
|
||||
model_config["connector_rope_type"] = rope_type
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
# @TODO(Wei): Better way to handle this?
|
||||
@@ -489,8 +533,20 @@ class TokenizerLoader(ComponentLoader):
|
||||
# in v0, this was same string as encoder_name "ClipTextModel"
|
||||
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
|
||||
# other method of config?
|
||||
padding_size="right",
|
||||
)
|
||||
padding_side = None
|
||||
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
|
||||
try:
|
||||
arch_config = fastvideo_args.pipeline_config.text_encoder_configs[
|
||||
0
|
||||
].arch_config
|
||||
padding_side = getattr(arch_config, "padding_side", None)
|
||||
except Exception:
|
||||
padding_side = None
|
||||
if padding_side:
|
||||
tokenizer.padding_side = padding_side
|
||||
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
logger.info("Loaded tokenizer: %s", tokenizer.__class__.__name__)
|
||||
return tokenizer
|
||||
|
||||
@@ -501,15 +557,12 @@ class VAELoader(ComponentLoader):
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
"""Load the VAE based on the model path, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
class_name = config.pop("_class_name")
|
||||
class_name = config.get("_class_name")
|
||||
assert class_name is not None, (
|
||||
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
)
|
||||
fastvideo_args.model_paths["vae"] = model_path
|
||||
|
||||
vae_config = fastvideo_args.pipeline_config.vae_config
|
||||
vae_config.update_model_arch(config)
|
||||
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
@@ -543,8 +596,29 @@ class VAELoader(ComponentLoader):
|
||||
vae.load_state_dict(sd, strict=False)
|
||||
return vae.eval()
|
||||
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(target_device)
|
||||
# LTX-2 uses CausalVideoAutoencoder with nested "vae" config
|
||||
if class_name == "CausalVideoAutoencoder" and "vae" in config:
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(config).to(target_device)
|
||||
if hasattr(vae, "set_tiling_config"):
|
||||
vae_config = fastvideo_args.pipeline_config.vae_config
|
||||
vae.set_tiling_config(
|
||||
spatial_tile_size_in_pixels=getattr(
|
||||
vae_config, "ltx2_spatial_tile_size_in_pixels", 512),
|
||||
spatial_tile_overlap_in_pixels=getattr(
|
||||
vae_config, "ltx2_spatial_tile_overlap_in_pixels", 64),
|
||||
temporal_tile_size_in_frames=getattr(
|
||||
vae_config, "ltx2_temporal_tile_size_in_frames", 64),
|
||||
temporal_tile_overlap_in_frames=getattr(
|
||||
vae_config,
|
||||
"ltx2_temporal_tile_overlap_in_frames", 24),
|
||||
)
|
||||
else:
|
||||
config.pop("_class_name", None)
|
||||
vae_config = fastvideo_args.pipeline_config.vae_config
|
||||
vae_config.update_model_arch(config)
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(target_device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -553,17 +627,101 @@ class VAELoader(ComponentLoader):
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
# Common case: a single `.safetensors` checkpoint file.
|
||||
# Some models may be sharded into multiple files; in that case we merge.
|
||||
if len(safetensors_list) == 1:
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
else:
|
||||
loaded = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
loaded = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
|
||||
# LTX-2 CausalVideoAutoencoder needs per_channel_statistics remapping
|
||||
if class_name == "CausalVideoAutoencoder" and "vae" in config:
|
||||
per_channel_prefixes = (
|
||||
"per_channel_statistics.",
|
||||
"vae.per_channel_statistics.",
|
||||
)
|
||||
remapped = {}
|
||||
for key, tensor in loaded.items():
|
||||
remapped[key] = tensor
|
||||
for prefix in per_channel_prefixes:
|
||||
if key.startswith(prefix):
|
||||
suffix = key[len(prefix):]
|
||||
remapped.setdefault(
|
||||
f"encoder.per_channel_statistics.{suffix}",
|
||||
tensor,
|
||||
)
|
||||
remapped.setdefault(
|
||||
f"decoder.per_channel_statistics.{suffix}",
|
||||
tensor,
|
||||
)
|
||||
break
|
||||
loaded = remapped
|
||||
|
||||
vae.load_state_dict(loaded, strict=False)
|
||||
|
||||
return vae.eval()
|
||||
|
||||
|
||||
class AudioDecoderLoader(ComponentLoader):
|
||||
"""Loader for LTX-2 audio decoder (audio_vae component)."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
config = get_diffusers_config(model=model_path)
|
||||
class_name = config.pop("_class_name", None) or "LTX2AudioDecoder"
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
target_device = get_local_torch_device()
|
||||
|
||||
precision = getattr(
|
||||
fastvideo_args.pipeline_config, "audio_decoder_precision", "bf16"
|
||||
)
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
|
||||
audio_decoder = model_cls(config).to(target_device)
|
||||
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(model_path), "*.safetensors")
|
||||
)
|
||||
loaded: dict[str, torch.Tensor] = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
|
||||
decoder_state = {}
|
||||
for name, tensor in loaded.items():
|
||||
if name.startswith("decoder."):
|
||||
decoder_state[name.replace("decoder.", "")] = tensor
|
||||
elif name.startswith("per_channel_statistics."):
|
||||
decoder_state[name] = tensor
|
||||
|
||||
target_module = getattr(audio_decoder, "model", audio_decoder)
|
||||
target_module.load_state_dict(decoder_state, strict=False)
|
||||
return audio_decoder.eval()
|
||||
|
||||
|
||||
class VocoderLoader(ComponentLoader):
|
||||
"""Loader for LTX-2 vocoder."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
config = get_diffusers_config(model=model_path)
|
||||
class_name = config.pop("_class_name", None) or "LTX2Vocoder"
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
target_device = get_local_torch_device()
|
||||
|
||||
precision = getattr(
|
||||
fastvideo_args.pipeline_config, "vocoder_precision", "bf16"
|
||||
)
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
|
||||
vocoder = model_cls(config).to(target_device)
|
||||
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(model_path), "*.safetensors")
|
||||
)
|
||||
loaded: dict[str, torch.Tensor] = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
|
||||
target_module = getattr(vocoder, "model", vocoder)
|
||||
target_module.load_state_dict(loaded, strict=False)
|
||||
return vocoder.eval()
|
||||
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
@@ -679,7 +837,18 @@ class TransformerLoader(ComponentLoader):
|
||||
model = model.eval()
|
||||
|
||||
if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload:
|
||||
enable_layerwise_offload(model)
|
||||
# Check if model has nn.ModuleList for layerwise offload compatibility
|
||||
has_module_list = any(
|
||||
isinstance(m, nn.ModuleList) for m in model.children()
|
||||
)
|
||||
if has_module_list:
|
||||
enable_layerwise_offload(model)
|
||||
else:
|
||||
logger.warning(
|
||||
"Layerwise offload requested but model %s does not have "
|
||||
"nn.ModuleList structure. Skipping layerwise offload.",
|
||||
cls_name
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -33,6 +33,7 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -54,6 +55,7 @@ _TEXT_ENCODER_MODELS = {
|
||||
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"Qwen2_5_VLForConditionalGeneration":
|
||||
("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
@@ -67,7 +69,14 @@ _VAE_MODELS = {
|
||||
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
|
||||
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
|
||||
}
|
||||
|
||||
_AUDIO_MODELS = {
|
||||
"LTX2AudioEncoder": ("audio", "ltx2_audio_vae", "LTX2AudioEncoder"),
|
||||
"LTX2AudioDecoder": ("audio", "ltx2_audio_vae", "LTX2AudioDecoder"),
|
||||
"LTX2Vocoder": ("audio", "ltx2_audio_vae", "LTX2Vocoder"),
|
||||
}
|
||||
|
||||
_SCHEDULERS = {
|
||||
@@ -91,6 +100,7 @@ _FAST_VIDEO_MODELS = {
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
**_AUDIO_MODELS,
|
||||
**_SCHEDULERS,
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,150 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 text-to-video pipeline.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (DecodingStage, InputValidationStage,
|
||||
LTX2AudioDecodingStage,
|
||||
LTX2DenoisingStage,
|
||||
LTX2LatentPreparationStage,
|
||||
TextEncodingStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LTX2Pipeline(ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"transformer",
|
||||
"vae",
|
||||
"audio_vae",
|
||||
"vocoder",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=LTX2LatentPreparationStage(
|
||||
transformer=self.get_module("transformer"), ),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=LTX2DenoisingStage(
|
||||
transformer=self.get_module("transformer"), ),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="audio_decoding_stage",
|
||||
stage=LTX2AudioDecodingStage(
|
||||
audio_decoder=self.get_module("audio_vae"),
|
||||
vocoder=self.get_module("vocoder"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
tokenizer = self.get_module("tokenizer")
|
||||
if tokenizer is not None:
|
||||
tokenizer.padding_side = "left"
|
||||
if tokenizer.pad_token is None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
def load_modules(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
loaded_modules: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
model_index = self._load_config(self.model_path)
|
||||
logger.info("Loading pipeline modules from config: %s", model_index)
|
||||
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
model_index.pop("workload_type", None)
|
||||
|
||||
if len(model_index) <= 1:
|
||||
raise ValueError(
|
||||
"model_index.json must contain at least one pipeline module")
|
||||
|
||||
required_modules = self.required_config_modules
|
||||
modules: dict[str, Any] = {}
|
||||
|
||||
for module_name, module_spec in model_index.items():
|
||||
if not isinstance(module_spec, list) or len(module_spec) < 1:
|
||||
continue
|
||||
transformers_or_diffusers = module_spec[0]
|
||||
if transformers_or_diffusers is None:
|
||||
if module_name in self.required_config_modules:
|
||||
self.required_config_modules.remove(module_name)
|
||||
continue
|
||||
if module_name not in required_modules:
|
||||
continue
|
||||
if loaded_modules is not None and module_name in loaded_modules:
|
||||
modules[module_name] = loaded_modules[module_name]
|
||||
continue
|
||||
|
||||
component_model_path = os.path.join(self.model_path, module_name)
|
||||
if module_name == "tokenizer" and not os.path.isdir(
|
||||
component_model_path):
|
||||
gemma_path = os.path.join(self.model_path, "text_encoder",
|
||||
"gemma")
|
||||
if os.path.isdir(gemma_path):
|
||||
component_model_path = gemma_path
|
||||
else:
|
||||
raise ValueError(
|
||||
"Tokenizer directory missing and Gemma weights were not found."
|
||||
)
|
||||
|
||||
module = PipelineComponentLoader.load_module(
|
||||
module_name=module_name,
|
||||
component_model_path=component_model_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
logger.info("Loaded module %s from %s", module_name,
|
||||
component_model_path)
|
||||
modules[module_name] = module
|
||||
|
||||
if "tokenizer" in required_modules and "tokenizer" not in modules:
|
||||
gemma_path = os.path.join(self.model_path, "text_encoder", "gemma")
|
||||
if os.path.isdir(gemma_path):
|
||||
modules["tokenizer"] = AutoTokenizer.from_pretrained(
|
||||
gemma_path, local_files_only=True)
|
||||
|
||||
for module_name in required_modules:
|
||||
if module_name not in modules or modules[module_name] is None:
|
||||
raise ValueError(
|
||||
f"Required module {module_name} was not loaded properly")
|
||||
|
||||
return modules
|
||||
|
||||
|
||||
EntryClass = LTX2Pipeline
|
||||
@@ -35,6 +35,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"LongCatPipeline": "longcat",
|
||||
"LongCatImageToVideoPipeline": "longcat",
|
||||
"LongCatVideoContinuationPipeline": "longcat",
|
||||
"LTX2Pipeline": "ltx2",
|
||||
}
|
||||
|
||||
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
|
||||
|
||||
@@ -22,6 +22,10 @@ from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.pipelines.stages.latent_preparation import (
|
||||
Cosmos25LatentPreparationStage, CosmosLatentPreparationStage,
|
||||
LatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.ltx2_audio_decoding import LTX2AudioDecodingStage
|
||||
from fastvideo.pipelines.stages.ltx2_denoising import LTX2DenoisingStage
|
||||
from fastvideo.pipelines.stages.ltx2_latent_preparation import (
|
||||
LTX2LatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.matrixgame_denoising import (
|
||||
MatrixGameCausalDenoisingStage)
|
||||
from fastvideo.pipelines.stages.stepvideo_encoding import (
|
||||
@@ -44,6 +48,8 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"CosmosLatentPreparationStage",
|
||||
"Cosmos25LatentPreparationStage",
|
||||
"LTX2LatentPreparationStage",
|
||||
"LTX2AudioDecodingStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
@@ -51,6 +57,7 @@ __all__ = [
|
||||
"MatrixGameCausalDenoisingStage",
|
||||
"CosmosDenoisingStage",
|
||||
"Cosmos25DenoisingStage",
|
||||
"LTX2DenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Audio decoding stage for LTX-2 pipelines.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.dits.ltx2 import DEFAULT_LTX2_VOCODER_OUTPUT_SAMPLE_RATE
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LTX2AudioDecodingStage(PipelineStage):
|
||||
"""Decode LTX-2 audio latents into a waveform."""
|
||||
|
||||
def __init__(self, audio_decoder, vocoder) -> None:
|
||||
super().__init__()
|
||||
self.audio_decoder = audio_decoder
|
||||
self.vocoder = vocoder
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
audio_latents = batch.extra.get("ltx2_audio_latents")
|
||||
if audio_latents is None:
|
||||
return batch
|
||||
|
||||
device = get_local_torch_device()
|
||||
self.audio_decoder = self.audio_decoder.to(device)
|
||||
self.vocoder = self.vocoder.to(device)
|
||||
audio_latents = audio_latents.to(device)
|
||||
|
||||
disable_autocast = os.getenv("LTX2_DISABLE_AUDIO_AUTOCAST", "1") == "1"
|
||||
with torch.no_grad(), torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=audio_latents.dtype,
|
||||
enabled=not disable_autocast,
|
||||
):
|
||||
decoded_spec = self.audio_decoder(audio_latents)
|
||||
audio_wave = self.vocoder(decoded_spec).squeeze(0).float()
|
||||
|
||||
# Move to CPU for pickling across process boundary
|
||||
batch.extra["audio"] = audio_wave.cpu()
|
||||
batch.extra[
|
||||
"audio_sample_rate"] = DEFAULT_LTX2_VOCODER_OUTPUT_SAMPLE_RATE
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("audio_latents", batch.extra.get("ltx2_audio_latents"),
|
||||
V.none_or_tensor)
|
||||
return result
|
||||
@@ -0,0 +1,308 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 denoising stage using the native sigma schedule.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import os
|
||||
|
||||
import torch
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.ltx2 import (
|
||||
AudioLatentShape, DEFAULT_LTX2_AUDIO_CHANNELS,
|
||||
DEFAULT_LTX2_AUDIO_DOWNSAMPLE, DEFAULT_LTX2_AUDIO_HOP_LENGTH,
|
||||
DEFAULT_LTX2_AUDIO_MEL_BINS, DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
|
||||
VideoLatentShape)
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
BASE_SHIFT_ANCHOR = 1024
|
||||
MAX_SHIFT_ANCHOR = 4096
|
||||
|
||||
# Official distilled sigma schedule (8 denoising steps)
|
||||
# From LTX-2/packages/ltx-pipelines/src/ltx_pipelines/utils/constants.py
|
||||
DISTILLED_SIGMA_VALUES = [
|
||||
1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0
|
||||
]
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _ltx2_sigmas(
|
||||
steps: int,
|
||||
latent: torch.Tensor | None,
|
||||
device: torch.device,
|
||||
max_shift: float = 2.05,
|
||||
base_shift: float = 0.95,
|
||||
stretch: bool = True,
|
||||
terminal: float = 0.1,
|
||||
) -> torch.Tensor:
|
||||
tokens = math.prod(
|
||||
latent.shape[2:]) if latent is not None else MAX_SHIFT_ANCHOR
|
||||
sigmas = torch.linspace(1.0,
|
||||
0.0,
|
||||
steps + 1,
|
||||
device=device,
|
||||
dtype=torch.float32)
|
||||
|
||||
mm = (max_shift - base_shift) / (MAX_SHIFT_ANCHOR - BASE_SHIFT_ANCHOR)
|
||||
b = base_shift - mm * BASE_SHIFT_ANCHOR
|
||||
sigma_shift = tokens * mm + b
|
||||
|
||||
numerator = math.exp(sigma_shift)
|
||||
sigmas = torch.where(
|
||||
sigmas != 0,
|
||||
numerator / (numerator + (1 / sigmas - 1)),
|
||||
torch.zeros_like(sigmas),
|
||||
)
|
||||
|
||||
if stretch:
|
||||
non_zero_mask = sigmas != 0
|
||||
non_zero_sigmas = sigmas[non_zero_mask]
|
||||
one_minus_z = 1.0 - non_zero_sigmas
|
||||
scale_factor = one_minus_z[-1] / (1.0 - terminal)
|
||||
stretched = 1.0 - (one_minus_z / scale_factor)
|
||||
sigmas = sigmas.clone()
|
||||
sigmas[non_zero_mask] = stretched
|
||||
|
||||
return sigmas
|
||||
|
||||
|
||||
class LTX2DenoisingStage(PipelineStage):
|
||||
"""Run the LTX-2 denoising loop over the sigma schedule."""
|
||||
|
||||
def __init__(self, transformer) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if batch.latents is None:
|
||||
raise ValueError("Latents must be provided before denoising.")
|
||||
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
prompt_mask = None
|
||||
|
||||
neg_prompt_embeds = None
|
||||
neg_prompt_mask = None
|
||||
# Only load negative prompts if CFG is actually enabled
|
||||
if batch.do_classifier_free_guidance:
|
||||
assert batch.negative_prompt_embeds is not None, (
|
||||
"CFG is enabled but negative_prompt_embeds is None")
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
|
||||
# Ensure text conditioning is on the same device as latents.
|
||||
if prompt_embeds.device != latents.device:
|
||||
prompt_embeds = prompt_embeds.to(latents.device)
|
||||
if prompt_mask is not None and prompt_mask.device != latents.device:
|
||||
prompt_mask = prompt_mask.to(latents.device)
|
||||
if neg_prompt_embeds is not None and neg_prompt_embeds.device != latents.device:
|
||||
neg_prompt_embeds = neg_prompt_embeds.to(latents.device)
|
||||
if neg_prompt_mask is not None and neg_prompt_mask.device != latents.device:
|
||||
neg_prompt_mask = neg_prompt_mask.to(latents.device)
|
||||
|
||||
target_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
disable_autocast = os.getenv("LTX2_DISABLE_AUTOCAST", "1") == "1"
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast and (
|
||||
not disable_autocast)
|
||||
|
||||
# Use official distilled sigma schedule for 8 steps (distilled models)
|
||||
use_distilled_sigmas = os.getenv("LTX2_USE_DISTILLED_SIGMAS",
|
||||
"1") == "1"
|
||||
if use_distilled_sigmas and batch.num_inference_steps == 8:
|
||||
sigmas = torch.tensor(
|
||||
DISTILLED_SIGMA_VALUES,
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
logger.info("[LTX2] Using official distilled sigma schedule")
|
||||
else:
|
||||
sigmas = _ltx2_sigmas(
|
||||
steps=batch.num_inference_steps,
|
||||
latent=None,
|
||||
device=latents.device,
|
||||
)
|
||||
if hasattr(self.transformer, "patchifier"):
|
||||
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
|
||||
token_count = self.transformer.patchifier.get_token_count(
|
||||
video_shape)
|
||||
else:
|
||||
token_count = 1
|
||||
timestep_template = torch.ones(
|
||||
(latents.shape[0], token_count),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
audio_prompt_embeds = batch.extra.get("ltx2_audio_prompt_embeds")
|
||||
audio_neg_embeds = batch.extra.get("ltx2_audio_negative_embeds")
|
||||
audio_context_p = audio_prompt_embeds[0] if audio_prompt_embeds else None
|
||||
audio_context_n = audio_neg_embeds[0] if audio_neg_embeds else None
|
||||
audio_latents = None
|
||||
audio_timestep_template = None
|
||||
if audio_context_p is not None:
|
||||
fps_value = batch.fps
|
||||
if isinstance(fps_value, list):
|
||||
fps_value = fps_value[0] if fps_value else None
|
||||
if fps_value is None:
|
||||
fps_value = 1.0
|
||||
duration = float(batch.num_frames) / float(fps_value)
|
||||
audio_shape = AudioLatentShape.from_duration(
|
||||
batch=latents.shape[0],
|
||||
duration=duration,
|
||||
channels=DEFAULT_LTX2_AUDIO_CHANNELS,
|
||||
mel_bins=DEFAULT_LTX2_AUDIO_MEL_BINS,
|
||||
sample_rate=DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
|
||||
hop_length=DEFAULT_LTX2_AUDIO_HOP_LENGTH,
|
||||
audio_latent_downsample_factor=DEFAULT_LTX2_AUDIO_DOWNSAMPLE,
|
||||
)
|
||||
audio_generator = None
|
||||
if fastvideo_args.ltx2_initial_latent_path and batch.seed is not None:
|
||||
audio_generator = torch.Generator(
|
||||
device=latents.device).manual_seed(batch.seed)
|
||||
elif batch.generator is not None:
|
||||
if isinstance(batch.generator, list):
|
||||
audio_generator = batch.generator[0]
|
||||
else:
|
||||
audio_generator = batch.generator
|
||||
if audio_generator is not None and audio_generator.device.type != latents.device.type:
|
||||
if batch.seed is None:
|
||||
audio_generator = torch.Generator(device=latents.device)
|
||||
else:
|
||||
audio_generator = torch.Generator(
|
||||
device=latents.device).manual_seed(batch.seed)
|
||||
audio_patch_shape = (
|
||||
audio_shape.batch,
|
||||
audio_shape.frames,
|
||||
audio_shape.channels * audio_shape.mel_bins,
|
||||
)
|
||||
audio_latents_patch = torch.randn(
|
||||
audio_patch_shape,
|
||||
generator=audio_generator,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
if hasattr(self.transformer, "audio_patchifier"):
|
||||
audio_latents = self.transformer.audio_patchifier.unpatchify(
|
||||
audio_latents_patch, audio_shape)
|
||||
else:
|
||||
audio_latents = audio_latents_patch.view(
|
||||
audio_shape.batch,
|
||||
audio_shape.frames,
|
||||
audio_shape.channels,
|
||||
audio_shape.mel_bins,
|
||||
).permute(0, 2, 1, 3).contiguous()
|
||||
audio_timestep_template = torch.ones(
|
||||
(latents.shape[0], audio_shape.frames),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
logger.info(
|
||||
"[LTX2] Denoising start: steps=%d dtype=%s guidance=%s "
|
||||
"sigmas_shape=%s latents_shape=%s",
|
||||
batch.num_inference_steps,
|
||||
target_dtype,
|
||||
batch.guidance_scale,
|
||||
tuple(sigmas.shape),
|
||||
tuple(latents.shape),
|
||||
)
|
||||
|
||||
for step_index in tqdm(range(len(sigmas) - 1)):
|
||||
sigma = sigmas[step_index]
|
||||
sigma_next = sigmas[step_index + 1]
|
||||
timestep = timestep_template * sigma
|
||||
audio_timestep = (audio_timestep_template * sigma
|
||||
if audio_timestep_template is not None else None)
|
||||
|
||||
with torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled,
|
||||
), set_forward_context(
|
||||
current_timestep=sigma.item(),
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
):
|
||||
pos_outputs = self.transformer(
|
||||
hidden_states=latents.to(target_dtype),
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
encoder_attention_mask=prompt_mask,
|
||||
timestep=timestep,
|
||||
audio_hidden_states=audio_latents,
|
||||
audio_encoder_hidden_states=audio_context_p,
|
||||
audio_timestep=audio_timestep,
|
||||
)
|
||||
if isinstance(pos_outputs, tuple):
|
||||
pos_denoised, pos_audio = pos_outputs
|
||||
else:
|
||||
pos_denoised = pos_outputs
|
||||
pos_audio = None
|
||||
|
||||
# Only run negative pass if CFG is enabled
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_outputs = self.transformer(
|
||||
hidden_states=latents.to(target_dtype),
|
||||
encoder_hidden_states=neg_prompt_embeds,
|
||||
encoder_attention_mask=neg_prompt_mask,
|
||||
timestep=timestep,
|
||||
audio_hidden_states=audio_latents,
|
||||
audio_encoder_hidden_states=audio_context_n,
|
||||
audio_timestep=audio_timestep,
|
||||
)
|
||||
if isinstance(neg_outputs, tuple):
|
||||
neg_denoised, neg_audio = neg_outputs
|
||||
else:
|
||||
neg_denoised = neg_outputs
|
||||
neg_audio = None
|
||||
pos_denoised = pos_denoised + (batch.guidance_scale - 1) * (
|
||||
pos_denoised - neg_denoised)
|
||||
if pos_audio is not None and neg_audio is not None:
|
||||
pos_audio = pos_audio + (batch.guidance_scale -
|
||||
1) * (pos_audio - neg_audio)
|
||||
|
||||
sigma_value = sigma.to(torch.float32) if isinstance(
|
||||
sigma, torch.Tensor) else torch.tensor(
|
||||
float(sigma),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
dt = sigma_next - sigma
|
||||
velocity = ((latents.float() - pos_denoised.float()) /
|
||||
sigma_value).to(latents.dtype)
|
||||
latents = (latents.float() + velocity.float() * dt).to(
|
||||
latents.dtype)
|
||||
if pos_audio is not None and audio_latents is not None:
|
||||
audio_velocity = ((audio_latents.float() - pos_audio.float()) /
|
||||
sigma_value).to(audio_latents.dtype)
|
||||
audio_latents = (audio_latents.float() +
|
||||
audio_velocity.float() * dt).to(
|
||||
audio_latents.dtype)
|
||||
|
||||
batch.latents = latents
|
||||
batch.extra["ltx2_audio_latents"] = audio_latents
|
||||
logger.info("[LTX2] Denoising done.")
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
return result
|
||||
@@ -0,0 +1,189 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Latent preparation stage for LTX-2 pipelines.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LTX2LatentPreparationStage(PipelineStage):
|
||||
"""Prepare initial LTX-2 latents without relying on a diffusers scheduler."""
|
||||
|
||||
def __init__(self, transformer) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
latent_num_frames = self._adjust_video_length(batch, fastvideo_args)
|
||||
|
||||
if not batch.prompt_embeds:
|
||||
batch_size = 1
|
||||
elif isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
if not batch.prompt_embeds:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
device = get_local_torch_device()
|
||||
dummy_prompt = torch.zeros(
|
||||
batch_size,
|
||||
0,
|
||||
self.transformer.hidden_size,
|
||||
device=device,
|
||||
dtype=transformer_dtype,
|
||||
)
|
||||
batch.prompt_embeds = [dummy_prompt]
|
||||
batch.negative_prompt_embeds = []
|
||||
batch.do_classifier_free_guidance = False
|
||||
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
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
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
latent_path = fastvideo_args.ltx2_initial_latent_path
|
||||
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
if height % spatial_ratio != 0 or width % spatial_ratio != 0:
|
||||
raise ValueError(
|
||||
f"Height and width must be divisible by {spatial_ratio} "
|
||||
f"but are {height} and {width}.")
|
||||
shape = (
|
||||
batch_size,
|
||||
self.transformer.num_channels_latents,
|
||||
num_frames,
|
||||
height // spatial_ratio,
|
||||
width // spatial_ratio,
|
||||
)
|
||||
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, "
|
||||
f"but requested an effective batch size of {batch_size}.")
|
||||
|
||||
if latents is None:
|
||||
if latent_path:
|
||||
loaded_latents = self._load_initial_latent(
|
||||
latent_path, device, dtype)
|
||||
if loaded_latents is not None:
|
||||
latents = loaded_latents
|
||||
else:
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self._save_initial_latent(latent_path, latents)
|
||||
else:
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = shape
|
||||
return batch
|
||||
|
||||
def _adjust_video_length(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> int | None:
|
||||
if not fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames:
|
||||
return None
|
||||
temporal_scale_factor = (fastvideo_args.pipeline_config.vae_config.
|
||||
arch_config.temporal_compression_ratio)
|
||||
video_length = batch.num_frames
|
||||
return int((video_length - 1) // temporal_scale_factor + 1)
|
||||
|
||||
def _load_initial_latent(
|
||||
self,
|
||||
latent_path: str,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor | None:
|
||||
path = Path(latent_path)
|
||||
if not path.exists():
|
||||
return None
|
||||
payload = torch.load(path, map_location=device)
|
||||
if isinstance(payload, dict):
|
||||
if "video_latent" in payload:
|
||||
latent = payload["video_latent"]
|
||||
elif "latent" in payload:
|
||||
latent = payload["latent"]
|
||||
else:
|
||||
latent = None
|
||||
else:
|
||||
latent = payload
|
||||
if not torch.is_tensor(latent):
|
||||
raise TypeError(f"Expected tensor for initial latent in {path}")
|
||||
logger.info("[LTX2] Loaded initial latent from %s", path)
|
||||
return latent.to(device=device, dtype=dtype)
|
||||
|
||||
def _save_initial_latent(self, latent_path: str,
|
||||
latents: torch.Tensor) -> None:
|
||||
path = Path(latent_path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if path.exists():
|
||||
return
|
||||
torch.save({"video_latent": latents.detach().cpu()}, path)
|
||||
logger.info("[LTX2] Saved initial latent to %s", path)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check(
|
||||
"prompt_or_embeds",
|
||||
None,
|
||||
lambda _: V.string_or_list_strings(batch.prompt) or not batch.
|
||||
prompt_embeds or V.list_not_empty(batch.prompt_embeds),
|
||||
)
|
||||
if batch.prompt_embeds:
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
V.list_of_tensors)
|
||||
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
|
||||
V.positive_int)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("latents", batch.latents, V.none_or_tensor)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple)
|
||||
return result
|
||||
@@ -11,14 +11,11 @@ from typing import Any
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class TextEncodingStage(PipelineStage):
|
||||
"""
|
||||
@@ -39,6 +36,7 @@ class TextEncodingStage(PipelineStage):
|
||||
super().__init__()
|
||||
self.tokenizers = tokenizers
|
||||
self.text_encoders = text_encoders
|
||||
self._last_audio_embeds: list[torch.Tensor] | None = None
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
@@ -70,6 +68,8 @@ class TextEncodingStage(PipelineStage):
|
||||
encoder_index=all_indices,
|
||||
return_attention_mask=True,
|
||||
)
|
||||
if self._last_audio_embeds is not None:
|
||||
batch.extra["ltx2_audio_prompt_embeds"] = self._last_audio_embeds
|
||||
|
||||
for pe in prompt_embeds_list:
|
||||
batch.prompt_embeds.append(pe)
|
||||
@@ -86,6 +86,9 @@ class TextEncodingStage(PipelineStage):
|
||||
encoder_index=all_indices,
|
||||
return_attention_mask=True,
|
||||
)
|
||||
if self._last_audio_embeds is not None:
|
||||
batch.extra[
|
||||
"ltx2_audio_negative_embeds"] = self._last_audio_embeds
|
||||
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
for ne in neg_embeds_list:
|
||||
@@ -184,10 +187,13 @@ class TextEncodingStage(PipelineStage):
|
||||
|
||||
embeds_list: list[torch.Tensor] = []
|
||||
attn_masks_list: list[torch.Tensor] = []
|
||||
audio_embeds_list: list[torch.Tensor] = []
|
||||
|
||||
preprocess_funcs = fastvideo_args.pipeline_config.preprocess_text_funcs
|
||||
postprocess_funcs = fastvideo_args.pipeline_config.postprocess_text_funcs
|
||||
encoder_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
|
||||
is_ltx2 = getattr(fastvideo_args.pipeline_config.dit_config, "prefix",
|
||||
"") == "ltx2"
|
||||
|
||||
if return_type not in ("list", "dict", "stack"):
|
||||
raise ValueError(
|
||||
@@ -259,6 +265,11 @@ class TextEncodingStage(PipelineStage):
|
||||
except Exception:
|
||||
prompt_embeds, attention_mask = postprocess_func(
|
||||
outputs, attention_mask)
|
||||
if is_ltx2 and getattr(outputs, "hidden_states", None):
|
||||
audio_embed = outputs.hidden_states[0]
|
||||
if dtype is not None:
|
||||
audio_embed = audio_embed.to(dtype=dtype)
|
||||
audio_embeds_list.append(audio_embed)
|
||||
|
||||
if dtype is not None:
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype)
|
||||
@@ -266,6 +277,7 @@ class TextEncodingStage(PipelineStage):
|
||||
if return_attention_mask:
|
||||
attn_masks_list.append(attention_mask)
|
||||
|
||||
self._last_audio_embeds = audio_embeds_list if is_ltx2 else None
|
||||
return self.return_embeds(embeds_list, attn_masks_list, return_type,
|
||||
return_attention_mask, indices)
|
||||
|
||||
|
||||
@@ -67,17 +67,23 @@ def run_test(pytest_command: str):
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="H100:1",
|
||||
image=image,
|
||||
timeout=1200,
|
||||
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_encoder_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_vae_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_transformer_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
|
||||
@app.function(
|
||||
gpu="L40S:4",
|
||||
@@ -89,13 +95,21 @@ def run_transformer_tests():
|
||||
def run_ssim_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:4",
|
||||
image=image,
|
||||
timeout=900,
|
||||
secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_training_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
|
||||
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:2",
|
||||
image=image,
|
||||
timeout=900,
|
||||
secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_training_lora_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
|
||||
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
|
||||
|
||||
@app.function(gpu="H100:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests_VSA():
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -80,9 +80,48 @@ WAN_I2V_PARAMS = {
|
||||
"text-encoder-precision": ("fp32",)
|
||||
}
|
||||
|
||||
# LTX-2 distilled one-stage params (no refine/upscale)
|
||||
# Official defaults: height=512, width=768, num_frames=121, fps=24, seed=10
|
||||
# Using num_frames=41 for faster CI (still valid: 41 = 8×5 + 1)
|
||||
LTX2_T2V_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "FastVideo/LTX2-Distilled-Diffusers",
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 41, # Shorter for CI; official default is 121
|
||||
"num_inference_steps": 8, # Distilled uses 8 steps
|
||||
"guidance_scale": 1.0, # No CFG for distilled
|
||||
"embedded_cfg_scale": 6,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
"fps": 24,
|
||||
"neg_prompt": (
|
||||
"blurry, out of focus, overexposed, underexposed, low contrast, washed out colors, "
|
||||
"excessive noise, grainy texture, poor lighting, flickering, motion blur, distorted "
|
||||
"proportions, unnatural skin tones, deformed facial features, asymmetrical face, "
|
||||
"missing facial features, extra limbs, disfigured hands, wrong hand count, artifacts "
|
||||
"around text, inconsistent perspective, camera shake, incorrect depth of field, "
|
||||
"background too sharp, background clutter, distracting reflections, harsh shadows, "
|
||||
"inconsistent lighting direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, wrong gender, "
|
||||
"exaggerated expressions, wrong gaze direction, mismatched lip sync, silent or muted "
|
||||
"audio, distorted voice, robotic voice, echo, background noise, off-sync audio, "
|
||||
"incorrect dialogue, added dialogue, repetitive speech, jittery movement, awkward "
|
||||
"pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, "
|
||||
"flat lighting, inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts."
|
||||
),
|
||||
"ltx2_vae_tiling": True,
|
||||
"ltx2_vae_spatial_tile_size_in_pixels": 512,
|
||||
"ltx2_vae_spatial_tile_overlap_in_pixels": 64,
|
||||
"ltx2_vae_temporal_tile_size_in_frames": 64,
|
||||
"ltx2_vae_temporal_tile_overlap_in_frames": 24,
|
||||
}
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"FastHunyuan-diffusers": HUNYUAN_PARAMS,
|
||||
"Wan2.1-T2V-1.3B-Diffusers": WAN_T2V_PARAMS,
|
||||
# "ltx2_diffusers": LTX2_T2V_PARAMS,
|
||||
}
|
||||
|
||||
I2V_MODEL_TO_PARAMS = {
|
||||
@@ -229,18 +268,26 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"flow_shift": BASE_PARAMS["flow_shift"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"use_fsdp_inference": True,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
}
|
||||
if "flow_shift" in BASE_PARAMS:
|
||||
init_kwargs["flow_shift"] = BASE_PARAMS["flow_shift"]
|
||||
if BASE_PARAMS.get("vae_sp"):
|
||||
init_kwargs["vae_sp"] = True
|
||||
init_kwargs["vae_tiling"] = True
|
||||
if "text-encoder-precision" in BASE_PARAMS:
|
||||
init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
|
||||
# LTX2-specific VAE tiling parameters
|
||||
if BASE_PARAMS.get("ltx2_vae_tiling"):
|
||||
init_kwargs["ltx2_vae_tiling"] = True
|
||||
init_kwargs["ltx2_vae_spatial_tile_size_in_pixels"] = BASE_PARAMS.get("ltx2_vae_spatial_tile_size_in_pixels", 512)
|
||||
init_kwargs["ltx2_vae_spatial_tile_overlap_in_pixels"] = BASE_PARAMS.get("ltx2_vae_spatial_tile_overlap_in_pixels", 64)
|
||||
init_kwargs["ltx2_vae_temporal_tile_size_in_frames"] = BASE_PARAMS.get("ltx2_vae_temporal_tile_size_in_frames", 64)
|
||||
init_kwargs["ltx2_vae_temporal_tile_overlap_in_frames"] = BASE_PARAMS.get("ltx2_vae_temporal_tile_overlap_in_frames", 24)
|
||||
|
||||
generation_kwargs = {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
|
||||
@@ -132,9 +132,13 @@ class MultiprocExecutor(Executor):
|
||||
else:
|
||||
logging_info = None
|
||||
|
||||
# Get extra dict (contains audio, etc.)
|
||||
extra = responses[0].get("extra", {})
|
||||
|
||||
result_batch = ForwardBatch(data_type=forward_batch.data_type,
|
||||
output=output,
|
||||
logging_info=logging_info)
|
||||
logging_info=logging_info,
|
||||
extra=extra)
|
||||
|
||||
return result_batch
|
||||
|
||||
@@ -648,7 +652,8 @@ class WorkerMultiprocProc:
|
||||
logging_info = output_batch.logging_info
|
||||
self.pipe.send({
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"logging_info": logging_info
|
||||
"logging_info": logging_info,
|
||||
"extra": output_batch.extra,
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(
|
||||
|
||||
@@ -29,6 +29,7 @@ dependencies = [
|
||||
"diffusers>=0.33.1",
|
||||
"torch>=2.9.1",
|
||||
"torchvision",
|
||||
"torchaudio",
|
||||
|
||||
# Acceleration & Optimization
|
||||
"accelerate==1.0.1",
|
||||
|
||||
@@ -0,0 +1,444 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Convert LTX-2 weights to FastVideo naming conventions and split by component.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
try:
|
||||
from huggingface_hub import snapshot_download
|
||||
except ImportError: # pragma: no cover - optional dependency
|
||||
snapshot_download = None
|
||||
|
||||
|
||||
PARAM_NAME_MAP: dict[str, str] = {
|
||||
r"^model\.diffusion_model\.(.*)$": r"\1",
|
||||
}
|
||||
|
||||
COMPONENT_PREFIXES: dict[str, tuple[str, ...]] = {
|
||||
"transformer": ("model.diffusion_model.",),
|
||||
"vae": ("vae.",),
|
||||
"audio_vae": ("audio_vae.",),
|
||||
"vocoder": ("vocoder.",),
|
||||
"text_embedding_projection": ("text_embedding_projection.", "model.text_embedding_projection."),
|
||||
}
|
||||
|
||||
|
||||
def _find_shards(model_path: Path) -> list[Path]:
|
||||
if model_path.is_file():
|
||||
return [model_path]
|
||||
|
||||
index_files = list(model_path.glob("*.safetensors.index.json"))
|
||||
if index_files:
|
||||
with index_files[0].open("r", encoding="utf-8") as f:
|
||||
index = json.load(f)
|
||||
return sorted({model_path / shard for shard in index["weight_map"].values()})
|
||||
return sorted(Path(p) for p in glob.glob(str(model_path / "*.safetensors")))
|
||||
|
||||
|
||||
def _apply_mapping(key: str) -> str:
|
||||
for pattern, replacement in PARAM_NAME_MAP.items():
|
||||
if re.match(pattern, key):
|
||||
return re.sub(pattern, replacement, key)
|
||||
return key
|
||||
|
||||
|
||||
def _load_weights(shards: list[Path]) -> dict[str, torch.Tensor]:
|
||||
weights: dict[str, torch.Tensor] = {}
|
||||
for shard in shards:
|
||||
weights.update(load_file(str(shard)))
|
||||
return weights
|
||||
|
||||
|
||||
def _read_metadata_config(path: Path) -> dict:
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
metadata = f.metadata()
|
||||
if not metadata or "config" not in metadata:
|
||||
return {}
|
||||
return json.loads(metadata["config"])
|
||||
|
||||
|
||||
def _filter_transformer_config(config: dict) -> dict:
|
||||
transformer = config.get("transformer", {})
|
||||
allowed = {
|
||||
"num_attention_heads",
|
||||
"attention_head_dim",
|
||||
"num_layers",
|
||||
"cross_attention_dim",
|
||||
"caption_channels",
|
||||
"norm_eps",
|
||||
"attention_type",
|
||||
"positional_embedding_theta",
|
||||
"positional_embedding_max_pos",
|
||||
"timestep_scale_multiplier",
|
||||
"use_middle_indices_grid",
|
||||
"rope_type",
|
||||
"frequencies_precision",
|
||||
"in_channels",
|
||||
"out_channels",
|
||||
"audio_num_attention_heads",
|
||||
"audio_attention_head_dim",
|
||||
"audio_in_channels",
|
||||
"audio_out_channels",
|
||||
"audio_cross_attention_dim",
|
||||
"audio_positional_embedding_max_pos",
|
||||
"av_ca_timestep_scale_multiplier",
|
||||
}
|
||||
filtered = {k: v for k, v in transformer.items() if k in allowed}
|
||||
if "frequencies_precision" in filtered:
|
||||
filtered["double_precision_rope"] = filtered["frequencies_precision"] == "float64"
|
||||
del filtered["frequencies_precision"]
|
||||
return filtered
|
||||
|
||||
|
||||
def _build_text_embedding_projection_config(
|
||||
gemma_model_path: str = "",
|
||||
) -> dict:
|
||||
return {
|
||||
"architectures": ["LTX2GemmaTextEncoderModel"],
|
||||
"hidden_size": 3840,
|
||||
"num_hidden_layers": 48,
|
||||
"num_attention_heads": 30,
|
||||
"text_len": 1024,
|
||||
"pad_token_id": 0,
|
||||
"eos_token_id": 2,
|
||||
"gemma_model_path": gemma_model_path,
|
||||
"gemma_dtype": "bfloat16",
|
||||
"padding_side": "left",
|
||||
"feature_extractor_in_features": 3840 * 49,
|
||||
"feature_extractor_out_features": 3840,
|
||||
"connector_num_attention_heads": 30,
|
||||
"connector_attention_head_dim": 128,
|
||||
"connector_num_layers": 2,
|
||||
"connector_positional_embedding_theta": 10000.0,
|
||||
"connector_positional_embedding_max_pos": [4096],
|
||||
"connector_rope_type": "split",
|
||||
"connector_double_precision_rope": True,
|
||||
"connector_num_learnable_registers": 128,
|
||||
}
|
||||
|
||||
|
||||
def _wrap_component_config(
|
||||
component_name: str,
|
||||
component_config: dict | None,
|
||||
class_name: str | None = None,
|
||||
) -> dict | None:
|
||||
if component_config is None:
|
||||
return None
|
||||
wrapped = {component_name: component_config}
|
||||
if class_name is not None:
|
||||
wrapped["_class_name"] = class_name
|
||||
return wrapped
|
||||
|
||||
|
||||
def _split_component_weights(weights: dict[str, torch.Tensor]) -> dict[str, OrderedDict]:
|
||||
components: dict[str, OrderedDict] = {name: OrderedDict() for name in COMPONENT_PREFIXES}
|
||||
for key, value in weights.items():
|
||||
if key.startswith("model.diffusion_model.audio_embeddings_connector."):
|
||||
new_key = key.replace("model.diffusion_model.audio_embeddings_connector.", "audio_embeddings_connector.")
|
||||
components["text_embedding_projection"][new_key] = value
|
||||
continue
|
||||
if key.startswith("model.diffusion_model.video_embeddings_connector."):
|
||||
new_key = key.replace("model.diffusion_model.video_embeddings_connector.", "embeddings_connector.")
|
||||
components["text_embedding_projection"][new_key] = value
|
||||
continue
|
||||
|
||||
matched = False
|
||||
for component, prefixes in COMPONENT_PREFIXES.items():
|
||||
for prefix in prefixes:
|
||||
if key.startswith(prefix):
|
||||
new_key = key[len(prefix):]
|
||||
components[component][new_key] = value
|
||||
matched = True
|
||||
break
|
||||
if matched:
|
||||
break
|
||||
return {name: weights for name, weights in components.items() if weights}
|
||||
|
||||
|
||||
def _write_component(
|
||||
output_dir: Path,
|
||||
name: str,
|
||||
weights: OrderedDict,
|
||||
config: dict | None,
|
||||
dir_name: str | None = None,
|
||||
) -> None:
|
||||
component_dir = output_dir / (dir_name or name)
|
||||
component_dir.mkdir(parents=True, exist_ok=True)
|
||||
output_file = component_dir / "model.safetensors"
|
||||
save_file(weights, str(output_file))
|
||||
print(f"Saved {name} weights to {output_file}")
|
||||
|
||||
if config is not None:
|
||||
config_path = component_dir / "config.json"
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(config, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f"Saved {name} config to {config_path}")
|
||||
|
||||
|
||||
def _build_model_index(
|
||||
transformer_class_name: str,
|
||||
vae_class_name: str,
|
||||
pipeline_class_name: str,
|
||||
diffusers_version: str,
|
||||
) -> dict:
|
||||
return {
|
||||
"_class_name": pipeline_class_name,
|
||||
"_diffusers_version": diffusers_version,
|
||||
"transformer": ["diffusers", transformer_class_name],
|
||||
"vae": ["diffusers", vae_class_name],
|
||||
"text_encoder": ["transformers", "LTX2GemmaTextEncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"audio_vae": ["diffusers", "LTX2AudioDecoder"],
|
||||
"vocoder": ["diffusers", "LTX2Vocoder"],
|
||||
}
|
||||
|
||||
|
||||
def _write_model_index(output_dir: Path, model_index: dict) -> None:
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
model_index_path = output_dir / "model_index.json"
|
||||
with model_index_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(model_index, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f"Saved model_index.json to {model_index_path}")
|
||||
|
||||
|
||||
def convert_components(
|
||||
source_path: Path,
|
||||
output_dir: Path,
|
||||
metadata_config: dict,
|
||||
transformer_class_name: str,
|
||||
components_to_write: set[str] | None = None,
|
||||
emit_diffusers_repo: bool = True,
|
||||
pipeline_class_name: str = "LTX2Pipeline",
|
||||
diffusers_version: str = "0.33.0.dev0",
|
||||
gemma_model_path: str = "",
|
||||
) -> None:
|
||||
shards = _find_shards(source_path)
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"No safetensors found in {source_path}")
|
||||
|
||||
weights = _load_weights(shards)
|
||||
split_weights = _split_component_weights(weights)
|
||||
if components_to_write is not None:
|
||||
split_weights = {name: weights for name, weights in split_weights.items() if name in components_to_write}
|
||||
|
||||
transformer_weights = split_weights.get("transformer", OrderedDict())
|
||||
converted_transformer = OrderedDict()
|
||||
for key, value in transformer_weights.items():
|
||||
new_key = _apply_mapping(f"model.diffusion_model.{key}")
|
||||
converted_transformer[new_key] = value
|
||||
split_weights["transformer"] = converted_transformer
|
||||
|
||||
transformer_config = _filter_transformer_config(metadata_config)
|
||||
if transformer_config:
|
||||
transformer_config["_class_name"] = transformer_class_name
|
||||
|
||||
component_configs: dict[str, dict | None] = {
|
||||
"transformer": transformer_config or None,
|
||||
"vae": _wrap_component_config(
|
||||
"vae",
|
||||
metadata_config.get("vae"),
|
||||
class_name="CausalVideoAutoencoder",
|
||||
),
|
||||
"audio_vae": _wrap_component_config(
|
||||
"audio_vae",
|
||||
metadata_config.get("audio_vae"),
|
||||
class_name="LTX2AudioDecoder",
|
||||
),
|
||||
"vocoder": _wrap_component_config(
|
||||
"vocoder",
|
||||
metadata_config.get("vocoder"),
|
||||
class_name="LTX2Vocoder",
|
||||
),
|
||||
"text_embedding_projection": _build_text_embedding_projection_config(
|
||||
gemma_model_path=gemma_model_path
|
||||
),
|
||||
}
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
for name, component_weights in split_weights.items():
|
||||
_write_component(output_dir, name, component_weights, component_configs.get(name))
|
||||
if emit_diffusers_repo and name == "text_embedding_projection":
|
||||
_write_component(
|
||||
output_dir,
|
||||
name,
|
||||
component_weights,
|
||||
component_configs.get(name),
|
||||
dir_name="text_encoder",
|
||||
)
|
||||
if emit_diffusers_repo:
|
||||
required_for_index = {
|
||||
"transformer",
|
||||
"vae",
|
||||
"audio_vae",
|
||||
"vocoder",
|
||||
"text_embedding_projection",
|
||||
}
|
||||
if components_to_write is not None and not required_for_index.issubset(components_to_write):
|
||||
print("Skipping model_index.json; not all diffusers components were written.")
|
||||
return
|
||||
if not required_for_index.issubset(split_weights.keys()):
|
||||
print("Skipping model_index.json; missing diffusers components in weights.")
|
||||
return
|
||||
vae_class_name = (component_configs.get("vae") or {}).get(
|
||||
"_class_name", "CausalVideoAutoencoder"
|
||||
)
|
||||
model_index = _build_model_index(
|
||||
transformer_class_name=transformer_class_name,
|
||||
vae_class_name=vae_class_name,
|
||||
pipeline_class_name=pipeline_class_name,
|
||||
diffusers_version=diffusers_version,
|
||||
)
|
||||
_write_model_index(output_dir, model_index)
|
||||
|
||||
|
||||
def update_transformer_config(config_path: Path, class_name: str) -> None:
|
||||
if not config_path.exists():
|
||||
print(f"Config file not found: {config_path}")
|
||||
return
|
||||
|
||||
with config_path.open("r", encoding="utf-8") as f:
|
||||
config = json.load(f)
|
||||
|
||||
config["_class_name"] = class_name
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(config, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f"Updated _class_name in {config_path} -> {class_name}")
|
||||
|
||||
|
||||
def maybe_download(repo_id: str, target_dir: Path, token: str | None, allow_patterns: str | None) -> Path:
|
||||
if snapshot_download is None:
|
||||
raise RuntimeError("huggingface_hub is required for --download")
|
||||
target_dir.mkdir(parents=True, exist_ok=True)
|
||||
snapshot_download(
|
||||
repo_id=repo_id,
|
||||
local_dir=str(target_dir),
|
||||
local_dir_use_symlinks=False,
|
||||
token=token,
|
||||
allow_patterns=allow_patterns,
|
||||
)
|
||||
return target_dir
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Convert LTX-2 weights to FastVideo format")
|
||||
parser.add_argument("--source", type=str, help="Path to transformer weights directory")
|
||||
parser.add_argument("--output", type=str, required=True, help="Output directory for converted weights")
|
||||
parser.add_argument("--download", type=str, help="HF repo id to download before conversion")
|
||||
parser.add_argument("--allow-patterns", type=str, help="Limit HF download to matching files")
|
||||
parser.add_argument("--token", type=str, default=os.getenv("HF_TOKEN"), help="HF token (or set HF_TOKEN)")
|
||||
parser.add_argument("--update-config", action="store_true", help="Update source config.json _class_name")
|
||||
parser.add_argument("--class-name", type=str, default="LTX2Transformer3DModel")
|
||||
parser.add_argument(
|
||||
"--diffusers-repo",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Emit a diffusers-style repo layout with model_index.json.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pipeline-class-name",
|
||||
type=str,
|
||||
default="LTX2Pipeline",
|
||||
help="Pipeline class name for model_index.json.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--diffusers-version",
|
||||
type=str,
|
||||
default="0.33.0.dev0",
|
||||
help="Diffusers version for model_index.json.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--transformer-only",
|
||||
action="store_true",
|
||||
help="Only convert transformer weights (no component split).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--components",
|
||||
type=str,
|
||||
default="",
|
||||
help=(
|
||||
"Comma-separated component list to write "
|
||||
"(transformer,vae,audio_vae,vocoder,text_embedding_projection)."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gemma-path",
|
||||
type=str,
|
||||
default="",
|
||||
help="Optional local Gemma model path to copy into the output repo.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.download:
|
||||
if args.source:
|
||||
raise ValueError("Use either --download or --source, not both.")
|
||||
source_dir = maybe_download(args.download, Path(args.output) / "download", args.token, args.allow_patterns)
|
||||
else:
|
||||
if not args.source:
|
||||
raise ValueError("--source is required when not using --download")
|
||||
source_dir = Path(args.source)
|
||||
|
||||
output_dir = Path(args.output)
|
||||
shards = _find_shards(source_dir)
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"No safetensors found in {source_dir}")
|
||||
metadata_path = shards[0]
|
||||
metadata_config = _read_metadata_config(metadata_path)
|
||||
components_to_write: set[str] | None = None
|
||||
if args.transformer_only:
|
||||
components_to_write = {"transformer"}
|
||||
elif args.components:
|
||||
components_to_write = {
|
||||
component.strip()
|
||||
for component in args.components.split(",")
|
||||
if component.strip()
|
||||
}
|
||||
|
||||
gemma_model_path = ""
|
||||
if args.gemma_path:
|
||||
gemma_src = Path(args.gemma_path)
|
||||
if not gemma_src.is_dir():
|
||||
raise ValueError(f"--gemma-path must be a directory: {gemma_src}")
|
||||
gemma_dest = output_dir / "text_encoder" / "gemma"
|
||||
if gemma_dest.exists():
|
||||
shutil.rmtree(gemma_dest)
|
||||
gemma_dest.parent.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copytree(gemma_src, gemma_dest)
|
||||
gemma_model_path = "gemma"
|
||||
|
||||
convert_components(
|
||||
source_dir,
|
||||
output_dir,
|
||||
metadata_config,
|
||||
args.class_name,
|
||||
components_to_write=components_to_write,
|
||||
emit_diffusers_repo=args.diffusers_repo,
|
||||
pipeline_class_name=args.pipeline_class_name,
|
||||
diffusers_version=args.diffusers_version,
|
||||
gemma_model_path=gemma_model_path,
|
||||
)
|
||||
|
||||
if args.update_config:
|
||||
if source_dir.is_dir():
|
||||
update_transformer_config(source_dir / "config.json", args.class_name)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,7 @@
|
||||
# Tests
|
||||
|
||||
- `tests/local_tests/` are local-only tests that require a checked-out `LTX-2/`
|
||||
directory under the repo root (`FastVideo/LTX-2`). Without that repo, they
|
||||
will skip or fail.
|
||||
- The CI-backed test suite still lives in `fastvideo/tests/`.
|
||||
- Eventually, all tests will move under `tests/`.
|
||||
@@ -0,0 +1,7 @@
|
||||
# Local LTX-2 Tests
|
||||
|
||||
These tests depend on a checked-out `LTX-2/` directory under the repo root
|
||||
(`FastVideo/LTX-2`). Without that local repo, the tests will skip or fail.
|
||||
|
||||
For the CI-backed test suite, see `fastvideo/tests/`. Those are the tests
|
||||
currently exercised in CI. (Eventually all tests will move to `tests/`.)
|
||||
@@ -0,0 +1,126 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.encoders import LTX2GemmaConfig
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
|
||||
def _load_connector_weights(path: str) -> dict[str, torch.Tensor]:
|
||||
weights = load_file(path)
|
||||
mapped: dict[str, torch.Tensor] = {}
|
||||
for name, tensor in weights.items():
|
||||
if name == "aggregate_embed.weight":
|
||||
mapped["feature_extractor_linear.aggregate_embed.weight"] = tensor
|
||||
elif name.startswith("embeddings_connector."):
|
||||
mapped[name] = tensor
|
||||
elif name.startswith("audio_embeddings_connector."):
|
||||
mapped[name] = tensor
|
||||
return mapped
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="LTX-2 Gemma encoder parity test requires CUDA.",
|
||||
)
|
||||
def test_ltx2_gemma_text_encoder_parity():
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
text_encoder_path = os.getenv(
|
||||
"LTX2_TEXT_ENCODER_PATH",
|
||||
str(diffusers_root / "text_encoder"),
|
||||
)
|
||||
gemma_model_path = str(Path(text_encoder_path) / "gemma")
|
||||
if not os.path.isdir(text_encoder_path):
|
||||
pytest.skip(f"LTX-2 text encoder weights not found at {text_encoder_path}")
|
||||
if not gemma_model_path or not os.path.isdir(gemma_model_path):
|
||||
pytest.skip("Gemma weights not found in text_encoder/gemma.")
|
||||
|
||||
try:
|
||||
from ltx_core.text_encoders.gemma.embeddings_connector import (
|
||||
Embeddings1DConnector,
|
||||
)
|
||||
from ltx_core.text_encoders.gemma.encoders.av_encoder import (
|
||||
AVGemmaTextEncoderModel,
|
||||
)
|
||||
from ltx_core.text_encoders.gemma.feature_extractor import (
|
||||
GemmaFeaturesExtractorProjLinear,
|
||||
)
|
||||
from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
|
||||
from transformers import Gemma3ForConditionalGeneration
|
||||
except Exception as exc:
|
||||
pytest.skip(f"LTX-2 Gemma import failed: {exc}")
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
|
||||
tokenizer = LTXVGemmaTokenizer(gemma_model_path, max_length=1024)
|
||||
gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
|
||||
gemma_model_path,
|
||||
local_files_only=True,
|
||||
torch_dtype=precision,
|
||||
).to(device)
|
||||
gemma_model.eval()
|
||||
|
||||
ref_model = AVGemmaTextEncoderModel(
|
||||
feature_extractor_linear=GemmaFeaturesExtractorProjLinear(),
|
||||
embeddings_connector=Embeddings1DConnector(),
|
||||
audio_embeddings_connector=Embeddings1DConnector(),
|
||||
tokenizer=tokenizer,
|
||||
model=gemma_model,
|
||||
dtype=precision,
|
||||
).to(device)
|
||||
ref_model.eval()
|
||||
|
||||
connector_weights = _load_connector_weights(
|
||||
os.path.join(text_encoder_path, "model.safetensors")
|
||||
)
|
||||
ref_model.load_state_dict(connector_weights, strict=False)
|
||||
|
||||
prompt = "A fast moving train in a snowy landscape."
|
||||
token_pairs = tokenizer.tokenize_with_weights(prompt)["gemma"]
|
||||
input_ids = torch.tensor(
|
||||
[[t[0] for t in token_pairs]], device=device, dtype=torch.long
|
||||
)
|
||||
attention_mask = torch.tensor(
|
||||
[[t[1] for t in token_pairs]], device=device, dtype=torch.long
|
||||
)
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=text_encoder_path,
|
||||
pipeline_config=PipelineConfig(
|
||||
text_encoder_configs=(LTX2GemmaConfig(),),
|
||||
text_encoder_precisions=("bf16",),
|
||||
),
|
||||
)
|
||||
loader = TextEncoderLoader()
|
||||
fastvideo_model = loader.load(text_encoder_path, args).to(device)
|
||||
fastvideo_model.eval()
|
||||
|
||||
with torch.no_grad():
|
||||
ref_video, ref_audio, ref_mask = ref_model(prompt, padding_side="left")
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
fastvideo_out = fastvideo_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
assert_close(ref_video, fastvideo_out.last_hidden_state, atol=1e-2, rtol=1e-2)
|
||||
assert_close(ref_audio, fastvideo_out.hidden_states[0], atol=1e-2, rtol=1e-2)
|
||||
assert torch.equal(ref_mask, fastvideo_out.attention_mask)
|
||||
@@ -0,0 +1,378 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.configs.models.encoders import LTX2GemmaConfig
|
||||
from fastvideo.models.encoders.gemma import LTX2GemmaTextEncoderModel
|
||||
from fastvideo.models.loader.component_loader import get_diffusers_config
|
||||
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
|
||||
def _init_log_paths() -> tuple[Path, Path]:
|
||||
base_dir = Path(os.getenv("LTX2_DEBUG_DIR", "ltx2_debug"))
|
||||
fastvideo_log = Path(os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", base_dir / "fastvideo_gemma.log"))
|
||||
reference_log = Path(os.getenv("LTX2_REFERENCE_GEMMA_LOG", base_dir / "reference_gemma.log"))
|
||||
fastvideo_log.parent.mkdir(parents=True, exist_ok=True)
|
||||
reference_log.parent.mkdir(parents=True, exist_ok=True)
|
||||
fastvideo_log.write_text("")
|
||||
reference_log.write_text("")
|
||||
return fastvideo_log, reference_log
|
||||
|
||||
|
||||
def _log_line(path: Path, message: str) -> None:
|
||||
with path.open("a", encoding="utf-8") as f:
|
||||
f.write(message + "\n")
|
||||
|
||||
|
||||
def _attach_encoder_logging(
|
||||
encoder: torch.nn.Module,
|
||||
log_path: Path,
|
||||
label: str,
|
||||
) -> None:
|
||||
def _format_sum(tensor: torch.Tensor | None) -> str:
|
||||
if tensor is None:
|
||||
return "None"
|
||||
return f"{tensor.float().sum().item():.6f}"
|
||||
|
||||
def _hook_factory(name: str):
|
||||
def _hook(_module, _inputs, outputs): # noqa: ANN001
|
||||
out = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
out_sum = _format_sum(out if torch.is_tensor(out) else None)
|
||||
_log_line(log_path, f"{label}:{name}:sum={out_sum}")
|
||||
return _hook
|
||||
|
||||
def _pre_hook_factory(name: str):
|
||||
def _hook(_module, inputs): # noqa: ANN001
|
||||
tensor = inputs[0] if inputs else None
|
||||
in_sum = _format_sum(tensor if torch.is_tensor(tensor) else None)
|
||||
_log_line(log_path, f"{label}:{name}:in_sum={in_sum}")
|
||||
return _hook
|
||||
|
||||
def _attach_block_detail(block: torch.nn.Module, prefix: str) -> None:
|
||||
block.register_forward_pre_hook(_pre_hook_factory(f"{prefix}:input"))
|
||||
block.register_forward_hook(_hook_factory(f"{prefix}:output"))
|
||||
if hasattr(block, "attn1"):
|
||||
block.attn1.register_forward_hook(
|
||||
_hook_factory(f"{prefix}:attn1"))
|
||||
if hasattr(block, "ff"):
|
||||
block.ff.register_forward_hook(_hook_factory(f"{prefix}:ff"))
|
||||
|
||||
if hasattr(encoder, "feature_extractor_linear"):
|
||||
encoder.feature_extractor_linear.register_forward_pre_hook(
|
||||
_pre_hook_factory("feature_extractor_linear"))
|
||||
encoder.feature_extractor_linear.register_forward_hook(_hook_factory("feature_extractor_linear"))
|
||||
if hasattr(encoder, "embeddings_connector"):
|
||||
encoder.embeddings_connector.register_forward_hook(_hook_factory("embeddings_connector"))
|
||||
for idx, block in enumerate(encoder.embeddings_connector.transformer_1d_blocks):
|
||||
_attach_block_detail(block, f"embeddings_block_{idx}")
|
||||
if hasattr(encoder, "audio_embeddings_connector"):
|
||||
encoder.audio_embeddings_connector.register_forward_hook(_hook_factory("audio_embeddings_connector"))
|
||||
for idx, block in enumerate(encoder.audio_embeddings_connector.transformer_1d_blocks):
|
||||
_attach_block_detail(block, f"audio_embeddings_block_{idx}")
|
||||
|
||||
|
||||
def _log_register_sums(encoder: torch.nn.Module, log_path: Path, label: str) -> None:
|
||||
def _sum_param(module: torch.nn.Module, name: str) -> float | None:
|
||||
if not hasattr(module, "learnable_registers"):
|
||||
return None
|
||||
param = getattr(module, "learnable_registers")
|
||||
if not torch.is_tensor(param):
|
||||
return None
|
||||
return param.float().sum().item()
|
||||
|
||||
if hasattr(encoder, "embeddings_connector"):
|
||||
reg_sum = _sum_param(encoder.embeddings_connector, "learnable_registers")
|
||||
if reg_sum is not None:
|
||||
_log_line(log_path, f"{label}:embeddings_registers:sum={reg_sum:.6f}")
|
||||
if hasattr(encoder, "audio_embeddings_connector"):
|
||||
reg_sum = _sum_param(encoder.audio_embeddings_connector, "learnable_registers")
|
||||
if reg_sum is not None:
|
||||
_log_line(log_path, f"{label}:audio_registers:sum={reg_sum:.6f}")
|
||||
|
||||
|
||||
def _log_param_sums(encoder: torch.nn.Module, log_path: Path, label: str) -> None:
|
||||
def _log_param(name: str, tensor: torch.Tensor | None) -> None:
|
||||
if tensor is None:
|
||||
_log_line(log_path, f"{label}:param:{name}:sum=None")
|
||||
return
|
||||
_log_line(log_path, f"{label}:param:{name}:sum={tensor.float().sum().item():.6f}")
|
||||
|
||||
if hasattr(encoder, "feature_extractor_linear"):
|
||||
_log_param(
|
||||
"feature_extractor_linear.aggregate_embed.weight",
|
||||
encoder.feature_extractor_linear.aggregate_embed.weight,
|
||||
)
|
||||
if hasattr(encoder, "embeddings_connector"):
|
||||
block0 = encoder.embeddings_connector.transformer_1d_blocks[0]
|
||||
_log_param("embeddings_block0.attn1.to_q.weight", block0.attn1.to_q.weight)
|
||||
_log_param("embeddings_block0.attn1.to_k.weight", block0.attn1.to_k.weight)
|
||||
_log_param("embeddings_block0.attn1.to_v.weight", block0.attn1.to_v.weight)
|
||||
_log_param("embeddings_block0.ff.net.0.proj.weight", block0.ff.net[0].proj.weight)
|
||||
if hasattr(encoder, "audio_embeddings_connector"):
|
||||
block0 = encoder.audio_embeddings_connector.transformer_1d_blocks[0]
|
||||
_log_param("audio_block0.attn1.to_q.weight", block0.attn1.to_q.weight)
|
||||
_log_param("audio_block0.attn1.to_k.weight", block0.attn1.to_k.weight)
|
||||
_log_param("audio_block0.attn1.to_v.weight", block0.attn1.to_v.weight)
|
||||
_log_param("audio_block0.ff.net.0.proj.weight", block0.ff.net[0].proj.weight)
|
||||
|
||||
|
||||
def _log_gemma_param_sums(
|
||||
gemma_model: torch.nn.Module | None,
|
||||
log_path: Path,
|
||||
label: str,
|
||||
) -> None:
|
||||
if gemma_model is None:
|
||||
_log_line(log_path, f"{label}:gemma_param:embed_tokens.weight:sum=None")
|
||||
return
|
||||
tokens = None
|
||||
if hasattr(gemma_model, "get_input_embeddings"):
|
||||
try:
|
||||
tokens = gemma_model.get_input_embeddings()
|
||||
except Exception:
|
||||
tokens = None
|
||||
if tokens is None:
|
||||
embed = getattr(gemma_model, "model", None)
|
||||
if embed is not None:
|
||||
tokens = getattr(embed, "embed_tokens", None)
|
||||
if tokens is None or not hasattr(tokens, "weight"):
|
||||
_log_line(log_path, f"{label}:gemma_param:embed_tokens.weight:sum=None")
|
||||
return
|
||||
_log_line(
|
||||
log_path,
|
||||
f"{label}:gemma_param:embed_tokens.weight:sum={tokens.weight.float().sum().item():.6f}",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="LTX-2 Gemma parity test requires CUDA.",
|
||||
)
|
||||
def test_ltx2_gemma_parity():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "TORCH_SDPA"
|
||||
torch.backends.cuda.enable_flash_sdp(False)
|
||||
torch.backends.cuda.enable_mem_efficient_sdp(False)
|
||||
torch.backends.cuda.enable_math_sdp(True)
|
||||
fastvideo_log, reference_log = _init_log_paths()
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
text_encoder_path = diffusers_root / "text_encoder"
|
||||
gemma_path = text_encoder_path / "gemma"
|
||||
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 weights not found at {official_path}")
|
||||
if not text_encoder_path.exists():
|
||||
pytest.skip(f"LTX-2 text encoder not found at {text_encoder_path}")
|
||||
if not gemma_path.exists():
|
||||
pytest.skip(f"LTX-2 Gemma weights not found at {gemma_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
from ltx_core.model.transformer.attention import Attention, AttentionFunction
|
||||
from ltx_core.text_encoders.gemma import (
|
||||
AV_GEMMA_TEXT_ENCODER_KEY_OPS,
|
||||
AVGemmaTextEncoderModelConfigurator,
|
||||
module_ops_from_gemma_root,
|
||||
)
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
|
||||
reference_builder = SingleGPUModelBuilder(
|
||||
model_path=str(official_path),
|
||||
model_class_configurator=AVGemmaTextEncoderModelConfigurator,
|
||||
model_sd_ops=AV_GEMMA_TEXT_ENCODER_KEY_OPS,
|
||||
module_ops=module_ops_from_gemma_root(str(gemma_path)),
|
||||
)
|
||||
reference_encoder = reference_builder.build(
|
||||
device=device, dtype=precision
|
||||
).to(device=device, dtype=precision)
|
||||
if hasattr(reference_encoder.model, "config"):
|
||||
if hasattr(reference_encoder.model.config, "attn_implementation"):
|
||||
reference_encoder.model.config.attn_implementation = "sdpa"
|
||||
if hasattr(reference_encoder.model.config, "_attn_implementation"):
|
||||
reference_encoder.model.config._attn_implementation = "sdpa"
|
||||
for module in reference_encoder.modules():
|
||||
if isinstance(module, Attention):
|
||||
module.attention_function = AttentionFunction.PYTORCH
|
||||
reference_encoder.eval()
|
||||
_attach_encoder_logging(reference_encoder, reference_log, "reference")
|
||||
_log_register_sums(reference_encoder, reference_log, "reference")
|
||||
_log_param_sums(reference_encoder, reference_log, "reference")
|
||||
_log_gemma_param_sums(reference_encoder.model, reference_log, "reference")
|
||||
_log_line(
|
||||
reference_log,
|
||||
"reference:gemma_config:attn_impl="
|
||||
f"{getattr(reference_encoder.model.config, 'attn_implementation', None)} "
|
||||
f"dtype={reference_encoder.model.dtype}",
|
||||
)
|
||||
|
||||
diffusers_config = get_diffusers_config(model=str(text_encoder_path))
|
||||
encoder_config = LTX2GemmaConfig()
|
||||
encoder_config.update_model_arch(diffusers_config)
|
||||
encoder_config.arch_config.gemma_model_path = str(gemma_path)
|
||||
fastvideo_encoder = LTX2GemmaTextEncoderModel(encoder_config).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
if hasattr(fastvideo_encoder.gemma_model, "config"):
|
||||
if hasattr(fastvideo_encoder.gemma_model.config, "attn_implementation"):
|
||||
fastvideo_encoder.gemma_model.config.attn_implementation = "sdpa"
|
||||
if hasattr(fastvideo_encoder.gemma_model.config, "_attn_implementation"):
|
||||
fastvideo_encoder.gemma_model.config._attn_implementation = "sdpa"
|
||||
|
||||
official_weights = load_file(str(official_path))
|
||||
fastvideo_weights: dict[str, torch.Tensor] = {}
|
||||
for name, tensor in official_weights.items():
|
||||
mapped_name = AV_GEMMA_TEXT_ENCODER_KEY_OPS.apply_to_key(name)
|
||||
if mapped_name is None:
|
||||
continue
|
||||
fastvideo_weights[mapped_name] = tensor
|
||||
fastvideo_encoder.load_weights(fastvideo_weights.items())
|
||||
fastvideo_encoder.eval()
|
||||
_attach_encoder_logging(fastvideo_encoder, fastvideo_log, "fastvideo")
|
||||
_log_register_sums(fastvideo_encoder, fastvideo_log, "fastvideo")
|
||||
_log_param_sums(fastvideo_encoder, fastvideo_log, "fastvideo")
|
||||
_log_gemma_param_sums(fastvideo_encoder.gemma_model, fastvideo_log, "fastvideo")
|
||||
_log_line(
|
||||
fastvideo_log,
|
||||
"fastvideo:gemma_config:attn_impl="
|
||||
f"{getattr(fastvideo_encoder.gemma_model.config, 'attn_implementation', None)} "
|
||||
f"dtype={fastvideo_encoder.gemma_model.dtype}",
|
||||
)
|
||||
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers."
|
||||
ref_tokenizer = reference_encoder.tokenizer
|
||||
if ref_tokenizer is None:
|
||||
pytest.skip("Reference tokenizer is not initialized.")
|
||||
token_pairs = ref_tokenizer.tokenize_with_weights(prompt)["gemma"]
|
||||
input_ids = torch.tensor([[t[0] for t in token_pairs]], device=device)
|
||||
attention_mask = torch.tensor([[w[1] for w in token_pairs]], device=device)
|
||||
|
||||
with torch.no_grad(), torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=False,
|
||||
enable_mem_efficient=False,
|
||||
enable_math=True,
|
||||
):
|
||||
ref_outputs = reference_encoder.model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
use_cache=False,
|
||||
)
|
||||
_log_line(
|
||||
reference_log,
|
||||
"reference:gemma_hidden_last:sum="
|
||||
f"{ref_outputs.hidden_states[-1].float().sum().item():.6f}",
|
||||
)
|
||||
ref_projected = reference_encoder._run_feature_extractor(
|
||||
ref_outputs.hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
padding_side="left",
|
||||
)
|
||||
# Compare rotary embeddings between implementations.
|
||||
from ltx_core.model.transformer.rope import precompute_freqs_cis
|
||||
from fastvideo.models.dits.ltx2 import precompute_ltx_freqs_cis
|
||||
|
||||
seq_len = ref_projected.shape[1]
|
||||
indices_grid = torch.arange(seq_len, device=device, dtype=torch.float32)[None, None, :]
|
||||
ref_cos, ref_sin = precompute_freqs_cis(
|
||||
indices_grid=indices_grid,
|
||||
dim=reference_encoder.embeddings_connector.inner_dim,
|
||||
out_dtype=ref_projected.dtype,
|
||||
theta=reference_encoder.embeddings_connector.positional_embedding_theta,
|
||||
max_pos=reference_encoder.embeddings_connector.positional_embedding_max_pos,
|
||||
num_attention_heads=reference_encoder.embeddings_connector.num_attention_heads,
|
||||
rope_type=reference_encoder.embeddings_connector.rope_type,
|
||||
)
|
||||
fast_cos, fast_sin = precompute_ltx_freqs_cis(
|
||||
indices_grid=indices_grid,
|
||||
dim=fastvideo_encoder.embeddings_connector.inner_dim,
|
||||
out_dtype=ref_projected.dtype,
|
||||
theta=fastvideo_encoder.embeddings_connector.positional_embedding_theta,
|
||||
max_pos=fastvideo_encoder.embeddings_connector.positional_embedding_max_pos,
|
||||
num_attention_heads=fastvideo_encoder.embeddings_connector.num_attention_heads,
|
||||
rope_type=fastvideo_encoder.embeddings_connector.rope_type,
|
||||
)
|
||||
_log_line(
|
||||
reference_log,
|
||||
"reference:rope:cos_sum="
|
||||
f"{ref_cos.float().sum().item():.6f} sin_sum={ref_sin.float().sum().item():.6f} "
|
||||
f"theta={reference_encoder.embeddings_connector.positional_embedding_theta} "
|
||||
f"max_pos={reference_encoder.embeddings_connector.positional_embedding_max_pos} "
|
||||
f"rope_type={reference_encoder.embeddings_connector.rope_type}"
|
||||
)
|
||||
_log_line(
|
||||
fastvideo_log,
|
||||
"fastvideo:rope:cos_sum="
|
||||
f"{fast_cos.float().sum().item():.6f} sin_sum={fast_sin.float().sum().item():.6f} "
|
||||
f"theta={fastvideo_encoder.embeddings_connector.positional_embedding_theta} "
|
||||
f"max_pos={fastvideo_encoder.embeddings_connector.positional_embedding_max_pos} "
|
||||
f"rope_type={fastvideo_encoder.embeddings_connector.rope_type}"
|
||||
)
|
||||
ref_video, ref_audio, _ = reference_encoder._run_connectors(
|
||||
ref_projected, attention_mask
|
||||
)
|
||||
fast_video_from_ref, fast_audio_from_ref, _ = fastvideo_encoder._run_connectors(
|
||||
ref_projected, attention_mask
|
||||
)
|
||||
_log_line(
|
||||
fastvideo_log,
|
||||
"fastvideo:connector_on_ref:video_sum="
|
||||
f"{fast_video_from_ref.float().sum().item():.6f}",
|
||||
)
|
||||
_log_line(
|
||||
fastvideo_log,
|
||||
"fastvideo:connector_on_ref:audio_sum="
|
||||
f"{fast_audio_from_ref.float().sum().item():.6f}",
|
||||
)
|
||||
|
||||
fast_outputs = fastvideo_encoder.gemma_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
use_cache=False,
|
||||
)
|
||||
_log_line(
|
||||
fastvideo_log,
|
||||
"fastvideo:gemma_hidden_last:sum="
|
||||
f"{fast_outputs.hidden_states[-1].float().sum().item():.6f}",
|
||||
)
|
||||
fast_projected = fastvideo_encoder._run_feature_extractor(
|
||||
fast_outputs.hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
padding_side=fastvideo_encoder.padding_side,
|
||||
)
|
||||
fast_video, fast_audio, _ = fastvideo_encoder._run_connectors(
|
||||
fast_projected, attention_mask
|
||||
)
|
||||
|
||||
assert ref_video.shape == fast_video.shape
|
||||
assert ref_audio.shape == fast_audio.shape
|
||||
assert torch.isfinite(ref_video).all(), "Reference Gemma produced non-finite video embeddings."
|
||||
assert torch.isfinite(ref_audio).all(), "Reference Gemma produced non-finite audio embeddings."
|
||||
assert torch.isfinite(fast_video).all(), "FastVideo Gemma produced non-finite video embeddings."
|
||||
assert torch.isfinite(fast_audio).all(), "FastVideo Gemma produced non-finite audio embeddings."
|
||||
assert_close(ref_video, fast_video, atol=3e-1, rtol=5e-2)
|
||||
assert_close(ref_audio, fast_audio, atol=3e-1, rtol=5e-2)
|
||||
@@ -0,0 +1,296 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.models.dits.ltx2 import (
|
||||
AudioLatentShape,
|
||||
DEFAULT_LTX2_AUDIO_CHANNELS,
|
||||
DEFAULT_LTX2_AUDIO_DOWNSAMPLE,
|
||||
DEFAULT_LTX2_AUDIO_HOP_LENGTH,
|
||||
DEFAULT_LTX2_AUDIO_MEL_BINS,
|
||||
DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
|
||||
)
|
||||
from fastvideo.models.loader.component_loader import PipelineComponentLoader
|
||||
|
||||
|
||||
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
|
||||
tensor_f32 = tensor.float()
|
||||
print(
|
||||
f"[LTX2 SMOKE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={tensor_f32.min().item():.6f} max={tensor_f32.max().item():.6f} "
|
||||
f"mean={tensor_f32.mean().item():.6f} sum={tensor_f32.sum().item():.6f}"
|
||||
)
|
||||
|
||||
def _truncate_debug_logs() -> None:
|
||||
for env_var in (
|
||||
"LTX2_PIPELINE_DEBUG_PATH",
|
||||
"LTX2_REFERENCE_DEBUG_PATH",
|
||||
"LTX2_PIPELINE_DEBUG_DETAIL_PATH",
|
||||
"LTX2_REFERENCE_DEBUG_DETAIL_PATH",
|
||||
):
|
||||
log_path = os.getenv(env_var, "")
|
||||
if not log_path:
|
||||
continue
|
||||
log_dir = os.path.dirname(log_path)
|
||||
if log_dir:
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
with open(log_path, "w", encoding="utf-8") as f:
|
||||
f.write("")
|
||||
|
||||
|
||||
def _run_audio_decode_smoke(
|
||||
diffusers_path: str,
|
||||
fastvideo_args,
|
||||
device: torch.device,
|
||||
num_frames: int,
|
||||
fps: float,
|
||||
) -> None:
|
||||
audio_decoder_path = os.path.join(diffusers_path, "audio_vae")
|
||||
vocoder_path = os.path.join(diffusers_path, "vocoder")
|
||||
if not os.path.isdir(audio_decoder_path):
|
||||
pytest.skip(f"Missing LTX-2 audio decoder at {audio_decoder_path}")
|
||||
if not os.path.isdir(vocoder_path):
|
||||
pytest.skip(f"Missing LTX-2 vocoder at {vocoder_path}")
|
||||
|
||||
audio_decoder = PipelineComponentLoader.load_module(
|
||||
module_name="audio_decoder",
|
||||
component_model_path=audio_decoder_path,
|
||||
transformers_or_diffusers="diffusers",
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
vocoder = PipelineComponentLoader.load_module(
|
||||
module_name="vocoder",
|
||||
component_model_path=vocoder_path,
|
||||
transformers_or_diffusers="diffusers",
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
|
||||
duration = float(num_frames) / float(fps)
|
||||
audio_shape = AudioLatentShape.from_duration(
|
||||
batch=1,
|
||||
duration=duration,
|
||||
channels=DEFAULT_LTX2_AUDIO_CHANNELS,
|
||||
mel_bins=DEFAULT_LTX2_AUDIO_MEL_BINS,
|
||||
sample_rate=DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
|
||||
hop_length=DEFAULT_LTX2_AUDIO_HOP_LENGTH,
|
||||
audio_latent_downsample_factor=DEFAULT_LTX2_AUDIO_DOWNSAMPLE,
|
||||
)
|
||||
audio_dtype = next(audio_decoder.parameters()).dtype
|
||||
audio_latents = torch.randn(
|
||||
audio_shape.to_torch_shape(),
|
||||
device=device,
|
||||
dtype=audio_dtype,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
decoded_spec = audio_decoder(audio_latents)
|
||||
audio_wave = vocoder(decoded_spec)
|
||||
|
||||
assert audio_wave.ndim == 3
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="LTX-2 pipeline smoke test requires CUDA.",
|
||||
)
|
||||
def test_ltx2_pipeline_smoke():
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
debug_dir = repo_root / "ltx2_debug"
|
||||
os.environ.setdefault(
|
||||
"LTX2_PIPELINE_DEBUG_PATH",
|
||||
str(debug_dir / "fastvideo_pipeline.log"),
|
||||
)
|
||||
os.environ.setdefault(
|
||||
"LTX2_REFERENCE_DEBUG_PATH",
|
||||
str(debug_dir / "reference_pipeline.log"),
|
||||
)
|
||||
os.environ.setdefault(
|
||||
"LTX2_PIPELINE_DEBUG_DETAIL_PATH",
|
||||
str(debug_dir / "fastvideo_pipeline_detail.log"),
|
||||
)
|
||||
os.environ.setdefault(
|
||||
"LTX2_REFERENCE_DEBUG_DETAIL_PATH",
|
||||
str(debug_dir / "reference_pipeline_detail.log"),
|
||||
)
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
os.environ.setdefault("LTX2_REFERENCE_ATTN", "pytorch")
|
||||
torch.backends.cuda.enable_flash_sdp(False)
|
||||
torch.backends.cuda.enable_mem_efficient_sdp(False)
|
||||
torch.backends.cuda.enable_math_sdp(True)
|
||||
_truncate_debug_logs()
|
||||
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
ltx_pipelines_path = repo_root / "LTX-2" / "packages" / "ltx-pipelines" / "src"
|
||||
if ltx_pipelines_path.exists() and str(ltx_pipelines_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_pipelines_path))
|
||||
os.environ["PYTHONPATH"] = str(repo_root)
|
||||
|
||||
diffusers_path = os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
gemma_model_path = os.path.join(diffusers_path, "text_encoder", "gemma")
|
||||
official_path = os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
|
||||
if not os.path.isdir(diffusers_path):
|
||||
pytest.skip(f"Missing LTX-2 diffusers repo at {diffusers_path}")
|
||||
if not os.path.isfile(os.path.join(diffusers_path, "model_index.json")):
|
||||
pytest.skip(f"Missing model_index.json in {diffusers_path}")
|
||||
|
||||
if not gemma_model_path or not os.path.isdir(gemma_model_path):
|
||||
pytest.skip("Gemma weights not found in text_encoder/gemma.")
|
||||
if not os.path.isfile(official_path):
|
||||
pytest.skip(f"Missing LTX-2 official weights at {official_path}")
|
||||
|
||||
try:
|
||||
from ltx_pipelines.ti2vid_one_stage import TI2VidOneStagePipeline
|
||||
from ltx_core.model.transformer import attention as ltx_attention
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 pipeline import failed: {exc}")
|
||||
ltx_attention.memory_efficient_attention = None
|
||||
ltx_attention.flash_attn_interface = None
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers."
|
||||
negative_prompt = "low quality, blurry, distorted, artifacts, jpeg compression"
|
||||
seed = 42
|
||||
height = 64
|
||||
width = 96
|
||||
num_frames = 9
|
||||
fps = 12.0
|
||||
steps = 4
|
||||
guidance_scale = 4.0
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
latent_path = str(Path(tmpdir) / "ltx2_initial_latent.pt")
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
diffusers_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
ltx2_vae_tiling=False,
|
||||
ltx2_initial_latent_path=latent_path,
|
||||
)
|
||||
result = generator.generate_video(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
output_path="outputs_video/ltx2_smoke",
|
||||
save_video=False,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
fps=fps,
|
||||
num_inference_steps=steps,
|
||||
guidance_scale=guidance_scale,
|
||||
seed=seed,
|
||||
)
|
||||
generator.shutdown()
|
||||
_run_audio_decode_smoke(
|
||||
diffusers_path=diffusers_path,
|
||||
fastvideo_args=generator.fastvideo_args,
|
||||
device=device,
|
||||
num_frames=num_frames,
|
||||
fps=fps,
|
||||
)
|
||||
|
||||
fastvideo_out = result["samples"]
|
||||
fastvideo_out = fastvideo_out.to(device=device, dtype=torch.float32)
|
||||
_log_tensor_stats("fastvideo_video", fastvideo_out)
|
||||
|
||||
ref_pipeline = TI2VidOneStagePipeline(
|
||||
checkpoint_path=official_path,
|
||||
gemma_root=gemma_model_path,
|
||||
loras=[],
|
||||
device=device,
|
||||
fp8transformer=False,
|
||||
)
|
||||
original_text_encoder = ref_pipeline.model_ledger.text_encoder
|
||||
original_transformer = ref_pipeline.model_ledger.transformer
|
||||
original_video_decoder = ref_pipeline.model_ledger.video_decoder
|
||||
|
||||
def _patched_text_encoder():
|
||||
encoder = original_text_encoder()
|
||||
try:
|
||||
from ltx_core.model.transformer.attention import ( # type: ignore
|
||||
Attention,
|
||||
AttentionFunction,
|
||||
)
|
||||
except ImportError:
|
||||
return encoder
|
||||
if hasattr(encoder, "model") and hasattr(encoder.model, "config"):
|
||||
if hasattr(encoder.model.config, "attn_implementation"):
|
||||
encoder.model.config.attn_implementation = "sdpa"
|
||||
if hasattr(encoder.model.config, "_attn_implementation"):
|
||||
encoder.model.config._attn_implementation = "sdpa"
|
||||
for module in encoder.modules():
|
||||
if isinstance(module, Attention):
|
||||
module.attention_function = AttentionFunction.PYTORCH
|
||||
return encoder
|
||||
|
||||
ref_pipeline.model_ledger.text_encoder = _patched_text_encoder
|
||||
if os.getenv("LTX2_DISABLE_VAE_NOISE", "1") == "1":
|
||||
def _patched_video_decoder():
|
||||
decoder = original_video_decoder()
|
||||
if hasattr(decoder, "decode_noise_scale"):
|
||||
decoder.decode_noise_scale = 0.0
|
||||
return decoder
|
||||
|
||||
ref_pipeline.model_ledger.video_decoder = _patched_video_decoder
|
||||
if os.getenv("LTX2_DEBUG_DETAIL", "0") == "1":
|
||||
from ..transformers.test_ltx2 import (
|
||||
_attach_block_detail_logging,
|
||||
)
|
||||
|
||||
def _patched_transformer():
|
||||
model = original_transformer()
|
||||
core = getattr(model, "velocity_model", model)
|
||||
_attach_block_detail_logging(
|
||||
core,
|
||||
Path(os.environ["LTX2_REFERENCE_DEBUG_DETAIL_PATH"]),
|
||||
"reference",
|
||||
True,
|
||||
)
|
||||
return model
|
||||
|
||||
ref_pipeline.model_ledger.transformer = _patched_transformer
|
||||
with torch.no_grad():
|
||||
ref_video_iter, _ = ref_pipeline(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
seed=seed,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
frame_rate=fps,
|
||||
num_inference_steps=steps,
|
||||
cfg_guidance_scale=guidance_scale,
|
||||
images=[],
|
||||
enhance_prompt=False,
|
||||
initial_video_latent_path=latent_path,
|
||||
)
|
||||
ref_chunks = list(ref_video_iter)
|
||||
ref_video = torch.cat(
|
||||
[chunk if torch.is_tensor(chunk) else torch.from_numpy(chunk) for chunk in ref_chunks],
|
||||
dim=0,
|
||||
)
|
||||
ref_video = ref_video.to(torch.float32) / 255.0
|
||||
ref_video = ref_video.permute(3, 0, 1, 2).unsqueeze(0)
|
||||
_log_tensor_stats("reference_video", ref_video)
|
||||
|
||||
assert ref_video.shape == fastvideo_out.shape
|
||||
assert_close(ref_video, fastvideo_out, atol=2 / 255, rtol=1e-3)
|
||||
@@ -0,0 +1,343 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.configs.models.dits import LTX2VideoConfig
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
|
||||
|
||||
def _read_transformer_config(config: dict) -> dict:
|
||||
transformer_config = config.get("transformer", {})
|
||||
if not transformer_config:
|
||||
raise ValueError("Missing transformer config in LTX-2 metadata.")
|
||||
return transformer_config
|
||||
|
||||
|
||||
def _infer_patch_params(in_channels: int) -> tuple[int, int]:
|
||||
patch_size = 1
|
||||
num_channels_latents = 128
|
||||
for candidate in (8, 16, 32, 64, 128):
|
||||
if in_channels % candidate != 0:
|
||||
continue
|
||||
patch_volume = in_channels // candidate
|
||||
root = int(round(patch_volume**0.5))
|
||||
if root * root == patch_volume:
|
||||
patch_size = root
|
||||
num_channels_latents = candidate
|
||||
break
|
||||
print(
|
||||
f"[LTX2 TEST] Inferred patch_size={patch_size}, "
|
||||
f"num_channels_latents={num_channels_latents}"
|
||||
)
|
||||
return patch_size, num_channels_latents
|
||||
|
||||
|
||||
def _attach_block_sum_logging(
|
||||
model: torch.nn.Module,
|
||||
log_path: Path,
|
||||
label: str,
|
||||
enabled: bool,
|
||||
) -> None:
|
||||
if not enabled:
|
||||
return
|
||||
|
||||
log_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if log_path.exists():
|
||||
log_path.unlink()
|
||||
|
||||
def _format_sum(tensor: torch.Tensor | None) -> str:
|
||||
if tensor is None:
|
||||
return "None"
|
||||
return f"{tensor.float().sum().item():.6f}"
|
||||
|
||||
def _hook(module, inputs, outputs): # noqa: ANN001
|
||||
if isinstance(outputs, tuple):
|
||||
video_args, audio_args = outputs
|
||||
video_sum = _format_sum(video_args.x if video_args is not None else None)
|
||||
audio_sum = _format_sum(audio_args.x if audio_args is not None else None)
|
||||
else:
|
||||
video_sum = _format_sum(outputs)
|
||||
audio_sum = "None"
|
||||
with log_path.open("a", encoding="utf-8") as f:
|
||||
f.write(f"{label}:{module.idx}:video_sum={video_sum},audio_sum={audio_sum}\n")
|
||||
|
||||
for block in model.transformer_blocks:
|
||||
block.register_forward_hook(_hook)
|
||||
|
||||
|
||||
def _attach_block_detail_logging(
|
||||
model: torch.nn.Module,
|
||||
log_path: Path,
|
||||
label: str,
|
||||
enabled: bool,
|
||||
) -> None:
|
||||
if not enabled:
|
||||
return
|
||||
|
||||
log_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if log_path.exists():
|
||||
log_path.unlink()
|
||||
|
||||
def _format_sum(tensor: torch.Tensor | None) -> str:
|
||||
if tensor is None:
|
||||
return "None"
|
||||
return f"{tensor.float().sum().item():.6f}"
|
||||
|
||||
def _hook_factory(block_idx: int, name: str):
|
||||
def _hook(_module, _inputs, outputs): # noqa: ANN001
|
||||
if isinstance(outputs, tuple):
|
||||
out = outputs[0]
|
||||
else:
|
||||
out = outputs
|
||||
out_sum = _format_sum(out if torch.is_tensor(out) else None)
|
||||
with log_path.open("a", encoding="utf-8") as f:
|
||||
f.write(f"{label}:{block_idx}:{name}:out_sum={out_sum}\n")
|
||||
return _hook
|
||||
|
||||
for block in model.transformer_blocks:
|
||||
idx = block.idx
|
||||
for name in (
|
||||
"attn1",
|
||||
"attn2",
|
||||
"ff",
|
||||
"audio_attn1",
|
||||
"audio_attn2",
|
||||
"audio_ff",
|
||||
"audio_to_video_attn",
|
||||
"video_to_audio_attn",
|
||||
):
|
||||
if hasattr(block, name):
|
||||
getattr(block, name).register_forward_hook(_hook_factory(idx, name))
|
||||
|
||||
def _output_hook(name: str):
|
||||
def _hook(_module, _inputs, outputs): # noqa: ANN001
|
||||
out = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
out_sum = _format_sum(out if torch.is_tensor(out) else None)
|
||||
with log_path.open("a", encoding="utf-8") as f:
|
||||
f.write(f"{label}:output:{name}:out_sum={out_sum}\n")
|
||||
return _hook
|
||||
|
||||
for name in ("proj_out", "audio_proj_out"):
|
||||
if hasattr(model, name):
|
||||
getattr(model, name).register_forward_hook(_output_hook(name))
|
||||
|
||||
|
||||
def test_ltx2_transformer_parity():
|
||||
torch.manual_seed(42)
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
fastvideo_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_FASTVIDEO_PATH",
|
||||
str(diffusers_root / "transformer"),
|
||||
)
|
||||
)
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 official weights not found at {official_path}")
|
||||
if not fastvideo_path.exists():
|
||||
pytest.skip(f"FastVideo converted weights not found at {fastvideo_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.components.patchifiers import VideoLatentPatchifier
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
from ltx_core.model.transformer import (LTXModelConfigurator,
|
||||
LTXV_MODEL_COMFY_RENAMING_MAP)
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.types import VideoLatentShape
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
config_loader = SafetensorsModelStateDictLoader()
|
||||
metadata = config_loader.metadata(str(official_path))
|
||||
transformer_config = _read_transformer_config(metadata)
|
||||
|
||||
config = LTX2VideoConfig()
|
||||
cfg = config.arch_config
|
||||
cfg.num_attention_heads = transformer_config.get("num_attention_heads",
|
||||
cfg.num_attention_heads)
|
||||
cfg.attention_head_dim = transformer_config.get("attention_head_dim",
|
||||
cfg.attention_head_dim)
|
||||
cfg.num_layers = transformer_config.get("num_layers", cfg.num_layers)
|
||||
cfg.cross_attention_dim = transformer_config.get(
|
||||
"cross_attention_dim", cfg.cross_attention_dim)
|
||||
cfg.caption_channels = transformer_config.get("caption_channels",
|
||||
cfg.caption_channels)
|
||||
cfg.norm_eps = transformer_config.get("norm_eps", cfg.norm_eps)
|
||||
cfg.attention_type = transformer_config.get("attention_type",
|
||||
cfg.attention_type)
|
||||
cfg.positional_embedding_theta = transformer_config.get(
|
||||
"positional_embedding_theta", cfg.positional_embedding_theta)
|
||||
cfg.positional_embedding_max_pos = transformer_config.get(
|
||||
"positional_embedding_max_pos", cfg.positional_embedding_max_pos)
|
||||
cfg.timestep_scale_multiplier = transformer_config.get(
|
||||
"timestep_scale_multiplier", cfg.timestep_scale_multiplier)
|
||||
cfg.use_middle_indices_grid = transformer_config.get(
|
||||
"use_middle_indices_grid", cfg.use_middle_indices_grid)
|
||||
cfg.rope_type = transformer_config.get("rope_type", cfg.rope_type)
|
||||
cfg.double_precision_rope = transformer_config.get(
|
||||
"double_precision_rope",
|
||||
transformer_config.get("frequencies_precision", "")
|
||||
== "float64",
|
||||
)
|
||||
cfg.audio_num_attention_heads = transformer_config.get(
|
||||
"audio_num_attention_heads", cfg.audio_num_attention_heads)
|
||||
cfg.audio_attention_head_dim = transformer_config.get(
|
||||
"audio_attention_head_dim", cfg.audio_attention_head_dim)
|
||||
cfg.audio_in_channels = transformer_config.get("audio_in_channels",
|
||||
cfg.audio_in_channels)
|
||||
cfg.audio_out_channels = transformer_config.get("audio_out_channels",
|
||||
cfg.audio_out_channels)
|
||||
cfg.audio_cross_attention_dim = transformer_config.get(
|
||||
"audio_cross_attention_dim", cfg.audio_cross_attention_dim)
|
||||
cfg.audio_positional_embedding_max_pos = transformer_config.get(
|
||||
"audio_positional_embedding_max_pos",
|
||||
cfg.audio_positional_embedding_max_pos,
|
||||
)
|
||||
cfg.av_ca_timestep_scale_multiplier = transformer_config.get(
|
||||
"av_ca_timestep_scale_multiplier", cfg.av_ca_timestep_scale_multiplier)
|
||||
cfg.in_channels = transformer_config.get("in_channels", cfg.in_channels)
|
||||
cfg.out_channels = transformer_config.get("out_channels", cfg.out_channels)
|
||||
|
||||
patch_size, num_channels_latents = _infer_patch_params(cfg.in_channels)
|
||||
cfg.patch_size = (1, patch_size, patch_size)
|
||||
cfg.num_channels_latents = num_channels_latents
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("LTX-2 transformer parity test requires CUDA for attention backends.")
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=str(fastvideo_path),
|
||||
dit_cpu_offload=True,
|
||||
use_fsdp_inference=False,
|
||||
pipeline_config=PipelineConfig(dit_config=config, dit_precision=precision_str),
|
||||
)
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
fastvideo_model = loader.load(str(fastvideo_path), args).to(device=device, dtype=precision)
|
||||
|
||||
reference_builder = SingleGPUModelBuilder(
|
||||
model_class_configurator=LTXModelConfigurator,
|
||||
model_path=str(official_path),
|
||||
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
)
|
||||
reference_model = reference_builder.build(
|
||||
device=device, dtype=precision).to(device=device, dtype=precision)
|
||||
reference_model.set_gradient_checkpointing(False)
|
||||
|
||||
fastvideo_model.eval()
|
||||
reference_model.eval()
|
||||
|
||||
debug_logs = os.getenv("LTX2_DEBUG_LOGS", "0") == "1"
|
||||
_attach_block_sum_logging(
|
||||
fastvideo_model.model,
|
||||
repo_root / "ltx2_debug" / "fastvideo.log",
|
||||
"fastvideo",
|
||||
debug_logs,
|
||||
)
|
||||
_attach_block_sum_logging(
|
||||
reference_model,
|
||||
repo_root / "ltx2_debug" / "reference.log",
|
||||
"reference",
|
||||
debug_logs,
|
||||
)
|
||||
_attach_block_detail_logging(
|
||||
fastvideo_model.model,
|
||||
repo_root / "ltx2_debug" / "fastvideo_detail.log",
|
||||
"fastvideo",
|
||||
os.getenv("LTX2_DEBUG_DETAIL", "0") == "1",
|
||||
)
|
||||
_attach_block_detail_logging(
|
||||
reference_model,
|
||||
repo_root / "ltx2_debug" / "reference_detail.log",
|
||||
"reference",
|
||||
os.getenv("LTX2_DEBUG_DETAIL", "0") == "1",
|
||||
)
|
||||
|
||||
patchifier = VideoLatentPatchifier(patch_size=cfg.patch_size[1])
|
||||
batch_size = 1
|
||||
frames = 4
|
||||
height = cfg.patch_size[1] * 4
|
||||
width = cfg.patch_size[2] * 4
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
cfg.num_channels_latents,
|
||||
frames,
|
||||
height,
|
||||
width,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
16,
|
||||
cfg.caption_channels,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
video_shape = VideoLatentShape.from_torch_shape(hidden_states.shape)
|
||||
positions = patchifier.get_patch_grid_bounds(video_shape, device=hidden_states.device)
|
||||
latents = patchifier.patchify(hidden_states)
|
||||
|
||||
video = Modality(
|
||||
enabled=True,
|
||||
latent=latents,
|
||||
timesteps=timestep,
|
||||
positions=positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out, _ = reference_model(
|
||||
video=video,
|
||||
audio=None,
|
||||
perturbations=BatchedPerturbationConfig.empty(batch_size),
|
||||
)
|
||||
ref_out = patchifier.unpatchify(ref_out, output_shape=video_shape)
|
||||
print(f"[LTX2 TEST] Reference model output shape: {ref_out.shape}")
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
fastvideo_out = fastvideo_model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
print(f"[LTX2 TEST] FastVideo model output shape: {fastvideo_out.shape}")
|
||||
assert ref_out.shape == fastvideo_out.shape
|
||||
assert ref_out.dtype == fastvideo_out.dtype
|
||||
assert_close(ref_out, fastvideo_out, atol=1e-4, rtol=1e-4)
|
||||
@@ -0,0 +1,280 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
# Force TORCH_SDPA backend for parity testing - both FastVideo and LTX-2 reference
|
||||
# will use PyTorch's scaled_dot_product_attention for consistent results
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.configs.models.dits import LTX2VideoConfig
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.dits.ltx2 import Modality as FastVideoModality
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from .test_ltx2 import (
|
||||
_attach_block_detail_logging,
|
||||
_attach_block_sum_logging,
|
||||
_infer_patch_params,
|
||||
_read_transformer_config,
|
||||
)
|
||||
|
||||
|
||||
def test_ltx2_transformer_audio_parity():
|
||||
torch.manual_seed(42)
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
fastvideo_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_FASTVIDEO_PATH",
|
||||
str(diffusers_root / "transformer"),
|
||||
)
|
||||
)
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 official weights not found at {official_path}")
|
||||
if not fastvideo_path.exists():
|
||||
pytest.skip(f"FastVideo converted weights not found at {fastvideo_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
from ltx_core.model.transformer import (LTXModelConfigurator,
|
||||
LTXV_MODEL_COMFY_RENAMING_MAP)
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.types import AudioLatentShape, VideoLatentShape
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
# Load config from metadata using same approach as test_ltx2.py
|
||||
config_loader = SafetensorsModelStateDictLoader()
|
||||
metadata = config_loader.metadata(str(official_path))
|
||||
transformer_config = _read_transformer_config(metadata)
|
||||
|
||||
config = LTX2VideoConfig()
|
||||
cfg = config.arch_config
|
||||
cfg.num_attention_heads = transformer_config.get("num_attention_heads",
|
||||
cfg.num_attention_heads)
|
||||
cfg.attention_head_dim = transformer_config.get("attention_head_dim",
|
||||
cfg.attention_head_dim)
|
||||
cfg.num_layers = transformer_config.get("num_layers", cfg.num_layers)
|
||||
cfg.cross_attention_dim = transformer_config.get(
|
||||
"cross_attention_dim", cfg.cross_attention_dim)
|
||||
cfg.caption_channels = transformer_config.get("caption_channels",
|
||||
cfg.caption_channels)
|
||||
cfg.norm_eps = transformer_config.get("norm_eps", cfg.norm_eps)
|
||||
cfg.attention_type = transformer_config.get("attention_type",
|
||||
cfg.attention_type)
|
||||
cfg.positional_embedding_theta = transformer_config.get(
|
||||
"positional_embedding_theta", cfg.positional_embedding_theta)
|
||||
cfg.positional_embedding_max_pos = transformer_config.get(
|
||||
"positional_embedding_max_pos", cfg.positional_embedding_max_pos)
|
||||
cfg.timestep_scale_multiplier = transformer_config.get(
|
||||
"timestep_scale_multiplier", cfg.timestep_scale_multiplier)
|
||||
cfg.use_middle_indices_grid = transformer_config.get(
|
||||
"use_middle_indices_grid", cfg.use_middle_indices_grid)
|
||||
cfg.rope_type = transformer_config.get("rope_type", cfg.rope_type)
|
||||
cfg.double_precision_rope = transformer_config.get(
|
||||
"double_precision_rope",
|
||||
transformer_config.get("frequencies_precision", "")
|
||||
== "float64",
|
||||
)
|
||||
cfg.audio_num_attention_heads = transformer_config.get(
|
||||
"audio_num_attention_heads", cfg.audio_num_attention_heads)
|
||||
cfg.audio_attention_head_dim = transformer_config.get(
|
||||
"audio_attention_head_dim", cfg.audio_attention_head_dim)
|
||||
cfg.audio_in_channels = transformer_config.get("audio_in_channels",
|
||||
cfg.audio_in_channels)
|
||||
cfg.audio_out_channels = transformer_config.get("audio_out_channels",
|
||||
cfg.audio_out_channels)
|
||||
cfg.audio_cross_attention_dim = transformer_config.get(
|
||||
"audio_cross_attention_dim", cfg.audio_cross_attention_dim)
|
||||
cfg.audio_positional_embedding_max_pos = transformer_config.get(
|
||||
"audio_positional_embedding_max_pos",
|
||||
cfg.audio_positional_embedding_max_pos,
|
||||
)
|
||||
cfg.av_ca_timestep_scale_multiplier = transformer_config.get(
|
||||
"av_ca_timestep_scale_multiplier", cfg.av_ca_timestep_scale_multiplier)
|
||||
cfg.in_channels = transformer_config.get("in_channels", cfg.in_channels)
|
||||
cfg.out_channels = transformer_config.get("out_channels", cfg.out_channels)
|
||||
|
||||
patch_size, num_channels_latents = _infer_patch_params(cfg.in_channels)
|
||||
cfg.patch_size = (1, patch_size, patch_size)
|
||||
cfg.num_channels_latents = num_channels_latents
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("LTX-2 transformer parity test requires CUDA for attention backends.")
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=str(fastvideo_path),
|
||||
dit_cpu_offload=True,
|
||||
use_fsdp_inference=False,
|
||||
pipeline_config=PipelineConfig(dit_config=config, dit_precision=precision_str),
|
||||
)
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
fastvideo_model = loader.load(str(fastvideo_path), args).to(device=device, dtype=precision)
|
||||
|
||||
# Use SingleGPUModelBuilder to load the reference model (same as test_ltx2.py)
|
||||
reference_builder = SingleGPUModelBuilder(
|
||||
model_class_configurator=LTXModelConfigurator,
|
||||
model_path=str(official_path),
|
||||
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
)
|
||||
reference_model = reference_builder.build(
|
||||
device=device, dtype=precision).to(device=device, dtype=precision)
|
||||
reference_model.set_gradient_checkpointing(False)
|
||||
|
||||
fastvideo_model.eval()
|
||||
reference_model.eval()
|
||||
|
||||
debug_logs = os.getenv("LTX2_DEBUG_LOGS", "0") == "1"
|
||||
_attach_block_sum_logging(
|
||||
fastvideo_model.model,
|
||||
repo_root / "ltx2_debug" / "fastvideo_audio.log",
|
||||
"fastvideo",
|
||||
debug_logs,
|
||||
)
|
||||
_attach_block_sum_logging(
|
||||
reference_model,
|
||||
repo_root / "ltx2_debug" / "reference_audio.log",
|
||||
"reference",
|
||||
debug_logs,
|
||||
)
|
||||
_attach_block_detail_logging(
|
||||
fastvideo_model.model,
|
||||
repo_root / "ltx2_debug" / "fastvideo_audio_detail.log",
|
||||
"fastvideo",
|
||||
os.getenv("LTX2_DEBUG_DETAIL", "0") == "1",
|
||||
)
|
||||
_attach_block_detail_logging(
|
||||
reference_model,
|
||||
repo_root / "ltx2_debug" / "reference_audio_detail.log",
|
||||
"reference",
|
||||
os.getenv("LTX2_DEBUG_DETAIL", "0") == "1",
|
||||
)
|
||||
|
||||
patchifier = VideoLatentPatchifier(patch_size=cfg.patch_size[1])
|
||||
audio_patchifier = AudioPatchifier(patch_size=1)
|
||||
batch_size = 1
|
||||
frames = 4
|
||||
height = cfg.patch_size[1] * 4
|
||||
width = cfg.patch_size[2] * 4
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
cfg.num_channels_latents,
|
||||
frames,
|
||||
height,
|
||||
width,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
16,
|
||||
cfg.caption_channels,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
video_shape = VideoLatentShape.from_torch_shape(hidden_states.shape)
|
||||
positions = patchifier.get_patch_grid_bounds(video_shape, device=hidden_states.device)
|
||||
latents = patchifier.patchify(hidden_states)
|
||||
|
||||
audio_frames = 16
|
||||
audio_channels = 8
|
||||
audio_mel_bins = 16
|
||||
audio_latents = torch.randn(
|
||||
batch_size,
|
||||
audio_channels,
|
||||
audio_frames,
|
||||
audio_mel_bins,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
audio_shape = AudioLatentShape.from_torch_shape(audio_latents.shape)
|
||||
audio_positions = audio_patchifier.get_patch_grid_bounds(audio_shape, device=audio_latents.device)
|
||||
audio_tokens = audio_patchifier.patchify(audio_latents)
|
||||
|
||||
video = Modality(
|
||||
enabled=True,
|
||||
latent=latents,
|
||||
timesteps=timestep,
|
||||
positions=positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
)
|
||||
audio = Modality(
|
||||
enabled=True,
|
||||
latent=audio_tokens,
|
||||
timesteps=timestep,
|
||||
positions=audio_positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
)
|
||||
|
||||
fastvideo_video = FastVideoModality(
|
||||
enabled=True,
|
||||
latent=latents,
|
||||
timesteps=timestep,
|
||||
positions=positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
)
|
||||
fastvideo_audio = FastVideoModality(
|
||||
enabled=True,
|
||||
latent=audio_tokens,
|
||||
timesteps=timestep,
|
||||
positions=audio_positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
_, ref_audio_out = reference_model(
|
||||
video=video,
|
||||
audio=audio,
|
||||
perturbations=BatchedPerturbationConfig.empty(batch_size),
|
||||
)
|
||||
ref_audio_out = audio_patchifier.unpatchify(ref_audio_out, output_shape=audio_shape)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
_, fastvideo_audio_out = fastvideo_model.model(
|
||||
video=fastvideo_video,
|
||||
audio=fastvideo_audio,
|
||||
)
|
||||
fastvideo_audio_out = audio_patchifier.unpatchify(fastvideo_audio_out, output_shape=audio_shape)
|
||||
|
||||
assert ref_audio_out.shape == fastvideo_audio_out.shape
|
||||
assert ref_audio_out.dtype == fastvideo_audio_out.dtype
|
||||
# With TORCH_SDPA backend for both, use same tolerance as video parity test
|
||||
assert_close(ref_audio_out, fastvideo_audio_out, atol=1e-4, rtol=1e-4)
|
||||
@@ -0,0 +1,272 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoder,
|
||||
LTX2AudioEncoder,
|
||||
LTX2Vocoder,
|
||||
)
|
||||
|
||||
|
||||
def _load_metadata(path: Path) -> dict:
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
meta = f.metadata()
|
||||
if not meta or "config" not in meta:
|
||||
raise KeyError("Missing config metadata in safetensors file.")
|
||||
return json.loads(meta["config"])
|
||||
|
||||
|
||||
def _load_weights(path: Path) -> dict[str, torch.Tensor]:
|
||||
print(f"[LTX2 AUDIO VAE TEST] Loading weights from {path}")
|
||||
return load_file(str(path))
|
||||
|
||||
|
||||
def _select_audio_vae_weights(
|
||||
weights: dict[str, torch.Tensor], prefix: str
|
||||
) -> dict[str, torch.Tensor]:
|
||||
filtered: dict[str, torch.Tensor] = {}
|
||||
alt_prefix = prefix.replace("audio_vae.", "")
|
||||
for name, tensor in weights.items():
|
||||
if name.startswith(prefix):
|
||||
filtered[name.replace(prefix, "")] = tensor
|
||||
elif alt_prefix and name.startswith(alt_prefix):
|
||||
filtered[name.replace(alt_prefix, "")] = tensor
|
||||
elif name.startswith("audio_vae.per_channel_statistics."):
|
||||
filtered[name.replace("audio_vae.", "")] = tensor
|
||||
elif name.startswith("per_channel_statistics."):
|
||||
filtered[name] = tensor
|
||||
print(f"[LTX2 AUDIO VAE TEST] Selected {len(filtered)} tensors for {prefix}")
|
||||
return filtered
|
||||
|
||||
|
||||
def _select_vocoder_weights(
|
||||
weights: dict[str, torch.Tensor]
|
||||
) -> dict[str, torch.Tensor]:
|
||||
if any(name.startswith("vocoder.") for name in weights):
|
||||
filtered = {
|
||||
name.replace("vocoder.", ""): tensor
|
||||
for name, tensor in weights.items()
|
||||
if name.startswith("vocoder.")
|
||||
}
|
||||
else:
|
||||
filtered = dict(weights)
|
||||
print(f"[LTX2 AUDIO VAE TEST] Selected {len(filtered)} tensors for vocoder.")
|
||||
return filtered
|
||||
|
||||
|
||||
def _load_into_model(
|
||||
model: torch.nn.Module, weights: dict[str, torch.Tensor]
|
||||
) -> tuple[int, list[str]]:
|
||||
model_state = model.state_dict()
|
||||
filtered = {
|
||||
k: v
|
||||
for k, v in weights.items()
|
||||
if k in model_state and model_state[k].shape == v.shape
|
||||
}
|
||||
missing = [k for k in model_state.keys() if k not in filtered]
|
||||
print(
|
||||
f"[LTX2 AUDIO VAE TEST] Loading {len(filtered)} / {len(model_state)} tensors "
|
||||
f"from {len(weights)} available"
|
||||
)
|
||||
if not filtered:
|
||||
return 0, missing
|
||||
model.load_state_dict(filtered, strict=False)
|
||||
return len(filtered), missing
|
||||
|
||||
|
||||
def test_ltx2_audio_vae_vocoder_parity():
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
audio_vae_path = Path(
|
||||
os.getenv("LTX2_AUDIO_VAE_PATH", str(diffusers_root / "audio_vae"))
|
||||
)
|
||||
vocoder_path = Path(
|
||||
os.getenv("LTX2_VOCODER_PATH", str(diffusers_root / "vocoder"))
|
||||
)
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 weights not found at {official_path}")
|
||||
if not audio_vae_path.exists():
|
||||
pytest.skip(f"LTX-2 audio VAE weights not found at {audio_vae_path}")
|
||||
if not vocoder_path.exists():
|
||||
pytest.skip(f"LTX-2 vocoder weights not found at {vocoder_path}")
|
||||
|
||||
config = _load_metadata(official_path)
|
||||
if "audio_vae" not in config or "vocoder" not in config:
|
||||
pytest.skip("Audio VAE or vocoder config not found in safetensors metadata.")
|
||||
|
||||
try:
|
||||
from ltx_core.model.audio_vae import (
|
||||
AudioDecoderConfigurator,
|
||||
AudioEncoderConfigurator,
|
||||
VocoderConfigurator,
|
||||
)
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
ref_weights = _load_weights(official_path)
|
||||
encoder_weights = _select_audio_vae_weights(ref_weights, "audio_vae.encoder.")
|
||||
decoder_weights = _select_audio_vae_weights(ref_weights, "audio_vae.decoder.")
|
||||
vocoder_weights = _select_vocoder_weights(ref_weights)
|
||||
if not encoder_weights or not decoder_weights or not vocoder_weights:
|
||||
pytest.skip("Audio VAE or vocoder weights not found in safetensors file.")
|
||||
|
||||
fastvideo_audio_weights_path = audio_vae_path / "model.safetensors"
|
||||
fastvideo_vocoder_weights_path = vocoder_path / "model.safetensors"
|
||||
if not fastvideo_audio_weights_path.exists():
|
||||
pytest.skip(
|
||||
f"FastVideo audio VAE weights not found at {fastvideo_audio_weights_path}"
|
||||
)
|
||||
if not fastvideo_vocoder_weights_path.exists():
|
||||
pytest.skip(
|
||||
f"FastVideo vocoder weights not found at {fastvideo_vocoder_weights_path}"
|
||||
)
|
||||
fastvideo_audio_weights = _load_weights(fastvideo_audio_weights_path)
|
||||
fastvideo_encoder_weights = _select_audio_vae_weights(
|
||||
fastvideo_audio_weights, "encoder."
|
||||
)
|
||||
fastvideo_decoder_weights = _select_audio_vae_weights(
|
||||
fastvideo_audio_weights, "decoder."
|
||||
)
|
||||
fastvideo_vocoder_weights = _select_vocoder_weights(
|
||||
_load_weights(fastvideo_vocoder_weights_path)
|
||||
)
|
||||
if (not fastvideo_encoder_weights or not fastvideo_decoder_weights
|
||||
or not fastvideo_vocoder_weights):
|
||||
pytest.skip("FastVideo audio VAE/vocoder weights not found in diffusers files.")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16 if torch.cuda.is_available() else torch.float32
|
||||
|
||||
fastvideo_encoder = LTX2AudioEncoder(config).to(device=device, dtype=precision)
|
||||
fastvideo_decoder = LTX2AudioDecoder(config).to(device=device, dtype=precision)
|
||||
fastvideo_vocoder = LTX2Vocoder(config).to(device=device, dtype=precision)
|
||||
|
||||
ref_encoder = AudioEncoderConfigurator.from_config(config).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
ref_decoder = AudioDecoderConfigurator.from_config(config).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
ref_vocoder = VocoderConfigurator.from_config(config).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
|
||||
loaded_fastvideo_encoder, missing_fastvideo_encoder = _load_into_model(
|
||||
fastvideo_encoder.model, fastvideo_encoder_weights
|
||||
)
|
||||
loaded_ref_encoder, missing_ref_encoder = _load_into_model(
|
||||
ref_encoder, encoder_weights
|
||||
)
|
||||
loaded_fastvideo_decoder, missing_fastvideo_decoder = _load_into_model(
|
||||
fastvideo_decoder.model, fastvideo_decoder_weights
|
||||
)
|
||||
loaded_ref_decoder, missing_ref_decoder = _load_into_model(
|
||||
ref_decoder, decoder_weights
|
||||
)
|
||||
loaded_fastvideo_vocoder, missing_fastvideo_vocoder = _load_into_model(
|
||||
fastvideo_vocoder.model, fastvideo_vocoder_weights
|
||||
)
|
||||
loaded_ref_vocoder, missing_ref_vocoder = _load_into_model(
|
||||
ref_vocoder, vocoder_weights
|
||||
)
|
||||
|
||||
if min(
|
||||
loaded_fastvideo_encoder,
|
||||
loaded_ref_encoder,
|
||||
loaded_fastvideo_decoder,
|
||||
loaded_ref_decoder,
|
||||
loaded_fastvideo_vocoder,
|
||||
loaded_ref_vocoder,
|
||||
) == 0:
|
||||
pytest.skip("Failed to load audio VAE or vocoder weights.")
|
||||
if (
|
||||
missing_fastvideo_encoder
|
||||
or missing_ref_encoder
|
||||
or missing_fastvideo_decoder
|
||||
or missing_ref_decoder
|
||||
or missing_fastvideo_vocoder
|
||||
or missing_ref_vocoder
|
||||
):
|
||||
print(
|
||||
f"[LTX2 AUDIO VAE TEST] Missing encoder keys: {len(missing_fastvideo_encoder)}"
|
||||
)
|
||||
print(
|
||||
f"[LTX2 AUDIO VAE TEST] Missing decoder keys: {len(missing_fastvideo_decoder)}"
|
||||
)
|
||||
print(
|
||||
f"[LTX2 AUDIO VAE TEST] Missing vocoder keys: {len(missing_fastvideo_vocoder)}"
|
||||
)
|
||||
pytest.skip("Missing audio VAE/vocoder keys; cannot ensure parity.")
|
||||
|
||||
fastvideo_encoder.model.eval()
|
||||
fastvideo_decoder.model.eval()
|
||||
fastvideo_vocoder.model.eval()
|
||||
ref_encoder.eval()
|
||||
ref_decoder.eval()
|
||||
ref_vocoder.eval()
|
||||
|
||||
ddconfig = config["audio_vae"]["model"]["params"]["ddconfig"]
|
||||
in_channels = ddconfig.get("in_channels", 2)
|
||||
resolution = ddconfig.get("resolution", 256)
|
||||
mel_bins = ddconfig.get("mel_bins", 64)
|
||||
batch_size = 1
|
||||
spectrogram = torch.randn(
|
||||
batch_size,
|
||||
in_channels,
|
||||
resolution,
|
||||
mel_bins,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_latents = ref_encoder(spectrogram)
|
||||
fast_latents = fastvideo_encoder(spectrogram)
|
||||
|
||||
assert ref_latents.shape == fast_latents.shape
|
||||
assert ref_latents.dtype == fast_latents.dtype
|
||||
assert torch.isfinite(ref_latents).all(), "Reference encoder produced non-finite latents."
|
||||
assert torch.isfinite(fast_latents).all(), "FastVideo encoder produced non-finite latents."
|
||||
assert_close(ref_latents, fast_latents, atol=1e-2, rtol=1e-2)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_decoded = ref_decoder(ref_latents)
|
||||
fast_decoded = fastvideo_decoder(ref_latents)
|
||||
|
||||
assert ref_decoded.shape == fast_decoded.shape
|
||||
assert ref_decoded.dtype == fast_decoded.dtype
|
||||
assert torch.isfinite(ref_decoded).all(), "Reference decoder produced non-finite output."
|
||||
assert torch.isfinite(fast_decoded).all(), "FastVideo decoder produced non-finite output."
|
||||
assert_close(ref_decoded, fast_decoded, atol=1e-2, rtol=1e-2)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_audio = ref_vocoder(ref_decoded)
|
||||
fast_audio = fastvideo_vocoder(ref_decoded)
|
||||
|
||||
assert ref_audio.shape == fast_audio.shape
|
||||
assert ref_audio.dtype == fast_audio.dtype
|
||||
assert torch.isfinite(ref_audio).all(), "Reference vocoder produced non-finite audio."
|
||||
assert torch.isfinite(fast_audio).all(), "FastVideo vocoder produced non-finite audio."
|
||||
assert_close(ref_audio, fast_audio, atol=1e-2, rtol=1e-2)
|
||||
@@ -0,0 +1,190 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.models.vaes.ltx2vae import LTX2VideoDecoder, LTX2VideoEncoder
|
||||
|
||||
|
||||
def _load_metadata(path: Path) -> dict:
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
meta = f.metadata()
|
||||
if not meta or "config" not in meta:
|
||||
raise KeyError("Missing config metadata in safetensors file.")
|
||||
return json.loads(meta["config"])
|
||||
|
||||
|
||||
def _load_weights(path: Path) -> dict[str, torch.Tensor]:
|
||||
print(f"[LTX2 VAE TEST] Loading weights from {path}")
|
||||
return load_file(str(path))
|
||||
|
||||
|
||||
def _select_vae_weights(weights: dict[str, torch.Tensor], prefix: str) -> dict[str, torch.Tensor]:
|
||||
filtered: dict[str, torch.Tensor] = {}
|
||||
alt_prefix = prefix.replace("vae.", "")
|
||||
for name, tensor in weights.items():
|
||||
if name.startswith(prefix):
|
||||
filtered[name.replace(prefix, "")] = tensor
|
||||
elif alt_prefix and name.startswith(alt_prefix):
|
||||
filtered[name.replace(alt_prefix, "")] = tensor
|
||||
elif name.startswith("vae.per_channel_statistics."):
|
||||
filtered[name.replace("vae.", "")] = tensor
|
||||
elif name.startswith("per_channel_statistics."):
|
||||
filtered[name] = tensor
|
||||
print(f"[LTX2 VAE TEST] Selected {len(filtered)} tensors for {prefix}")
|
||||
return filtered
|
||||
|
||||
|
||||
def _load_into_model(model: torch.nn.Module, weights: dict[str, torch.Tensor]) -> tuple[int, list[str]]:
|
||||
model_state = model.state_dict()
|
||||
filtered = {
|
||||
k: v
|
||||
for k, v in weights.items()
|
||||
if k in model_state and model_state[k].shape == v.shape
|
||||
}
|
||||
missing = [k for k in model_state.keys() if k not in filtered]
|
||||
print(
|
||||
f"[LTX2 VAE TEST] Loading {len(filtered)} / {len(model_state)} tensors "
|
||||
f"from {len(weights)} available"
|
||||
)
|
||||
if not filtered:
|
||||
return 0, missing
|
||||
model.load_state_dict(filtered, strict=False)
|
||||
return len(filtered), missing
|
||||
|
||||
|
||||
def test_ltx2_vae_parity():
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
fastvideo_path = Path(
|
||||
os.getenv("LTX2_VAE_PATH", str(diffusers_root / "vae"))
|
||||
)
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 weights not found at {official_path}")
|
||||
if not fastvideo_path.exists():
|
||||
pytest.skip(f"LTX-2 diffusers VAE not found at {fastvideo_path}")
|
||||
|
||||
config = _load_metadata(official_path)
|
||||
if "vae" not in config:
|
||||
pytest.skip("VAE config not found in safetensors metadata.")
|
||||
|
||||
try:
|
||||
from ltx_core.model.video_vae import VideoDecoderConfigurator, VideoEncoderConfigurator
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
ref_weights = _load_weights(official_path)
|
||||
encoder_weights = _select_vae_weights(ref_weights, "vae.encoder.")
|
||||
decoder_weights = _select_vae_weights(ref_weights, "vae.decoder.")
|
||||
if not encoder_weights or not decoder_weights:
|
||||
pytest.skip("VAE weights not found in safetensors file.")
|
||||
|
||||
fastvideo_weights_path = fastvideo_path / "model.safetensors"
|
||||
if not fastvideo_weights_path.exists():
|
||||
pytest.skip(f"FastVideo VAE weights not found at {fastvideo_weights_path}")
|
||||
fastvideo_weights = _load_weights(fastvideo_weights_path)
|
||||
fastvideo_encoder_weights = _select_vae_weights(
|
||||
fastvideo_weights, "encoder."
|
||||
)
|
||||
fastvideo_decoder_weights = _select_vae_weights(
|
||||
fastvideo_weights, "decoder."
|
||||
)
|
||||
if not fastvideo_encoder_weights or not fastvideo_decoder_weights:
|
||||
pytest.skip("FastVideo VAE weights not found in diffusers file.")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16 if torch.cuda.is_available() else torch.float32
|
||||
|
||||
fastvideo_encoder = LTX2VideoEncoder(config).to(device=device, dtype=precision)
|
||||
fastvideo_decoder = LTX2VideoDecoder(config).to(device=device, dtype=precision)
|
||||
ref_encoder = VideoEncoderConfigurator.from_config(config).to(device=device, dtype=precision)
|
||||
ref_decoder = VideoDecoderConfigurator.from_config(config).to(device=device, dtype=precision)
|
||||
|
||||
loaded_fastvideo_encoder, missing_fastvideo_encoder = _load_into_model(
|
||||
fastvideo_encoder.model, fastvideo_encoder_weights
|
||||
)
|
||||
loaded_ref_encoder, missing_ref_encoder = _load_into_model(ref_encoder, encoder_weights)
|
||||
loaded_fastvideo_decoder, missing_fastvideo_decoder = _load_into_model(
|
||||
fastvideo_decoder.model, fastvideo_decoder_weights
|
||||
)
|
||||
loaded_ref_decoder, missing_ref_decoder = _load_into_model(ref_decoder, decoder_weights)
|
||||
|
||||
if min(
|
||||
loaded_fastvideo_encoder,
|
||||
loaded_ref_encoder,
|
||||
loaded_fastvideo_decoder,
|
||||
loaded_ref_decoder,
|
||||
) == 0:
|
||||
pytest.skip("Failed to load VAE weights into one or more models.")
|
||||
if (
|
||||
missing_fastvideo_encoder
|
||||
or missing_ref_encoder
|
||||
or missing_fastvideo_decoder
|
||||
or missing_ref_decoder
|
||||
):
|
||||
print(f"[LTX2 VAE TEST] Missing encoder keys: {len(missing_fastvideo_encoder)}")
|
||||
print(f"[LTX2 VAE TEST] Missing decoder keys: {len(missing_fastvideo_decoder)}")
|
||||
pytest.skip("Missing VAE keys; cannot ensure parity.")
|
||||
|
||||
fastvideo_encoder.model.eval()
|
||||
fastvideo_decoder.model.eval()
|
||||
ref_encoder.eval()
|
||||
ref_decoder.eval()
|
||||
|
||||
fastvideo_decoder.model.decode_noise_scale = 0.0
|
||||
ref_decoder.decode_noise_scale = 0.0
|
||||
|
||||
batch_size = 1
|
||||
frames = 9
|
||||
height = 64
|
||||
width = 64
|
||||
video = torch.randn(
|
||||
batch_size,
|
||||
3,
|
||||
frames,
|
||||
height,
|
||||
width,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_latents = ref_encoder(video)
|
||||
fast_latents = fastvideo_encoder(video)
|
||||
|
||||
assert ref_latents.shape == fast_latents.shape
|
||||
assert ref_latents.dtype == fast_latents.dtype
|
||||
assert torch.isfinite(ref_latents).all(), "Reference encoder produced non-finite latents."
|
||||
assert torch.isfinite(fast_latents).all(), "FastVideo encoder produced non-finite latents."
|
||||
assert_close(ref_latents, fast_latents, atol=1e-2, rtol=1e-2)
|
||||
|
||||
timestep = torch.tensor([0.05], device=device, dtype=precision)
|
||||
with torch.no_grad():
|
||||
ref_decoded = ref_decoder(ref_latents, timestep=timestep)
|
||||
fast_decoded = fastvideo_decoder(fast_latents, timestep=timestep)
|
||||
|
||||
assert ref_decoded.shape == fast_decoded.shape
|
||||
assert ref_decoded.dtype == fast_decoded.dtype
|
||||
assert torch.isfinite(ref_decoded).all(), "Reference decoder produced non-finite output."
|
||||
assert torch.isfinite(fast_decoded).all(), "FastVideo decoder produced non-finite output."
|
||||
assert_close(ref_decoded, fast_decoded, atol=1e-2, rtol=1e-2)
|
||||
@@ -0,0 +1,132 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.configs.models.vaes import LTX2VAEConfig
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.loader.component_loader import VAELoader
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="LTX-2 VAE parity test requires CUDA.",
|
||||
)
|
||||
def test_ltx2_vae_parity_official():
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
fastvideo_path = Path(
|
||||
os.getenv("LTX2_VAE_PATH", str(diffusers_root / "vae"))
|
||||
)
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 weights not found at {official_path}")
|
||||
if not fastvideo_path.exists():
|
||||
pytest.skip(f"LTX-2 diffusers VAE not found at {fastvideo_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
from ltx_core.model.video_vae import (
|
||||
VAE_DECODER_COMFY_KEYS_FILTER,
|
||||
VAE_ENCODER_COMFY_KEYS_FILTER,
|
||||
VideoDecoderConfigurator,
|
||||
VideoEncoderConfigurator,
|
||||
)
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=str(fastvideo_path),
|
||||
vae_cpu_offload=False,
|
||||
pipeline_config=PipelineConfig(
|
||||
vae_config=LTX2VAEConfig(),
|
||||
vae_precision=precision_str,
|
||||
),
|
||||
)
|
||||
|
||||
loader = VAELoader()
|
||||
fastvideo_vae = loader.load(str(fastvideo_path), args).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
|
||||
encoder_builder = SingleGPUModelBuilder(
|
||||
model_class_configurator=VideoEncoderConfigurator,
|
||||
model_path=str(official_path),
|
||||
model_sd_ops=VAE_ENCODER_COMFY_KEYS_FILTER,
|
||||
)
|
||||
decoder_builder = SingleGPUModelBuilder(
|
||||
model_class_configurator=VideoDecoderConfigurator,
|
||||
model_path=str(official_path),
|
||||
model_sd_ops=VAE_DECODER_COMFY_KEYS_FILTER,
|
||||
)
|
||||
ref_encoder = encoder_builder.build(
|
||||
device=device, dtype=precision
|
||||
).to(device=device, dtype=precision)
|
||||
ref_decoder = decoder_builder.build(
|
||||
device=device, dtype=precision
|
||||
).to(device=device, dtype=precision)
|
||||
|
||||
fastvideo_vae.encoder.eval()
|
||||
fastvideo_vae.decoder.eval()
|
||||
ref_encoder.eval()
|
||||
ref_decoder.eval()
|
||||
|
||||
if hasattr(fastvideo_vae.decoder, "decode_noise_scale"):
|
||||
fastvideo_vae.decoder.decode_noise_scale = 0.0
|
||||
if hasattr(ref_decoder, "decode_noise_scale"):
|
||||
ref_decoder.decode_noise_scale = 0.0
|
||||
|
||||
batch_size = 1
|
||||
frames = 9
|
||||
height = 64
|
||||
width = 64
|
||||
video = torch.randn(
|
||||
batch_size,
|
||||
3,
|
||||
frames,
|
||||
height,
|
||||
width,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_latents = ref_encoder(video)
|
||||
fast_latents = fastvideo_vae.encoder(video)
|
||||
|
||||
assert ref_latents.shape == fast_latents.shape
|
||||
assert ref_latents.dtype == fast_latents.dtype
|
||||
assert torch.isfinite(ref_latents).all(), "Reference encoder produced non-finite latents."
|
||||
assert torch.isfinite(fast_latents).all(), "FastVideo encoder produced non-finite latents."
|
||||
assert_close(ref_latents, fast_latents, atol=1e-2, rtol=1e-2)
|
||||
|
||||
timestep = torch.tensor([0.05], device=device, dtype=precision)
|
||||
with torch.no_grad():
|
||||
ref_decoded = ref_decoder(ref_latents, timestep=timestep)
|
||||
fast_decoded = fastvideo_vae.decoder(fast_latents, timestep=timestep)
|
||||
|
||||
assert ref_decoded.shape == fast_decoded.shape
|
||||
assert ref_decoded.dtype == fast_decoded.dtype
|
||||
assert torch.isfinite(ref_decoded).all(), "Reference decoder produced non-finite output."
|
||||
assert torch.isfinite(fast_decoded).all(), "FastVideo decoder produced non-finite output."
|
||||
assert_close(ref_decoded, fast_decoded, atol=1e-2, rtol=1e-2)
|
||||
Reference in New Issue
Block a user