Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
141a1140f6 | ||
|
|
6294015389 | ||
|
|
e31b6c9e90 | ||
|
|
bf0ff21eeb | ||
|
|
6f937102ad | ||
|
|
0164e93019 | ||
|
|
67e457aa92 | ||
|
|
d795f0c443 | ||
|
|
02452dd6e7 | ||
|
|
91ef24bc14 | ||
|
|
e76e9fda15 | ||
|
|
3b17f5a621 | ||
|
|
d758878705 | ||
|
|
689e629420 | ||
|
|
873dc9695f | ||
|
|
bfc0f46d61 | ||
|
|
39907dbe4d | ||
|
|
abdd0c9b6a | ||
|
|
f32a12200d | ||
|
|
f1d2c9e6b7 | ||
|
|
450579cb42 | ||
|
|
26d7d6cc08 | ||
|
|
44f0124eaa | ||
|
|
d3ace51394 | ||
|
|
58954c660b |
@@ -1,34 +0,0 @@
|
||||
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()
|
||||
Executable
+129
@@ -0,0 +1,129 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Change to FastVideo root directory (3 levels up from this script)
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
FASTVIDEO_ROOT="$(cd "$SCRIPT_DIR/../../.." && pwd)"
|
||||
cd "$FASTVIDEO_ROOT"
|
||||
|
||||
# Add FastVideo root to PYTHONPATH so Python can find the fastvideo package
|
||||
export PYTHONPATH="$FASTVIDEO_ROOT${PYTHONPATH:+:$PYTHONPATH}"
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
RL_DATASET_DIR="data/ocr/" # Path to RL prompt dataset directory (should contain train.txt and test.txt)
|
||||
VALIDATION_DATASET_FILE="$SCRIPT_DIR/validation.json"
|
||||
NUM_GPUS=1
|
||||
|
||||
# use GPU 3
|
||||
export CUDA_VISIBLE_DEVICES=3
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_grpo"
|
||||
--output_dir "checkpoints/wan_t2v_grpo"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 4
|
||||
# --train_sp_batch_size 4
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 5
|
||||
--num_height 240
|
||||
--num_width 416
|
||||
--num_frames 33
|
||||
--lora_rank 32
|
||||
--lora_training True
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
# --use-fsdp-inference False
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments (for RL prompt dataset)
|
||||
dataset_args=(
|
||||
--data_path $RL_DATASET_DIR # Used as fallback if rl_dataset_path not set
|
||||
--rl_dataset_path $RL_DATASET_DIR # RL prompt dataset directory
|
||||
--rl_dataset_type "text" # "text" or "geneval"
|
||||
--rl_num_image_per_prompt 4 # k parameter (number of samples per prompt)
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation True
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 5
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 10
|
||||
--training_state_checkpointing_steps 10
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# RL-specific arguments
|
||||
rl_args=(
|
||||
--inference_mode False
|
||||
--rl_mode True
|
||||
--rl_algorithm "grpo"
|
||||
--rl_kl_beta 0.004 # KL regularization coefficient
|
||||
--rl_policy_clip_range 0.2 # Policy clipping range for GRPO
|
||||
--rl_kl_reward 0.0 # KL reward coefficient (typically 0)
|
||||
--rl_global_std False # Use per-prompt std (recommended for GRPO)
|
||||
--rl_per_prompt_stat_tracking True # Enable per-prompt stat tracking
|
||||
--rl_warmup_steps 0 # Number of warmup steps (SFT before RL)
|
||||
--reward-models "{\"paddle_ocr\": 1.0}" # use video_ocr reward function
|
||||
)
|
||||
|
||||
# CFG arguments
|
||||
cfg_args=(
|
||||
--guidance_scale 1.0 # use guidance_scale > 1.0 to enable CFG
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0 # No CFG during training (CFG used in sampling)
|
||||
--dit_precision "fp32"
|
||||
# --dit_precision "bf16"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --resume_from_checkpoint "checkpoints/wan_t2v_grpo/checkpoint-XXX"
|
||||
--enable-gradient-checkpointing-type "full"
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port 29501 \
|
||||
"$FASTVIDEO_ROOT/fastvideo/training/wan_rl_training_pipeline.py" \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${rl_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -2,16 +2,5 @@ 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",
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.configs.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
@@ -1,31 +0,0 @@
|
||||
# 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,12 +3,11 @@ 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", "LTX2VideoConfig"
|
||||
"LongCatVideoConfig"
|
||||
]
|
||||
|
||||
@@ -1,84 +0,0 @@
|
||||
# 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,11 +8,10 @@ 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", "LTX2GemmaConfig"
|
||||
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config"
|
||||
]
|
||||
|
||||
@@ -1,48 +0,0 @@
|
||||
# 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,7 +2,6 @@ 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
|
||||
|
||||
@@ -13,5 +12,4 @@ __all__ = [
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -1,45 +0,0 @@
|
||||
# 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,7 +4,6 @@ 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
|
||||
@@ -17,6 +16,5 @@ __all__ = [
|
||||
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -1,50 +0,0 @@
|
||||
# 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,7 +9,6 @@ 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 (
|
||||
@@ -65,9 +64,6 @@ 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,
|
||||
@@ -106,8 +102,6 @@ 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
|
||||
}
|
||||
|
||||
@@ -129,7 +123,6 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig,
|
||||
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
|
||||
"ltx2": LTX2T2VConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
# 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,7 +10,6 @@ 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 (
|
||||
@@ -41,36 +40,48 @@ 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":
|
||||
@@ -91,9 +102,12 @@ 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":
|
||||
@@ -103,10 +117,6 @@ 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
|
||||
}
|
||||
|
||||
@@ -134,8 +144,6 @@ 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
|
||||
}
|
||||
|
||||
@@ -156,7 +164,6 @@ 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
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset,
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.dataset.rl_prompt_dataset import build_rl_prompt_dataloader
|
||||
|
||||
|
||||
def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
@@ -47,5 +48,6 @@ def gettextdataset(args) -> TextDataset:
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
"VideoCaptionMergedDataset", "TextDataset",
|
||||
"build_rl_prompt_dataloader"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import torch
|
||||
from torch.utils.data import Dataset, DataLoader, Sampler
|
||||
import json
|
||||
import os
|
||||
|
||||
|
||||
class TextPromptDataset(Dataset):
|
||||
"""Dataset for loading text prompts from a simple text file (one prompt per line)."""
|
||||
|
||||
def __init__(self, dataset, split='train'):
|
||||
self.file_path = os.path.join(dataset, f'{split}.txt')
|
||||
with open(self.file_path, 'r') as f:
|
||||
self.prompts = [line.strip() for line in f.readlines()]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.prompts)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {"prompt": self.prompts[idx], "metadata": {}}
|
||||
|
||||
@staticmethod
|
||||
def collate_fn(examples):
|
||||
prompts = [example["prompt"] for example in examples]
|
||||
metadatas = [example["metadata"] for example in examples]
|
||||
return prompts, metadatas
|
||||
|
||||
|
||||
class GenevalPromptDataset(Dataset):
|
||||
"""Dataset for loading prompts with metadata from JSONL files (e.g., GenEval format)."""
|
||||
|
||||
def __init__(self, dataset, split='train'):
|
||||
self.file_path = os.path.join(dataset, f'{split}_metadata.jsonl')
|
||||
with open(self.file_path, 'r', encoding='utf-8') as f:
|
||||
self.metadatas = [json.loads(line) for line in f]
|
||||
self.prompts = [item['prompt'] for item in self.metadatas]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.prompts)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {"prompt": self.prompts[idx], "metadata": self.metadatas[idx]}
|
||||
|
||||
@staticmethod
|
||||
def collate_fn(examples):
|
||||
prompts = [example["prompt"] for example in examples]
|
||||
metadatas = [example["metadata"] for example in examples]
|
||||
return prompts, metadatas
|
||||
|
||||
|
||||
class KRepeatSampler(Sampler):
|
||||
"""Sampler that repeats each sample k times, ensuring synchronized random selection. For single-node training, set num_replicas=1 and rank=0."""
|
||||
|
||||
def __init__(self, dataset, batch_size, k, num_replicas, rank, seed=0):
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size # Batch size per GPU/card
|
||||
self.k = k # Number of repetitions per sample
|
||||
self.num_replicas = num_replicas # Total number of GPUs/cards
|
||||
self.rank = rank # Current GPU/card rank
|
||||
self.seed = seed # Random seed for synchronization
|
||||
|
||||
# Calculate the number of unique samples needed for each iteration
|
||||
self.total_samples = self.num_replicas * self.batch_size
|
||||
assert self.total_samples % self.k == 0, f"k can not div n*b, k{k}-num_replicas{num_replicas}-batch_size{batch_size}"
|
||||
self.m = self.total_samples // self.k # different number of samples
|
||||
self.step = 0
|
||||
|
||||
def __iter__(self):
|
||||
while True:
|
||||
# Generate a deterministic random sequence to ensure all cards are synchronized
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.seed + self.step)
|
||||
|
||||
# Randomly select m unique samples
|
||||
indices = torch.randperm(len(self.dataset), generator=g)[:self.m].tolist()
|
||||
|
||||
# Repeat each sample k times to generate a total of n*b samples
|
||||
repeated_indices = [idx for idx in indices for _ in range(self.k)]
|
||||
|
||||
# Shuffle the order to ensure even distribution
|
||||
shuffled_indices = torch.randperm(len(repeated_indices), generator=g).tolist()
|
||||
shuffled_samples = [repeated_indices[i] for i in shuffled_indices]
|
||||
|
||||
# Split samples among all cards
|
||||
per_card_samples = []
|
||||
for i in range(self.num_replicas):
|
||||
start = i * self.batch_size
|
||||
end = start + self.batch_size
|
||||
per_card_samples.append(shuffled_samples[start:end])
|
||||
|
||||
# Return the sample indices for the current card
|
||||
yield per_card_samples[self.rank]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dataset) // self.batch_size
|
||||
|
||||
def set_step(self, step):
|
||||
"""Used to synchronize the random state for different epochs."""
|
||||
self.step = step
|
||||
|
||||
|
||||
def build_rl_prompt_dataloader(
|
||||
dataset_path: str,
|
||||
dataset_type: str = "text",
|
||||
split: str = "train",
|
||||
train_batch_size: int = 8,
|
||||
test_batch_size: int = 8,
|
||||
k: int = 1,
|
||||
seed: int = 42,
|
||||
train_num_workers: int = 1,
|
||||
test_num_workers: int = 8,
|
||||
num_replicas: int = 1,
|
||||
rank: int = 0,
|
||||
) -> tuple[DataLoader, DataLoader]:
|
||||
"""
|
||||
Factory function to create train and test dataloaders for RL prompt datasets.
|
||||
|
||||
Args:
|
||||
dataset_path: Path to dataset directory
|
||||
dataset_type: "text" for TextPromptDataset or "geneval" for GenevalPromptDataset
|
||||
split: Dataset split ("train" or "test")
|
||||
train_batch_size: Batch size per GPU for training
|
||||
test_batch_size: Batch size for testing
|
||||
k: Number of times to repeat each sample (num_image_per_prompt)
|
||||
seed: Random seed for sampler synchronization
|
||||
train_num_workers: Number of workers for training dataloader
|
||||
test_num_workers: Number of workers for test dataloader
|
||||
num_replicas: Number of replicas (default 1 for single-node)
|
||||
rank: Rank of current process (default 0 for single-node)
|
||||
|
||||
Returns:
|
||||
Tuple of (train_dataloader, test_dataloader)
|
||||
"""
|
||||
# Create datasets based on type
|
||||
if dataset_type == "text":
|
||||
train_dataset = TextPromptDataset(dataset_path, 'train')
|
||||
test_dataset = TextPromptDataset(dataset_path, 'test')
|
||||
collate_fn = TextPromptDataset.collate_fn
|
||||
elif dataset_type == "geneval":
|
||||
train_dataset = GenevalPromptDataset(dataset_path, 'train')
|
||||
test_dataset = GenevalPromptDataset(dataset_path, 'test')
|
||||
collate_fn = GenevalPromptDataset.collate_fn
|
||||
else:
|
||||
raise ValueError(f"Unknown dataset_type: {dataset_type}. Must be 'text' or 'geneval'")
|
||||
|
||||
# Create infinite-loop training sampler
|
||||
train_sampler = KRepeatSampler(
|
||||
dataset=train_dataset,
|
||||
batch_size=train_batch_size,
|
||||
k=k,
|
||||
num_replicas=num_replicas,
|
||||
rank=rank,
|
||||
seed=seed
|
||||
)
|
||||
|
||||
# Create training dataloader with batch_sampler (infinite loop)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_sampler=train_sampler,
|
||||
num_workers=train_num_workers,
|
||||
collate_fn=collate_fn,
|
||||
)
|
||||
|
||||
# Create standard test dataloader
|
||||
test_dataloader = DataLoader(
|
||||
test_dataset,
|
||||
batch_size=test_batch_size,
|
||||
collate_fn=collate_fn,
|
||||
shuffle=False,
|
||||
num_workers=test_num_workers,
|
||||
)
|
||||
|
||||
return train_dataloader, test_dataloader, train_dataset, test_dataset
|
||||
|
||||
@@ -18,8 +18,6 @@ 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
|
||||
@@ -391,11 +389,6 @@ 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
|
||||
@@ -403,7 +396,6 @@ 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,
|
||||
@@ -413,98 +405,6 @@ 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:
|
||||
|
||||
+311
-83
@@ -166,14 +166,6 @@ 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: {
|
||||
@@ -211,44 +203,8 @@ 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
|
||||
@@ -369,44 +325,6 @@ 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",
|
||||
@@ -822,6 +740,271 @@ def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
return _current_fastvideo_args
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class RLArgs:
|
||||
"""
|
||||
Reinforcement Learning (RL) specific arguments
|
||||
"""
|
||||
# ============================================================================
|
||||
# SHARED RL CONFIGURATION
|
||||
rl_mode: bool = False # Enable RL training mode
|
||||
rl_algorithm: str = "grpo" # RL algorithm to use: "grpo", "ppo", "dpo"
|
||||
|
||||
# Trajectory collection
|
||||
num_rollouts: int = 4 # Number of rollouts to collect per training step
|
||||
rollout_steps: str = "20,30" # Random intermediate steps for sampling (comma-separated)
|
||||
noise_injection_min: int = 10 # Minimum timestep for noise injection
|
||||
noise_injection_max: int = 40 # Maximum timestep for noise injection
|
||||
use_sde_sampling: bool = True # Use SDE sampling (Flow-GRPO-Fast)
|
||||
num_denoising_steps: int = 2 # Number of denoising steps per trajectory (1-2 for fast)
|
||||
|
||||
# Advantage estimation
|
||||
gamma: float = 0.99 # Discount factor for returns
|
||||
lambda_param: float = 0.95 # GAE lambda parameter
|
||||
use_gae: bool = True # Use Generalized Advantage Estimation
|
||||
normalize_advantages: bool = True # Normalize advantages before policy update
|
||||
|
||||
# Reward models
|
||||
reward_models: dict[str, float] = field(default_factory=lambda: {"dummy": 1.0}) # reward models (names, weight)
|
||||
value_model_path: str = "" # Path to value model (can be empty to train from scratch)
|
||||
value_model_share_backbone: bool = False # Share transformer backbone between policy and value
|
||||
|
||||
# Training schedule
|
||||
warmup_steps: int = 1000 # Collect SFT-style data before starting RL
|
||||
collect_on_policy: bool = True # Collect fresh rollouts each step (on-policy)
|
||||
timestep_fraction: float = 0.99 # Fraction of timesteps to train on
|
||||
num_inner_epochs: int = 1 # Number of inner epochs per outer epoch
|
||||
|
||||
# KL regularization
|
||||
kl_beta: float = 0.004 # KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)
|
||||
kl_reward: float = 0.0 # KL reward coefficient (alternative to KL loss, typically 0)
|
||||
|
||||
# SFT integration
|
||||
sft_weight: float = 0.0 # SFT loss weight for supervised learning in RL training
|
||||
sft_batch_size: int = 3 # Batch size for SFT data
|
||||
|
||||
# CFG
|
||||
guidance_scale = 1.0 # use guidance_scale > 1.0 to enable CFG
|
||||
|
||||
# Statistics tracking
|
||||
global_std: bool = False # Use global std across all samples vs per-group std
|
||||
per_prompt_stat_tracking: bool = True # Track statistics per prompt
|
||||
|
||||
# Training options
|
||||
use_diffusion_loss: bool = True # Use diffusion loss in training
|
||||
|
||||
# ============================================================================
|
||||
# GRPO-SPECIFIC CONFIGURATION
|
||||
|
||||
# Policy optimization
|
||||
grpo_policy_clip_range: float = 0.001 # PPO-style clipping range for policy ratio
|
||||
grpo_value_clip_range: float = 0.2 # Value function clipping range
|
||||
grpo_num_policy_epochs: int = 1 # Number of policy update epochs (GRPO typically uses 1)
|
||||
grpo_num_value_epochs: int = 1 # Number of value function update epochs
|
||||
grpo_target_kl: float = 0.01 # Target KL divergence for early stopping
|
||||
grpo_entropy_coef: float = 0.0 # Entropy coefficient for exploration
|
||||
grpo_value_loss_coef: float = 0.5 # Value loss coefficient
|
||||
|
||||
# GRPO-Guard safety mechanisms
|
||||
grpo_use_grpo_guard: bool = True # Enable GRPO-Guard safety mechanisms
|
||||
grpo_ratio_norm_correction: bool = True # RatioNorm: correct importance ratio bias
|
||||
grpo_gradient_reweighting: bool = True # Reweight gradients across denoising steps
|
||||
grpo_max_importance_ratio: float = 10.0 # Clip importance ratios above this value
|
||||
|
||||
# ============================================================================
|
||||
# DPO-SPECIFIC CONFIGURATION
|
||||
|
||||
dpo_beta: float = 100.0 # DPO regularization parameter (typically much larger than GRPO beta)
|
||||
dpo_ref_update_step: int = 10000000 # Reference model update frequency for OnlineDPO
|
||||
dpo_label_smoothing: float = 0.0 # Label smoothing for DPO loss
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
"""Add RL-specific CLI arguments to the parser."""
|
||||
# RL (Reinforcement Learning) arguments
|
||||
parser.add_argument("--rl-mode",
|
||||
action=StoreBoolean,
|
||||
help="Enable RL training mode")
|
||||
parser.add_argument("--rl-algorithm",
|
||||
type=str,
|
||||
default=RLArgs.rl_algorithm,
|
||||
choices=["grpo", "ppo", "dpo"],
|
||||
help="RL algorithm to use (grpo, ppo, dpo)")
|
||||
|
||||
# Trajectory collection (Flow-GRPO-Fast)
|
||||
parser.add_argument("--rl-num-rollouts",
|
||||
type=int,
|
||||
default=RLArgs.num_rollouts,
|
||||
help="Number of rollouts to collect per training step")
|
||||
parser.add_argument("--rl-rollout-steps",
|
||||
type=str,
|
||||
default=RLArgs.rollout_steps,
|
||||
help="Random intermediate steps for sampling (comma-separated)")
|
||||
parser.add_argument("--rl-noise-injection-min",
|
||||
type=int,
|
||||
default=RLArgs.noise_injection_min,
|
||||
help="Minimum timestep for noise injection")
|
||||
parser.add_argument("--rl-noise-injection-max",
|
||||
type=int,
|
||||
default=RLArgs.noise_injection_max,
|
||||
help="Maximum timestep for noise injection")
|
||||
parser.add_argument("--rl-use-sde-sampling",
|
||||
action=StoreBoolean,
|
||||
help="Use SDE sampling (Flow-GRPO-Fast)")
|
||||
parser.add_argument("--rl-num-denoising-steps",
|
||||
type=int,
|
||||
default=RLArgs.num_denoising_steps,
|
||||
help="Number of denoising steps per trajectory (1-2 for fast)")
|
||||
|
||||
# Advantage estimation
|
||||
parser.add_argument("--rl-gamma",
|
||||
type=float,
|
||||
default=RLArgs.gamma,
|
||||
help="Discount factor for returns")
|
||||
parser.add_argument("--rl-lambda",
|
||||
type=float,
|
||||
default=RLArgs.lambda_param,
|
||||
help="GAE lambda parameter")
|
||||
parser.add_argument("--rl-use-gae",
|
||||
action=StoreBoolean,
|
||||
help="Use Generalized Advantage Estimation")
|
||||
parser.add_argument("--rl-normalize-advantages",
|
||||
action=StoreBoolean,
|
||||
help="Normalize advantages before policy update")
|
||||
|
||||
# Policy optimization (GRPO/PPO)
|
||||
parser.add_argument("--rl-policy-clip-range",
|
||||
type=float,
|
||||
default=RLArgs.grpo_policy_clip_range,
|
||||
dest="grpo_policy_clip_range", # Map to RLArgs field name
|
||||
help="PPO-style clipping range for policy ratio")
|
||||
parser.add_argument("--rl-value-clip-range",
|
||||
type=float,
|
||||
default=RLArgs.grpo_value_clip_range,
|
||||
help="Value function clipping range")
|
||||
parser.add_argument("--rl-num-policy-epochs",
|
||||
type=int,
|
||||
default=RLArgs.grpo_num_policy_epochs,
|
||||
help="Number of policy update epochs (GRPO typically uses 1)")
|
||||
parser.add_argument("--rl-num-value-epochs",
|
||||
type=int,
|
||||
default=RLArgs.grpo_num_value_epochs,
|
||||
help="Number of value function update epochs")
|
||||
parser.add_argument("--rl-target-kl",
|
||||
type=float,
|
||||
default=RLArgs.grpo_target_kl,
|
||||
help="Target KL divergence for early stopping")
|
||||
parser.add_argument("--rl-entropy-coef",
|
||||
type=float,
|
||||
default=RLArgs.grpo_entropy_coef,
|
||||
help="Entropy coefficient for exploration")
|
||||
parser.add_argument("--rl-value-loss-coef",
|
||||
type=float,
|
||||
default=RLArgs.grpo_value_loss_coef,
|
||||
help="Value loss coefficient")
|
||||
|
||||
# GRPO-Guard (safety mechanisms)
|
||||
parser.add_argument("--rl-use-grpo-guard",
|
||||
action=StoreBoolean,
|
||||
help="Enable GRPO-Guard safety mechanisms")
|
||||
parser.add_argument("--rl-ratio-norm-correction",
|
||||
action=StoreBoolean,
|
||||
help="RatioNorm: correct importance ratio bias")
|
||||
parser.add_argument("--rl-gradient-reweighting",
|
||||
action=StoreBoolean,
|
||||
help="Reweight gradients across denoising steps")
|
||||
parser.add_argument("--rl-max-importance-ratio",
|
||||
type=float,
|
||||
default=RLArgs.grpo_max_importance_ratio,
|
||||
help="Clip importance ratios above this value")
|
||||
|
||||
# Reward models
|
||||
parser.add_argument("--reward-models",
|
||||
type=str,
|
||||
default='{"dummy": 1.0}',
|
||||
help="Reward models as JSON dict (e.g., '{\"video_ocr\": 1.0, \"pickscore\": 0.5}')")
|
||||
parser.add_argument("--value-model-path",
|
||||
type=str,
|
||||
default=RLArgs.value_model_path,
|
||||
help="Path to value model (can be empty to train from scratch)")
|
||||
parser.add_argument("--value-model-share-backbone",
|
||||
action=StoreBoolean,
|
||||
help="Share transformer backbone between policy and value")
|
||||
|
||||
# Training schedule
|
||||
parser.add_argument("--rl-warmup-steps",
|
||||
type=int,
|
||||
default=RLArgs.warmup_steps,
|
||||
help="Collect SFT-style data before starting RL")
|
||||
parser.add_argument("--rl-collect-on-policy",
|
||||
action=StoreBoolean,
|
||||
help="Collect fresh rollouts each step (on-policy)")
|
||||
parser.add_argument("--rl-timestep-fraction",
|
||||
type=float,
|
||||
default=RLArgs.timestep_fraction,
|
||||
help="Fraction of timesteps to train on")
|
||||
parser.add_argument("--rl-num-inner-epochs",
|
||||
type=int,
|
||||
default=RLArgs.num_inner_epochs,
|
||||
help="Number of inner epochs per outer epoch")
|
||||
|
||||
# KL regularization
|
||||
parser.add_argument("--rl-kl-beta",
|
||||
type=float,
|
||||
default=RLArgs.kl_beta,
|
||||
dest="kl_beta", # Map CLI arg to RLArgs field name
|
||||
help="KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)")
|
||||
parser.add_argument("--rl-kl-reward",
|
||||
type=float,
|
||||
default=RLArgs.kl_reward,
|
||||
help="KL reward coefficient (alternative to KL loss, typically 0)")
|
||||
|
||||
# SFT integration
|
||||
parser.add_argument("--rl-sft-weight",
|
||||
type=float,
|
||||
default=RLArgs.sft_weight,
|
||||
help="SFT loss weight for supervised learning in RL training")
|
||||
parser.add_argument("--rl-sft-batch-size",
|
||||
type=int,
|
||||
default=RLArgs.sft_batch_size,
|
||||
help="Batch size for SFT data")
|
||||
|
||||
# CFG settings
|
||||
parser.add_argument("--guidance-scale",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Guidance scale for CFG")
|
||||
|
||||
# Statistics tracking
|
||||
parser.add_argument("--rl-global-std",
|
||||
action=StoreBoolean,
|
||||
help="Use global std across all samples vs per-group std")
|
||||
parser.add_argument("--rl-per-prompt-stat-tracking",
|
||||
action=StoreBoolean,
|
||||
help="Track statistics per prompt")
|
||||
|
||||
# Training options
|
||||
parser.add_argument("--rl-use-diffusion-loss",
|
||||
action=StoreBoolean,
|
||||
help="Use diffusion loss in training")
|
||||
|
||||
# DPO-specific
|
||||
parser.add_argument("--dpo-beta",
|
||||
type=float,
|
||||
default=RLArgs.dpo_beta,
|
||||
help="DPO regularization parameter (typically much larger than GRPO beta)")
|
||||
parser.add_argument("--dpo-ref-update-step",
|
||||
type=int,
|
||||
default=RLArgs.dpo_ref_update_step,
|
||||
help="Reference model update frequency for OnlineDPO")
|
||||
parser.add_argument("--dpo-label-smoothing",
|
||||
type=float,
|
||||
default=RLArgs.dpo_label_smoothing,
|
||||
help="Label smoothing for DPO loss")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TrainingArgs(FastVideoArgs):
|
||||
"""
|
||||
@@ -834,6 +1017,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
num_height: int = 0
|
||||
num_width: int = 0
|
||||
num_frames: int = 0
|
||||
|
||||
# RL dataset configuration (for RL prompt datasets)
|
||||
rl_dataset_path: str = "" # Path to RL prompt dataset directory (defaults to data_path if not set)
|
||||
rl_dataset_type: str = "text" # "text" or "geneval"
|
||||
rl_num_image_per_prompt: int = 4 # k parameter for KRepeatSampler (num_image_per_prompt)
|
||||
|
||||
train_batch_size: int = 0
|
||||
num_latent_t: int = 0
|
||||
@@ -944,6 +1132,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
|
||||
# Nested RL configuration
|
||||
rl_args: RLArgs = dataclasses.field(default_factory=RLArgs)
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
provided_args = clean_cli_args(args)
|
||||
@@ -968,6 +1159,25 @@ class TrainingArgs(FastVideoArgs):
|
||||
kwargs[attr] = WorkloadType.from_string(
|
||||
workload_type_value) if isinstance(
|
||||
workload_type_value, str) else workload_type_value
|
||||
elif attr == 'rl_args':
|
||||
# Construct nested RLArgs from CLI arguments
|
||||
rl_kwargs = {}
|
||||
for rl_field in dataclasses.fields(RLArgs):
|
||||
rl_attr = rl_field.name
|
||||
if hasattr(args, rl_attr):
|
||||
value = getattr(args, rl_attr)
|
||||
# Special handling for reward_models: parse JSON string to dict
|
||||
if rl_attr == 'reward_models' and isinstance(value, str):
|
||||
rl_kwargs[rl_attr] = json.loads(value) if value else {}
|
||||
else:
|
||||
rl_kwargs[rl_attr] = value
|
||||
else:
|
||||
# Use default value from RLArgs
|
||||
if rl_field.default_factory is not dataclasses.MISSING:
|
||||
rl_kwargs[rl_attr] = rl_field.default_factory()
|
||||
elif rl_field.default is not dataclasses.MISSING:
|
||||
rl_kwargs[rl_attr] = rl_field.default
|
||||
kwargs[attr] = RLArgs(**rl_kwargs)
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
# Get the field to check its default value
|
||||
@@ -997,11 +1207,26 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--data-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to parquet files")
|
||||
help="Path to parquet files (or RL prompt dataset directory for RL training)")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of workers for dataloader")
|
||||
|
||||
# RL dataset arguments (optional, defaults to data_path)
|
||||
parser.add_argument("--rl-dataset-path",
|
||||
type=str,
|
||||
default="",
|
||||
help="Path to RL prompt dataset directory (defaults to --data-path if not set)")
|
||||
parser.add_argument("--rl-dataset-type",
|
||||
type=str,
|
||||
default="text",
|
||||
choices=["text", "geneval"],
|
||||
help="RL dataset type: 'text' for TextPromptDataset or 'geneval' for GenevalPromptDataset")
|
||||
parser.add_argument("--rl-num-image-per-prompt",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Number of times to repeat each prompt (k parameter for KRepeatSampler)")
|
||||
parser.add_argument("--num-height",
|
||||
type=int,
|
||||
required=True,
|
||||
@@ -1366,6 +1591,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
|
||||
# RL (Reinforcement Learning) arguments
|
||||
RLArgs.add_cli_args(parser)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -1,9 +0,0 @@
|
||||
# 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
@@ -1,563 +0,0 @@
|
||||
# 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,9 +80,6 @@ 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"),
|
||||
@@ -245,47 +242,6 @@ 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?
|
||||
@@ -533,20 +489,8 @@ 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
|
||||
|
||||
@@ -557,12 +501,15 @@ 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.get("_class_name")
|
||||
class_name = config.pop("_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:
|
||||
@@ -596,29 +543,8 @@ class VAELoader(ComponentLoader):
|
||||
vae.load_state_dict(sd, strict=False)
|
||||
return vae.eval()
|
||||
|
||||
# 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)
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(target_device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -627,101 +553,17 @@ 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.
|
||||
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
|
||||
|
||||
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))
|
||||
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."""
|
||||
|
||||
@@ -837,18 +679,7 @@ class TransformerLoader(ComponentLoader):
|
||||
model = model.eval()
|
||||
|
||||
if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload:
|
||||
# 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
|
||||
)
|
||||
enable_layerwise_offload(model)
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -153,6 +153,7 @@ def maybe_load_fsdp_model(
|
||||
f"Unexpected param or buffer {n} on meta device.")
|
||||
# Avoid unintended computation graph accumulation during inference
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
|
||||
p.requires_grad = False
|
||||
|
||||
compile_in_loader = enable_torch_compile and training_mode
|
||||
|
||||
@@ -33,7 +33,6 @@ _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 = {
|
||||
@@ -55,7 +54,6 @@ _TEXT_ENCODER_MODELS = {
|
||||
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"Qwen2_5_VLForConditionalGeneration":
|
||||
("encoders", "reason1", "Reason1TextEncoder"),
|
||||
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
@@ -69,14 +67,7 @@ _VAE_MODELS = {
|
||||
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"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"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
|
||||
}
|
||||
|
||||
_SCHEDULERS = {
|
||||
@@ -100,7 +91,6 @@ _FAST_VIDEO_MODELS = {
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
**_AUDIO_MODELS,
|
||||
**_SCHEDULERS,
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,150 +0,0 @@
|
||||
# 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
|
||||
@@ -201,8 +201,6 @@ class ComposedPipelineBase(ABC):
|
||||
# fwd, bwd, and other operations' precision.
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
pipe = cls(model_path,
|
||||
fastvideo_args,
|
||||
required_config_modules=required_config_modules,
|
||||
|
||||
@@ -67,6 +67,21 @@ class ForwardBatch:
|
||||
execution, allowing methods to update specific components without needing
|
||||
to manage numerous individual parameters.
|
||||
"""
|
||||
|
||||
@dataclass
|
||||
class RLData:
|
||||
"""RL-specific data collection options and outputs."""
|
||||
enabled: bool = False
|
||||
collect_log_probs: bool = True
|
||||
collect_kl: bool = False
|
||||
kl_reward: float = 0.0
|
||||
store_trajectory: bool = True
|
||||
keep_trajectory_on_cpu: bool = False
|
||||
log_probs: torch.Tensor | None = None
|
||||
kl: torch.Tensor | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
|
||||
# TODO(will): double check that args are separate from fastvideo_args
|
||||
# properly. Also maybe think about providing an abstraction for pipeline
|
||||
# specific arguments.
|
||||
@@ -197,6 +212,9 @@ class ForwardBatch:
|
||||
logging_info: PipelineLoggingInfo = field(
|
||||
default_factory=PipelineLoggingInfo)
|
||||
|
||||
# RL data collection
|
||||
rl_data: "ForwardBatch.RLData" = field(default_factory=RLData)
|
||||
|
||||
def __post_init__(self):
|
||||
"""Initialize dependent fields after dataclass initialization."""
|
||||
|
||||
@@ -267,6 +285,36 @@ class TrainingBatch:
|
||||
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# RL/GRPO-specific attributes
|
||||
reward_scores: torch.Tensor | None = None # Computed rewards from reward models
|
||||
log_probs: torch.Tensor | None = None # Current policy log probabilities [B, num_steps] or [B]
|
||||
old_log_probs: torch.Tensor | None = None # Old policy log probs (for importance ratio) [B, num_steps] or [B]
|
||||
advantages: torch.Tensor | None = None # GAE advantages [B, num_steps] or [B]
|
||||
returns: torch.Tensor | None = None # TD returns (advantages + values) [B, num_steps] or [B]
|
||||
values: torch.Tensor | None = None # Value function predictions [B]
|
||||
old_values: torch.Tensor | None = None # Old value predictions (for clipping) [B]
|
||||
|
||||
# GRPO sampling-specific attributes
|
||||
kl: torch.Tensor | None = None # KL divergences from sampling [B, num_steps] (if kl_reward > 0)
|
||||
prompt_ids: torch.Tensor | None = None # Prompt token IDs for stat tracking [B, seq_len]
|
||||
prompt_embeds: torch.Tensor | None = None # Prompt embeddings used in sampling [B, seq_len, hidden_dim]
|
||||
negative_prompt_embeds: torch.Tensor | None = None # Negative prompt embeddings for CFG [B, seq_len, hidden_dim]
|
||||
|
||||
# RL loss components
|
||||
policy_loss: float = 0.0 # GRPO/PPO policy loss
|
||||
value_loss: float = 0.0 # Value function loss
|
||||
kl_divergence: float = 0.0 # KL(new_policy || old_policy)
|
||||
importance_ratio: float = 1.0 # exp(log_prob - old_log_prob)
|
||||
clip_fraction: float = 0.0 # Fraction of ratios that were clipped
|
||||
|
||||
# RL metrics
|
||||
advantage_mean: float = 0.0 # Mean advantage (should be ~0 after normalization)
|
||||
advantage_std: float = 1.0 # Std of advantages
|
||||
reward_mean: float = 0.0 # Mean reward across batch
|
||||
reward_std: float = 0.0 # Std of rewards
|
||||
value_mean: float = 0.0 # Mean value prediction
|
||||
entropy: float = 0.0 # Policy entropy (for exploration)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PreprocessBatch(ForwardBatch):
|
||||
|
||||
@@ -35,7 +35,6 @@ _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,10 +22,6 @@ 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 (
|
||||
@@ -48,8 +44,6 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"CosmosLatentPreparationStage",
|
||||
"Cosmos25LatentPreparationStage",
|
||||
"LTX2LatentPreparationStage",
|
||||
"LTX2AudioDecodingStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
@@ -57,7 +51,6 @@ __all__ = [
|
||||
"MatrixGameCausalDenoisingStage",
|
||||
"CosmosDenoisingStage",
|
||||
"Cosmos25DenoisingStage",
|
||||
"LTX2DenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
|
||||
@@ -4,11 +4,14 @@ Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import math
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from contextlib import nullcontext
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.attention import get_attn_backend
|
||||
@@ -52,6 +55,84 @@ except ImportError:
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def sde_step_with_logprob(
|
||||
scheduler,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: float | torch.FloatTensor,
|
||||
sample: torch.FloatTensor,
|
||||
prev_sample: torch.FloatTensor | None = None,
|
||||
generator: torch.Generator | None = None,
|
||||
deterministic: bool = False,
|
||||
return_pixel_log_prob: bool = False,
|
||||
return_dt_and_std_dev_t: bool = False
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ...]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE and
|
||||
compute log probabilities for the transition.
|
||||
"""
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
if timestep.ndim == 0:
|
||||
timestep = timestep.unsqueeze(0)
|
||||
step_indices = [
|
||||
scheduler.index_for_timestep(t.item()) for t in timestep
|
||||
]
|
||||
else:
|
||||
step_indices = [scheduler.index_for_timestep(timestep)]
|
||||
|
||||
prev_step_indices = [step + 1 for step in step_indices]
|
||||
|
||||
sigmas = scheduler.sigmas.to(sample.device, sample.dtype)
|
||||
sigma = sigmas[step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_prev = sigmas[prev_step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_max = sigmas[0].item()
|
||||
sigma_min = sigmas[-1].item()
|
||||
|
||||
dt = sigma_prev - sigma
|
||||
|
||||
std_dev_t = sigma_min + (sigma_max - sigma_min) * sigma
|
||||
prev_sample_mean = (sample * (1 + std_dev_t**2 / (2 * sigma) * dt) +
|
||||
model_output * (1 + std_dev_t**2 * (1 - sigma) /
|
||||
(2 * sigma)) * dt)
|
||||
|
||||
if prev_sample is not None and generator is not None:
|
||||
raise ValueError(
|
||||
"Cannot pass both generator and prev_sample. Please make sure that either `generator` or"
|
||||
" `prev_sample` stays `None`.")
|
||||
|
||||
if prev_sample is None:
|
||||
variance_noise = randn_tensor(
|
||||
model_output.shape,
|
||||
generator=generator,
|
||||
device=model_output.device,
|
||||
dtype=model_output.dtype,
|
||||
)
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
prev_sample = prev_sample_mean + std_dev_t * sqrt_dt * variance_noise
|
||||
else:
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
if deterministic:
|
||||
prev_sample = sample + dt * model_output
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
if return_pixel_log_prob:
|
||||
raise NotImplementedError(
|
||||
"Pixel-level log prob is not supported in this helper.")
|
||||
|
||||
std_dev_sqrt_dt = std_dev_t * sqrt_dt
|
||||
log_prob = (
|
||||
-((prev_sample.detach() - prev_sample_mean)**2) /
|
||||
(2 *
|
||||
(std_dev_sqrt_dt**2)) - torch.log(std_dev_sqrt_dt + 1e-8) - torch.log(
|
||||
torch.sqrt(2 * torch.as_tensor(math.pi, device=sample.device))))
|
||||
|
||||
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
|
||||
|
||||
if return_dt_and_std_dev_t:
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt
|
||||
|
||||
|
||||
class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
Stage for running the denoising loop in diffusion pipelines.
|
||||
@@ -203,9 +284,11 @@ class DenoisingStage(PipelineStage):
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
rl_data = batch.rl_data if batch.rl_data and batch.rl_data.enabled else None
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
assert latent_model_input.shape[
|
||||
0] == 1, "TI2V task only supports batch size 1"
|
||||
# TI2V directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
@@ -243,6 +326,12 @@ class DenoisingStage(PipelineStage):
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
rl_timesteps: list[torch.Tensor] = []
|
||||
rl_latents: list[torch.Tensor] = []
|
||||
rl_log_probs: list[torch.Tensor] = []
|
||||
rl_kl: list[torch.Tensor] = []
|
||||
if rl_data is not None and rl_data.store_trajectory:
|
||||
rl_latents.append(latents)
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -329,6 +418,24 @@ class DenoisingStage(PipelineStage):
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
def run_transformer(model, encoder_hidden_states, cond_kwargs,
|
||||
is_cfg_negative: bool):
|
||||
batch.is_cfg_negative = is_cfg_negative
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
return model(
|
||||
latent_model_input,
|
||||
encoder_hidden_states,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
@@ -390,40 +497,13 @@ class DenoisingStage(PipelineStage):
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred = current_model(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
noise_pred = run_transformer(current_model, prompt_embeds,
|
||||
pos_cond_kwargs, False)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
batch.is_cfg_negative = True
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
noise_pred_uncond = current_model(
|
||||
latent_model_input,
|
||||
neg_prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
noise_pred_uncond = run_transformer(
|
||||
current_model, neg_prompt_embeds, neg_cond_kwargs,
|
||||
True)
|
||||
|
||||
noise_pred_text = noise_pred
|
||||
noise_pred = noise_pred_uncond + current_guidance_scale * (
|
||||
@@ -438,11 +518,58 @@ class DenoisingStage(PipelineStage):
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
)
|
||||
# Compute the previous noisy sample
|
||||
prev_latents = latents
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
if rl_data is not None:
|
||||
if rl_data.collect_log_probs:
|
||||
_, log_prob, prev_latents_mean, std_dev_t, _ = sde_step_with_logprob(
|
||||
self.scheduler,
|
||||
noise_pred.float(),
|
||||
t,
|
||||
prev_latents.float(),
|
||||
prev_sample=latents.float(),
|
||||
deterministic=False,
|
||||
return_dt_and_std_dev_t=True,
|
||||
)
|
||||
rl_log_probs.append(log_prob)
|
||||
|
||||
if rl_data.collect_kl and rl_data.kl_reward > 0:
|
||||
adapter_ctx = nullcontext()
|
||||
if hasattr(current_model, "disable_adapter"):
|
||||
adapter_ctx = current_model.disable_adapter()
|
||||
with adapter_ctx:
|
||||
noise_pred_ref = run_transformer(
|
||||
current_model, prompt_embeds,
|
||||
pos_cond_kwargs, False)
|
||||
if batch.do_classifier_free_guidance:
|
||||
noise_pred_uncond_ref = run_transformer(
|
||||
current_model, neg_prompt_embeds,
|
||||
neg_cond_kwargs, True)
|
||||
noise_pred_text_ref = noise_pred_ref
|
||||
noise_pred_ref = noise_pred_uncond_ref + current_guidance_scale * (
|
||||
noise_pred_text_ref -
|
||||
noise_pred_uncond_ref)
|
||||
_, _, prev_latents_mean_ref, std_dev_t_ref, _ = sde_step_with_logprob(
|
||||
self.scheduler,
|
||||
noise_pred_ref.float(),
|
||||
t,
|
||||
prev_latents.float(),
|
||||
prev_sample=latents.float(),
|
||||
deterministic=False,
|
||||
return_dt_and_std_dev_t=True,
|
||||
)
|
||||
if not torch.allclose(std_dev_t, std_dev_t_ref):
|
||||
logger.warning(
|
||||
"std_dev_t mismatch in RL KL computation at step %s",
|
||||
i)
|
||||
kl = (prev_latents_mean -
|
||||
prev_latents_mean_ref)**2 / (2 * std_dev_t**2)
|
||||
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
|
||||
rl_kl.append(kl)
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
@@ -452,6 +579,15 @@ class DenoisingStage(PipelineStage):
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(t)
|
||||
trajectory_latents.append(latents)
|
||||
if rl_data is not None:
|
||||
rl_timesteps.append(t)
|
||||
if rl_data.store_trajectory:
|
||||
rl_latents.append(latents)
|
||||
if rl_data.collect_kl and rl_data.kl_reward <= 0:
|
||||
rl_kl.append(
|
||||
torch.zeros(latents.shape[0],
|
||||
device=latents.device,
|
||||
dtype=latents.dtype))
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
@@ -472,6 +608,25 @@ class DenoisingStage(PipelineStage):
|
||||
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
|
||||
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
if rl_data is not None:
|
||||
if rl_timesteps:
|
||||
rl_data.trajectory_timesteps = torch.stack(rl_timesteps, dim=0)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.trajectory_timesteps = rl_data.trajectory_timesteps.cpu(
|
||||
)
|
||||
if rl_data.store_trajectory and rl_latents:
|
||||
rl_data.trajectory_latents = torch.stack(rl_latents, dim=1)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.trajectory_latents = rl_data.trajectory_latents.cpu(
|
||||
)
|
||||
if rl_log_probs:
|
||||
rl_data.log_probs = torch.stack(rl_log_probs, dim=1)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.log_probs = rl_data.log_probs.cpu()
|
||||
if rl_kl:
|
||||
rl_data.kl = torch.stack(rl_kl, dim=1)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.kl = rl_data.kl.cpu()
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
# 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
|
||||
@@ -1,308 +0,0 @@
|
||||
# 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
|
||||
@@ -1,189 +0,0 @@
|
||||
# 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,11 +11,14 @@ 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):
|
||||
"""
|
||||
@@ -36,7 +39,6 @@ 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(
|
||||
@@ -68,8 +70,6 @@ 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,9 +86,6 @@ 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:
|
||||
@@ -187,13 +184,10 @@ 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(
|
||||
@@ -265,11 +259,6 @@ 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)
|
||||
@@ -277,7 +266,6 @@ 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,23 +67,17 @@ 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", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
@app.function(gpu="H100:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_encoder_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
|
||||
run_test("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", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_vae_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
|
||||
run_test("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", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_transformer_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
|
||||
@app.function(
|
||||
gpu="L40S:4",
|
||||
@@ -95,21 +89,13 @@ 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", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
|
||||
run_test("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", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_lora_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
|
||||
run_test("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,48 +80,9 @@ 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 = {
|
||||
@@ -268,26 +229,18 @@ 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,
|
||||
|
||||
@@ -17,10 +17,11 @@ from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
os.environ["MASTER_PORT"] = "29701"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
@@ -121,4 +122,24 @@ def test_wan_transformer():
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from fastvideo.distributed import (
|
||||
cleanup_dist_env_and_memory,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
)
|
||||
# Allow running this test file directly without pytest.
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
try:
|
||||
test_wan_transformer()
|
||||
logger.info("test_wan_transformer finished successfully.")
|
||||
finally:
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -1,5 +1,12 @@
|
||||
from .distillation_pipeline import DistillationPipeline
|
||||
from .training_pipeline import TrainingPipeline
|
||||
from .wan_training_pipeline import WanTrainingPipeline
|
||||
from fastvideo.training.rl import RLPipeline, create_rl_pipeline
|
||||
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
|
||||
__all__ = [
|
||||
"TrainingPipeline",
|
||||
"WanTrainingPipeline",
|
||||
"DistillationPipeline",
|
||||
"RLPipeline",
|
||||
"create_rl_pipeline",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from .rl_pipeline import RLPipeline, create_rl_pipeline
|
||||
|
||||
__all__ = [
|
||||
"RLPipeline",
|
||||
"create_rl_pipeline",
|
||||
]
|
||||
@@ -0,0 +1,11 @@
|
||||
from .rewards import (
|
||||
create_reward_models,
|
||||
MultiRewardAggregator,
|
||||
ValueModel
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"create_reward_models",
|
||||
"MultiRewardAggregator",
|
||||
"ValueModel",
|
||||
]
|
||||
@@ -0,0 +1,63 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Abstract base class for VIDEO reward models.
|
||||
|
||||
All VIDEO reward models should inherit from this class and implement
|
||||
the compute_reward() method.
|
||||
|
||||
IMPORTANT: Reward models must process FULL VIDEO SEQUENCES, not individual frames.
|
||||
Input shape is [B, T, C, H, W] where T is the temporal (frame) dimension.
|
||||
|
||||
For video-specific rewards, consider:
|
||||
- Temporal coherence across frames
|
||||
- Motion quality and smoothness
|
||||
- Video-text alignment (not just frame-text)
|
||||
- Multi-frame aesthetic quality
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
class BaseRewardModel(ABC, nn.Module):
|
||||
def __init__(self, model_path: str | None = None, device: str = "cuda"):
|
||||
super().__init__()
|
||||
self.model_path = model_path
|
||||
self.device = device
|
||||
|
||||
@abstractmethod
|
||||
def compute_reward(
|
||||
self,
|
||||
videos: torch.Tensor, # [B, T, C, H, W] decoded video sequences
|
||||
prompts: list[str] | None, # Text prompts
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute rewards for generated VIDEO sequences.
|
||||
|
||||
IMPORTANT: This method must process the FULL temporal sequence [B, T, C, H, W].
|
||||
Do NOT evaluate individual frames independently and average.
|
||||
|
||||
Args:
|
||||
videos: Decoded video tensors [B, T, C, H, W] in range [0, 1]
|
||||
B = batch size
|
||||
T = number of frames (temporal dimension)
|
||||
C = channels (typically 3 for RGB)
|
||||
H, W = height, width
|
||||
prompts: List of text prompts (length B) describing each video
|
||||
**kwargs: Additional model-specific arguments
|
||||
|
||||
Returns:
|
||||
rewards: Tensor of shape [B] with reward scores for each video sequence
|
||||
|
||||
Example:
|
||||
>>> videos = torch.rand(4, 17, 3, 256, 256) # 4 videos, 17 frames each
|
||||
>>> prompts = ["A cat jumping", "A dog running", ...]
|
||||
>>> rewards = model.compute_reward(videos, prompts)
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement compute_reward()")
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(model_path={self.model_path})"
|
||||
@@ -0,0 +1,206 @@
|
||||
from paddleocr import PaddleOCR
|
||||
import torch
|
||||
import numpy as np
|
||||
from Levenshtein import distance
|
||||
from typing import Any
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.training.rl.rewards.base import BaseRewardModel
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class OcrScorerVideo(BaseRewardModel):
|
||||
"""
|
||||
OCR reward model for multi-frame video OCR evaluation.
|
||||
|
||||
This model evaluates multiple frames across the video sequence,
|
||||
sampling frames at a specified interval and averaging the OCR scores.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
model_path: str | None = None,
|
||||
device: str = "cpu",
|
||||
frame_interval: int = 4):
|
||||
"""
|
||||
OCR reward calculator for videos
|
||||
|
||||
Args:
|
||||
model_path: Not used for PaddleOCR (kept for BaseRewardModel compatibility)
|
||||
device: Device string (used to determine use_gpu if not explicitly set)
|
||||
frame_interval: Sample every Nth frame (default: 4)
|
||||
"""
|
||||
super().__init__(model_path=model_path, device=device)
|
||||
|
||||
self.frame_interval = frame_interval
|
||||
self.ocr = PaddleOCR(
|
||||
use_angle_cls=False,
|
||||
lang="en",
|
||||
use_gpu=False,
|
||||
show_log=False # Disable unnecessary log output
|
||||
)
|
||||
|
||||
logger.info("Initialized OcrScorerVideo (device=%s, frame_interval=%d)",
|
||||
device, frame_interval)
|
||||
|
||||
def _process_single_video(self, video_tensor: torch.Tensor,
|
||||
prompt: str) -> float:
|
||||
"""
|
||||
Process a single video tensor and return its OCR reward.
|
||||
|
||||
Args:
|
||||
video_tensor: Video tensor of shape [C, T, H, W]
|
||||
prompt: Text prompt containing target OCR text in quotes
|
||||
|
||||
Returns:
|
||||
Average reward across positive-scoring frames
|
||||
"""
|
||||
# Extract target text from prompt
|
||||
try:
|
||||
target_text = prompt.split('"')[1].replace(' ', '').lower()
|
||||
except IndexError:
|
||||
logger.warning("Failed to extract quoted text from prompt: %s",
|
||||
prompt)
|
||||
target_text = prompt.replace(' ', '').lower()
|
||||
|
||||
if not target_text:
|
||||
return 0.0
|
||||
|
||||
# video_tensor is [C, T, H, W]
|
||||
C, T, H, W = video_tensor.shape
|
||||
|
||||
# Convert to numpy and move to CPU if needed
|
||||
video_np = video_tensor.detach().cpu().numpy()
|
||||
|
||||
# Convert from [C, T, H, W] to [T, H, W, C] for easier frame extraction
|
||||
video_np = np.transpose(video_np, (1, 2, 3, 0)) # [T, H, W, C]
|
||||
logger.info(f"in ocr 1.5, video_np[0][0]: {video_np[0][0]}")
|
||||
|
||||
# Normalize to [0, 255] uint8 if needed
|
||||
if video_np.max() <= 1.0:
|
||||
video_np = (video_np * 255).astype(np.uint8)
|
||||
else:
|
||||
video_np = video_np.astype(np.uint8)
|
||||
|
||||
frame_rewards = []
|
||||
|
||||
# Sample frames at specified interval
|
||||
for frame_idx in range(0, T, self.frame_interval):
|
||||
frame = video_np[frame_idx] # [H, W, C]
|
||||
logger.info(f"in ocr 2, frame.shape: {frame.shape}")
|
||||
# Run OCR
|
||||
try:
|
||||
result = self.ocr.ocr(frame, cls=False)
|
||||
logger.info(f"in ocr 3, result: {result}")
|
||||
if result and result[0]:
|
||||
recognized_text = "".join(
|
||||
[line[1][0] for line in result[0] if line[1][1] > 0])
|
||||
else:
|
||||
recognized_text = ""
|
||||
except Exception as e:
|
||||
logger.info("OCR failed on frame %d: %s", frame_idx, str(e))
|
||||
recognized_text = ''
|
||||
|
||||
logger.info(f"in ocr 4, recognized_text: {recognized_text}")
|
||||
|
||||
recognized_text = recognized_text.replace(' ', '').lower()
|
||||
if target_text in recognized_text:
|
||||
dist = 0
|
||||
else:
|
||||
dist = distance(recognized_text, target_text)
|
||||
dist = min(dist, len(target_text))
|
||||
reward = 1.0 - dist / len(target_text)
|
||||
|
||||
logger.info(f"in ocr 5, reward: {reward}")
|
||||
if reward > 0:
|
||||
frame_rewards.append(reward)
|
||||
|
||||
logger.info(f"in ocr 6, frame_rewards: {frame_rewards}")
|
||||
|
||||
return sum([reward / len(frame_rewards)
|
||||
for reward in frame_rewards]) if frame_rewards else 0.0
|
||||
|
||||
@torch.no_grad()
|
||||
def compute_reward(self, videos: torch.Tensor, prompts: list[str],
|
||||
**kwargs: Any) -> torch.Tensor:
|
||||
"""
|
||||
Calculate OCR reward by evaluating sampled frames across the video.
|
||||
|
||||
Args:
|
||||
videos: Video tensor of shape [B, C, T, H, W]
|
||||
B = batch size
|
||||
C = channels (typically 3 for RGB)
|
||||
T = number of frames (temporal dimension)
|
||||
H, W = height, width
|
||||
prompts: List of text prompts containing target OCR text in quotes (length B)
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
Reward tensor [B] with averaged OCR similarity scores across frames
|
||||
"""
|
||||
# Ensure videos is a torch tensor with correct shape
|
||||
assert isinstance(
|
||||
videos,
|
||||
torch.Tensor), f"videos must be torch.Tensor, got {type(videos)}"
|
||||
assert videos.ndim == 5, f"videos must have 5 dimensions [B, C, T, H, W], got shape {videos.shape}"
|
||||
|
||||
logger.info(f"in ocr 1, videos.shape: {videos.shape}")
|
||||
|
||||
B, C, T, H, W = videos.shape
|
||||
assert len(
|
||||
prompts
|
||||
) == B, f"Number of prompts ({len(prompts)}) must match batch size ({B})"
|
||||
|
||||
rewards = []
|
||||
for b in range(B):
|
||||
# Extract single video: [C, T, H, W]
|
||||
video = videos[b]
|
||||
reward = self._process_single_video(video, prompts[b])
|
||||
rewards.append(reward)
|
||||
|
||||
logger.info(f"in ocr 7, rewards: {rewards}")
|
||||
|
||||
rewards = torch.tensor(rewards, dtype=torch.float32, device=self.device)
|
||||
|
||||
logger.info(f"in ocr 8, rewards: {rewards}")
|
||||
|
||||
# Check for NaN or Inf values
|
||||
if torch.isnan(rewards).any() or torch.isinf(rewards).any():
|
||||
logger.warning(
|
||||
"NaN or Inf detected in OCR rewards, returning zero tensor")
|
||||
return torch.zeros_like(rewards)
|
||||
|
||||
return rewards
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
example_image_path = "flowgrpo_cmd.png"
|
||||
example_image = Image.open(example_image_path)
|
||||
example_prompt = '/f1ow_grpo$'
|
||||
|
||||
# Convert image to RGB if needed
|
||||
if example_image.mode != 'RGB':
|
||||
example_image = example_image.convert('RGB')
|
||||
|
||||
# Convert PIL Image to numpy array [H, W, C]
|
||||
image_np = np.array(example_image)
|
||||
|
||||
# Normalize to [0, 1] range and convert to float32
|
||||
image_np = image_np.astype(np.float32) / 255.0
|
||||
|
||||
# Convert to torch tensor and reshape: [H, W, C] -> [C, H, W]
|
||||
image_tensor = torch.from_numpy(image_np).permute(2, 0, 1)
|
||||
|
||||
# Add temporal dimension: [C, H, W] -> [C, T, H, W] where T=1
|
||||
video_tensor = image_tensor.unsqueeze(1) # [C, 1, H, W]
|
||||
|
||||
# Add batch dimension: [C, T, H, W] -> [B, C, T, H, W] where B=1
|
||||
video_tensor = video_tensor.unsqueeze(0) # [1, C, 1, H, W]
|
||||
|
||||
# Instantiate scorer
|
||||
scorer = OcrScorerVideo(device="cpu")
|
||||
|
||||
# Call compute_reward method with video tensor
|
||||
reward = scorer.compute_reward(video_tensor, [example_prompt])
|
||||
print(f"OCR Reward: {reward.item()}")
|
||||
@@ -0,0 +1,338 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Base infrastructure for VIDEO reward models in RL/GRPO training.
|
||||
|
||||
IMPORTANT: This module is designed exclusively for VIDEO generation models.
|
||||
All reward models must operate on video sequences [B, T, C, H, W], not single frames.
|
||||
|
||||
This module provides:
|
||||
1. Multi-reward aggregation for video
|
||||
2. Value model wrapper
|
||||
3. Integration with FastVideo video generation infrastructure
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.training.rl.rewards.ocr import OcrScorerVideo
|
||||
from fastvideo.training.rl.rewards.base import BaseRewardModel
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MultiRewardAggregator(nn.Module):
|
||||
"""
|
||||
Aggregates multiple reward models with configurable weights.
|
||||
|
||||
This implements the multi-reward aggregation strategy from flow_grpo,
|
||||
allowing combination of different reward signals (aesthetic quality,
|
||||
text-video alignment, compositional understanding, etc.)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reward_models: list[BaseRewardModel],
|
||||
reward_weights: list[float] | None = None,
|
||||
normalize_rewards: bool = True
|
||||
):
|
||||
"""
|
||||
Initialize multi-reward aggregator.
|
||||
|
||||
Args:
|
||||
reward_models: List of reward model instances
|
||||
reward_weights: Weights for each reward model (default: uniform)
|
||||
normalize_rewards: Whether to normalize rewards before aggregation
|
||||
"""
|
||||
super().__init__()
|
||||
self.reward_models = nn.ModuleList(reward_models)
|
||||
|
||||
if reward_weights is None:
|
||||
reward_weights = [1.0 / len(reward_models)] * len(reward_models)
|
||||
|
||||
assert len(reward_weights) == len(reward_models), \
|
||||
f"Number of weights ({len(reward_weights)}) must match number of models ({len(reward_models)})"
|
||||
|
||||
assert abs(sum(reward_weights) - 1.0) < 1e-6, \
|
||||
f"Reward weights must sum to 1.0, got {sum(reward_weights)}"
|
||||
|
||||
self.reward_weights = reward_weights
|
||||
self.normalize_rewards = normalize_rewards
|
||||
|
||||
logger.info(
|
||||
"Initialized MultiRewardAggregator with %d models: %s",
|
||||
len(reward_models),
|
||||
[(type(m).__name__, w) for m, w in zip(reward_models, reward_weights, strict=False)]
|
||||
)
|
||||
|
||||
def compute_reward(
|
||||
self,
|
||||
videos: torch.Tensor,
|
||||
prompts: list[str],
|
||||
return_individual: bool = False,
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor | dict[str, torch.Tensor]:
|
||||
"""
|
||||
Compute aggregated reward from multiple models.
|
||||
|
||||
Args:
|
||||
videos: Decoded video tensors [B, C, T, H, W]
|
||||
prompts: List of text prompts
|
||||
return_individual: If True, return dict with individual rewards
|
||||
**kwargs: Additional arguments passed to reward models
|
||||
|
||||
Returns:
|
||||
If return_individual=False: aggregated_rewards [B]
|
||||
If return_individual=True: dict with "aggregated" and individual model rewards
|
||||
"""
|
||||
batch_size = videos.shape[0]
|
||||
individual_rewards: dict[str, torch.Tensor] = {}
|
||||
|
||||
# Collect rewards from all models
|
||||
all_rewards = []
|
||||
for i, (model, weight) in enumerate(zip(self.reward_models, self.reward_weights, strict=False)):
|
||||
reward = model.compute_reward(videos, prompts, **kwargs)
|
||||
assert reward.shape == (batch_size,), \
|
||||
f"Reward model {i} returned shape {reward.shape}, expected ({batch_size},)"
|
||||
|
||||
# Optionally normalize individual rewards
|
||||
if self.normalize_rewards:
|
||||
reward = (reward - reward.mean()) / (reward.std() + 1e-8)
|
||||
|
||||
individual_rewards[f"reward_{type(model).__name__}"] = reward
|
||||
all_rewards.append(weight * reward)
|
||||
|
||||
# Aggregate with weights
|
||||
aggregated = sum(all_rewards)
|
||||
|
||||
if return_individual:
|
||||
individual_rewards["aggregated"] = aggregated
|
||||
return individual_rewards
|
||||
|
||||
return aggregated
|
||||
|
||||
def __repr__(self) -> str:
|
||||
models_str = ", ".join([
|
||||
f"{type(m).__name__}(w={w:.3f})"
|
||||
for m, w in zip(self.reward_models, self.reward_weights, strict=False)
|
||||
])
|
||||
return f"MultiRewardAggregator({models_str})"
|
||||
|
||||
|
||||
class ValueModel(nn.Module):
|
||||
"""
|
||||
Value function model wrapper for RL training.
|
||||
|
||||
The value model can either:
|
||||
1. Share the transformer backbone with the policy (memory efficient)
|
||||
2. Use a separate transformer (more flexible)
|
||||
|
||||
For now, this is a placeholder that will be expanded based on
|
||||
the chosen architecture strategy.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer: nn.Module,
|
||||
share_backbone: bool = False,
|
||||
hidden_size: int | None = None
|
||||
):
|
||||
"""
|
||||
Initialize value model.
|
||||
|
||||
Args:
|
||||
transformer: Transformer model (policy or separate)
|
||||
share_backbone: Whether to share backbone with policy
|
||||
hidden_size: Hidden size for value head (inferred if None)
|
||||
"""
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.share_backbone = share_backbone
|
||||
|
||||
# Value head will be added later based on transformer architecture
|
||||
# For now, just store the transformer reference
|
||||
logger.info(
|
||||
"Initialized ValueModel (share_backbone=%s)",
|
||||
share_backbone
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass to compute value predictions.
|
||||
|
||||
Args:
|
||||
hidden_states: Latent states [B, C, T, H, W]
|
||||
encoder_hidden_states: Text embeddings [B, L, D]
|
||||
timestep: Timesteps [B]
|
||||
**kwargs: Additional transformer arguments
|
||||
|
||||
Returns:
|
||||
values: Value predictions [B]
|
||||
"""
|
||||
# TODO: Implement value prediction
|
||||
# For now, return dummy values
|
||||
batch_size = hidden_states.shape[0]
|
||||
return torch.zeros(batch_size, device=hidden_states.device)
|
||||
|
||||
|
||||
class DummyRewardModel(BaseRewardModel):
|
||||
"""
|
||||
Dummy VIDEO reward model for testing and development.
|
||||
|
||||
Returns random rewards in the range [0, 1] for VIDEO inputs.
|
||||
This is a placeholder for testing the RL pipeline before real video reward models
|
||||
are implemented.
|
||||
|
||||
NOTE: This does NOT actually evaluate video quality - it's just for testing!
|
||||
"""
|
||||
|
||||
def __init__(self, mean: float = 0.5, std: float = 0.1):
|
||||
super().__init__(model_path=None)
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
logger.info("Initialized DummyRewardModel (VIDEO) - mean=%.2f, std=%.2f", mean, std)
|
||||
logger.warning(
|
||||
"DummyRewardModel is for TESTING ONLY - does not evaluate actual video quality!"
|
||||
)
|
||||
|
||||
def compute_reward(
|
||||
self,
|
||||
videos: torch.Tensor, # [B, T, C, H, W]
|
||||
prompts: list[str],
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Return random rewards for testing.
|
||||
|
||||
Args:
|
||||
videos: Video sequences [B, T, C, H, W]
|
||||
prompts: Text prompts
|
||||
|
||||
Returns:
|
||||
Random rewards [B] in range [0, 1]
|
||||
"""
|
||||
batch_size = videos.shape[0]
|
||||
num_frames = videos.shape[1]
|
||||
|
||||
logger.debug(
|
||||
"DummyRewardModel processing %d videos with %d frames each",
|
||||
batch_size,
|
||||
num_frames
|
||||
)
|
||||
|
||||
# Generate random rewards (not based on actual video content!)
|
||||
rewards = torch.randn(batch_size, device=videos.device) * self.std + self.mean
|
||||
return rewards.clamp(0.0, 1.0)
|
||||
|
||||
def load_model(self) -> None:
|
||||
"""No model to load for dummy."""
|
||||
pass
|
||||
|
||||
|
||||
def create_reward_models(
|
||||
reward_models: dict,
|
||||
device: str = "cuda"
|
||||
) -> MultiRewardAggregator:
|
||||
"""
|
||||
Factory function to create VIDEO reward models from configuration strings.
|
||||
|
||||
IMPORTANT: Only creates VIDEO reward models. Image-only reward models
|
||||
(PickScore, ImageReward, GenEval, etc.) are NOT supported.
|
||||
|
||||
Args:
|
||||
reward_models: dictionary of reward model names to weights
|
||||
Example: {"paddle_ocr": 0.5, "video_score": 0.5}
|
||||
device: Device to load models on
|
||||
|
||||
Returns:
|
||||
MultiRewardAggregator with loaded VIDEO reward models
|
||||
|
||||
Supported VIDEO Reward Types:
|
||||
- "paddle_ocr": PaddleOCR multi-frame video text recognition
|
||||
- "video_score": Video aesthetic quality (multi-frame) - TODO
|
||||
- "video_text_alignment": CLIP-based video-text similarity - TODO
|
||||
- "temporal_coherence": Frame-to-frame consistency - TODO
|
||||
- "motion_quality": Motion smoothness and realism - TODO
|
||||
- "dummy": Random rewards for testing (VIDEO-aware)
|
||||
|
||||
NOT Supported (Image-Only):
|
||||
- "pickscore": Image aesthetic (use "video_score" instead)
|
||||
- "imagereward": Image quality (use "video_score" instead)
|
||||
- "geneval": Image compositional (no video equivalent yet)
|
||||
- Any single-frame reward models
|
||||
|
||||
Example:
|
||||
>>> models = create_reward_models(
|
||||
... reward_models={
|
||||
... "paddle_ocr": 0.5,
|
||||
... "video_text_alignment": 0.5
|
||||
... },
|
||||
... device="cuda"
|
||||
... )
|
||||
"""
|
||||
|
||||
|
||||
assert reward_models, "No reward models specified. Please select at least 1 reward model"
|
||||
|
||||
types = [t.strip() for t in reward_models.keys()]
|
||||
weights = list(reward_models.values())
|
||||
|
||||
assert len(types) == len(weights), \
|
||||
f"Number of models ({len(types)}) must match number of weights ({len(weights)})"
|
||||
|
||||
# Create reward models based on types
|
||||
models_list: list[BaseRewardModel] = []
|
||||
for reward_type in types:
|
||||
if reward_type == "dummy":
|
||||
model = DummyRewardModel()
|
||||
|
||||
elif reward_type == "paddle_ocr":
|
||||
logger.info("Creating PaddleOCR reward model")
|
||||
model = OcrScorerVideo(device=device)
|
||||
|
||||
elif reward_type == "video_score":
|
||||
# TODO: Implement VideoScore reward model (Phase 2)
|
||||
logger.warning(
|
||||
"VideoScore reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
elif reward_type == "video_text_alignment":
|
||||
# TODO: Implement VideoTextAlignment reward model (Phase 2)
|
||||
logger.warning(
|
||||
"VideoTextAlignment reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
elif reward_type == "temporal_coherence":
|
||||
# TODO: Implement TemporalCoherence reward model (Phase 2)
|
||||
logger.warning(
|
||||
"TemporalCoherence reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
elif reward_type == "motion_quality":
|
||||
# TODO: Implement MotionQuality reward model (Phase 2)
|
||||
logger.warning(
|
||||
"MotionQuality reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
else:
|
||||
logger.warning(
|
||||
"Unknown VIDEO reward type '%s', using DummyRewardModel",
|
||||
reward_type
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
|
||||
models_list.append(model)
|
||||
|
||||
logger.info(
|
||||
"Created MultiRewardAggregator with %d VIDEO reward models",
|
||||
len(models_list)
|
||||
)
|
||||
|
||||
return MultiRewardAggregator(models_list, weights, normalize_rewards=True)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,385 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Utility functions for RL/GRPO training.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def compute_gae(
|
||||
rewards: torch.Tensor,
|
||||
values: torch.Tensor,
|
||||
next_values: torch.Tensor,
|
||||
dones: torch.Tensor | None = None,
|
||||
gamma: float = 0.99,
|
||||
lambda_: float = 0.95
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Compute Generalized Advantage Estimation (GAE-lambda).
|
||||
|
||||
GAE reduces variance in advantage estimation while allowing some bias.
|
||||
This is a key component of modern policy gradient methods like PPO and GRPO.
|
||||
|
||||
Args:
|
||||
rewards: Rewards at each step [B, T] or [B]
|
||||
values: Value predictions at each step [B, T] or [B]
|
||||
next_values: Value predictions at next step [B, T] or [B]
|
||||
dones: Episode termination flags [B, T] or [B] (1 if done, 0 otherwise)
|
||||
gamma: Discount factor
|
||||
lambda_: GAE lambda parameter (0=TD(0), 1=Monte Carlo)
|
||||
|
||||
Returns:
|
||||
advantages: GAE advantages [B, T] or [B]
|
||||
returns: TD(lambda) returns [B, T] or [B]
|
||||
|
||||
Reference:
|
||||
Schulman et al. "High-Dimensional Continuous Control Using Generalized Advantage Estimation"
|
||||
https://arxiv.org/abs/1506.02438
|
||||
"""
|
||||
if dones is None:
|
||||
dones = torch.zeros_like(rewards)
|
||||
|
||||
# Compute TD residuals: delta_t = r_t + gamma * V(s_{t+1}) - V(s_t)
|
||||
deltas = rewards + gamma * next_values * (1.0 - dones) - values
|
||||
|
||||
# If single step (no time dimension), return directly
|
||||
if deltas.dim() == 1:
|
||||
advantages = deltas
|
||||
returns = advantages + values
|
||||
return advantages, returns
|
||||
|
||||
# Multi-step: compute GAE recursively
|
||||
batch_size, num_steps = deltas.shape
|
||||
advantages = torch.zeros_like(deltas)
|
||||
gae = torch.zeros(batch_size, device=deltas.device)
|
||||
|
||||
# Backward pass to compute GAE
|
||||
for t in reversed(range(num_steps)):
|
||||
gae = deltas[:, t] + gamma * lambda_ * (1.0 - dones[:, t]) * gae
|
||||
advantages[:, t] = gae
|
||||
|
||||
# Returns are advantages + values
|
||||
returns = advantages + values
|
||||
|
||||
return advantages, returns
|
||||
|
||||
|
||||
def normalize_advantages(
|
||||
advantages: torch.Tensor,
|
||||
epsilon: float = 1e-8
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Normalize advantages to have zero mean and unit variance.
|
||||
|
||||
This is a common practice in PPO and GRPO to stabilize training.
|
||||
|
||||
Args:
|
||||
advantages: Raw advantages [B, ...]
|
||||
epsilon: Small constant for numerical stability
|
||||
|
||||
Returns:
|
||||
normalized_advantages: Normalized advantages [B, ...]
|
||||
"""
|
||||
mean = advantages.mean()
|
||||
std = advantages.std()
|
||||
return (advantages - mean) / (std + epsilon)
|
||||
|
||||
#TODO(jiali): refactor into algorithm
|
||||
def compute_grpo_policy_loss(
|
||||
log_probs: torch.Tensor,
|
||||
old_log_probs: torch.Tensor,
|
||||
advantages: torch.Tensor,
|
||||
clip_range: float = 0.2,
|
||||
use_ratio_norm: bool = True,
|
||||
max_importance_ratio: float = 10.0
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""
|
||||
Compute GRPO policy loss with importance sampling and clipping.
|
||||
|
||||
This implements the core GRPO objective with safety mechanisms from GRPO-Guard:
|
||||
- Importance ratio clipping (PPO-style)
|
||||
- RatioNorm correction (GRPO-Guard)
|
||||
- Ratio clamping for extreme values
|
||||
|
||||
Args:
|
||||
log_probs: Log probabilities from current policy [B]
|
||||
old_log_probs: Log probabilities from old policy [B]
|
||||
advantages: Advantages [B]
|
||||
clip_range: Clipping range for importance ratios
|
||||
use_ratio_norm: Apply RatioNorm correction (GRPO-Guard)
|
||||
max_importance_ratio: Maximum importance ratio before clamping
|
||||
|
||||
Returns:
|
||||
loss: Policy loss (scalar)
|
||||
info: Dictionary with diagnostic information
|
||||
|
||||
Reference:
|
||||
- PPO: Schulman et al. "Proximal Policy Optimization Algorithms"
|
||||
- GRPO-Guard: RatioNorm and gradient reweighting
|
||||
"""
|
||||
# Compute importance ratio: r_t = pi_new(a|s) / pi_old(a|s)
|
||||
log_ratio = log_probs - old_log_probs
|
||||
ratio = torch.exp(log_ratio)
|
||||
|
||||
# Clamp extreme ratios for numerical stability
|
||||
ratio = torch.clamp(ratio, 1.0 / max_importance_ratio, max_importance_ratio)
|
||||
|
||||
# RatioNorm correction (GRPO-Guard)
|
||||
# Corrects bias in importance sampling when ratio >> 1
|
||||
if use_ratio_norm:
|
||||
ratio_mean = ratio.mean()
|
||||
ratio = ratio / (ratio_mean + 1e-8)
|
||||
|
||||
# Clipped surrogate objective
|
||||
ratio_clipped = torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range)
|
||||
surrogate1 = ratio * advantages
|
||||
surrogate2 = ratio_clipped * advantages
|
||||
policy_loss = -torch.min(surrogate1, surrogate2).mean()
|
||||
|
||||
# Compute diagnostics
|
||||
with torch.no_grad():
|
||||
# Clip fraction: how often ratios were clipped
|
||||
clip_fraction = ((ratio < 1.0 - clip_range) | (ratio > 1.0 + clip_range)).float().mean()
|
||||
|
||||
# KL divergence (approximate)
|
||||
kl_div = log_ratio.mean()
|
||||
|
||||
# Importance ratio stats
|
||||
importance_ratio_mean = ratio.mean()
|
||||
importance_ratio_std = ratio.std()
|
||||
|
||||
info = {
|
||||
"policy_loss": policy_loss.item(),
|
||||
"clip_fraction": clip_fraction.item(),
|
||||
"kl_divergence": kl_div.item(),
|
||||
"importance_ratio_mean": importance_ratio_mean.item(),
|
||||
"importance_ratio_std": importance_ratio_std.item(),
|
||||
}
|
||||
|
||||
return policy_loss, info
|
||||
|
||||
|
||||
def compute_value_loss(
|
||||
values: torch.Tensor,
|
||||
returns: torch.Tensor,
|
||||
old_values: torch.Tensor | None = None,
|
||||
clip_range: float = 0.2,
|
||||
use_clipping: bool = True
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""
|
||||
Compute value function loss with optional clipping.
|
||||
|
||||
Args:
|
||||
values: Value predictions from current model [B]
|
||||
returns: Target returns (from GAE) [B]
|
||||
old_values: Value predictions from old model [B] (for clipping)
|
||||
clip_range: Clipping range for value updates
|
||||
use_clipping: Whether to use clipped value loss (PPO-style)
|
||||
|
||||
Returns:
|
||||
loss: Value loss (scalar)
|
||||
info: Dictionary with diagnostic information
|
||||
"""
|
||||
# Standard MSE loss
|
||||
value_loss_unclipped = F.mse_loss(values, returns, reduction="none")
|
||||
|
||||
# Clipped value loss (PPO-style)
|
||||
if use_clipping and old_values is not None:
|
||||
values_clipped = old_values + torch.clamp(
|
||||
values - old_values,
|
||||
-clip_range,
|
||||
clip_range
|
||||
)
|
||||
value_loss_clipped = F.mse_loss(values_clipped, returns, reduction="none")
|
||||
value_loss = torch.max(value_loss_unclipped, value_loss_clipped).mean()
|
||||
else:
|
||||
value_loss = value_loss_unclipped.mean()
|
||||
|
||||
# Compute diagnostics
|
||||
with torch.no_grad():
|
||||
explained_variance = 1.0 - (returns - values).var() / (returns.var() + 1e-8)
|
||||
|
||||
info = {
|
||||
"value_loss": value_loss.item(),
|
||||
"explained_variance": explained_variance.item(),
|
||||
"value_mean": values.mean().item(),
|
||||
"value_std": values.std().item(),
|
||||
}
|
||||
|
||||
return value_loss, info
|
||||
|
||||
|
||||
def compute_policy_entropy(log_probs: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Compute policy entropy for exploration bonus.
|
||||
|
||||
Args:
|
||||
log_probs: Log probabilities [B]
|
||||
|
||||
Returns:
|
||||
entropy: Mean entropy across batch (scalar)
|
||||
"""
|
||||
# For continuous actions: H = -log_prob (assuming Gaussian)
|
||||
# For discrete: H = -sum(p * log(p))
|
||||
# Here we use a simple approximation
|
||||
entropy = -log_probs.mean()
|
||||
return entropy
|
||||
|
||||
|
||||
def apply_gradient_reweighting(
|
||||
gradients: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
num_train_timesteps: int = 1000
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply GRPO-Guard gradient reweighting across denoising steps.
|
||||
|
||||
This reweights gradients based on the timestep to balance learning
|
||||
across different noise levels.
|
||||
|
||||
Args:
|
||||
gradients: Gradients to reweight [B, ...]
|
||||
timesteps: Timesteps at which gradients were computed [B]
|
||||
num_train_timesteps: Total number of training timesteps
|
||||
|
||||
Returns:
|
||||
reweighted_gradients: Reweighted gradients [B, ...]
|
||||
"""
|
||||
# Compute timestep weights (higher weight for later timesteps)
|
||||
# This is a simple linear weighting, can be made more sophisticated
|
||||
timestep_weights = 1.0 + (timesteps.float() / num_train_timesteps)
|
||||
timestep_weights = timestep_weights.view(-1, *([1] * (gradients.dim() - 1)))
|
||||
|
||||
return gradients * timestep_weights
|
||||
|
||||
|
||||
def sample_random_timesteps(
|
||||
batch_size: int,
|
||||
min_timestep: int,
|
||||
max_timestep: int,
|
||||
device: torch.device,
|
||||
generator: torch.Generator | None = None
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Sample random timesteps for noise injection (Flow-GRPO-Fast).
|
||||
|
||||
Args:
|
||||
batch_size: Number of samples
|
||||
min_timestep: Minimum timestep
|
||||
max_timestep: Maximum timestep
|
||||
device: Device for tensor
|
||||
generator: Random generator for reproducibility
|
||||
|
||||
Returns:
|
||||
timesteps: Random timesteps [B]
|
||||
"""
|
||||
if generator is not None:
|
||||
timesteps = torch.randint(
|
||||
min_timestep,
|
||||
max_timestep + 1,
|
||||
(batch_size,),
|
||||
device=device,
|
||||
generator=generator
|
||||
)
|
||||
else:
|
||||
timesteps = torch.randint(
|
||||
min_timestep,
|
||||
max_timestep + 1,
|
||||
(batch_size,),
|
||||
device=device
|
||||
)
|
||||
|
||||
return timesteps
|
||||
|
||||
|
||||
def compute_reward_statistics(
|
||||
rewards: torch.Tensor
|
||||
) -> dict[str, float]:
|
||||
"""
|
||||
Compute statistics for reward distribution.
|
||||
|
||||
Args:
|
||||
rewards: Reward values [B]
|
||||
|
||||
Returns:
|
||||
stats: Dictionary with mean, std, min, max
|
||||
"""
|
||||
return {
|
||||
"reward_mean": rewards.mean().item(),
|
||||
"reward_std": rewards.std().item(),
|
||||
"reward_min": rewards.min().item(),
|
||||
"reward_max": rewards.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def check_early_stopping(
|
||||
kl_divergence: float,
|
||||
target_kl: float
|
||||
) -> bool:
|
||||
"""
|
||||
Check if training should stop early based on KL divergence.
|
||||
|
||||
Args:
|
||||
kl_divergence: Current KL divergence
|
||||
target_kl: Target KL threshold
|
||||
|
||||
Returns:
|
||||
should_stop: True if KL exceeds target
|
||||
"""
|
||||
if kl_divergence > target_kl:
|
||||
logger.warning(
|
||||
"Early stopping triggered: KL divergence %.4f > target %.4f",
|
||||
kl_divergence,
|
||||
target_kl
|
||||
)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def compute_log_probs_from_model_output(
|
||||
model_output: torch.Tensor,
|
||||
target: torch.Tensor,
|
||||
noise_level: float = 0.1
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute log probabilities from model predictions.
|
||||
|
||||
For diffusion models, we approximate log probabilities using the
|
||||
negative squared error (assuming Gaussian likelihood).
|
||||
|
||||
Args:
|
||||
model_output: Model predictions [B, C, T, H, W]
|
||||
target: Target values [B, C, T, H, W]
|
||||
noise_level: Assumed noise level (std) for Gaussian likelihood
|
||||
|
||||
Returns:
|
||||
log_probs: Log probabilities [B]
|
||||
"""
|
||||
# Compute mean squared error per sample
|
||||
mse = ((model_output - target) ** 2).flatten(1).mean(dim=1)
|
||||
|
||||
# Log probability under Gaussian: log p(x) = -0.5 * (x - mu)^2 / sigma^2 + const
|
||||
log_probs = -0.5 * mse / (noise_level ** 2)
|
||||
|
||||
return log_probs
|
||||
|
||||
|
||||
def check_for_nan_inf(tensor: torch.Tensor, name: str) -> None:
|
||||
"""
|
||||
Check tensor for NaN or Inf values and raise error if found.
|
||||
|
||||
Args:
|
||||
tensor: Tensor to check
|
||||
name: Name for error message
|
||||
"""
|
||||
if torch.isnan(tensor).any():
|
||||
raise ValueError(f"{name} contains NaN values")
|
||||
if torch.isinf(tensor).any():
|
||||
raise ValueError(f"{name} contains Inf values")
|
||||
@@ -0,0 +1,189 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Per-prompt statistics tracking for GRPO training.
|
||||
|
||||
This module ports the PerPromptStatTracker from FlowGRPO to FastVideo.
|
||||
It tracks reward statistics per unique prompt and computes normalized advantages.
|
||||
|
||||
Ported from:
|
||||
- flow_grpo/flow_grpo/stat_tracking.py
|
||||
|
||||
Key adaptations:
|
||||
1. Uses FastVideo's logging instead of print statements
|
||||
2. Works with single GPU (no distributed logic)
|
||||
3. Supports numpy arrays and torch tensors
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PerPromptStatTracker:
|
||||
"""
|
||||
Tracks reward statistics per unique prompt for advantage normalization.
|
||||
|
||||
This class maintains running statistics (mean, std) for each unique prompt
|
||||
and computes normalized advantages using either per-prompt or global statistics.
|
||||
|
||||
Used in GRPO training to normalize advantages within groups of samples
|
||||
generated from the same prompt, which helps stabilize training when different
|
||||
prompts have different reward scales.
|
||||
"""
|
||||
|
||||
def __init__(self, global_std: bool = False):
|
||||
"""
|
||||
Initialize the per-prompt stat tracker.
|
||||
|
||||
Args:
|
||||
global_std: If True, use global std across all rewards for normalization.
|
||||
If False, use per-prompt std (default, recommended for GRPO).
|
||||
"""
|
||||
self.global_std = global_std
|
||||
self.stats: dict[str, list] = {} # Maps prompt -> list of rewards
|
||||
self.history_prompts: set[int] = set() # Set of hashed prompts seen
|
||||
|
||||
def update(
|
||||
self,
|
||||
prompts: Union[list[str], np.ndarray],
|
||||
rewards: Union[list[float], np.ndarray, torch.Tensor],
|
||||
type: str = 'grpo'
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Update statistics and compute normalized advantages.
|
||||
|
||||
Args:
|
||||
prompts: List or array of prompt strings (one per sample)
|
||||
rewards: Array or tensor of reward values (one per sample)
|
||||
type: Advantage computation type:
|
||||
- 'grpo': Normalize by (reward - mean) / std (default)
|
||||
- 'rwr': Return rewards as-is (reward-weighted regression)
|
||||
- 'sft': Binary advantages (1 for max, 0 otherwise)
|
||||
- 'dpo': DPO-style advantages (1 for max, -1 for min)
|
||||
|
||||
Returns:
|
||||
advantages: Normalized advantages array [num_samples] or [num_samples, ...]
|
||||
Shape matches rewards shape
|
||||
"""
|
||||
# Convert to numpy arrays
|
||||
prompts = np.array(prompts)
|
||||
if isinstance(rewards, torch.Tensor):
|
||||
rewards = rewards.detach().cpu().numpy()
|
||||
rewards = np.array(rewards, dtype=np.float64)
|
||||
|
||||
# Ensure rewards are 1D (one reward per sample)
|
||||
# FlowGRPO expects rewards to be aggregated per sample
|
||||
if rewards.ndim > 1:
|
||||
# If multi-dimensional, flatten or take mean
|
||||
# For [B, num_steps] shape, we typically want one reward per sample
|
||||
# So we take the mean across timesteps
|
||||
if rewards.ndim == 2:
|
||||
# Assume shape is [B, num_steps] - take mean across timesteps
|
||||
rewards = rewards.mean(axis=1)
|
||||
else:
|
||||
# Flatten and take mean for higher dimensions
|
||||
rewards = rewards.reshape(len(prompts), -1).mean(axis=1)
|
||||
|
||||
# Ensure prompts and rewards have matching lengths
|
||||
assert len(prompts) == len(rewards), \
|
||||
f"Prompts ({len(prompts)}) and rewards ({len(rewards)}) must have same length"
|
||||
|
||||
unique_prompts = np.unique(prompts)
|
||||
advantages = np.zeros_like(rewards, dtype=np.float64)
|
||||
|
||||
# First pass: collect rewards for each prompt
|
||||
for prompt in unique_prompts:
|
||||
prompt_mask = prompts == prompt
|
||||
prompt_rewards = rewards[prompt_mask]
|
||||
|
||||
# Store rewards in stats
|
||||
if prompt not in self.stats:
|
||||
self.stats[prompt] = []
|
||||
self.stats[prompt].extend(prompt_rewards.tolist())
|
||||
self.history_prompts.add(hash(prompt))
|
||||
|
||||
# Second pass: compute statistics and advantages
|
||||
for prompt in unique_prompts:
|
||||
prompt_mask = prompts == prompt
|
||||
prompt_rewards = rewards[prompt_mask]
|
||||
|
||||
# Stack all historical rewards for this prompt
|
||||
if len(self.stats[prompt]) > 0:
|
||||
all_prompt_rewards = np.array(self.stats[prompt])
|
||||
else:
|
||||
all_prompt_rewards = prompt_rewards
|
||||
|
||||
# Compute mean and std
|
||||
mean = np.mean(all_prompt_rewards, axis=0, keepdims=True)
|
||||
|
||||
if self.global_std:
|
||||
# Use global std across all rewards
|
||||
std = np.std(rewards, axis=0, keepdims=True) + 1e-4
|
||||
else:
|
||||
# Use per-prompt std
|
||||
std = np.std(all_prompt_rewards, axis=0, keepdims=True) + 1e-4
|
||||
|
||||
# Compute advantages based on type
|
||||
if type == 'grpo':
|
||||
# GRPO: normalize by (reward - mean) / std
|
||||
advantages[prompt_mask] = (prompt_rewards - mean) / std
|
||||
elif type == 'rwr':
|
||||
# Reward-weighted regression: use rewards as-is
|
||||
advantages[prompt_mask] = prompt_rewards
|
||||
elif type == 'sft':
|
||||
# Supervised fine-tuning: binary (1 for max, 0 otherwise)
|
||||
max_reward = np.max(prompt_rewards)
|
||||
advantages[prompt_mask] = (prompt_rewards == max_reward).astype(np.float64)
|
||||
elif type == 'dpo':
|
||||
# DPO-style: 1 for max, -1 for min
|
||||
prompt_rewards_tensor = torch.tensor(prompt_rewards)
|
||||
max_idx = torch.argmax(prompt_rewards_tensor)
|
||||
min_idx = torch.argmin(prompt_rewards_tensor)
|
||||
|
||||
# If all rewards are the same, use first two indices
|
||||
if max_idx == min_idx:
|
||||
min_idx = torch.tensor(0)
|
||||
max_idx = torch.tensor(1) if len(prompt_rewards_tensor) > 1 else torch.tensor(0)
|
||||
|
||||
result = torch.zeros_like(prompt_rewards_tensor, dtype=torch.float64)
|
||||
result[max_idx] = 1.0
|
||||
result[min_idx] = -1.0
|
||||
advantages[prompt_mask] = result.numpy()
|
||||
else:
|
||||
raise ValueError(f"Unknown advantage type: {type}. Must be one of: 'grpo', 'rwr', 'sft', 'dpo'")
|
||||
|
||||
return advantages
|
||||
|
||||
def get_stats(self) -> tuple[float, int]:
|
||||
"""
|
||||
Get statistics about tracked prompts.
|
||||
|
||||
Returns:
|
||||
avg_group_size: Average number of samples per unique prompt
|
||||
history_prompts: Number of unique prompts seen (across all updates)
|
||||
"""
|
||||
if not self.stats:
|
||||
avg_group_size = 0.0
|
||||
else:
|
||||
total_samples = sum(len(v) for v in self.stats.values())
|
||||
avg_group_size = total_samples / len(self.stats)
|
||||
|
||||
history_prompts = len(self.history_prompts)
|
||||
|
||||
return avg_group_size, history_prompts
|
||||
|
||||
def clear(self) -> None:
|
||||
"""
|
||||
Clear all statistics (but keep history_prompts for tracking).
|
||||
|
||||
This is typically called after each epoch to reset per-epoch statistics
|
||||
while maintaining a record of all prompts seen during training.
|
||||
"""
|
||||
self.stats = {}
|
||||
logger.debug("Cleared per-prompt statistics (kept %d unique prompts in history)",
|
||||
len(self.history_prompts))
|
||||
@@ -0,0 +1,877 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
GRPO utilities for Wan model in FastVideo.
|
||||
|
||||
This module ports the SDE step and pipeline functions from FlowGRPO to work with
|
||||
FastVideo's scheduler and pipeline interfaces.
|
||||
|
||||
Ported from:
|
||||
- flow_grpo/flow_grpo/diffusers_patch/wan_pipeline_with_logprob.py
|
||||
|
||||
Key adaptations:
|
||||
1. Uses FastVideo's FlowUniPCMultistepScheduler instead of diffusers' UniPCMultistepScheduler
|
||||
2. Works with FastVideo's WanPipeline (ComposedPipelineBase) instead of diffusers' WanPipeline
|
||||
3. Direct module access via pipeline.get_module() instead of pipeline attributes
|
||||
4. Simplified prompt encoding (direct text encoder usage instead of pipeline stages)
|
||||
"""
|
||||
|
||||
import math
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
|
||||
# for test_wan_transformer2
|
||||
import os
|
||||
from diffusers import WanTransformer3DModel
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def test_wan_transformer():
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(
|
||||
dit_config=WanVideoConfig(),
|
||||
dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = WanTransformer3DModel.from_pretrained(
|
||||
TRANSFORMER_PATH, device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(data_type="dummy", )
|
||||
|
||||
with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
'''
|
||||
INFO 01-19 22:53:46 [wan_grpo_utils.py:74] Model 1 weight sum: 395834.3506456231████ | 1/2 [00:00<00:00, 7.84it/s]
|
||||
INFO 01-19 22:53:46 [wan_grpo_utils.py:75] Model 1 weight mean: 0.0002789536598289884
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:83] Model 2 weight sum: 395834.3506456231
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:84] Model 2 weight mean: 0.0002789536598289884
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:87] Weight sum difference: 0.0
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:89] Weight mean difference: 0.0
|
||||
INFO 01-19 22:53:54 [wan_grpo_utils.py:145] Max Diff: 0.08203125
|
||||
INFO 01-19 22:53:54 [wan_grpo_utils.py:146] Mean Diff: 0.01129150390625
|
||||
'''
|
||||
|
||||
|
||||
def test_wan_transformer2(model2):
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
|
||||
logger.info("loading model1 transformer weight")
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
model1 = WanTransformer3DModel.from_pretrained(
|
||||
TRANSFORMER_PATH,
|
||||
device=device,
|
||||
torch_dtype=precision,
|
||||
).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(data_type="dummy", )
|
||||
|
||||
with torch.amp.autocast("cuda", dtype=precision):
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
|
||||
# Print basic stats for debugging (cast to float32 for stability)
|
||||
out1 = output1.detach().float()
|
||||
out2 = output2.detach().float()
|
||||
logger.info(
|
||||
"output1 stats: min=%s max=%s mean=%s std=%s",
|
||||
out1.min().item(),
|
||||
out1.max().item(),
|
||||
out1.mean().item(),
|
||||
out1.std(unbiased=False).item(),
|
||||
)
|
||||
logger.info(
|
||||
"output2 stats: min=%s max=%s mean=%s std=%s",
|
||||
out2.min().item(),
|
||||
out2.max().item(),
|
||||
out2.mean().item(),
|
||||
out2.std(unbiased=False).item(),
|
||||
)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert (output1.shape == output2.shape
|
||||
), f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert (output1.dtype == output2.dtype
|
||||
), f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
'''
|
||||
when --dit_precision "bf16", use_fsdp hardcoded to False:
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:65] Model 1 weight sum: 395834.3506456231████████████████████████████████████████████████████████| 2/2 [00:00<00:00, 3.25it/s]
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:66] Model 1 weight mean: 0.0002789536598289884
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:75] Model 2 weight sum: 395125.463677882
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:76] Model 2 weight mean: 0.0002739000890162162
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:79] Weight sum difference: 708.8869677411276
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:81] Weight mean difference: 5.053570812772192e-06
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:139] output1 stats: min=-2.28125 max=1.921875 mean=-0.16638492047786713 std=0.458170622587204
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:146] output2 stats: min=-2.296875 max=1.90625 mean=-0.166452556848526 std=0.4579130709171295
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:165] Max Diff: 0.08984375
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:166] Mean Diff: 0.0120849609375
|
||||
|
||||
when --dit_precision "fp32", use_fsdp not changed:
|
||||
|
||||
'''
|
||||
|
||||
|
||||
def sde_step_with_logprob(
|
||||
scheduler: FlowUniPCMultistepScheduler,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: float | torch.FloatTensor,
|
||||
sample: torch.FloatTensor,
|
||||
prev_sample: torch.FloatTensor | None = None,
|
||||
generator: torch.Generator | None = None,
|
||||
deterministic: bool = False,
|
||||
return_pixel_log_prob: bool = False,
|
||||
return_dt_and_std_dev_t: bool = False
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ...]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE.
|
||||
This function propagates the flow process from the learned model outputs
|
||||
(most often the predicted velocity) and computes log probabilities.
|
||||
|
||||
Ported from FlowGRPO's sde_step_with_logprob to work with FastVideo's
|
||||
FlowUniPCMultistepScheduler.
|
||||
|
||||
Args:
|
||||
scheduler: FastVideo FlowUniPCMultistepScheduler instance
|
||||
model_output: The direct output from learned flow model
|
||||
timestep: The current discrete timestep in the diffusion chain
|
||||
sample: A current instance of a sample created by the diffusion process
|
||||
prev_sample: Optional previous sample (if provided, used instead of sampling)
|
||||
generator: Optional random number generator
|
||||
deterministic: If True, no noise is added (deterministic sampling)
|
||||
return_pixel_log_prob: If True, return pixel-level log probabilities (not used)
|
||||
return_dt_and_std_dev_t: If True, return dt and std_dev_t separately
|
||||
|
||||
Returns:
|
||||
If return_dt_and_std_dev_t=True:
|
||||
(prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt)
|
||||
Otherwise:
|
||||
(prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt)
|
||||
"""
|
||||
|
||||
# # Convert all variables to fp32 for numerical stability
|
||||
# model_output = model_output.float()
|
||||
# sample = sample.float()
|
||||
# if prev_sample is not None:
|
||||
# prev_sample = prev_sample.float()
|
||||
|
||||
# Get step indices for current and previous timesteps
|
||||
# Handle both single timestep and batch of timesteps
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
if timestep.ndim == 0:
|
||||
timestep = timestep.unsqueeze(0)
|
||||
step_indices = [
|
||||
scheduler.index_for_timestep(t.item()) for t in timestep
|
||||
]
|
||||
else:
|
||||
step_indices = [scheduler.index_for_timestep(timestep)]
|
||||
|
||||
prev_step_indices = [step + 1 for step in step_indices]
|
||||
|
||||
# Move sigmas to sample device
|
||||
sigmas = scheduler.sigmas.to(sample.device)
|
||||
# myregion debug: hardcode sigmas to flow_grpo's
|
||||
sigmas = torch.Tensor([
|
||||
0.9997, 0.9824, 0.9639, 0.9441, 0.9227, 0.8996, 0.8746, 0.8475, 0.8178,
|
||||
0.7853, 0.7496, 0.7102, 0.6663, 0.6173, 0.5621, 0.4997, 0.4283, 0.3459,
|
||||
0.2498, 0.1362, 0.0000
|
||||
]).to(sample.device, sample.dtype)
|
||||
# end region
|
||||
|
||||
# Get sigma values for current and previous steps
|
||||
sigma = sigmas[step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_prev = sigmas[prev_step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_max = sigmas[0].item() # First sigma (highest)
|
||||
sigma_min = sigmas[-1].item() # Last sigma (lowest)
|
||||
|
||||
dt = sigma_prev - sigma
|
||||
|
||||
# myregion debug
|
||||
print(f"[DEBUG]: sigma_max: {sigma_max}, sigma_min: {sigma_min}, dt: {dt}")
|
||||
print(f"[DEBUG]: in sde_step_with_logprob(), timestep: {timestep}")
|
||||
print(f"[DEBUG]: in sde_step_with_logprob(), sigmas: {sigmas}")
|
||||
print(f"[DEBUG]: in sde_step_with_logprob(), step_indices: {step_indices}")
|
||||
print(
|
||||
f"[DEBUG]: in sde_step_with_logprob(), prev_step_indices: {prev_step_indices}"
|
||||
)
|
||||
'''
|
||||
[DEBUG]: in sde_step_with_logprob(), timestep: tensor([428, 428, 428, 428], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
|
||||
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
|
||||
0.2499, 0.1363, 0.0000], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), step_indices: [16, 16, 16, 16]
|
||||
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [17, 17, 17, 17]
|
||||
|
||||
DEBUG]: in sde_step_with_logprob(), timestep: tensor([249], device='cuda:0')███████████████▎ | 18/20 [00:04<00:00, 3.87step/s, step_time=0.26s, timestep=346.0]
|
||||
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
|
||||
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
|
||||
0.2499, 0.1363, 0.0000], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), step_indices: [18]
|
||||
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [19]
|
||||
|
||||
[DEBUG]: in sde_step_with_logprob(), timestep: tensor([617, 617, 617, 617], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
|
||||
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
|
||||
0.2499, 0.1363, 0.0000], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), step_indices: [13, 13, 13, 13]
|
||||
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [14, 14, 14, 14]
|
||||
'''
|
||||
# endregion
|
||||
|
||||
# Compute std_dev_t and prev_sample_mean using SDE formulation
|
||||
std_dev_t = sigma_min + (sigma_max - sigma_min) * sigma
|
||||
prev_sample_mean = (sample * (1 + std_dev_t**2 / (2 * sigma) * dt) +
|
||||
model_output * (1 + std_dev_t**2 * (1 - sigma) /
|
||||
(2 * sigma)) * dt)
|
||||
|
||||
if prev_sample is not None and generator is not None:
|
||||
raise ValueError(
|
||||
"Cannot pass both generator and prev_sample. Please make sure that either `generator` or"
|
||||
" `prev_sample` stays `None`.")
|
||||
|
||||
# Sample prev_sample if not provided
|
||||
if prev_sample is None:
|
||||
variance_noise = randn_tensor(
|
||||
model_output.shape,
|
||||
generator=generator,
|
||||
device=model_output.device,
|
||||
dtype=model_output.dtype,
|
||||
)
|
||||
sqrt_dt = torch.sqrt(-1 * dt) # dt is negative (going backwards)
|
||||
prev_sample = prev_sample_mean + std_dev_t * sqrt_dt * variance_noise
|
||||
else:
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
# No noise is added during evaluation (deterministic)
|
||||
if deterministic:
|
||||
prev_sample = sample + dt * model_output
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
# Compute log probability: log p(prev_sample | sample, model_output)
|
||||
# Assuming Gaussian distribution: N(prev_sample_mean, (std_dev_t * sqrt_dt)^2)
|
||||
std_dev_sqrt_dt = std_dev_t * sqrt_dt
|
||||
log_prob = (
|
||||
-((prev_sample.detach() - prev_sample_mean)**2) /
|
||||
(2 * (std_dev_sqrt_dt**2)) - torch.log(
|
||||
std_dev_sqrt_dt + 1e-8) # Add small epsilon for numerical stability
|
||||
- torch.log(
|
||||
torch.sqrt(2 * torch.as_tensor(math.pi, device=sample.device))))
|
||||
|
||||
# Mean along all but batch dimension
|
||||
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
|
||||
|
||||
if return_dt_and_std_dev_t:
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt
|
||||
|
||||
|
||||
def wan_pipeline_with_logprob(
|
||||
pipeline,
|
||||
prompt: str | list[str] = None,
|
||||
negative_prompt: str | list[str] = None,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 5.0,
|
||||
num_videos_per_prompt: int | None = 1,
|
||||
generator: torch.Generator | list[torch.Generator] | None = None,
|
||||
latents: torch.Tensor | None = None,
|
||||
prompt_embeds: torch.Tensor | None = None,
|
||||
negative_prompt_embeds: torch.Tensor | None = None,
|
||||
output_type: str | None = "pt",
|
||||
return_dict: bool = False,
|
||||
attention_kwargs: dict[str, Any] | None = None,
|
||||
max_sequence_length: int = 512,
|
||||
deterministic: bool = False,
|
||||
kl_reward: float = 0.0,
|
||||
return_pixel_log_prob: bool = False,
|
||||
) -> tuple[torch.Tensor, list[torch.Tensor], list[torch.Tensor],
|
||||
list[torch.Tensor], torch.Tensor | None]:
|
||||
"""
|
||||
Wan pipeline with log probability computation for GRPO training.
|
||||
|
||||
Ported from FlowGRPO's wan_pipeline_with_logprob to work with FastVideo's WanPipeline.
|
||||
This function generates videos and computes log probabilities at each denoising step.
|
||||
|
||||
Args:
|
||||
pipeline: FastVideo WanPipeline instance
|
||||
prompt: Text prompt(s) for generation
|
||||
negative_prompt: Negative prompt(s) for classifier-free guidance
|
||||
height: Height of generated video
|
||||
width: Width of generated video
|
||||
num_frames: Number of frames in generated video
|
||||
num_inference_steps: Number of denoising steps
|
||||
guidance_scale: Classifier-free guidance scale
|
||||
num_videos_per_prompt: Number of videos to generate per prompt
|
||||
generator: Random generator for reproducibility
|
||||
latents: Optional initial latents
|
||||
prompt_embeds: Optional pre-computed prompt embeddings
|
||||
negative_prompt_embeds: Optional pre-computed negative prompt embeddings
|
||||
output_type: Output type ("pt" for PyTorch tensor, "np" for numpy, "latent" for latents only)
|
||||
return_dict: Whether to return dict (not used, always returns tuple)
|
||||
attention_kwargs: Optional attention kwargs
|
||||
max_sequence_length: Maximum sequence length for text encoding
|
||||
deterministic: If True, use deterministic sampling (no noise)
|
||||
kl_reward: KL reward coefficient (if > 0, computes KL divergence)
|
||||
return_pixel_log_prob: If True, return pixel-level log probabilities (not used)
|
||||
|
||||
Returns:
|
||||
Tuple of:
|
||||
- video: Generated video tensor [B, C, T, H, W] or latents if output_type="latent"
|
||||
- all_latents: List of latents at each step [num_steps+1] of shape [B, C, T, H, W]
|
||||
- all_log_probs: List of log probabilities at each step [num_steps] of shape [B]
|
||||
- all_kl: List of KL divergences at each step [num_steps] of shape [B] (if kl_reward > 0)
|
||||
- prompt_ids: Tokenized prompt IDs [B, seq_len] (None if prompt_embeds were provided)
|
||||
"""
|
||||
# Get device from transformer
|
||||
transformer = pipeline.get_module("transformer")
|
||||
|
||||
# myregion debug: test transformer output
|
||||
logger.info("testing transformer, running test_wan_transformer2")
|
||||
test_wan_transformer()
|
||||
# test_wan_transformer2(transformer)
|
||||
# endregion
|
||||
|
||||
# hardcode dtype for debug
|
||||
# transformer_dtype = torch.float32
|
||||
# use get_compute_dtype() to get dtype based on mixed precision
|
||||
transformer_dtype = get_compute_dtype()
|
||||
logger.info(f"[DEBUG]: transformer_dtype: {transformer_dtype}")
|
||||
|
||||
# Get scheduler and other modules
|
||||
scheduler = pipeline.get_module("scheduler")
|
||||
vae = pipeline.get_module("vae")
|
||||
|
||||
# Determine batch size
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
elif prompt_embeds is not None:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
else:
|
||||
raise ValueError("Either prompt or prompt_embeds must be provided")
|
||||
|
||||
# Encode prompts if not provided
|
||||
prompt_ids = None
|
||||
if prompt_embeds is None:
|
||||
# Encode prompts directly using text encoder and tokenizer
|
||||
# This is a simplified encoding - for full pipeline encoding, use TextEncodingStage
|
||||
text_encoder = pipeline.get_module("text_encoder")
|
||||
tokenizer = pipeline.get_module("tokenizer")
|
||||
|
||||
# Normalize to list
|
||||
if isinstance(prompt, str):
|
||||
prompts_list = [prompt]
|
||||
else:
|
||||
prompts_list = prompt
|
||||
|
||||
# Tokenize prompts
|
||||
text_inputs = tokenizer(prompts_list,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(pipeline.device)
|
||||
|
||||
# Store prompt_ids for return
|
||||
prompt_ids = text_inputs["input_ids"]
|
||||
|
||||
# Encode with text encoder
|
||||
with torch.no_grad():
|
||||
outputs = text_encoder(
|
||||
text_inputs["input_ids"],
|
||||
attention_mask=text_inputs["attention_mask"],
|
||||
output_hidden_states=True,
|
||||
)
|
||||
# Get last hidden state (Wan typically uses last hidden state)
|
||||
prompt_embeds = outputs.last_hidden_state
|
||||
|
||||
# Encode negative prompts if CFG is enabled
|
||||
if guidance_scale > 1.0:
|
||||
if negative_prompt is None:
|
||||
negative_prompt = [""] * len(prompts_list)
|
||||
elif isinstance(negative_prompt, str):
|
||||
negative_prompt = [negative_prompt]
|
||||
|
||||
neg_text_inputs = tokenizer(negative_prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(pipeline.device)
|
||||
|
||||
with torch.no_grad():
|
||||
neg_outputs = text_encoder(
|
||||
neg_text_inputs["input_ids"],
|
||||
attention_mask=neg_text_inputs["attention_mask"],
|
||||
output_hidden_states=True,
|
||||
)
|
||||
negative_prompt_embeds = neg_outputs.last_hidden_state
|
||||
else:
|
||||
negative_prompt_embeds = None
|
||||
|
||||
# myregion Debug: Print shapes of prompt embeddings
|
||||
logger.info(
|
||||
f"After encoding - prompt_embeds shape: {prompt_embeds.shape if prompt_embeds is not None else None}"
|
||||
)
|
||||
logger.info(
|
||||
f"After encoding - negative_prompt_embeds shape: {negative_prompt_embeds.shape if negative_prompt_embeds is not None else None}"
|
||||
)
|
||||
logger.info(
|
||||
f"After encoding - prompt_embeds dtype: {prompt_embeds.dtype if prompt_embeds is not None else None}"
|
||||
)
|
||||
logger.info(
|
||||
f"After encoding - negative_prompt_embeds dtype: {negative_prompt_embeds.dtype if negative_prompt_embeds is not None else None}"
|
||||
)
|
||||
'''
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:290] After encoding - prompt_embeds shape: torch.Size([4, 512, 4096])
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:291] After encoding - negative_prompt_embeds shape: None
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:292] After encoding - prompt_embeds dtype: torch.float32
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:293] After encoding - negative_prompt_embeds dtype: None
|
||||
'''
|
||||
# endregion
|
||||
# logger.info("wan_pipeline_with_logprob's transformer class type: %s", type(transformer))
|
||||
# logger.info("Variables in transformer: %s", str(dir(transformer)))
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
if negative_prompt_embeds is not None:
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
|
||||
|
||||
# Prepare timesteps
|
||||
scheduler.set_timesteps(num_inference_steps, device=pipeline.device)
|
||||
timesteps = scheduler.timesteps
|
||||
|
||||
# Prepare latent variables
|
||||
num_channels_latents = transformer.config.in_channels
|
||||
vae = pipeline.get_module("vae")
|
||||
# Get VAE scale factors
|
||||
vae_scale_factor_spatial = vae.spatial_compression_ratio
|
||||
vae_scale_factor_temporal = vae.temporal_compression_ratio
|
||||
|
||||
if latents is None:
|
||||
# Generate random latents
|
||||
# Note: num_frames in latents accounts for temporal compression
|
||||
num_latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
|
||||
latents_shape = (
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
num_latent_frames,
|
||||
height // vae_scale_factor_spatial,
|
||||
width // vae_scale_factor_spatial,
|
||||
)
|
||||
if generator is not None:
|
||||
if isinstance(generator, list):
|
||||
latents = [
|
||||
torch.randn(
|
||||
latents_shape[1:],
|
||||
generator=gen,
|
||||
device=pipeline.device,
|
||||
dtype=transformer_dtype,
|
||||
) for gen in generator
|
||||
]
|
||||
latents = torch.stack(latents, dim=0)
|
||||
else:
|
||||
latents = torch.randn(
|
||||
latents_shape,
|
||||
generator=generator,
|
||||
device=pipeline.device,
|
||||
dtype=transformer_dtype,
|
||||
)
|
||||
else:
|
||||
latents = torch.randn(latents_shape,
|
||||
device=pipeline.device,
|
||||
dtype=transformer_dtype)
|
||||
else:
|
||||
latents = latents.to(device=pipeline.device, dtype=transformer_dtype)
|
||||
|
||||
|
||||
# myregion Debug: Print latents shape, dtype, and value range
|
||||
logger.info("=" * 80)
|
||||
logger.info("Latents Debug Information:")
|
||||
logger.info(f" Shape: {latents.shape}")
|
||||
logger.info(f" Dtype: {latents.dtype}")
|
||||
logger.info(f" Min value: {latents.min().item():.6f}")
|
||||
logger.info(f" Max value: {latents.max().item():.6f}")
|
||||
logger.info(f" Mean value: {latents.mean().item():.6f}")
|
||||
logger.info(f" Std value: {latents.std().item():.6f}")
|
||||
logger.info(f" Device: {latents.device}")
|
||||
logger.info("=" * 80)
|
||||
'''
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:355] ================================================================================
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:356] Latents Debug Information:
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:357] Shape: torch.Size([4, 16, 9, 30, 52])
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:358] Dtype: torch.bfloat16
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:359] Min value: -4.500000
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:360] Max value: 4.656250
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:361] Mean value: 0.000111
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:362] Std value: 1.000000
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:363] Device: cuda:0
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:364] ================================================================================
|
||||
'''
|
||||
# endregion
|
||||
|
||||
all_latents = [latents]
|
||||
all_log_probs = []
|
||||
all_kl = []
|
||||
|
||||
# myregion Debug
|
||||
logger.info("Tensor type issue debugging:")
|
||||
logger.info(f"latents: {type(latents)}")
|
||||
logger.info(f"prompt_embeds: {type(prompt_embeds)}")
|
||||
logger.info(
|
||||
f"[DEBUG]: before denoising loop: type(timesteps): {type(timesteps)}")
|
||||
logger.info(
|
||||
f"[DEBUG]: before denoising loop: timesteps.shape: {timesteps.shape}")
|
||||
# endregion
|
||||
|
||||
# Progress bar for denoising loop
|
||||
progress_bar = tqdm(enumerate(timesteps),
|
||||
total=len(timesteps),
|
||||
desc="Denoising steps",
|
||||
unit="step")
|
||||
|
||||
for i, t in progress_bar:
|
||||
step_start_time = time.time()
|
||||
latents_ori = latents.clone()
|
||||
timestep = t.expand(latents.shape[0]) if isinstance(
|
||||
t, torch.Tensor) else torch.tensor([t] * latents.shape[0],
|
||||
device=pipeline.device)
|
||||
|
||||
logger.info(
|
||||
f"[DEBUG]: before set_forward_context: current_timestep=i:{i}")
|
||||
# Predict noise with transformer
|
||||
with set_forward_context(
|
||||
current_timestep=t.item(),
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
noise_pred = transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_pred.to(prompt_embeds.dtype)
|
||||
|
||||
# Classifier-free guidance
|
||||
if guidance_scale > 1.0:
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
noise_uncond = transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_uncond + guidance_scale * (noise_pred -
|
||||
noise_uncond)
|
||||
|
||||
# SDE step with log probability
|
||||
latents, log_prob, prev_latents_mean, std_dev_t = sde_step_with_logprob(
|
||||
scheduler,
|
||||
noise_pred, #.float(),
|
||||
t.unsqueeze(0) if isinstance(t, torch.Tensor) else t,
|
||||
latents, #.float(),
|
||||
deterministic=deterministic,
|
||||
return_pixel_log_prob=return_pixel_log_prob)
|
||||
# sde_step_with_logprob returns fp32
|
||||
# latents = latents.to(transformer_dtype)
|
||||
prev_latents = latents.clone()
|
||||
|
||||
all_latents.append(latents)
|
||||
all_log_probs.append(log_prob)
|
||||
|
||||
# Compute KL divergence if kl_reward > 0 (for KL reward in sampling)
|
||||
if kl_reward > 0 and not deterministic:
|
||||
# Use reference model (disable adapter if using LoRA)
|
||||
latent_model_input_ref = torch.cat(
|
||||
[latents_ori] * 2) if guidance_scale > 1.0 else latents_ori
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
with transformer.disable_adapter() if hasattr(
|
||||
transformer, 'disable_adapter') else torch.no_grad():
|
||||
noise_pred_ref = transformer(
|
||||
hidden_states=latent_model_input_ref,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred_ref = noise_pred_ref.to(prompt_embeds.dtype)
|
||||
|
||||
# Perform guidance for reference model
|
||||
if guidance_scale > 1.0:
|
||||
noise_pred_uncond_ref, noise_pred_text_ref = noise_pred_ref.chunk(
|
||||
2)
|
||||
noise_pred_ref = noise_pred_uncond_ref + guidance_scale * (
|
||||
noise_pred_text_ref - noise_pred_uncond_ref)
|
||||
|
||||
# Compute reference log prob
|
||||
_, ref_log_prob, ref_prev_latents_mean, ref_std_dev_t = sde_step_with_logprob(
|
||||
scheduler,
|
||||
noise_pred_ref.float(),
|
||||
t.unsqueeze(0) if isinstance(t, torch.Tensor) else t,
|
||||
latents_ori.float(),
|
||||
prev_sample=prev_latents.float(),
|
||||
deterministic=deterministic,
|
||||
)
|
||||
|
||||
# Compute KL divergence: KL = (mean_diff)^2 / (2 * std^2)
|
||||
assert torch.allclose(
|
||||
std_dev_t, ref_std_dev_t
|
||||
), "std_dev_t should match between current and reference"
|
||||
kl = (prev_latents_mean - ref_prev_latents_mean)**2 / (2 *
|
||||
std_dev_t**2)
|
||||
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
|
||||
all_kl.append(kl)
|
||||
else:
|
||||
# No KL reward, set to zero
|
||||
all_kl.append(torch.zeros(len(latents), device=latents.device))
|
||||
|
||||
# Update progress bar with timing information
|
||||
step_time = time.time() - step_start_time
|
||||
progress_bar.set_postfix({
|
||||
"step_time":
|
||||
f"{step_time:.2f}s",
|
||||
"timestep":
|
||||
f"{t.item() if isinstance(t, torch.Tensor) else t:.1f}"
|
||||
})
|
||||
|
||||
# Decode latents to video if needed
|
||||
if output_type != "latent":
|
||||
latents = latents.to(vae.dtype)
|
||||
|
||||
# Apply VAE normalization (Wan VAE specific)
|
||||
# Wan VAE requires denormalization before decoding
|
||||
if hasattr(vae, 'config') and hasattr(vae.config,
|
||||
'latents_mean') and hasattr(
|
||||
vae.config, 'latents_std'):
|
||||
# Get z_dim from config or VAE
|
||||
z_dim = getattr(vae.config, 'z_dim', latents.shape[1])
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype).view(
|
||||
1, z_dim, 1, 1, 1))
|
||||
latents_std = (
|
||||
1.0 / torch.tensor(vae.config.latents_std,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype).view(1, z_dim, 1, 1, 1))
|
||||
latents = latents / latents_std + latents_mean
|
||||
elif hasattr(vae, 'latents_mean') and hasattr(vae, 'latents_std'):
|
||||
# Alternative: check if latents_mean/std are direct attributes
|
||||
z_dim = latents.shape[1]
|
||||
latents_mean = (torch.tensor(vae.latents_mean,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype).view(
|
||||
1, z_dim, 1, 1, 1))
|
||||
latents_std = (1.0 / torch.tensor(
|
||||
vae.latents_std, device=latents.device,
|
||||
dtype=latents.dtype).view(1, z_dim, 1, 1, 1))
|
||||
latents = latents / latents_std + latents_mean
|
||||
|
||||
# Decode using VAE
|
||||
with torch.no_grad():
|
||||
video = vae.decode(latents.float(), return_dict=False)[0]
|
||||
# VAE.decode returns tensor directly (not tuple)
|
||||
|
||||
# Postprocess video: convert from [-1, 1] to [0, 1]
|
||||
# FastVideo VAE typically outputs in [-1, 1] range
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
return video, all_latents, all_log_probs, all_kl, prompt_ids
|
||||
@@ -174,17 +174,18 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
training_args.train_batch_size,
|
||||
parquet_schema=self.train_dataset_schema,
|
||||
num_data_workers=training_args.dataloader_num_workers,
|
||||
cfg_rate=training_args.training_cfg_rate,
|
||||
drop_last=True,
|
||||
text_padding_length=training_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.
|
||||
text_len, # type: ignore[attr-defined]
|
||||
seed=self.seed)
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
training_args.train_batch_size,
|
||||
parquet_schema=self.train_dataset_schema,
|
||||
num_data_workers=training_args.dataloader_num_workers,
|
||||
cfg_rate=training_args.training_cfg_rate,
|
||||
drop_last=True,
|
||||
text_padding_length=training_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.
|
||||
text_len, # type: ignore[attr-defined]
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
if self.training_args.boundary_ratio is not None:
|
||||
@@ -192,19 +193,21 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
else:
|
||||
self.boundary_timestep = None
|
||||
|
||||
logger.info("train_dataloader length: %s", len(self.train_dataloader))
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
logger.info("train_dataloader length: %s", len(self.train_dataloader))
|
||||
logger.info("train_sp_batch_size: %s",
|
||||
training_args.train_sp_batch_size)
|
||||
logger.info("gradient_accumulation_steps: %s",
|
||||
training_args.gradient_accumulation_steps)
|
||||
logger.info("sp_size: %s", training_args.sp_size)
|
||||
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
training_args.train_sp_batch_size)
|
||||
self.num_train_epochs = math.ceil(training_args.max_train_steps /
|
||||
self.num_update_steps_per_epoch)
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
training_args.train_sp_batch_size)
|
||||
self.num_train_epochs = math.ceil(training_args.max_train_steps /
|
||||
self.num_update_steps_per_epoch)
|
||||
|
||||
# TODO(will): is there a cleaner way to track epochs?
|
||||
self.current_epoch = 0
|
||||
@@ -575,8 +578,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
else:
|
||||
self.noise_random_generator = torch.Generator(device=self.device).manual_seed(
|
||||
self.seed)
|
||||
self.noise_gen_cuda = torch.Generator(
|
||||
device=current_platform.device_name).manual_seed(self.seed)
|
||||
self.validation_random_generator = torch.Generator(
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
|
||||
from fastvideo.training.rl.rl_pipeline import RLPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanRLTrainingPipeline(RLPipeline):
|
||||
"""
|
||||
A training pipeline for Wan with RL/GRPO support.
|
||||
|
||||
This pipeline extends RLPipeline with Wan-specific initialization.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _create_inference_pipeline(self, training_args: TrainingArgs,
|
||||
dit_cpu_offload: bool):
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
loaded_modules = {
|
||||
"transformer": self.get_module("transformer"),
|
||||
}
|
||||
transformer_2 = self.get_module("transformer_2", None)
|
||||
if transformer_2 is not None:
|
||||
loaded_modules["transformer_2"] = transformer_2
|
||||
text_encoder = self.get_module("text_encoder", None)
|
||||
if text_encoder is not None:
|
||||
loaded_modules["text_encoder"] = text_encoder
|
||||
tokenizer = self.get_module("tokenizer", None)
|
||||
if tokenizer is not None:
|
||||
loaded_modules["tokenizer"] = tokenizer
|
||||
vae = self.get_module("vae", None)
|
||||
if vae is not None:
|
||||
loaded_modules["vae"] = vae
|
||||
|
||||
return WanPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules=loaded_modules,
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=dit_cpu_offload)
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
self.validation_pipeline = self._create_inference_pipeline(
|
||||
training_args, dit_cpu_offload=True)
|
||||
|
||||
def _build_sampling_pipeline(self, training_args: TrainingArgs):
|
||||
return self._create_inference_pipeline(training_args,
|
||||
dit_cpu_offload=False)
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting RL training pipeline...")
|
||||
|
||||
pipeline = WanRLTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("RL training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
# Enable RL mode
|
||||
args.rl_mode = True
|
||||
main(args)
|
||||
@@ -132,13 +132,9 @@ 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,
|
||||
extra=extra)
|
||||
logging_info=logging_info)
|
||||
|
||||
return result_batch
|
||||
|
||||
@@ -652,8 +648,7 @@ class WorkerMultiprocProc:
|
||||
logging_info = output_batch.logging_info
|
||||
self.pipe.send({
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"logging_info": logging_info,
|
||||
"extra": output_batch.extra,
|
||||
"logging_info": logging_info
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(
|
||||
|
||||
@@ -29,7 +29,6 @@ dependencies = [
|
||||
"diffusers>=0.33.1",
|
||||
"torch>=2.9.1",
|
||||
"torchvision",
|
||||
"torchaudio",
|
||||
|
||||
# Acceleration & Optimization
|
||||
"accelerate==1.0.1",
|
||||
|
||||
@@ -1,444 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,7 +0,0 @@
|
||||
# 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/`.
|
||||
@@ -1,7 +0,0 @@
|
||||
# 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/`.)
|
||||
@@ -1,126 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,378 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,296 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,343 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,280 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,272 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,190 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,132 +0,0 @@
|
||||
# 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