Compare commits

..
25 Commits
Author SHA1 Message Date
Will Lin 141a1140f6 refactor sampling pipeline 2026-01-20 15:40:53 -08:00
Shijie Wang 6294015389 Debug transformer output misalignment 2026-01-20 14:38:47 -08:00
Shijie Wang e31b6c9e90 Fix OCR Rewards 2026-01-20 14:38:22 -08:00
Tamoghno Kandar bf0ff21eeb Fix OCR Rewards 2026-01-20 14:38:22 -08:00
Shijie Wang 6f937102ad Enable validation videos 2026-01-20 14:38:22 -08:00
Tamoghno Kandar 0164e93019 Add Validation Loop 2026-01-20 14:38:21 -08:00
Shijie Wang 67e457aa92 resolved cuda OOM error 2026-01-20 14:38:21 -08:00
Shijie Wang d795f0c443 remove additional sampling pipeline 2026-01-20 14:38:21 -08:00
loaydatrain 02452dd6e7 fixed dtype mismatch 2026-01-20 14:38:21 -08:00
Shijie Wang 91ef24bc14 update run script 2026-01-20 14:38:20 -08:00
Shijie Wang e76e9fda15 minor fix 2026-01-20 14:38:20 -08:00
Shijie Wang 3b17f5a621 fix trajectory collection & reward computation 2026-01-20 14:38:20 -08:00
Shijie Wang d758878705 minor fix 2026-01-20 14:38:19 -08:00
Shijie Wang 689e629420 Add entry point script 2026-01-20 14:38:19 -08:00
Shijie Wang 873dc9695f Complete train_one_step and grpo policy loss 2026-01-20 14:38:19 -08:00
Shijie Wang bfc0f46d61 Implement trajectories collection, reward and advantage computing 2026-01-20 14:38:18 -08:00
Shijie Wang 39907dbe4d Port per-prompt stat tracker 2026-01-20 14:38:18 -08:00
Shijie Wang abdd0c9b6a Implement SDE step & SDE pipeline with log prob 2026-01-20 14:38:18 -08:00
Shijie (Jacob) Wang f32a12200d Refactor and trim down unnecessary RL args 2026-01-20 14:38:18 -08:00
Shijie Wang f1d2c9e6b7 Add RL dataset & dataloader 2026-01-20 14:38:18 -08:00
Jiali Chen 450579cb42 init algorithm backbone and refactor rl_pipeline 2026-01-20 14:38:17 -08:00
Jiali Chen 26d7d6cc08 minor bug fix 2026-01-20 14:38:17 -08:00
Jiali Chen 44f0124eaa refactor and add ocr reward model 2026-01-20 14:38:17 -08:00
Jiali Chen d3ace51394 Phase 1 minor fixes 2026-01-20 14:38:17 -08:00
Jiali Chen 58954c660b implement Phase 1 backbone code 2026-01-20 14:38:16 -08:00
75 changed files with 4535 additions and 9607 deletions
-34
View File
@@ -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()
+129
View File
@@ -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[@]}"
+1 -12
View File
@@ -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"]))
+1 -2
View File
@@ -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"
]
-84
View File
@@ -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",
]
-45
View File
@@ -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
+1 -3
View File
@@ -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"
]
-50
View File
@@ -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
-7
View File
@@ -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
}
-20
View File
@@ -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 = ""
+30 -23
View File
@@ -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
}
+3 -1
View File
@@ -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"
]
+174
View File
@@ -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
-100
View File
@@ -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
View File
@@ -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
-9
View File
@@ -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
-563
View File
@@ -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
+14 -183
View File
@@ -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
+1
View File
@@ -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
+1 -11
View File
@@ -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):
-1
View File
@@ -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] = {
-7
View File
@@ -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",
+188 -33
View File
@@ -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
+3 -15
View File
@@ -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)
+10 -24
View File
@@ -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():
@@ -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,
+23 -2
View File
@@ -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()
+8 -1
View File
@@ -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",
]
+6
View File
@@ -0,0 +1,6 @@
from .rl_pipeline import RLPipeline, create_rl_pipeline
__all__ = [
"RLPipeline",
"create_rl_pipeline",
]
+11
View File
@@ -0,0 +1,11 @@
from .rewards import (
create_reward_models,
MultiRewardAggregator,
ValueModel
)
__all__ = [
"create_reward_models",
"MultiRewardAggregator",
"ValueModel",
]
+63
View File
@@ -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})"
+206
View File
@@ -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()}")
+338
View File
@@ -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
+385
View File
@@ -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")
+189
View File
@@ -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))
+877
View File
@@ -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
+27 -20
View File
@@ -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)
+2 -7
View File
@@ -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(
-1
View File
@@ -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()
-7
View File
@@ -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/`.
View File
-7
View File
@@ -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/`.)
View File
@@ -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)
-343
View File
@@ -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)
View File
@@ -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)
-190
View File
@@ -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)