Compare commits

...
Author SHA1 Message Date
JerryZhou54 38f9c41e46 refactor hy15 sf distill 2026-01-26 20:40:11 +00:00
JerryZhou54 b1f88ad5d5 refactor hy15 causal denoising 2026-01-26 15:54:20 +00:00
JerryZhou54 90a598bd9e fix lint 2026-01-26 03:09:37 +00:00
JerryZhou54 90d86d5a79 compatible with wan sf 2026-01-26 01:25:34 +00:00
JerryZhou54 7d373cd2c4 Add context forcing, ode_init to sf training 2026-01-26 01:01:47 +00:00
JerryZhou54 0806218156 small change 2026-01-26 01:01:47 +00:00
JerryZhou54 8c6056fbe2 ckpt 2026-01-26 01:01:44 +00:00
JerryZhou54 e4705349d0 Ode init running for hy15 2026-01-26 00:57:22 +00:00
JerryZhou54 50e63840c6 Ode runnable for hy15 2026-01-26 00:57:18 +00:00
JerryZhou54 f1ec0cde18 Add support for ode_init inference for hy15 & support multiple timesteps for hy15 2026-01-26 00:51:50 +00:00
Matthew Noto 1b503554d1 [bugfix] fix torchvision import (#1039) 2026-01-24 22:37:14 -08:00
Shreejith SGandgemini-code-assist[bot] 351ceb7c59 [bugfix]: handle architectural differences while lora extraction (#1035)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-24 15:42:28 -08:00
KyleShao 10875e0d7b [bugfix] Fix NCCL all_gather contiguity + correct ParallelTiledVAE decode tiling threshold (#1037) 2026-01-24 15:37:35 -08:00
alexzms 1eaae8a10b [ci] Increase ci test error threshold (#1038) 2026-01-24 15:36:10 -08:00
Mingjia Huo 59e00f6164 [feat] Add HY-World1.5-Bidirectional-480P-I2V (#1027)
VAE requires further improvement, will raise PR in near future.
2026-01-23 14:18:04 -08:00
745cc05b10 [bugfix] Allow update timesteps for hy1.5 model. (#1033)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-22 22:04:53 -08:00
William Lin c5dc244871 [bugfix] add omegaconf as dep. (#1032) 2026-01-22 11:59:28 -08:00
alexzms dbf3917bf4 [fastvideo-kernel] replace map to index with Triton implementation + add vsa benchmark (#1029) 2026-01-22 11:35:02 -08:00
XOR-op 050f189c95 fix: SP for hunyuanvideo 1.5 (#1026) 2026-01-21 14:40:06 -08:00
Shao DuanandWill Lin 029216029f Added LTX-2 Distilled T2V Generation (#1016)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-21 14:11:39 -08:00
126 changed files with 17689 additions and 530 deletions
@@ -0,0 +1,169 @@
#!/bin/bash
#SBATCH --job-name=dmd_3333
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_3333_output/dmd_3333_%j.out
#SBATCH --error=dmd_3333_output/dmd_3333_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv-distill
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation with Hunyuan1.5:
GENERATOR_MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v" # Teacher model
FAKE_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
# REAL_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled"
# FAKE_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
# DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
DATA_DIR="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/vidprom_16k_text_embed"
DATA_DIR_2="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/fv-ode-preprocessing-16k-hy15-121"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFhy1.5_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_3333_1e-5_1e-5_cfg6_corrected_scheduler"
--wandb_run_name "self_forcing_3333_context_forcing"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 848
--enable_gradient_checkpointing_type "full"
--simulate_generator_forward
--log_visualization
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
# --resume-from-checkpoint "/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_1333/checkpoint-300"
--init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_worldplay_df_init_3333/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_distilled_ode_init_3333/checkpoint-2400/transformer/diffusion_pytorch_model.safetensors
)
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
# --data_path_2 "$DATA_DIR_2"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "8"
--validation_guidance_scale "1.0" # not used for dmd inference
--text-encoder-cpu-offload
# --vae_cpu_offload True
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 100
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
)
dmd_args=(
--warp_denoising_step
--dmd_denoising_steps '1000,875,750,625,500,375,250,125'
# --dmd_denoising_steps '1000,760,520,280'
# --dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.5
--fake_score_learning_rate 5e-6
--fake_score_betas '0.0,0.999'
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
# --use-context-forcing True
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/hy15_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
+24 -4
View File
@@ -2,14 +2,14 @@ from fastvideo import VideoGenerator
import json
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_hy15"
OUTPUT_PATH = "video_samples_hy15_t2v_distilled"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
@@ -18,15 +18,35 @@ def main():
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
# init_weights_from_safetensors="/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_1333_5e-5_2e-6_cfg3.5/checkpoint-500/ema/generator_ema.safetensors"
)
# json_path = "/mnt/weka/home/hao.zhang/wei/FastVideo/data/mixkit_i2v_full_720p.json"
# with open(json_path, 'r') as f:
# data_list = json.load(f)["data"]
# # Now you can index into data_list however you like
# # For example: data_list[0], data_list[1:3], etc.
# for data in data_list:
# prompt = data["prompt"]
# image_path = data["image_path"]
# generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path, num_frames=121, fps=24)
# return
# prompt = (
# "In a brightly lit studio, a photographer wearing a denim jacket focuses intently, capturing shots with a professional camera. Facing him, a model stands gracefully, adjusting her long, flowing hair with delicate movements. The scene is characterized by strong contrasts; the model's soft pink attire and gentle gestures complement the rugged, precise demeanor of the photographer. Positioned against a minimalist backdrop, the pair work seamlessly, with the camera\u2019s lens pointed directly at the model, capturing her elegance. The soft, diffused lighting casts a gentle glow on both subjects, creating an airy and ethereal atmosphere perfect for a high-fashion photo shoot."
# )
# video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=121, fps=24, image_path="data/1.png")
# return
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
@@ -35,7 +55,7 @@ def main():
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
+224
View File
@@ -0,0 +1,224 @@
# SPDX-License-Identifier: Apache-2.0
"""
Basic example for HYWorld (HY-WorldPlay) video generation using FastVideo.
This example replicates the same functionality as HY-WorldPlay/run.sh,
demonstrating image-to-video generation with camera trajectory control.
"""
import time
import math
import numpy as np
import imageio
import torchvision
from einops import rearrange
from fastvideo import VideoGenerator
from fastvideo.pipelines import ForwardBatch
from fastvideo.utils import shallow_asdict, align_to
from fastvideo.logger import init_logger
from fastvideo.models.dits.hyworld.pose import pose_to_input, compute_latent_num
from fastvideo.models.dits.hyworld.resolution_utils import get_resolution_from_image
logger = init_logger(__name__)
class HYWorldVideoGenerator(VideoGenerator):
"""Extended VideoGenerator that adds HYWorld-specific parameters to batch.extra."""
def _generate_single_video(self, prompt: str, sampling_param=None, **kwargs):
"""Override to add viewmats, Ks, and action to batch.extra."""
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
if sampling_param is None:
from fastvideo.configs.sample import SamplingParam
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
# Update sampling param with kwargs
if kwargs:
for key, value in kwargs.items():
if hasattr(sampling_param, key):
setattr(sampling_param, key, value)
# Get pose string from sampling_param or kwargs
pose = kwargs.get('pose', getattr(sampling_param, 'POSE', 'w-31'))
num_frames = kwargs.get('num_frames', getattr(sampling_param, 'num_frames', 125))
# Calculate number of latents
latent_num = compute_latent_num(num_frames)
# Convert pose to viewmats, Ks, and action
viewmats, Ks, action = pose_to_input(pose, latent_num)
# Convert to tensors and add batch dimension
viewmats = viewmats.unsqueeze(0) # (1, T, 4, 4)
Ks = Ks.unsqueeze(0) # (1, T, 3, 3)
action = action.unsqueeze(0) # (1, T)
# Validate inputs
prompt = prompt.strip()
sampling_param = sampling_param.__class__(**shallow_asdict(sampling_param))
output_path = kwargs.get("output_path", sampling_param.output_path)
sampling_param.prompt = prompt
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip()
# Validate dimensions
if (sampling_param.height <= 0 or sampling_param.width <= 0 or
sampling_param.num_frames <= 0):
raise ValueError(
f"Height, width, and num_frames must be positive integers")
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1
else:
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
if use_temporal_scaling_frames:
new_num_frames = (orig_latent_num_frames - 1) * temporal_scale_factor + 1
else:
divisor = math.lcm(3, num_gpus)
orig_latent_num_frames = (
(orig_latent_num_frames + divisor - 1) // divisor) * divisor
new_num_frames = orig_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
sampling_param.num_frames, new_num_frames, fastvideo_args.num_gpus)
sampling_param.num_frames = new_num_frames
# Calculate sizes
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
# Calculate latent sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
# Add HYWorld-specific parameters to batch.extra
batch.extra['viewmats'] = viewmats
batch.extra['Ks'] = Ks
batch.extra['action'] = action
batch.extra['chunk_latent_frames'] = 16 # For bidirectional model
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
if batch.return_frames:
return frames
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
# Default prompt from HY-WorldPlay run.sh
DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a calm body of water. Lush green trees and foliage line the path and the far bank of the water. A traditional-style pavilion with a tiered, reddish-brown roof sits on the far shore. The water reflects the surrounding greenery and the sky. The scene is bathed in soft, natural light, creating a tranquil and serene atmosphere. The pathway is composed of large, rectangular stones, and the bridge is constructed of light gray stone. The overall composition emphasizes the peaceful and harmonious nature of the landscape.'
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
def main():
import argparse
# pose: (a, w, s, d) - (15, 31)
# num_frames: (61, 125)
parser = argparse.ArgumentParser(description="HYWorld video generation with FastVideo")
parser.add_argument("--prompt", type=str, default=DEFAULT_PROMPT, help="Text prompt for video generation")
parser.add_argument("--image", type=str, default=DEFAULT_IMAGE, help="Path or URL to input image")
parser.add_argument("--pose", type=str, default='w-31', help="Pose string (e.g., 'a-31', 'w-31', 's-31', 'd-31')")
parser.add_argument("--output_path", type=str, default='video_samples_hyworld', help="Output video path")
parser.add_argument("--num-frames", type=int, default=125, help="Number of frames")
parser.add_argument("--seed", type=int, default=1, help="Random seed")
parser.add_argument("--resolution", type=str, default="480p", help="Only support 480p for now")
args = parser.parse_args()
# Automatically determine resolution from input image
HEIGHT, WIDTH = get_resolution_from_image(args.image, args.resolution)
print(f"Image: {args.image}")
print(f"Pose: {args.pose}")
print(f"Resolution: {HEIGHT}x{WIDTH} (from {args.resolution} buckets)")
print(f"Num frames: {args.num_frames}")
print(f"Output path: {args.output_path}")
# Initialize generator
print("\nInitializing VideoGenerator for HYWorld...")
generator = HYWorldVideoGenerator.from_pretrained(
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
image_encoder_cpu_offload=True,
)
# Generate video
print("\nGenerating video...")
start_time = time.time()
video = generator.generate_video(
prompt=args.prompt,
image_path=args.image,
output_path=args.output_path,
save_video=True,
negative_prompt="",
num_frames=args.num_frames,
fps=24,
height=HEIGHT,
width=WIDTH,
seed=args.seed,
pose=args.pose,
)
elapsed = time.time() - start_time
print(f"\nVideo generated successfully!")
print(f"Saved to: {args.output_path}")
print(f"Time: {elapsed:.2f}s")
if __name__ == "__main__":
main()
+34
View File
@@ -0,0 +1,34 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,36 +1,38 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_MODE=offline
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
DATA_DIR="data/ode-preprocessing-hy15-test/"
VALIDATION_DATASET_FILE="data/validation_64.json"
NUM_GPUS=1
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--output_dir "ode_init_hy15_test"
--wandb_run_name "vidprom_bz128_1e-5"
# --resume_from_checkpoint "ode_init_diffusers/"
# --warp_denoising_step
# --log_visualization
--max_train_steps 6001
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_latent_t 19
--num_height 480
--num_width 832
--num_frames 77
--warp_denoising_step
--num_frames 73
--dmd_denoising_steps "1000,750,500,250"
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--num_gpus 1
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
@@ -51,18 +53,17 @@ dataset_args=(
# Validation arguments
validation_args=(
--log_validation
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
@@ -78,6 +79,7 @@ miscellaneous_args=(
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
# --enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
@@ -0,0 +1,132 @@
#!/bin/bash
#SBATCH --job-name=hy15_ode_3333
#SBATCH --partition=main
#SBATCH --nodes=8
#SBATCH --ntasks=8
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=hy15_ode_3333_output/hy15_ode_3333_%j.out
#SBATCH --error=hy15_ode_3333_output/hy15_ode_3333_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv-distill
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v"
DATA_DIR="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/fv-ode-preprocessing-16k-hy15-121"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "hy15_ode_init"
--output_dir "/mnt/weka/home/hao.zhang/wei/hy15_ode_init_1333_new"
--wandb_run_name "hy15_ode_init_1333_new"
--warp_denoising_step
--dmd_denoising_steps '1000,760,520,280,0'
--log_visualization
--visualization_steps 100
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 31
--num_height 480
--num_width 848
--num_frames 121
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--text-encoder-cpu-offload
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 400
--weight_decay 0.01
--max_grad_norm 1.0
)
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 10
--validation_sampling_steps "4"
--validation_guidance_scale "1.0" # not used for dmd inference
# --vae_cpu_offload True
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 6
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
# --enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
+11
View File
@@ -40,6 +40,17 @@ out = video_sparse_attn(q, k, v, block_sizes, block_sizes, topk=5)
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
```
## Benchmark
### VSA (block-sparse) TFLOPs
After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
```
### TurboDiffusion Kernels
This package also includes kernels from [TurboDiffusion](https://github.com/thu-ml/TurboDiffusion), including INT8 GEMM, Quantization, RMSNorm and LayerNorm.
+166
View File
@@ -0,0 +1,166 @@
#!/usr/bin/env python3
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
from __future__ import annotations
import argparse
import os
import random
from typing import Tuple, Callable
import numpy as np
import torch
try:
from triton.testing import do_bench
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def parse_arguments() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Benchmark FastVideo VSA block-sparse attention")
p.add_argument("--batch_size", type=int, default=1)
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--kv_seq_lens", type=int, nargs="+", default=None, help="KV sequence lengths (defaults to q_seq_len)")
p.add_argument("--warmup", type=int, default=5)
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--force_triton", action="store_true", help="Force wrapper to use Triton path (if supported by shapes).")
return p.parse_args()
def create_qkv(batch: int, heads: int, q_len: int, kv_len: int, d: int, dtype: torch.dtype) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = torch.randn(batch, heads, q_len, d, dtype=dtype, device="cuda")
k = torch.randn(batch, heads, kv_len, d, dtype=dtype, device="cuda")
v = torch.randn(batch, heads, kv_len, d, dtype=dtype, device="cuda")
return q, k, v
def make_block_map(bs: int, h: int, num_q_blocks: int, num_kv_blocks: int, topk: int) -> torch.Tensor:
# block_map: [bs, h, num_q_blocks, num_kv_blocks] bool
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device="cuda")
topk = min(max(1, topk), num_kv_blocks)
idx = torch.topk(scores, topk, dim=-1).indices
block_map = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device="cuda")
block_map.scatter_(-1, idx, True)
return block_map
def flops_sparse_attention(bs: int, h: int, d: int, q_len: int, topk_blocks: int, block_n: int) -> float:
# Approx: QK^T + PV, each is ~2*bs*h*q_len*(topk_blocks*block_n)*d
return 4.0 * bs * h * d * q_len * (topk_blocks * block_n)
def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
bs, h, d = args.batch_size, args.num_heads, args.head_dim
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
if len(kv_seq_lens) != len(args.q_seq_lens):
raise ValueError("kv_seq_lens must have the same number of entries as q_seq_lens (or be omitted).")
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
continue
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
print("\n" + "=" * 80)
print(f"q_len={q_len}, kv_len={kv_len}, num_q_blocks={num_q_blocks}, num_kv_blocks={num_kv_blocks}, topk={topk}")
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks,), BLOCK_N, dtype=torch.int32, device="cuda")
def _fwd():
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
# Backward benchmark (wrapper autograd). We build the graph once, then repeatedly run backward
# on the retained graph so bwd timing excludes the forward compute.
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
for _ in range(max(1, args.warmup // 2)):
torch.autograd.grad(loss, (q_, k_, v_), retain_graph=True)
torch.cuda.synchronize()
bwd_ms = bench_ms(
lambda: torch.autograd.grad(loss, (q_, k_, v_), retain_graph=True),
warmup=0,
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
print(f"fwd(wrapper): {fwd_ms:.3f} ms | {fwd_tflops:.2f} TFLOPs (approx)")
print(f"bwd(wrapper): {bwd_ms:.3f} ms | {bwd_tflops:.2f} TFLOPs (approx)")
if __name__ == "__main__":
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this benchmark.")
main()
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.4"
version = "0.2.5"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -30,13 +30,12 @@ def _force_triton() -> bool:
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Pure-torch (no triton) conversion:
block_map: [B, H, Q, KV] bool (or [H, Q, KV] which will be treated as B=1)
returns:
index: [B, H, Q, KV] int32 (packed KV indices, -1 padding)
num: [B, H, Q] int32 (#kv blocks per q block)
Preferred map->index conversion used by the wrapper.
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
"""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
@@ -45,20 +44,17 @@ def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Te
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
B, H, Q, KV = block_map.shape
index = torch.full((B, H, Q, KV), -1, dtype=torch.int32, device=block_map.device)
num = torch.zeros((B, H, Q), dtype=torch.int32, device=block_map.device)
if not block_map.is_cuda:
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
# Small sizes in practice (B=1, H<=16, Q/KV<=64), so a Python loop is fine.
for b in range(B):
for h in range(H):
for q in range(Q):
kv_idx = torch.nonzero(block_map[b, h, q], as_tuple=False).flatten().to(torch.int32)
n = int(kv_idx.numel())
if n:
index[b, h, q, :n] = kv_idx
num[b, h, q] = n
return index, num
try:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
raise ImportError(
"Triton map_to_index is required but not available. "
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
) from e
return triton_map_to_index(block_map)
@torch.library.custom_op(
@@ -77,7 +73,7 @@ def block_sparse_attn_triton(
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
q2k_idx, q2k_num = _map_to_index(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_forward,
@@ -87,6 +83,7 @@ def block_sparse_attn_triton(
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
@@ -117,8 +114,8 @@ def block_sparse_attn_backward_triton(
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
q2k_idx, q2k_num = _map_to_index(block_map)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_backward,
@@ -182,7 +179,7 @@ def block_sparse_attn_sm90(
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
q2k_idx, q2k_num = _map_to_index(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
@@ -224,7 +221,7 @@ def block_sparse_attn_backward_sm90(
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
dq, dk, dv = block_sparse_bwd(
q_padded,
@@ -1 +1 @@
__version__ = "0.2.4"
__version__ = "0.2.5"
+12 -1
View File
@@ -2,5 +2,16 @@ from fastvideo.configs.models.base import ModelConfig
from fastvideo.configs.models.dits.base import DiTConfig
from fastvideo.configs.models.encoders.base import EncoderConfig
from fastvideo.configs.models.vaes.base import VAEConfig
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig)
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
__all__ = [
"ModelConfig",
"VAEConfig",
"DiTConfig",
"EncoderConfig",
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -0,0 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.configs.models.audio.ltx2_audio_vae import (
LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
)
__all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -0,0 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 audio VAE and vocoder configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig, ModelConfig
@dataclass
class LTX2AudioArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=list)
@dataclass
class LTX2AudioEncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioEncoder"]))
@dataclass
class LTX2AudioDecoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioDecoder"]))
@dataclass
class LTX2VocoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2Vocoder"]))
+3 -1
View File
@@ -3,11 +3,13 @@ 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
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig"
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
]
@@ -26,6 +26,9 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
param_names_mapping: dict = field(
default_factory=lambda: {
r"^cond_type_embed\.(.*)$":
r"cond_type_embed.\1",
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
@@ -55,6 +58,16 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
r"^image_embedder\.linear_1\.(.*)$":
r"image_embedder.linear_1.\1",
r"^image_embedder\.linear_2\.(.*)$":
r"image_embedder.linear_2.\1",
r"^image_embedder\.norm_in\.(.*)$":
r"image_embedder.norm_in.\1",
r"^image_embedder\.norm_out\.(.*)$":
r"image_embedder.norm_out.\1",
# 2. txt_in_2 mapping:
r"^context_embedder_2\.(.*)$":
@@ -144,6 +157,12 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
# Causal HunyuanVideo1.5
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 31
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
+207
View File
@@ -0,0 +1,207 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
# def is_refiner_block(n: str, m) -> bool:
# return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HYWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. txt_in submodules (text embedder, refiner blocks):
r"^txt_in\.t_embedder\.mlp\.0\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^txt_in\.t_embedder\.mlp\.2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^txt_in\.c_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^txt_in\.c_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_proj\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 2. time_in mappings (HYWorld uses TimestepEmbedder directly,
# but FastVideo model inherits HunyuanVideo15TimeEmbedding with timestep_embedder):
r"^time_in\.mlp\.0\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_out.\1",
# 3. action_in mappings:
r"^action_in\.mlp\.0\.(.*)$":
r"action_in.mlp.fc_in.\1",
r"^action_in\.mlp\.2\.(.*)$":
r"action_in.mlp.fc_out.\1",
# 4. byt5_in -> txt_in_2 mappings:
r"^byt5_in\.layernorm\.(.*)$":
r"txt_in_2.norm.\1",
r"^byt5_in\.fc1\.(.*)$":
r"txt_in_2.linear_1.\1",
r"^byt5_in\.fc2\.(.*)$":
r"txt_in_2.linear_2.\1",
r"^byt5_in\.fc3\.(.*)$":
r"txt_in_2.linear_3.\1",
# 5. cond_type_embedding -> cond_type_embed:
r"^cond_type_embedding\.(.*)$":
r"cond_type_embed.\1",
# 6. vision_in -> image_embedder mappings:
r"^vision_in\.proj\.0\.(.*)$":
r"image_embedder.norm_in.\1",
r"^vision_in\.proj\.1\.(.*)$":
r"image_embedder.linear_1.\1",
r"^vision_in\.proj\.3\.(.*)$":
r"image_embedder.linear_2.\1",
r"^vision_in\.proj\.4\.(.*)$":
r"image_embedder.norm_out.\1",
# 7. double_blocks mapping:
r"^double_blocks\.(\d+)\.img_attn_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.img_attn_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.img_attn_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.txt_attn_q\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.txt_attn_k\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.txt_attn_v\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.img_mlp\.fc1\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.img_mlp\.fc2\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc1\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc2\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 8. Final layer mapping:
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Parameters from HY-WorldPlay config.json (loaded from checkpoint)
patch_size: list | tuple | int = field(default_factory=lambda: [1, 1, 1])
# Base latent channels - will be expanded in __post_init__ if concat_condition=True
in_channels: int = 32
concat_condition: bool = True
out_channels: int = 32
hidden_size: int = 2048
heads_num: int = 16
mlp_width_ratio: float = 4.0
mlp_act_type: str = "gelu_tanh"
mm_double_blocks_depth: int = 54
mm_single_blocks_depth: int = 0
rope_dim_list: list | tuple = field(default_factory=lambda: [16, 56, 56])
qkv_bias: bool = True
qk_norm: bool | str = True
qk_norm_type: str = "rms"
guidance_embed: bool = False
use_meanflow: bool = False
text_projection: str = "single_refiner"
use_attention_mask: bool = True
text_states_dim: int = 3584
text_states_dim_2: int | None = None
text_pool_type: str | None = None
rope_theta: float = 256.0
attn_mode: str = "flash"
attn_param: str | None = None
glyph_byT5_v2: bool = True
vision_projection: str = "linear"
vision_states_dim: int = 1152
is_reshape_temporal_channels: bool = False
use_cond_type_embedding: bool = True
ideal_resolution: str = "480p"
ideal_task: str = "i2v"
task_type: str = "i2v"
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
# Convert HY-WorldPlay naming to FastVideo naming conventions
self.num_attention_heads: int = self.heads_num
self.attention_head_dim: int = self.hidden_size // self.heads_num
self.num_layers: int = self.mm_double_blocks_depth
self.num_single_layers: int = self.mm_single_blocks_depth
self.num_refiner_layers: int = 2 # Default for HYWorld
self.mlp_ratio: float = float(self.mlp_width_ratio)
self.text_embed_dim: int = self.text_states_dim
self.text_embed_2_dim: int = self.text_states_dim_2 if self.text_states_dim_2 else 1472
self.image_embed_dim: int = self.vision_states_dim
self.rope_axes_dim: tuple[int, ...] = tuple(self.rope_dim_list)
self.num_channels_latents: int = self.out_channels
self.target_size: int = 640
# Handle concat_condition: when True, actual in_channels = base * 2 + 1
# (base latent + condition latent + mask channel)
# config.json has base in_channels (32), but img_in needs full (65)
if self.concat_condition and self.in_channels == 32:
if self.is_reshape_temporal_channels:
self.in_channels = self.in_channels + self.in_channels // 2 + 1
else:
self.in_channels = self.in_channels * 2 + 1 # 32 * 2 + 1 = 65
# Handle patch_size (can be list/tuple or int)
if isinstance(self.patch_size, list | tuple):
self.patch_size_t: int = self.patch_size[0]
# assume square patch size for height and width
patch_size_hw: int = self.patch_size[1]
object.__setattr__(self, 'patch_size', patch_size_hw)
else:
self.patch_size_t = 1
# Convert qk_norm to string format
if isinstance(self.qk_norm, bool):
if self.qk_norm:
self.qk_norm = "rms_norm" if self.qk_norm_type == "rms" else self.qk_norm_type
else:
self.qk_norm = "none"
@dataclass
class HYWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HYWorldArchConfig)
prefix: str = "HYWorld"
+84
View File
@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 Transformer configuration for native FastVideo integration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_ltx2_blocks(name: str, _module) -> bool:
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
@dataclass
class LTX2VideoArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 video transformer."""
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_ltx2_blocks])
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo)
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.diffusion_model\.(.*)$": r"model.\1",
r"^diffusion_model\.(.*)$": r"model.\1",
r"^model\.(.*)$": r"model.\1",
r"^(.*)$": r"model.\1",
})
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Core transformer settings (defaults from LTX-2 metadata)
num_attention_heads: int = 32
attention_head_dim: int = 128
num_layers: int = 48
cross_attention_dim: int = 4096
caption_channels: int = 3840
norm_eps: float = 1e-6
attention_type: str = "default"
rope_type: str = "split"
double_precision_rope: bool = True
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20, 2048, 2048])
timestep_scale_multiplier: int = 1000
use_middle_indices_grid: bool = True
# Patchification (video-only path)
patch_size: tuple[int, int, int] = (1, 1, 1)
num_channels_latents: int = 128
in_channels: int | None = None
out_channels: int | None = None
# Audio defaults (reserved for joint AV ports)
audio_num_attention_heads: int = 32
audio_attention_head_dim: int = 64
audio_in_channels: int = 128
audio_out_channels: int = 128
audio_cross_attention_dim: int = 2048
audio_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20])
av_ca_timestep_scale_multiplier: int = 1
def __post_init__(self):
super().__post_init__()
patch_volume = self.patch_size[0] * self.patch_size[
1] * self.patch_size[2]
if self.in_channels is None:
self.in_channels = self.num_channels_latents * patch_volume
if self.out_channels is None:
self.out_channels = self.in_channels
@dataclass
class LTX2VideoConfig(DiTConfig):
"""Main configuration for LTX-2 transformer."""
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
prefix: str = "ltx2"
@@ -7,11 +7,14 @@ from fastvideo.configs.models.encoders.clip import (
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.siglip import SiglipVisionConfig
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config"
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig",
"SiglipVisionConfig"
]
@@ -0,0 +1,48 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class LTX2GemmaArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(
default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
hidden_size: int = 3840
num_hidden_layers: int = 48
num_attention_heads: int = 30
text_len: int = 1024
pad_token_id: int = 0
eos_token_id: int = 2
gemma_model_path: str = ""
gemma_dtype: str = "bfloat16"
padding_side: str = "left"
feature_extractor_in_features: int = 3840 * 49
feature_extractor_out_features: int = 3840
connector_num_attention_heads: int = 30
connector_attention_head_dim: int = 128
connector_num_layers: int = 2
connector_positional_embedding_theta: float = 10000.0
connector_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [4096])
connector_rope_type: str = "split"
connector_double_precision_rope: bool = False
connector_num_learnable_registers: int | None = 128
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LTX2GemmaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=LTX2GemmaArchConfig)
prefix: str = "ltx2_gemma"
@@ -0,0 +1,53 @@
# SPDX-License-Identifier: Apache-2.0
"""SigLIP vision encoder configuration for FastVideo."""
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig)
@dataclass
class SiglipVisionArchConfig(ImageEncoderArchConfig):
"""Architecture configuration for SigLIP vision encoder.
Fields match the config.json from HuggingFace SigLIP checkpoints.
"""
# From config.json
architectures: list[str] = field(
default_factory=lambda: ["SiglipVisionModel"])
attention_dropout: float = 0.0
dtype: str | None = None
hidden_act: str = "gelu_pytorch_tanh"
hidden_size: int = 1152
image_size: int = 384
intermediate_size: int = 4304
layer_norm_eps: float = 1e-6
model_type: str = "siglip_vision_model"
num_attention_heads: int = 16
num_channels: int = 3
num_hidden_layers: int = 27
patch_size: int = 14
# FastVideo specific - QKV fusion mapping
stacked_params_mapping: list = field(default_factory=lambda: [
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
@dataclass
class SiglipVisionConfig(ImageEncoderConfig):
"""Configuration for SigLIP vision encoder."""
arch_config: ImageEncoderArchConfig = field(
default_factory=SiglipVisionArchConfig)
# FastVideo specific
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = False
prefix: str = "siglip"
@@ -2,6 +2,7 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
@@ -12,4 +13,5 @@ __all__ = [
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
]
+45
View File
@@ -0,0 +1,45 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 VAE configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class LTX2VAEArchConfig(VAEArchConfig):
# Mirrors LTX-2 safetensors metadata config under "vae"
_class_name: str = "CausalVideoAutoencoder"
dims: int = 3
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 128
encoder_blocks: list = field(default_factory=list)
decoder_blocks: list = field(default_factory=list)
patch_size: int = 4
norm_layer: str = "pixel_norm"
latent_log_var: str = "uniform"
encoder_spatial_padding_mode: str = "zeros"
decoder_spatial_padding_mode: str = "reflect"
causal_decoder: bool = False
timestep_conditioning: bool = True
use_quant_conv: bool = False
scaling_factor: float = 1.0
normalize_latent_channels: bool = False
# Match FastVideo naming for compression ratios (LTX-2 default)
temporal_compression_ratio: int = 8
spatial_compression_ratio: int = 32
@dataclass
class LTX2VAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
ltx2_spatial_tile_size_in_pixels: int = 512
ltx2_spatial_tile_overlap_in_pixels: int = 64
ltx2_temporal_tile_size_in_frames: int = 64
ltx2_temporal_tile_overlap_in_frames: int = 24
+4 -1
View File
@@ -4,6 +4,8 @@ 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.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -16,5 +18,6 @@ __all__ = [
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"get_pipeline_config_cls_from_name"
]
+17 -1
View File
@@ -127,13 +127,29 @@ class Hunyuan15T2V480PConfig(PipelineConfig):
vae_tiling: bool = True
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Hunyuan15DistilledI2V480PConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 7
@dataclass
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 9
@dataclass
class SelfForcingHunyuan15T2V480PConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 5
is_causal: bool = True
# dmd_denoising_steps: list[int] | None = field(
# default_factory=lambda: [1000, 750, 500, 250])
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 875, 750, 625, 500, 375, 250, 125])
warp_denoising_step: bool = True
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, EncoderConfig
from fastvideo.configs.models.dits import HYWorldConfig as HYWorldDiTConfig
from fastvideo.configs.models.encoders import SiglipVisionConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig
@dataclass
class HYWorldConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HYWorld pipeline architecture."""
# HYWorldConfig-specific parameters with defaults
dit_config: DiTConfig = field(default_factory=HYWorldDiTConfig)
# SigLIP image encoder for I2V
image_encoder_config: EncoderConfig = field(
default_factory=SiglipVisionConfig)
image_encoder_precision: str = "fp16"
# vae_precision: str = "fp32"
# Text encoding
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp32"))
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
+50
View File
@@ -0,0 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
LTX2AudioDecoderConfig, LTX2VocoderConfig,
VAEConfig)
from fastvideo.configs.models.dits import LTX2VideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, LTX2GemmaConfig
from fastvideo.configs.models.vaes import LTX2VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def ltx2_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
return outputs.last_hidden_state
@dataclass
class LTX2T2VConfig(PipelineConfig):
"""Configuration for LTX-2 T2V pipeline."""
dit_config: DiTConfig = field(default_factory=LTX2VideoConfig)
vae_config: VAEConfig = field(default_factory=LTX2VAEConfig)
vae_tiling: bool = True
vae_sp: bool = False
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LTX2GemmaConfig(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(ltx2_postprocess_text, ))
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", ))
audio_decoder_config: ModelConfig = field(
default_factory=LTX2AudioDecoderConfig)
vocoder_config: ModelConfig = field(default_factory=LTX2VocoderConfig)
audio_decoder_precision: str = "bf16"
vocoder_precision: str = "bf16"
def __post_init__(self) -> None:
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
+18 -1
View File
@@ -8,7 +8,9 @@ from fastvideo.configs.pipelines.base import PipelineConfig
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.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig, SelfForcingHunyuan15T2V480PConfig, Hunyuan15DistilledI2V480PConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.turbodiffusion import (
@@ -33,10 +35,15 @@ logger = init_logger(__name__)
PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v":
SelfForcingHunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled":
Hunyuan15DistilledI2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15T2V720PConfig,
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorldConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
@@ -64,6 +71,9 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
# LTX-2 models
"Lightricks/LTX-2": LTX2T2VConfig,
"converted/ltx2_diffusers": LTX2T2VConfig,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers": TurboDiffusionT2V_1_3B_Config,
"loayrashid/TurboWan2.1-T2V-14B-Diffusers": TurboDiffusionT2V_14B_Config,
@@ -83,6 +93,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"hyworld":
lambda id: "hyworld" in id.lower(),
"matrixgame":
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
"wanpipeline":
@@ -102,6 +114,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "cosmos25" in id.lower(),
"turbodiffusion":
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -116,6 +130,8 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"matrixgame": MatrixGameI2V480PConfig,
"hunyuan15":
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"hyworld":
HYWorldConfig, # HYWorld-specific config as fallback for any HYWorld variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
@@ -123,6 +139,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig,
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
"ltx2": LTX2T2VConfig,
# Other fallbacks by architecture
}
+15 -2
View File
@@ -17,11 +17,24 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
# sigmas: list[float] | None = field(
# default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
def __post_init__(self):
super().__post_init__()
self.sigmas = list(
np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
@dataclass
class Hunyuan15_480P_Distilled_SamplingParam(Hunyuan15_480P_SamplingParam):
num_inference_steps: int = 8
fps: int = 24
guidance_scale: float = 1.0
@dataclass
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
+26
View File
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
import numpy as np
@dataclass
class HYWorld_SamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 480
width: int = 832
fps: int = 24
# Camera trajectory: pose string (e.g., 'w-31' means generating [1 + 31] latents) or JSON file path
pose: str = 'w-31'
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
+20
View File
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
"""
seed: int = 10
num_frames: int = 121
height: int = 1024
width: int = 1536
fps: int = 24
num_inference_steps: int = 8
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
+34 -32
View File
@@ -5,11 +5,13 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam, Hunyuan15_480P_Distilled_SamplingParam
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
from fastvideo.configs.sample.cosmos2_5 import Cosmos_Predict2_5_2B_Diffusers_SamplingParam
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
@@ -40,48 +42,41 @@ 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,
"weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled":
Hunyuan15_480P_Distilled_SamplingParam,
"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/HY-WorldPlay-Bidirectional-Diffusers": HYWorld_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
WanI2V_14B_720P_SamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
Wan2_2_I2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
FastWanT2V480P_SamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480P_SamplingParam,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
@@ -102,12 +97,9 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
# MatrixGame2.0 models
"FastVideo/Matrix-Game-2.0-Base-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers":
@@ -117,6 +109,10 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers":
TurboDiffusionI2V_A14B_SamplingParam,
# LTX-2 models
"Lightricks/LTX-2": LTX2SamplingParam,
"FastVideo/LTX2-Distilled-Diffusers": LTX2SamplingParam,
# Add other specific weight variants
}
@@ -126,6 +122,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"hyworld":
lambda id: "hyworld" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
@@ -144,6 +142,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "cosmos2_5" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower() and "2_5" not in id.lower(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -153,6 +153,8 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"hunyuan15":
Hunyuan15_480P_SamplingParam, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"hyworld":
HYWorld_SamplingParam, # HYWorld-specific config as fallback for any HYWorld variant
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
@@ -164,13 +166,13 @@ 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
}
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
"""Get the appropriate sampling param for specific pretrained weights."""
# First try exact match for specific weights
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
+24 -2
View File
@@ -127,7 +127,10 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
def ode_text_only_record_creator(
video_name: str, text_embedding: np.ndarray, caption: str,
trajectory_latents: np.ndarray,
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
trajectory_timesteps: np.ndarray,
text_mask: np.ndarray | None = None,
text_embedding_2: np.ndarray | None = None,
text_mask_2: np.ndarray | None = None) -> dict[str, Any]:
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
Args:
@@ -165,6 +168,25 @@ def ode_text_only_record_creator(
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
if text_embedding_2 is not None:
record.update({
"text_embedding_2_bytes": text_embedding_2.tobytes(),
"text_embedding_2_shape": list(text_embedding_2.shape),
"text_embedding_2_dtype": str(text_embedding_2.dtype),
})
if text_mask is not None:
record.update({
"text_mask_bytes": text_mask.tobytes(),
"text_mask_shape": list(text_mask.shape),
"text_mask_dtype": str(text_mask.dtype),
})
if text_mask_2 is not None:
record.update({
"text_mask_2_bytes": text_mask_2.tobytes(),
"text_mask_2_shape": list(text_mask_2.shape),
"text_mask_2_dtype": str(text_mask_2.dtype),
})
return record
@@ -187,4 +209,4 @@ def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
"text_embedding_dtype": str(text_embedding.dtype),
"caption": caption,
}
return record
return record
+10 -1
View File
@@ -90,6 +90,15 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_embedding_2_bytes", pa.binary()),
pa.field("text_embedding_2_shape", pa.list_(pa.int64())),
pa.field("text_embedding_2_dtype", pa.string()),
pa.field("text_mask_bytes", pa.binary()),
pa.field("text_mask_shape", pa.list_(pa.int64())),
pa.field("text_mask_dtype", pa.string()),
pa.field("text_mask_2_bytes", pa.binary()),
pa.field("text_mask_2_shape", pa.list_(pa.int64())),
pa.field("text_mask_2_dtype", pa.string()),
# --- ODE Trajectory ---
pa.field("trajectory_latents_bytes", pa.binary()),
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
@@ -115,4 +124,4 @@ pyarrow_schema_text_only = pa.schema([
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("caption", pa.string()),
])
])
+3 -2
View File
@@ -17,6 +17,7 @@ from PIL import Image
from transformers import AutoTokenizer
from fastvideo.logger import init_logger
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -655,7 +656,7 @@ class TextDataset(torch.utils.data.IterableDataset,
self.seed = seed
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer_path = os.path.join(maybe_download_model(args.model_path), "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
@@ -758,4 +759,4 @@ class TextDataset(torch.utils.data.IterableDataset,
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
self.processed_batches = state_dict["processed_batches"]
+9 -3
View File
@@ -154,8 +154,14 @@ def collate_rows_from_parquet_schema(rows,
) if rng else random.random()) < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
else:
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
if row[f"{tensor_name}_dtype"] == "float32":
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
elif row[f"{tensor_name}_dtype"] == "int64":
data = np.frombuffer(
bytes_data, dtype=np.int64).reshape(shape).copy()
else:
raise ValueError(f"Unsupported dtype: {row[f"{tensor_name}_dtype"]}")
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
@@ -168,7 +174,7 @@ def collate_rows_from_parquet_schema(rows,
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_name == 'text_embedding':
if tensor_name == 'null':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
@@ -57,6 +57,9 @@ class DistributedAutograd:
ctx.dim = dim
ctx.input_shape = input_.shape
# NCCL all_gather_into_tensor requires contiguous tensors.
if not input_.is_contiguous():
input_ = input_.contiguous()
input_size = input_.size()
output_size = (input_size[0] * world_size, ) + input_size[1:]
output_tensor = torch.empty(output_size,
+100
View File
@@ -18,6 +18,8 @@ import numpy as np
import torch
import torchvision
from einops import rearrange
import shutil
import tempfile
from fastvideo.configs.sample import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
@@ -389,6 +391,11 @@ class VideoGenerator:
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if (audio is not None and audio_sample_rate is not None and
not self._mux_audio(output_path, audio, audio_sample_rate)):
logger.warning("Audio mux failed; saved video without audio.")
if batch.return_frames:
return frames
@@ -396,6 +403,7 @@ class VideoGenerator:
return {
"samples": samples,
"frames": frames,
"audio": output_batch.extra.get("audio"),
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
@@ -405,6 +413,98 @@ class VideoGenerator:
"trajectory_decoded": output_batch.trajectory_decoded,
}
@staticmethod
def _mux_audio(
video_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
) -> bool:
"""Mux audio into video using PyAV."""
try:
import av
except ImportError:
logger.warning("PyAV not installed; cannot mux audio. "
"Install with: pip install av")
return False
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
logger.warning("Unexpected audio shape %s; skipping mux.",
audio_np.shape)
return False
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
num_channels = audio_int16.shape[1]
layout = "stereo" if num_channels == 2 else "mono"
try:
import wave
with tempfile.TemporaryDirectory() as tmpdir:
out_path = os.path.join(tmpdir, "muxed.mp4")
wav_path = os.path.join(tmpdir, "audio.wav")
# Write audio to WAV file
with wave.open(wav_path, "wb") as wav_file:
wav_file.setnchannels(num_channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(audio_int16.tobytes())
# Open input video and audio
input_video = av.open(video_path)
input_audio = av.open(wav_path)
# Create output with both streams
output = av.open(out_path, mode="w")
# Add video stream (copy codec from input)
in_video_stream = input_video.streams.video[0]
out_video_stream = output.add_stream(
codec_name=in_video_stream.codec_context.name,
rate=in_video_stream.average_rate,
)
out_video_stream.width = in_video_stream.width
out_video_stream.height = in_video_stream.height
out_video_stream.pix_fmt = in_video_stream.pix_fmt
# Add audio stream (AAC)
out_audio_stream = output.add_stream("aac", rate=sample_rate)
out_audio_stream.layout = layout
# Remux video (decode and re-encode to be safe)
for frame in input_video.decode(video=0):
for packet in out_video_stream.encode(frame):
output.mux(packet)
for packet in out_video_stream.encode():
output.mux(packet)
# Encode audio
for frame in input_audio.decode(audio=0):
frame.pts = None # Let encoder assign PTS
for packet in out_audio_stream.encode(frame):
output.mux(packet)
for packet in out_audio_stream.encode():
output.mux(packet)
input_video.close()
input_audio.close()
output.close()
shutil.move(out_path, video_path)
return True
except Exception as e:
logger.warning("Audio mux failed: %s", e)
return False
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
+104
View File
@@ -166,6 +166,14 @@ class FastVideoArgs:
# Prompt text file for batch processing
prompt_txt: str | None = None
# LTX-2 VAE tiling overrides
ltx2_vae_tiling: bool | None = None
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
ltx2_vae_temporal_tile_size_in_frames: int | None = None
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
ltx2_initial_latent_path: str | None = None
# model paths for correct deallocation
model_paths: dict[str, str] = field(default_factory=dict)
model_loaded: dict[str, bool] = field(default_factory=lambda: {
@@ -203,8 +211,44 @@ class FastVideoArgs:
logger.error("Failed to load V-MoBA config from %s: %s",
self.moba_config_path, e)
raise
self._apply_ltx2_vae_overrides()
self.check_fastvideo_args()
def _apply_ltx2_vae_overrides(self) -> None:
if self.pipeline_config is None:
return
vae_config = self.pipeline_config.vae_config
has_any = any(value is not None for value in (
self.ltx2_vae_spatial_tile_size_in_pixels,
self.ltx2_vae_spatial_tile_overlap_in_pixels,
self.ltx2_vae_temporal_tile_size_in_frames,
self.ltx2_vae_temporal_tile_overlap_in_frames,
))
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
"vae_tiling"):
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
self.pipeline_config.vae_tiling = True
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
vae_config.ltx2_spatial_tile_size_in_pixels = (
self.ltx2_vae_spatial_tile_size_in_pixels)
if hasattr(
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
self.ltx2_vae_spatial_tile_overlap_in_pixels)
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
vae_config.ltx2_temporal_tile_size_in_frames = (
self.ltx2_vae_temporal_tile_size_in_frames)
if hasattr(
vae_config, "ltx2_temporal_tile_overlap_in_frames"
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
vae_config.ltx2_temporal_tile_overlap_in_frames = (
self.ltx2_vae_temporal_tile_overlap_in_frames)
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
# Model and path configuration
@@ -325,6 +369,44 @@ class FastVideoArgs:
"Path to a text file containing prompts (one per line) for batch processing",
)
# LTX-2 VAE tiling overrides
parser.add_argument(
"--ltx2-vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.ltx2_vae_tiling,
help="Enable LTX-2 VAE tiling overrides.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-size-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
help="LTX-2 VAE spatial tile size in pixels.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-overlap-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
help="LTX-2 VAE spatial tile overlap in pixels.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-size-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
help="LTX-2 VAE temporal tile size in frames.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-overlap-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
help="LTX-2 VAE temporal tile overlap in frames.",
)
parser.add_argument(
"--ltx2-initial-latent-path",
type=str,
default=FastVideoArgs.ltx2_initial_latent_path,
help="Path to load/save a precomputed LTX-2 initial latent.",
)
# LoRA parameters (inference-time adapter loading)
parser.add_argument(
"--lora-path",
@@ -748,6 +830,7 @@ class TrainingArgs(FastVideoArgs):
precedence.
"""
data_path: str = ""
data_path_2: str | None = None
dataloader_num_workers: int = 0
num_height: int = 0
num_width: int = 0
@@ -777,6 +860,7 @@ class TrainingArgs(FastVideoArgs):
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
visualization_steps: float = 0.0
log_validation: bool = False
trackers: list[str] = dataclasses.field(default_factory=list)
tracker_project_name: str = ""
@@ -792,6 +876,7 @@ class TrainingArgs(FastVideoArgs):
num_train_epochs: int = 0
max_train_steps: int = 0
gradient_accumulation_steps: int = 0
optimizer_type: str = "adamw"
learning_rate: float = 0.0
scale_lr: bool = False
lr_scheduler: str = "constant"
@@ -861,6 +946,8 @@ class TrainingArgs(FastVideoArgs):
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
use_context_forcing: bool = False
use_ode_init: bool = False
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -916,6 +1003,10 @@ class TrainingArgs(FastVideoArgs):
type=str,
required=True,
help="Path to parquet files")
parser.add_argument("--data-path-2",
type=str,
required=False,
help="Path to parquet files")
parser.add_argument("--dataloader-num-workers",
type=int,
required=True,
@@ -1009,6 +1100,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--validation-steps",
type=float,
help="Number of validation steps")
parser.add_argument("--visualization-steps",
type=float,
help="Number of visualization steps")
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
@@ -1057,6 +1151,10 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--gradient-accumulation-steps",
type=int,
help="Number of steps to accumulate gradients")
parser.add_argument("--optimizer-type",
type=str,
choices=["adamw", "muon"],
help="Optimizer type")
parser.add_argument("--learning-rate",
type=float,
required=True,
@@ -1283,6 +1381,12 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
parser.add_argument("--use-context-forcing",
action=StoreBoolean,
help="Whether to use context forcing")
parser.add_argument("--use-ode-init",
action=StoreBoolean,
help="Whether to use ODE init")
return parser
+23 -5
View File
@@ -168,9 +168,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor | int, shift: torch.Tensor,
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
def forward(
self,
residual: torch.Tensor,
x: torch.Tensor,
gate: torch.Tensor | int,
shift: torch.Tensor,
scale: torch.Tensor,
convert_modulation_dtype: bool = False
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
scale/shift in a single fused operation.
@@ -205,6 +211,11 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
# Apply normalization
normalized = self.norm(residual_output)
if convert_modulation_dtype:
scale = scale.to(normalized.dtype)
shift = shift.to(normalized.dtype)
# Apply scale and shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
@@ -254,14 +265,21 @@ class LayerNormScaleShift(nn.Module):
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, x: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> torch.Tensor:
def forward(self,
x: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor,
convert_modulation_dtype: bool = False) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
# x.shape: [batch_size, seq_len, inner_dim]
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
normalized = normalized.float()
if convert_modulation_dtype:
scale = scale.to(normalized.dtype)
shift = shift.to(normalized.dtype)
if scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
+9
View File
@@ -0,0 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.models.audio.ltx2_audio_vae import (
LTX2AudioDecoder,
LTX2AudioEncoder,
LTX2Vocoder,
)
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,808 @@
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Any, Dict, Optional, List
import math
import torch
# import torch._dynamo
# torch._dynamo.config.cache_size_limit = 128
# try:
# torch._dynamo.config.recompile_limit = 128
# except AttributeError:
# pass
import torch.nn as nn
import torch.nn.functional as F
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
from torch.nn.attention.flex_attention import BlockMask
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="default")
import torch.distributed as dist
from fastvideo.attention import LocalAttention
from fastvideo.forward_context import set_forward_context
from fastvideo.configs.models.dits import HunyuanVideo15Config
from fastvideo.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
# TODO(will-PY-refactor): RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed, _apply_rotary_emb
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
unpatchify)
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.models.dits.hunyuanvideo15 import (
HunyuanRMSNorm,
HunyuanVideo15TimeEmbedding,
HunyuanVideo15ByT5TextProjection,
HunyuanVideo15ImageProjection,
SingleTokenRefiner,
FinalLayer)
logger = init_logger(__name__)
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
local_attn_size: int = -1,
sink_size: int = 0,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
self.head_dim = hidden_size // num_attention_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
self.max_attention_size = 21 * 1590 if local_attn_size == -1 else local_attn_size * 1590
self.attn = LocalAttention(
num_heads=self.num_attention_heads,
head_size=self.head_dim,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn"
)
def forward_txt(
self,
txt: torch.Tensor,
vec: torch.Tensor,
cache_txt: bool = False,
):
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
t_kv = {}
if cache_txt:
t_kv["k_txt"] = txt_k
t_kv["v_txt"] = txt_v
txt_attn = self.attn(txt_q, txt_k, txt_v)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return txt, t_kv
def forward_vision(
self,
img: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple,
block_mask: BlockMask,
kv_cache: dict | None = None,
txt_kv_cache: list | None = None,
current_start: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
if vec.dim() == 3:
img_mod_outputs = self.img_mod(vec).unflatten(dim=-1, sizes=(6, -1))
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=2)
else:
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings
cos, sin = freqs_cis
img_q = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False)
img_k = _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
# Apply flex_attention
# Does not support SP padding for now
if kv_cache is None:
q = img_q
k = torch.cat([img_k, txt_kv_cache["k_txt"]], dim=1)
v = torch.cat([img_v, txt_kv_cache["v_txt"]], dim=1)
# Padding for flex attention
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_kv_length = math.ceil(k.shape[1] / 128) * 128 - k.shape[1]
padded_roped_query = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_roped_key = torch.cat(
[k, torch.zeros([k.shape[0], padded_kv_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_kv_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
img_attn = flex_attention(
query=padded_roped_query.transpose(2, 1),
key=padded_roped_key.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask
)[:, :, :-padded_length].transpose(2, 1)
assert img_attn.shape[1] == image_seq_len
updated_kv_cache = None
else:
current_end = current_start + img_q.shape[1]
num_new_tokens = img_q.shape[1]
sink_tokens = self.sink_size * 1590
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = self.max_attention_size
# Clone cache to avoid in-place modification during gradient checkpointing
k_cache = kv_cache["k"].clone()
v_cache = kv_cache["v"].clone()
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_size - num_new_tokens - sink_tokens
k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
assert local_end_index == self.max_attention_size
else:
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
assert local_start_index >= 0
q = img_q
k = torch.cat([k_cache[:, :local_start_index], img_k, txt_kv_cache["k_txt"]], dim=1)
v = torch.cat([v_cache[:, :local_start_index], img_v, txt_kv_cache["v_txt"]], dim=1)
img_attn = self.attn(q, k, v)
k_cache[:, local_start_index:local_end_index] = img_k
v_cache[:, local_start_index:local_end_index] = img_v
updated_kv_cache = {
"k": k_cache,
"v": v_cache,
"global_end_index": torch.tensor([current_end], dtype=torch.long, device=k_cache.device),
"local_end_index": torch.tensor([local_end_index], dtype=torch.long, device=k_cache.device)
}
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
return img, updated_kv_cache
def forward(
self,
txt_inference=False,
vision_inference=False,
**kwargs
):
if txt_inference:
return self.forward_txt(**kwargs)
elif vision_inference:
return self.forward_vision(**kwargs)
else:
raise ValueError("txt_inference and vision_inference cannot be both False")
class CausalHunyuanVideo15Transformer3DModel(CachableDiT):
r"""
A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5).
"""
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = HunyuanVideo15Config()._fsdp_shard_conditions
_compile_conditions = HunyuanVideo15Config()._compile_conditions
_supported_attention_backends = HunyuanVideo15Config(
)._supported_attention_backends
param_names_mapping = HunyuanVideo15Config().param_names_mapping
reverse_param_names_mapping = HunyuanVideo15Config(
).reverse_param_names_mapping
lora_param_names_mapping = HunyuanVideo15Config().lora_param_names_mapping
def __init__(
self,
config: HunyuanVideo15Config,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
self.out_channels = config.out_channels or config.in_channels
self.patch_size = (config.patch_size_t, config.patch_size, config.patch_size)
# 1. Latent and condition embedders
self.img_in = PatchEmbed(self.patch_size,
config.in_channels,
self.hidden_size,
prefix=f"{config.prefix}.img_in")
self.image_embedder = HunyuanVideo15ImageProjection(config.image_embed_dim, self.hidden_size)
self.txt_in = SingleTokenRefiner(config.text_embed_dim,
self.hidden_size,
config.num_attention_heads,
depth=config.num_refiner_layers,
dtype=None,
prefix=f"{config.prefix}.txt_in")
self.txt_in_2 = HunyuanVideo15ByT5TextProjection(config.text_embed_2_dim, 2048, self.hidden_size)
self.time_in = HunyuanVideo15TimeEmbedding(self.hidden_size, use_meanflow=config.use_meanflow)
self.cond_type_embed = nn.Embedding(3, self.hidden_size)
# 3. Dual stream transformer blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
local_attn_size=config.local_attn_size,
sink_size=config.sink_size,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.num_layers)
]
)
# 5. Output projection
self.final_layer = FinalLayer(self.hidden_size,
self.patch_size,
self.out_channels,
prefix=f"{config.prefix}.final_layer")
self.gradient_checkpointing = False
self.num_frame_per_block = config.num_frames_per_block
self.local_attn_size = config.local_attn_size
self.block_mask = None
self.__post_init__()
def get_text_and_mask(
self,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_2: torch.Tensor,
encoder_attention_mask: torch.Tensor,
encoder_attention_mask_2: torch.Tensor,
encoder_hidden_states_image: torch.Tensor,
timestep: torch.Tensor,
):
batch_size, txt_seq_len = encoder_hidden_states.shape[0], encoder_hidden_states.shape[1]
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
assert encoder_hidden_states.shape[0] == 1
return encoder_hidden_states, encoder_attention_mask
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1,
text_seq_len: int = 0
) -> BlockMask:
"""
we will divide the token sequence into the following format
[1 latent frame] [1 latent frame] ... [1 latent frame]
We use flexattention to construct the attention mask
"""
total_length = num_frames * frame_seqlen
total_kv_length = total_length + text_seq_len
total_length_tensor = torch.tensor(total_length, device=device)
total_kv_length_tensor = torch.tensor(total_kv_length, device=device)
# we do right padding to get to a multiple of 128
padded_length = math.ceil(total_length / 128) * 128 - total_length
kv_padded_length = math.ceil(total_kv_length / 128) * 128 - total_kv_length
ends = torch.zeros(total_length + padded_length,
device=device, dtype=torch.long)
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
frame_indices = torch.arange(
# start=frame_seqlen,
start=0,
end=total_length,
step=frame_seqlen * num_frame_per_block,
device=device
)
# frame_indices = torch.cat([torch.tensor([0], device=device), frame_indices])
for i, tmp in enumerate(frame_indices):
# if i == 0:
# ends[tmp:tmp + frame_seqlen] = tmp + frame_seqlen
# else:
ends[tmp:tmp + frame_seqlen * num_frame_per_block] = tmp + \
frame_seqlen * num_frame_per_block
def attention_mask(b, h, q_idx, kv_idx):
if local_attn_size == -1:
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx) | ((kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor))
else:
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx) | ((kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor))
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
KV_LEN=total_kv_length + kv_padded_length, _compile=False, device=device)
# if not dist.is_initialized() or dist.get_rank() == 0:
# print(
# f" cache a block wise causal mask with block size of {num_frame_per_block} frames")
# print(block_mask)
# import imageio
# import numpy as np
# from torch.nn.attention.flex_attention import create_mask
# mask = create_mask(attention_mask, B=None, H=None, Q_LEN=total_length +
# padded_length, KV_LEN=total_length + padded_length, device=device)
# import cv2
# mask = cv2.resize(mask[0, 0].cpu().float().numpy(), (1024, 1024))
# imageio.imwrite("mask_%d.jpg" % (0), np.uint8(255. * mask))
return block_mask
def forward_txt(
self,
encoder_hidden_states: List[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[List[torch.Tensor]] = None,
encoder_attention_mask: Optional[List[torch.Tensor]] = None,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
cache_txt: bool = False,
**kwargs
):
# Check that the timestep is only consisted of 0s
assert torch.all(timestep == 0), "Timestep for txt must be only consisted of 0s"
if cache_txt:
_kv_cache_new = []
transformer_num_layers = len(self.double_blocks)
for _ in range(transformer_num_layers):
_kv_cache_new.append(
{"k_vision": None, "v_vision": None, "k_txt": None, "v_txt": None}
)
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
# 2. Conditional embeddings
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
encoder_hidden_states, encoder_attention_mask = self.get_text_and_mask(
encoder_hidden_states,
encoder_hidden_states_2,
encoder_attention_mask,
encoder_attention_mask_2,
encoder_hidden_states_image,
timestep
)
encoder_hidden_states = encoder_hidden_states[encoder_attention_mask.bool().to(encoder_hidden_states.device)].unsqueeze(0)
# 4. Transformer blocks
for index, block in enumerate(self.double_blocks):
encoder_hidden_states, t_kv = block(
txt_inference=True,
vision_inference=False,
txt=encoder_hidden_states,
vec=temb,
cache_txt=cache_txt,
)
if cache_txt:
_kv_cache_new[index]["k_txt"] = t_kv["k_txt"]
_kv_cache_new[index]["v_txt"] = t_kv["v_txt"]
if cache_txt:
return _kv_cache_new
def forward_vision(
self,
hidden_states: torch.Tensor,
timestep: torch.LongTensor,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
kv_cache: dict | None = None,
txt_kv_cache: list | None = None,
current_start: int = 0,
rope_start_idx: int = 0,
):
assert txt_kv_cache is not None, "txt_kv_cache must be provided"
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width), self.hidden_size,
self.num_attention_heads, self.config.rope_axes_dim, self.config.rope_theta, start_frame=rope_start_idx)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
hidden_states = self.img_in(hidden_states)
# Prepare block-wise causal attention mask
if kv_cache is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
text_seq_len=txt_kv_cache[0]["k_txt"].shape[1]
)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block_index, block in enumerate(self.double_blocks):
hidden_states, new_cache = self._gradient_checkpointing_func(
block,
txt_inference=False,
vision_inference=True,
img=hidden_states,
vec=temb,
freqs_cis=freqs_cis,
block_mask=self.block_mask,
kv_cache=kv_cache[block_index] if kv_cache is not None else None,
txt_kv_cache=txt_kv_cache[block_index],
current_start=current_start
)
if new_cache is not None and kv_cache is not None:
for k in new_cache.keys():
kv_cache[block_index][k] = new_cache[k].clone()
else:
for block_index, block in enumerate(self.double_blocks):
hidden_states, new_cache = block(
txt_inference=False,
vision_inference=True,
img=hidden_states,
vec=temb,
freqs_cis=freqs_cis,
block_mask=self.block_mask,
kv_cache=kv_cache[block_index] if kv_cache is not None else None,
txt_kv_cache=txt_kv_cache[block_index],
current_start=current_start
)
if new_cache is not None and kv_cache is not None:
for k in new_cache.keys():
kv_cache[block_index][k] = new_cache[k].clone()
# Final layer processing
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states, kv_cache
def forward(
self,
txt_inference=False,
vision_inference=False,
**kwargs,
):
if txt_inference:
return self.forward_txt(**kwargs)
elif vision_inference:
return self.forward_vision(**kwargs)
else:
raise ValueError("txt_inference and vision_inference cannot be both False")
+20 -3
View File
@@ -127,8 +127,9 @@ class HunyuanVideo15TimeEmbedding(nn.Module):
self,
timestep: torch.Tensor,
timestep_r: Optional[torch.Tensor] = None,
timestep_seq_len: int | None = None,
) -> torch.Tensor:
timesteps_emb = self.timestep_embedder(timestep)
timesteps_emb = self.timestep_embedder(timestep, timestep_seq_len)
if timestep_r is not None:
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
@@ -473,6 +474,7 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
**kwargs
):
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
@@ -494,7 +496,12 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
@@ -626,6 +633,8 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
)
# Final layer processing
if get_sp_world_size() > 1:
hidden_states = hidden_states.contiguous()
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
@@ -696,6 +705,7 @@ class SingleTokenRefiner(nn.Module):
else:
mask_float = mask.float().unsqueeze(-1) # [B, L, 1]
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
context_aware_representations = context_aware_representations.to(original_dtype)
context_aware_representations = self.c_embedder(
context_aware_representations)
@@ -848,6 +858,13 @@ class FinalLayer(nn.Module):
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
if c.dim() == 3:
# [bs, seq_len, inner_dim]
num_frames = scale.shape[1]
frame_seqlen = x.shape[1] // num_frames
x = (self.norm_final(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale.unsqueeze(2)) + shift.unsqueeze(2)).flatten(1, 2)
else:
# [bs, inner_dim]
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
+24
View File
@@ -0,0 +1,24 @@
# SPDX-License-Identifier: Apache-2.0
"""
HYWorld (HY-WorldPlay) model components for FastVideo.
This module provides:
- HYWorldTransformer3DModel: The main transformer model with ProPE and action conditioning
- HYWorldVideoGenerator: Extended VideoGenerator for HYWorld inference
- Utilities for pose processing and camera trajectory generation
"""
from .hyworld import HYWorldTransformer3DModel, HYWorldDoubleStreamBlock
# Inference utilities (used by examples)
from .resolution_utils import (
get_resolution_from_image,
)
__all__ = [
# Model (used by model registry)
"HYWorldTransformer3DModel",
"HYWorldDoubleStreamBlock",
# Inference utilities (used by examples)
"get_resolution_from_image",
]
@@ -0,0 +1,261 @@
# HY-WorldPlay/hyvideo/prope/camera_rope.py
# MIT License
#
# Copyright (c) Authors of
# "PRoPE: Projective Positional Encoding for Multiview Transformers"
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
# How to use PRoPE attention for self-attention:
#
# 1. Easiest way (fast):
# attn = PropeDotProductAttention(...)
# o = attn(q, k, v, viewmats, Ks)
#
# 2. More flexible way (fast):
# attn = PropeDotProductAttention(...)
# attn._precompute_and_cache_apply_fns(viewmats, Ks)
# q = attn._apply_to_q(q)
# k = attn._apply_to_kv(k)
# v = attn._apply_to_kv(v)
# o = F.scaled_dot_product_attention(q, k, v, **kwargs)
# o = attn._apply_to_o(o)
#
# 3. The most flexible way (but slower because repeated computation of RoPE coefficients):
# o = prope_dot_product_attention(q, k, v, ...)
#
# How to use PRoPE attention for cross-attention:
#
# attn_src = PropeDotProductAttention(...)
# attn_tgt = PropeDotProductAttention(...)
# attn_src._precompute_and_cache_apply_fns(viewmats_src, Ks_src)
# attn_tgt._precompute_and_cache_apply_fns(viewmats_tgt, Ks_tgt)
# q_src = attn_src._apply_to_q(q_src)
# k_tgt = attn_tgt._apply_to_kv(k_tgt)
# v_tgt = attn_tgt._apply_to_kv(v_tgt)
# o_src = F.scaled_dot_product_attention(q_src, k_tgt, v_tgt, **kwargs)
# o_src = attn_src._apply_to_o(o_src)
from functools import partial
from typing import Callable, Optional, Tuple, List
import torch
import torch.nn.functional as F
def prope_qkv(
q: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
k: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
v: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
*,
viewmats: torch.Tensor, # (batch, cameras, 4, 4)
Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3)
patches_x: int = None, # How many patches wide is each image?
patches_y: int = None, # How many patches tall is each image?
image_width: int = None, # Width of the image. Used to normalize intrinsics.
image_height: int = None, # Height of the image. Used to normalize intrinsics.
coeffs_x: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
coeffs_y: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
mask: Optional[torch.Tensor] = None,
kv_cache=None,
is_cache: bool = False,
**kwargs,
) -> torch.Tensor:
"""Similar to torch.nn.functional.scaled_dot_product_attention, but applies PRoPE-style
positional encoding.
Currently, we assume that the sequence length is equal to:
cameras * patches_x * patches_y
And token ordering allows the `(seqlen,)` axis to be reshaped into
`(cameras, patches_x, patches_y)`.
"""
# We're going to assume self-attention: all inputs are the same shape.
(batch, num_heads, seqlen, head_dim) = q.shape
cameras = viewmats.shape[1]
assert q.shape == k.shape == v.shape
assert viewmats.shape == (batch, cameras, 4, 4)
assert Ks is None or Ks.shape == (batch, cameras, 3, 3)
# assert seqlen == cameras * patches_x * patches_y
apply_fn_q, apply_fn_kv, apply_fn_o = _prepare_apply_fns_all_dim(
head_dim=head_dim,
viewmats=viewmats,
Ks=Ks,
patches_x=patches_x,
patches_y=patches_y,
image_width=image_width,
image_height=image_height,
coeffs_x=coeffs_x,
coeffs_y=coeffs_y,
)
query = apply_fn_q(q)
key = apply_fn_kv(k)
value = apply_fn_kv(v)
return query, key, value, apply_fn_o
def _prepare_apply_fns_all_dim(
head_dim: int, # Q/K/V will have this last dimension
viewmats: torch.Tensor, # (batch, cameras, 4, 4)
Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3)
patches_x: int, # How many patches wide is each image?
patches_y: int, # How many patches tall is each image?
image_width: int, # Width of the image. Used to normalize intrinsics.
image_height: int, # Height of the image. Used to normalize intrinsics.
coeffs_x: Optional[torch.Tensor] = None,
coeffs_y: Optional[torch.Tensor] = None,
) -> Tuple[
Callable[[torch.Tensor], torch.Tensor],
Callable[[torch.Tensor], torch.Tensor],
Callable[[torch.Tensor], torch.Tensor],
]:
"""Prepare transforms for PRoPE-style positional encoding."""
device = viewmats.device
(batch, cameras, _, _) = viewmats.shape
# Normalize camera intrinsics.
if Ks is not None:
Ks_norm = torch.zeros_like(Ks)
Ks_norm[..., 0, 0] = Ks[..., 0, 0]
Ks_norm[..., 1, 1] = Ks[..., 1, 1]
Ks_norm[..., 0, 2] = 0
Ks_norm[..., 1, 2] = 0
Ks_norm[..., 2, 2] = 1.0
Ks_norm = Ks_norm.to(dtype=Ks.dtype)
del Ks
# Compute the camera projection matrices we use in PRoPE.
# - K is an `image<-camera` transform.
# - viewmats is a `camera<-world` transform.
# - P = lift(K) @ viewmats is an `image<-world` transform.
P = torch.einsum("...ij,...jk->...ik", _lift_K(Ks_norm), viewmats)
P_T = P.transpose(-1, -2).to(dtype=viewmats.dtype)
P_inv = torch.einsum(
"...ij,...jk->...ik",
_invert_SE3(viewmats),
_lift_K(_invert_K(Ks_norm)),
).to(dtype=viewmats.dtype)
else:
# GTA formula. P is `camera<-world` transform.
P = viewmats
P_T = P.transpose(-1, -2)
P_inv = _invert_SE3(viewmats)
assert P.shape == P_inv.shape == (batch, cameras, 4, 4)
# Block-diagonal transforms to the inputs and outputs of the attention operator.
assert head_dim % 4 == 0
transforms_q = [
(partial(_apply_tiled_projmat, matrix=P_T), head_dim),
]
transforms_kv = [
(partial(_apply_tiled_projmat, matrix=P_inv), head_dim),
]
transforms_o = [
(partial(_apply_tiled_projmat, matrix=P), head_dim),
]
apply_fn_q = partial(_apply_block_diagonal, func_size_pairs=transforms_q)
apply_fn_kv = partial(_apply_block_diagonal, func_size_pairs=transforms_kv)
apply_fn_o = partial(_apply_block_diagonal, func_size_pairs=transforms_o)
return apply_fn_q, apply_fn_kv, apply_fn_o
def _apply_tiled_projmat(
feats: torch.Tensor, # (batch, num_heads, seqlen, feat_dim)
matrix: torch.Tensor, # (batch, cameras, D, D)
) -> torch.Tensor:
"""Apply projection matrix to features."""
# - seqlen => (cameras, patches_x * patches_y)
# - feat_dim => (feat_dim // 4, 4)
(batch, num_heads, seqlen, feat_dim) = feats.shape
cameras = matrix.shape[1]
assert seqlen >= cameras and seqlen % cameras == 0
D = matrix.shape[-1]
assert matrix.shape == (batch, cameras, D, D)
assert feat_dim % D == 0
return torch.einsum(
"bcij,bncpkj->bncpki",
matrix,
feats.reshape((batch, num_heads, cameras, -1, feat_dim // D, D)),
).reshape(feats.shape)
def _apply_block_diagonal(
feats: torch.Tensor, # (..., dim)
func_size_pairs: List[Tuple[Callable[[torch.Tensor], torch.Tensor], int]],
) -> torch.Tensor:
"""Apply a block-diagonal function to an input array.
Each function is specified as a tuple with form:
((Tensor) -> Tensor, int)
Where the integer is the size of the input to the function.
"""
funcs, block_sizes = zip(*func_size_pairs)
assert feats.shape[-1] == sum(block_sizes)
x_blocks = torch.split(feats, block_sizes, dim=-1)
out = torch.cat(
[f(x_block) for f, x_block in zip(funcs, x_blocks)],
dim=-1,
)
assert out.shape == feats.shape, "Input/output shapes should match."
return out
def _invert_SE3(transforms: torch.Tensor) -> torch.Tensor:
"""Invert a 4x4 SE(3) matrix."""
assert transforms.shape[-2:] == (4, 4)
Rinv = transforms[..., :3, :3].transpose(-1, -2)
out = torch.zeros_like(transforms)
out[..., :3, :3] = Rinv
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", Rinv, transforms[..., :3, 3])
out[..., 3, 3] = 1.0
out = out.to(dtype=transforms.dtype)
return out
def _lift_K(Ks: torch.Tensor) -> torch.Tensor:
"""Lift 3x3 matrices to homogeneous 4x4 matrices."""
assert Ks.shape[-2:] == (3, 3)
out = torch.zeros(Ks.shape[:-2] + (4, 4), device=Ks.device)
out[..., :3, :3] = Ks
out[..., 3, 3] = 1.0
out = out.to(dtype=Ks.dtype)
return out
def _invert_K(Ks: torch.Tensor) -> torch.Tensor:
"""Invert 3x3 intrinsics matrices. Assumes no skew."""
assert Ks.shape[-2:] == (3, 3)
out = torch.zeros_like(Ks)
out[..., 0, 0] = 1.0 / Ks[..., 0, 0]
out[..., 1, 1] = 1.0 / Ks[..., 1, 1]
out[..., 0, 2] = -Ks[..., 0, 2] / Ks[..., 0, 0]
out[..., 1, 2] = -Ks[..., 1, 2] / Ks[..., 1, 1]
out[..., 2, 2] = 1.0
out = out.to(dtype=Ks.dtype)
return out
@@ -0,0 +1,76 @@
# HY-WorldPlay/hyvideo/utils/data_utils.py
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
import numpy as np
from PIL import Image
def resize_and_center_crop(image, target_width, target_height):
if target_height == image.shape[0] and target_width == image.shape[1]:
return image
pil_image = Image.fromarray(image)
original_width, original_height = pil_image.size
scale_factor = max(target_width / original_width, target_height / original_height)
resized_width = int(round(original_width * scale_factor))
resized_height = int(round(original_height * scale_factor))
resized_image = pil_image.resize((resized_width, resized_height), Image.LANCZOS)
left = (resized_width - target_width) / 2
top = (resized_height - target_height) / 2
right = (resized_width + target_width) / 2
bottom = (resized_height + target_height) / 2
cropped_image = resized_image.crop((left, top, right, bottom))
return np.array(cropped_image)
def get_closest_ratio(height: float, width: float, ratios: list, buckets: list):
"""
Get the closest ratio in the buckets.
Args:
height (float): video height
width (float): video width
ratios (list): video aspect ratio
buckets (list): buckets generated by `generate_crop_size_list`
Returns:
the closest size in the buckets and the corresponding ratio
"""
aspect_ratio = float(height) / float(width)
ratios_array = np.array(ratios)
closest_ratio_id = np.abs(ratios_array - aspect_ratio).argmin()
closest_size = buckets[closest_ratio_id]
closest_ratio = ratios_array[closest_ratio_id]
return closest_size, closest_ratio
def generate_crop_size_list(base_size=256, patch_size=16, max_ratio=4.0):
num_patches = round((base_size / patch_size) ** 2)
assert max_ratio >= 1.0
crop_size_list = []
wp, hp = num_patches, 1
while wp > 0:
if max(wp, hp) / min(wp, hp) <= max_ratio:
crop_size_list.append((wp * patch_size, hp * patch_size))
if (hp + 1) * wp <= num_patches:
hp += 1
else:
wp -= 1
return crop_size_list
+569
View File
@@ -0,0 +1,569 @@
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
from typing import Any, Optional
import torch
import torch.nn as nn
from einops import rearrange, repeat
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.configs.models.dits import HYWorldConfig
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import TimestepEmbedder, unpatchify
from fastvideo.models.dits.hunyuanvideo15 import (
MMDoubleStreamBlock,
HunyuanVideo15Transformer3DModel,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.forward_context import set_forward_context
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
from .camera_rope import prope_qkv
logger = init_logger(__name__)
class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
"""
Extended MMDoubleStreamBlock with ProPE (Projective Positional Encoding) support
for camera-aware attention in HY-World/WorldPlay models.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = "",
):
super().__init__(
hidden_size=hidden_size,
num_attention_heads=num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
supported_attention_backends=supported_attention_backends,
prefix=prefix,
)
self.hidden_size = hidden_size
# Add ProPE projection layer for camera-aware attention
self.img_attn_prope_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_prope_proj"
)
# Zero-initialize ProPE projection (starts as identity)
nn.init.zeros_(self.img_attn_prope_proj.weight)
if self.img_attn_prope_proj.bias is not None:
nn.init.zeros_(self.img_attn_prope_proj.bias)
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
encoder_attention_mask: torch.Tensor,
vec: torch.Tensor,
vec_txt: torch.Tensor,
freqs_cis: tuple,
seq_attention_mask: torch.Tensor,
viewmats: torch.Tensor,
Ks: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Forward pass with ProPE camera conditioning.
Args:
img: Image/video tokens
txt: Text tokens
encoder_attention_mask: Text attention mask
vec: Modulation vector
freqs_cis: Rotary embedding frequencies
seq_attention_mask: Sequence attention mask
viewmats: Camera view matrices for ProPE [B, T, 4, 4]
Ks: Camera intrinsics for ProPE [B, T, 3, 3]
Returns:
Tuple of (img, txt) output tokens
"""
# Process modulation vectors (inherited from parent)
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec_txt)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale, convert_modulation_dtype=True)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale, convert_modulation_dtype=True)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# begin hyworld: add camera pose through prope
img_q_prope, img_k_prope, img_v_prope, apply_fn_o = prope_qkv(
img_q.permute(0, 2, 1, 3),
img_k.permute(0, 2, 1, 3),
img_v.permute(0, 2, 1, 3),
viewmats=viewmats,
Ks=Ks,
) # [batch, num_heads, seqlen, head_dim]
img_q_prope = img_q_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
img_k_prope = img_k_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
img_v_prope = img_v_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
# end hyworld
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis, attention_mask=seq_attention_mask)
# begin hyworld
# attention with prope
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_prope = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# NOTE: Do NOT pass freqs_cis to prope attention - HY-WorldPlay does not apply RoPE to prope
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata_prope):
img_attn_prope, _ = self.attn(
img_q_prope, img_k_prope, img_v_prope, txt_q, txt_k, txt_v,
freqs_cis=None, attention_mask=seq_attention_mask # No RoPE for prope attention
)
img_attn_prope = img_attn_prope.reshape(batch_size, image_seq_len, -1)
img_attn_prope = rearrange(
img_attn_prope, "B L (H D) -> B H L D", H=self.num_attention_heads
)
img_attn_prope = apply_fn_o(img_attn_prope) # [batch, num_heads, seqlen, head_dim]
img_attn_prope = rearrange(img_attn_prope, "B H L D -> B L (H D)")
# add prope to img_attn
img_attn_out, _ = self.img_attn_proj(img_attn.view(batch_size, image_seq_len, -1))
img_attn_prope_out, _ = self.img_attn_prope_proj(img_attn_prope)
img_attn_out = img_attn_out + img_attn_prope_out
# end hyworld
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale, convert_modulation_dtype=True)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale, convert_modulation_dtype=True)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class HYWorldFinalLayer(nn.Module):
"""
Final layer for HYWorld that uses modulate() to handle per-token conditioning.
This matches HY-WorldPlay's FinalLayer behavior.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.visual_embedding import ModulateProjection
from fastvideo.layers.layernorm import LayerNormScaleShift
self.norm_final = LayerNormScaleShift(
hidden_size,
norm_type="layer",
eps=1e-6,
elementwise_affine=False,
dtype=dtype,
prefix=f"{prefix}.norm_final")
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation projection to get shift/scale
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x, shift, scale, convert_modulation_dtype=True)
x, _ = self.linear(x)
return x
class HYWorldTransformer3DModel(HunyuanVideo15Transformer3DModel):
r"""
HY-World Transformer extending HunyuanVideo15 with:
- ProPE (Projective Positional Encoding) for camera-aware attention
- Action conditioning for interactive video generation
"""
# Class attributes for weight loading - use HYWorld-specific mapping
_fsdp_shard_conditions = HYWorldConfig().arch_config._fsdp_shard_conditions
_compile_conditions = HYWorldConfig().arch_config._compile_conditions
param_names_mapping = HYWorldConfig().arch_config.param_names_mapping
reverse_param_names_mapping = HYWorldConfig().arch_config.reverse_param_names_mapping
def __init__(
self,
config: HYWorldConfig,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
# Replace double_blocks with HY-World version that supports ProPE
self.double_blocks = nn.ModuleList([
HYWorldDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=self.num_attention_heads,
mlp_ratio=config.arch_config.mlp_ratio,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.arch_config.num_layers)
])
# Add action conditioning module
self.action_in = TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=None,
prefix=f"{config.prefix}.action_in"
)
# Zero-initialize action embedding (starts with no effect)
nn.init.zeros_(self.action_in.mlp.fc_out.weight)
if self.action_in.mlp.fc_out.bias is not None:
nn.init.zeros_(self.action_in.mlp.fc_out.bias)
# Override final_layer with HYWorld version that uses per-token modulate()
self.final_layer = HYWorldFinalLayer(
hidden_size=self.hidden_size,
patch_size=self.patch_size,
out_channels=self.out_channels,
dtype=None,
prefix=f"{config.prefix}.final_layer"
)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: list[torch.Tensor],
encoder_attention_mask: list[torch.Tensor],
action: torch.Tensor,
viewmats: torch.Tensor,
Ks: torch.Tensor,
timestep_txt: torch.LongTensor,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[dict[str, Any]] = None,
):
"""
Forward pass with action and camera conditioning.
Args:
action: Action tensor for action conditioning [B, T] or [B*T]
viewmats: Camera view matrices [B, T, 4, 4]
Ks: Camera intrinsics [B, T, 3, 3]
... (other args same as parent)
"""
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
self.config.rope_axes_dim,
self.config.rope_theta
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
# NOTE: freqs_cis does NOT need sharding because FastVideo's DistributedAttention
# uses all-to-all to gather the full sequence before applying RoPE
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
temb_txt = self.time_in(timestep_txt, timestep_r=timestep_r)
# Add action conditioning if provided
# temb shape: [B*T, C] where T = num_frames
temb = temb + self.action_in(action.reshape(-1))
# Broadcast timestep embedding for transformer blocks (one per spatial token)
# [B*T, C] -> [B, T*H*W, C] -> [B*T*H*W, C]
temb = repeat(temb, "(B T) C -> B (T H W) C", B=batch_size, H=post_patch_height, W=post_patch_width)
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
seq_attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
seq_attention_mask = None
viewmats_seq = repeat(
viewmats, "B T M N->B (T H W) M N",
H=post_patch_height,
W=post_patch_width
)
Ks_seq = repeat(
Ks, "B T M N->B (T H W) M N",
H=post_patch_height,
W=post_patch_width
)
# Shard viewmats, Ks, and temb for sequence parallelism (shard along sequence dim=1)
# Note that temb in HY1.5 does not need sharding because it is per-sample modulation
# In HYWorld, temb is per-token modulation.
if sp_world_size > 1:
viewmats_seq, _ = sequence_model_parallel_shard(viewmats_seq, dim=1)
Ks_seq, _ = sequence_model_parallel_shard(Ks_seq, dim=1)
temb, _ = sequence_model_parallel_shard(temb, dim=1)
# Rearrange temb after sharding to match expected shape
temb = rearrange(temb, "B S C -> (B S) C")
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep_txt, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats_seq, # hyworld
Ks_seq, # hyworld
)
else:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats=viewmats_seq, # hyworld
Ks=Ks_seq, # hyworld
)
# Final layer processing (per-token conditioning via HYWorldFinalLayer)
# Apply final_layer on sharded data first, then gather
hidden_states = self.final_layer(hidden_states, temb)
# Gather the output from all ranks
if get_sp_world_size() > 1:
hidden_states = hidden_states.contiguous()
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states
+413
View File
@@ -0,0 +1,413 @@
# Some functions from HY-WorldPlay/hyvideo/generate.py
"""
Pose processing utilities for HYWorld video generation.
This module provides functions to convert camera poses to model input tensors,
including viewmats, intrinsics, and action labels.
Adapted from HY-WorldPlay: https://github.com/Tencent-Hunyuan/HY-WorldPlay
"""
import json
import numpy as np
import torch
from scipy.spatial.transform import Rotation as R
from typing import Union, Optional
from .trajectory import generate_camera_trajectory_local
# Mapping from one-hot action encoding to single label
mapping = {
(0, 0, 0, 0): 0,
(1, 0, 0, 0): 1,
(0, 1, 0, 0): 2,
(0, 0, 1, 0): 3,
(0, 0, 0, 1): 4,
(1, 0, 1, 0): 5,
(1, 0, 0, 1): 6,
(0, 1, 1, 0): 7,
(0, 1, 0, 1): 8,
}
# Default camera intrinsic matrix (for 1920x1080 resolution)
DEFAULT_INTRINSIC = [
[969.6969696969696, 0.0, 960.0],
[0.0, 969.6969696969696, 540.0],
[0.0, 0.0, 1.0],
]
# Default movement speeds
DEFAULT_FORWARD_SPEED = 0.08 # units per frame
DEFAULT_YAW_SPEED = np.deg2rad(3) # radians per frame
DEFAULT_PITCH_SPEED = np.deg2rad(3) # radians per frame
def one_hot_to_one_dimension(one_hot: torch.Tensor) -> torch.Tensor:
"""Convert one-hot action encoding to single dimension labels."""
return torch.tensor([mapping[tuple(row.tolist())] for row in one_hot])
def parse_pose_string(
pose_string: str,
forward_speed: float = DEFAULT_FORWARD_SPEED,
yaw_speed: float = DEFAULT_YAW_SPEED,
pitch_speed: float = DEFAULT_PITCH_SPEED,
) -> list[dict]:
"""
Parse pose string to motions list.
Format: "w-3, right-0.5, d-4"
- w: forward movement
- s: backward movement
- a: left movement
- d: right movement
- up: pitch up rotation
- down: pitch down rotation
- left: yaw left rotation
- right: yaw right rotation
- number after dash: duration in frames/latents
Args:
pose_string: Comma-separated pose commands
forward_speed: Movement amount per frame
yaw_speed: Yaw rotation amount per frame (radians)
pitch_speed: Pitch rotation amount per frame (radians)
Returns:
List of motion dictionaries for generate_camera_trajectory_local
"""
motions = []
commands = [cmd.strip() for cmd in pose_string.split(",")]
for cmd in commands:
if not cmd:
continue
parts = cmd.split("-")
if len(parts) != 2:
raise ValueError(
f"Invalid pose command: {cmd}. Expected format: 'action-duration'"
)
action = parts[0].strip()
try:
duration = float(parts[1].strip())
except ValueError:
raise ValueError(f"Invalid duration in command: {cmd}")
num_frames = int(duration)
# Parse action and create motion dicts
if action == "w":
# Forward
for _ in range(num_frames):
motions.append({"forward": forward_speed})
elif action == "s":
# Backward
for _ in range(num_frames):
motions.append({"forward": -forward_speed})
elif action == "a":
# Left
for _ in range(num_frames):
motions.append({"right": -forward_speed})
elif action == "d":
# Right
for _ in range(num_frames):
motions.append({"right": forward_speed})
elif action == "up":
# Pitch up
for _ in range(num_frames):
motions.append({"pitch": pitch_speed})
elif action == "down":
# Pitch down
for _ in range(num_frames):
motions.append({"pitch": -pitch_speed})
elif action == "left":
# Yaw left
for _ in range(num_frames):
motions.append({"yaw": -yaw_speed})
elif action == "right":
# Yaw right
for _ in range(num_frames):
motions.append({"yaw": yaw_speed})
else:
raise ValueError(
f"Unknown action: {action}. "
f"Supported actions: w, s, a, d, up, down, left, right"
)
return motions
def pose_string_to_json(
pose_string: str,
intrinsic: Optional[list[list[float]]] = None,
) -> dict:
"""
Convert pose string to pose JSON format.
Args:
pose_string: Comma-separated pose commands
intrinsic: Camera intrinsic matrix (default: DEFAULT_INTRINSIC from trajectory)
Returns:
Dict with frame indices as keys, containing extrinsic and K (intrinsic) matrices
"""
if intrinsic is None:
intrinsic = DEFAULT_INTRINSIC
motions = parse_pose_string(pose_string)
poses = generate_camera_trajectory_local(motions)
pose_json = {}
for i, p in enumerate(poses):
pose_json[str(i)] = {"extrinsic": p.tolist(), "K": intrinsic}
return pose_json
def pose_to_input(
pose_data: Union[str, dict],
latent_num: int,
tps: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Convert pose data to model input tensors.
Args:
pose_data: One of:
- str ending with '.json': path to JSON file
- str: pose string (e.g., "w-3, right-0.5, d-4")
- dict: pose JSON data
latent_num: Number of latents (frames in latent space)
tps: Third-person mode flag
Returns:
Tuple of (viewmats, intrinsics, action_labels):
- viewmats: World-to-camera matrices [T, 4, 4]
- intrinsics: Normalized camera intrinsics [T, 3, 3]
- action_labels: Action labels for each frame [T]
"""
# Handle different input types
if isinstance(pose_data, str):
if pose_data.endswith(".json"):
# Load from JSON file
with open(pose_data, "r") as f:
pose_json = json.load(f)
else:
# Parse pose string
pose_json = pose_string_to_json(pose_data)
elif isinstance(pose_data, dict):
pose_json = pose_data
else:
raise ValueError(
f"Invalid pose_data type: {type(pose_data)}. Expected str or dict."
)
pose_keys = list(pose_json.keys())
latent_num_from_pose = len(pose_keys)
assert latent_num_from_pose == latent_num, (
f"pose corresponds to {latent_num_from_pose * 4 - 3} frames, num_frames "
f"must be set to {latent_num_from_pose * 4 - 3} to ensure alignment."
)
intrinsic_list = []
w2c_list = []
for i in range(latent_num):
t_key = pose_keys[i]
c2w = np.array(pose_json[t_key]["extrinsic"])
w2c = np.linalg.inv(c2w)
w2c_list.append(w2c)
# Normalize intrinsics
intrinsic = np.array(pose_json[t_key]["K"])
intrinsic[0, 0] /= intrinsic[0, 2] * 2
intrinsic[1, 1] /= intrinsic[1, 2] * 2
intrinsic[0, 2] = 0.5
intrinsic[1, 2] = 0.5
intrinsic_list.append(intrinsic)
w2c_list = np.array(w2c_list)
intrinsic_list = torch.tensor(np.array(intrinsic_list))
# Compute relative camera-to-world transforms
c2ws = np.linalg.inv(w2c_list)
C_inv = np.linalg.inv(c2ws[:-1])
relative_c2w = np.zeros_like(c2ws)
relative_c2w[0, ...] = c2ws[0, ...]
relative_c2w[1:, ...] = C_inv @ c2ws[1:, ...]
# Initialize one-hot action encodings
trans_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
rotate_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
move_norm_valid = 0.0001
for i in range(1, relative_c2w.shape[0]):
move_dirs = relative_c2w[i, :3, 3] # direction vector
move_norms = np.linalg.norm(move_dirs)
if move_norms > move_norm_valid: # threshold for movement
move_norm_dirs = move_dirs / move_norms
angles_rad = np.arccos(move_norm_dirs.clip(-1.0, 1.0))
trans_angles_deg = angles_rad * (180.0 / np.pi) # convert to degrees
else:
trans_angles_deg = np.zeros(3)
R_rel = relative_c2w[i, :3, :3]
r = R.from_matrix(R_rel)
rot_angles_deg = r.as_euler("xyz", degrees=True)
# Determine movement and rotation actions
if move_norms > move_norm_valid: # threshold for movement
if (not tps) or (
tps and abs(rot_angles_deg[1]) < 5e-2 and abs(rot_angles_deg[0]) < 5e-2
):
if trans_angles_deg[2] < 60:
trans_one_hot[i, 0] = 1 # forward
elif trans_angles_deg[2] > 120:
trans_one_hot[i, 1] = 1 # backward
if trans_angles_deg[0] < 60:
trans_one_hot[i, 2] = 1 # right
elif trans_angles_deg[0] > 120:
trans_one_hot[i, 3] = 1 # left
if rot_angles_deg[1] > 5e-2:
rotate_one_hot[i, 0] = 1 # right
elif rot_angles_deg[1] < -5e-2:
rotate_one_hot[i, 1] = 1 # left
if rot_angles_deg[0] > 5e-2:
rotate_one_hot[i, 2] = 1 # up
elif rot_angles_deg[0] < -5e-2:
rotate_one_hot[i, 3] = 1 # down
trans_one_hot = torch.tensor(trans_one_hot)
rotate_one_hot = torch.tensor(rotate_one_hot)
# Convert one-hot to single-dimension labels
trans_one_label = one_hot_to_one_dimension(trans_one_hot)
rotate_one_label = one_hot_to_one_dimension(rotate_one_hot)
action_one_label = trans_one_label * 9 + rotate_one_label
return (
torch.as_tensor(w2c_list),
torch.as_tensor(intrinsic_list),
action_one_label,
)
def camera_center_normalization(w2c: np.ndarray) -> np.ndarray:
"""Normalize camera centers relative to the first camera."""
c2w = np.linalg.inv(w2c)
C0_inv = np.linalg.inv(c2w[0])
c2w_aligned = np.array([C0_inv @ C for C in c2w])
return np.linalg.inv(c2w_aligned)
def parse_pose_string_to_actions(pose_string: str, fps: int = 24) -> list[dict]:
"""
Parse pose string to frame-level action timeline.
Format: pose string uses latent counts, where:
- 1 latent = 4 frames
- Special rule: first frame of entire video is extra (frame 0)
- Example: "w-4,d-4" means:
- w-4: forward for frames 0-16 (17 frames total: 1 extra + 4*4)
- d-4: right for frames 17-32 (16 frames total: 4*4)
Args:
pose_string: Comma-separated pose commands (e.g., "w-4,d-4")
fps: Frames per second for video (default: 24)
Returns:
List of dicts with action values for each frame
"""
commands = [cmd.strip() for cmd in pose_string.split(",")]
frame_actions = []
is_first_command = True
for cmd in commands:
if not cmd:
continue
parts = cmd.split("-")
if len(parts) != 2:
raise ValueError(
f"Invalid pose command: {cmd}. Expected format: 'action-duration'"
)
action = parts[0].strip()
try:
num_latents = int(parts[1].strip())
except ValueError:
raise ValueError(f"Invalid duration in command: {cmd}")
# Convert latents to frames
# First command gets 1 extra frame (the special frame 0)
if is_first_command:
num_frames = 1 + num_latents * 4
is_first_command = False
else:
num_frames = num_latents * 4
# Map action to action values
action_values = {"forward": 0, "left": 0, "yaw": 0, "pitch": 0}
if action == "w":
action_values["forward"] = 1
elif action == "s":
action_values["forward"] = -1
elif action == "a":
action_values["left"] = 1
elif action == "d":
action_values["left"] = -1
elif action == "up":
action_values["pitch"] = 1
elif action == "down":
action_values["pitch"] = -1
elif action == "left":
action_values["yaw"] = -1
elif action == "right":
action_values["yaw"] = 1
else:
raise ValueError(f"Unknown action: {action}")
# Add frame-level actions
for _ in range(num_frames):
frame_actions.append(action_values.copy())
return frame_actions
def compute_latent_num(num_frames: int) -> int:
"""
Compute the number of latents from number of frames.
Formula: num_frames = (latent_num - 1) * 4 + 1
So: latent_num = (num_frames - 1) // 4 + 1
Args:
num_frames: Number of video frames
Returns:
Number of latents
"""
return (num_frames - 1) // 4 + 1
def compute_num_frames(latent_num: int) -> int:
"""
Compute the number of frames from number of latents.
Formula: num_frames = (latent_num - 1) * 4 + 1
Args:
latent_num: Number of latents
Returns:
Number of video frames
"""
return (latent_num - 1) * 4 + 1
@@ -0,0 +1,64 @@
import numpy as np
from PIL import Image
import requests
from io import BytesIO
from fastvideo.models.dits.hyworld.data_utils import generate_crop_size_list
# Target resolution configs (matching HY-WorldPlay)
TARGET_SIZE_CONFIG = {
"360p": {"bucket_hw_base_size": 480, "bucket_hw_bucket_stride": 16},
"480p": {"bucket_hw_base_size": 640, "bucket_hw_bucket_stride": 16},
"720p": {"bucket_hw_base_size": 960, "bucket_hw_bucket_stride": 16},
"1080p": {"bucket_hw_base_size": 1440, "bucket_hw_bucket_stride": 16},
}
def get_closest_resolution(image_height, image_width, target_resolution="480p"):
"""
Get closest supported resolution for given image dimensions.
Args:
image_height: Height of input image
image_width: Width of input image
target_resolution: Target resolution string (e.g., "480p", "720p")
Returns:
tuple[int, int]: (height, width) of closest supported resolution
"""
config = TARGET_SIZE_CONFIG[target_resolution]
bucket_hw_base_size = config["bucket_hw_base_size"]
bucket_hw_bucket_stride = config["bucket_hw_bucket_stride"]
crop_size_list = generate_crop_size_list(bucket_hw_base_size, bucket_hw_bucket_stride)
aspect_ratios = np.array([round(float(h) / float(w), 5) for h, w in crop_size_list])
# Find closest aspect ratio
image_ratio = float(image_height) / float(image_width)
closest_idx = np.abs(aspect_ratios - image_ratio).argmin()
closest_size = crop_size_list[closest_idx]
return closest_size[0], closest_size[1] # (height, width)
def get_resolution_from_image(image_path, target_resolution="480p"):
"""
Automatically determine resolution from input image.
Args:
image_path: Path or URL to input image
target_resolution: Target resolution tier ("480p", "720p", etc.)
Returns:
tuple[int, int]: (height, width) matching HY-WorldPlay's bucket selection
"""
# Handle URL inputs
if isinstance(image_path, str) and image_path.startswith(('http://', 'https://')):
response = requests.get(image_path)
response.raise_for_status()
img = Image.open(BytesIO(response.content))
else:
img = Image.open(image_path)
img_width, img_height = img.size
return get_closest_resolution(img_height, img_width, target_resolution)
@@ -0,0 +1,316 @@
# HY-WorldPlay/hyvideo/utils/retrieval_context.py
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
import torch
import numpy as np
from typing import List, Tuple, Dict
import math
def generate_points_in_sphere(n_points: int, radius: float) -> torch.Tensor:
"""
Uniformly sample points within a sphere of a specified radius.
:param n_points: The number of points to generate.
:param radius: The radius of the sphere.
:return: A tensor of shape (n_points, 3), representing the (x, y, z) coordinates of the points.
"""
samples_r = torch.rand(n_points)
samples_phi = torch.rand(n_points)
samples_u = torch.rand(n_points)
r = radius * torch.pow(samples_r, 1 / 3)
phi = 2 * math.pi * samples_phi
theta = torch.acos(1 - 2 * samples_u)
# transfer the coordinates from spherical to cartesian
x = r * torch.sin(theta) * torch.cos(phi)
y = r * torch.sin(theta) * torch.sin(phi)
z = r * torch.cos(theta)
points = torch.stack((x, y, z), dim=1)
return points
def rotation_matrix_to_angles(R: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Estimate the Pitch and Yaw angles from a 3x3 rotation matrix R in the camera coordinate system.
Assumed Camera Coordinate System: X=Right, Y=Up, Z=Backward
(or NeRF style: X=Right, Y=Down, Z=Forward).
Here we adopt the common Computer Vision convention: Z-axis is Forward.
Note: The angle calculations here are directly based on the conventions of your `is_inside_fov_3d_hv` function:
- Yaw/Azimuth angle is in the XZ plane (atan2(x, z)).
- Pitch/Elevation angle is relative to the horizontal plane (atan2(y, sqrt(x^2 + z^2))).
For the third column R[:, 2] of the W2C matrix R (the direction of the World Z-axis in the Camera frame),
this typically corresponds to the direction the camera is looking
(the representation of the world Z-axis in the camera frame).
To simplify and match your `is_inside_fov` logic, we directly use the camera's Z-axis vector:
Camera Z-axis direction in World Frame (Forward Vector): fwd = R_w2c_inv @ [0, 0, 1]
More simply, the Z-axis vector of the C2W matrix is the camera's forward vector in the world frame.
C2W = W2C_inv
"""
R_c2w = R.T
fwd = R_c2w[:, 2]
x = fwd[0]
y = fwd[1]
z = fwd[2]
# compute yaw and pitch
yaw_rad = torch.atan2(x, z)
yaw_deg = yaw_rad * (180.0 / math.pi)
pitch_rad = torch.atan2(y, torch.sqrt(x**2 + z**2))
pitch_deg = pitch_rad * (180.0 / math.pi)
return pitch_deg, yaw_deg
def is_inside_fov_3d_hv(
points: torch.Tensor,
center: torch.Tensor,
center_pitch: torch.Tensor,
center_yaw: torch.Tensor,
fov_half_h: torch.Tensor,
fov_half_v: torch.Tensor,
) -> torch.Tensor:
"""
Check whether points are inside a 3D view frustum defined by a center coordinate, pitch angle, and yaw angle.
:param points: Tensor of shape (N, 3) or (N, B, 3) representing the coordinates of the sampled points.
:param center: Tensor of shape (3) or (B, 3) representing the camera center coordinates.
:param center_pitch: Tensor of shape (1) or (B) representing the pitch angle of center view direction.
:param center_yaw: Tensor of shape (1) or (B) representing the yaw angle of the center view direction.
:param fov_half_h: The horizontal half field-of-view angle (in degrees).
:param fov_half_v: The vertical half field-of-view angle (in degrees).
:return: Boolean tensor of shape (N) or (N, B), indicating whether each point is inside the FOV.
"""
if points.ndim == 2: # N, 3
vectors = points - center[None, :]
C = 1
elif points.ndim == 3: # N, B, 3
vectors = points - center[None, ...]
center_pitch = center_pitch[None, :] if center_pitch.ndim == 1 else center_pitch
center_yaw = center_yaw[None, :] if center_yaw.ndim == 1 else center_yaw
else:
raise ValueError("points' shape should be (N, 3) or (N, B, 3)")
x = vectors[..., 0]
y = vectors[..., 1]
z = vectors[..., 2]
# Calculate the horizontal angle (yaw/azimuth), assuming the Z-axis is forward.
azimuth = torch.atan2(x, z) * (180 / math.pi)
# Calculate the vertical angle (pitch/elevation).
elevation = torch.atan2(y, torch.sqrt(x**2 + z**2)) * (180 / math.pi)
# Calculate the angular difference from the center view direction (handling angle wrapping).
diff_azimuth = azimuth - center_yaw
diff_azimuth = torch.remainder(diff_azimuth + 180, 360) - 180
diff_elevation = elevation - center_pitch
diff_elevation = torch.remainder(diff_elevation + 180, 360) - 180
# Check if within FOV
in_fov_h = diff_azimuth.abs() < fov_half_h
in_fov_v = diff_elevation.abs() < fov_half_v
return in_fov_h & in_fov_v
def calculate_fov_overlap_similarity(
w2c_matrix_curr: torch.Tensor,
w2c_matrix_hist: torch.Tensor,
fov_h_deg: float = 105.0,
fov_v_deg: float = 75.0,
device=None,
points_local=None,
) -> float:
"""
Calculate the Field-of-View (FOV) overlap similarity between two W2C poses using Monte Carlo sampling.
Similarity = (Number of points in Curr_FOV ∩ Hist_FOV) / (Number of points in Curr_FOV).
:param w2c_matrix_curr: The (4, 4) W2C matrix for the current frame.
:param w2c_matrix_hist: The (4, 4) W2C matrix for the historical frame.
:param num_samples, radius, fov_h_deg, fov_v_deg: Sampling and FOV parameters.
:return: The overlap ratio (a float between 0.0 and 1.0).
"""
w2c_matrix_curr = torch.tensor(w2c_matrix_curr, device=device)
w2c_matrix_hist = torch.tensor(w2c_matrix_hist, device=device)
c2w_matrix_curr = torch.linalg.inv(w2c_matrix_curr)
c2w_matrix_hist = torch.linalg.inv(w2c_matrix_hist)
C_inv = w2c_matrix_curr
w2c_matrix_curr = torch.linalg.inv(C_inv @ c2w_matrix_curr)
w2c_matrix_hist = torch.linalg.inv(C_inv @ c2w_matrix_hist)
R_curr, t_curr = w2c_matrix_curr[:3, :3], w2c_matrix_curr[:3, 3]
R_hist, t_hist = w2c_matrix_hist[:3, :3], w2c_matrix_hist[:3, 3]
P_w_curr = -R_curr.T @ t_curr
P_w_hist = -R_hist.T @ t_hist
# pitch, yaw
pitch_curr, yaw_curr = rotation_matrix_to_angles(R_curr)
pitch_hist, yaw_hist = rotation_matrix_to_angles(R_hist)
fov_half_h = torch.tensor(fov_h_deg / 2.0, device=device)
fov_half_v = torch.tensor(fov_v_deg / 2.0, device=device)
# move to P_w_curr (N, 3)
points_world = points_local + P_w_curr[None, :]
in_fov_curr = is_inside_fov_3d_hv(
points_world,
P_w_curr[None, :],
pitch_curr[None],
yaw_curr[None],
fov_half_h,
fov_half_v,
)
# compute based on angle
in_fov_hist = is_inside_fov_3d_hv(
points_world,
P_w_hist[None, :],
pitch_hist[None],
yaw_hist[None],
fov_half_h,
fov_half_v,
)
# compute based on distance
dist = torch.norm(points_world - P_w_hist.reshape(1, -1), dim=1) < 8.0
in_fov_hist = in_fov_hist.bool() & dist.reshape(1, -1).bool()
overlap_count = (in_fov_curr.bool() & in_fov_hist.bool()).sum().float()
fov_curr_count = in_fov_curr.sum().float()
if fov_curr_count == 0:
return 0.0
overlap_ratio = overlap_count / fov_curr_count
return overlap_ratio.item()
def select_aligned_memory_frames(
w2c_list: List[np.ndarray],
current_frame_idx: int,
memory_frames: int,
temporal_context_size: int,
pred_latent_size: int,
pos_weight: float = 1.0,
ang_weight: float = 1.0,
device=None,
points_local=None,
) -> List[int]:
"""
Selects memory and context frames for a given frame based on a four-frame segment distance calculation.
:param w2c_list: List of all N 4x4 World-to-Camera (W2C) extrinsic matrices (np.ndarray).
:param current_frame_idx: The index of the current frame to be processed.
:param memory_frames: The total number of memory frames to select.
:param context_size: The total number of context frames to select.
:param pos_weight: The weight applied to the spatial (position) distance component.
:param ang_weight: The weight applied to the angular distance component.
:return: List[int]: A list containing the indices of the selected memory frames and context frames.
"""
if current_frame_idx <= memory_frames:
return list(range(0, current_frame_idx))
num_total_frames = len(w2c_list)
if current_frame_idx >= num_total_frames or current_frame_idx < 3:
raise ValueError(
f"The current frame index must be within the valid range of w2c_list and must be at least 3."
f"{current_frame_idx}, {len(w2c_list)}"
)
start_context_idx = max(0, current_frame_idx - temporal_context_size)
context_frames_indices = list(range(start_context_idx, current_frame_idx))
candidate_distances = []
query_clip_indices = list(
range(
current_frame_idx,
(
current_frame_idx + pred_latent_size
if current_frame_idx + pred_latent_size <= num_total_frames
else num_total_frames
),
)
)
historical_clip_indices = list(
range(4, current_frame_idx - temporal_context_size, 4)
)
memory_frames_indices = [0, 1, 2, 3] # add the first chunk as context
memory_frames = memory_frames - temporal_context_size
for hist_idx in historical_clip_indices:
total_dist = 0
hist_w2c_1 = w2c_list[hist_idx]
hist_w2c_2 = w2c_list[hist_idx + 2]
for query_idx in query_clip_indices:
dist_1_for_query_idx = 1.0 - calculate_fov_overlap_similarity(
w2c_list[query_idx],
hist_w2c_1,
fov_h_deg=60.0,
fov_v_deg=35.0,
device=device,
points_local=points_local,
)
dist_2_for_query_idx = 1.0 - calculate_fov_overlap_similarity(
w2c_list[query_idx],
hist_w2c_2,
fov_h_deg=60.0,
fov_v_deg=35.0,
device=device,
points_local=points_local,
)
dist_for_query_idx = (dist_1_for_query_idx + dist_2_for_query_idx) / 2.0
total_dist += dist_for_query_idx
final_clip_distance = total_dist / len(query_clip_indices)
candidate_distances.append((hist_idx, final_clip_distance))
candidate_distances.sort(key=lambda x: x[1])
for start_idx, _ in candidate_distances:
# check the memory frame number
if len(memory_frames_indices) >= memory_frames:
break
if start_idx not in memory_frames_indices:
memory_frames_indices.extend(range(start_idx, start_idx + 4))
# exclude the repeated frames
selected_frames_set = set(context_frames_indices)
selected_frames_set.update(memory_frames_indices)
final_selected_frames = sorted(list(selected_frames_set))
return final_selected_frames
+112
View File
@@ -0,0 +1,112 @@
# HY-WorldPlay/hyvideo/generate_custom_trajectory.py
import numpy as np
import json
def rot_x(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[1, 0, 0], [0, c, -s], [0, s, c]])
def rot_y(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]])
def rot_z(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[c, -s, 0], [s, c, 0], [0, 0, 1]])
def generate_camera_trajectory_local(motions):
"""
motions: list of dict
{"forward": 1.0}, {"yaw": np.pi/2}, {"pitch": np.pi/6}, {"right": 1.0}
- forward: Translation (Forward or Backward)
- yaw: Rotate (Left or Right)
- pitch: Rotate (Up or Down)
- right: Translation (Right or Left)
- third_yaw: Third Perspective Rotate (Left or Right)
"""
poses = []
T = np.eye(4)
poses.append(T.copy())
for move in motions:
# Rotate (Left or Right)
if "yaw" in move:
R = rot_y(move["yaw"])
T[:3, :3] = T[:3, :3] @ R
# Rotate (Up or Down)
if "pitch" in move:
R = rot_x(move["pitch"])
T[:3, :3] = T[:3, :3] @ R
# Translation (Z-direction of the camera's local coordinate system)
forward = move.get("forward", 0.0)
if forward != 0:
local_t = np.array([0, 0, forward])
world_t = T[:3, :3] @ local_t
T[:3, 3] += world_t
# Translation (Z-direction of the camera's local coordinate system)
right = move.get("right", 0.0)
if right != 0:
local_t = np.array([right, 0, 0])
world_t = T[:3, :3] @ local_t
T[:3, 3] += world_t
# Third Perspective Rotate (Left or Right)
third_yaw = move.get("third_yaw", 0.0)
if third_yaw != 0:
theta = -third_yaw
C = np.array([[1, 0.0, 0, 0], [0, 1, 0, 0], [0, 0, 1, -1.0], [0, 0, 0, 1]])
c_origin = C.copy()
# Rotation around the Y-axis
R_y = np.array(
[
[np.cos(theta), 0, np.sin(theta)],
[0, 1, 0],
[-np.sin(theta), 0, np.cos(theta)],
]
)
# Translation
C[:3, :3] = C[:3, :3] @ R_y
C[:3, 3] = R_y @ C[:3, 3]
c_inv = np.linalg.inv(c_origin)
c_relative = c_inv @ C
T = T @ c_relative
poses.append(T.copy())
return poses
if __name__ == "__main__":
# Examples: Forward 0.08 * 16 -> Right Rotate 3 degree * 16
motions = []
for i in range(15):
motions.append({"forward": 0.08})
for i in range(16):
motions.append({"yaw": np.deg2rad(3)})
intrinsic = [
[969.6969696969696, 0.0, 960.0],
[0.0, 969.6969696969696, 540.0],
[0.0, 0.0, 1.0],
]
poses = generate_camera_trajectory_local(motions)
custom_c2w = {}
for i, p in enumerate(poses):
custom_c2w[str(i)] = {"extrinsic": p.tolist(), "K": intrinsic}
json.dump(
custom_c2w,
open("./assets/pose/pose.json", "w"),
indent=4,
ensure_ascii=False,
)
File diff suppressed because it is too large Load Diff
+563
View File
@@ -0,0 +1,563 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
import os
from typing import Iterable
import torch
from torch import nn
from transformers import Gemma3ForConditionalGeneration
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.dits.ltx2 import (
FeedForward,
LTXRopeType,
apply_ltx_rotary_emb,
generate_ltx_freq_grid_np,
generate_ltx_freq_grid_pytorch,
precompute_ltx_freqs_cis,
)
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.platforms import AttentionBackendEnum
def _debug_log_line(message: str) -> None:
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") != "1":
return
log_path = os.getenv("LTX2_PIPELINE_DEBUG_PATH", "")
if not log_path:
return
log_dir = os.path.dirname(log_path)
if log_dir:
os.makedirs(log_dir, exist_ok=True)
with open(log_path, "a", encoding="utf-8") as f:
f.write(message + "\n")
def _debug_gemma_log_line(message: str) -> None:
log_path = os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", "")
if not log_path:
return
log_dir = os.path.dirname(log_path)
if log_dir:
os.makedirs(log_dir, exist_ok=True)
with open(log_path, "a", encoding="utf-8") as f:
f.write(message + "\n")
@dataclass(frozen=True)
class GemmaConnectorConfig:
num_attention_heads: int
attention_head_dim: int
num_layers: int
positional_embedding_theta: float
positional_embedding_max_pos: list[int]
rope_type: LTXRopeType
double_precision_rope: bool
num_learnable_registers: int | None
class GemmaFeaturesExtractorProjLinear(nn.Module):
"""Linear projection that aggregates stacked Gemma hidden states."""
def __init__(self, in_features: int, out_features: int) -> None:
super().__init__()
self.aggregate_embed = nn.Linear(in_features, out_features, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.aggregate_embed(x)
class _BasicTransformerBlock1D(nn.Module):
"""1D transformer block for connector processing."""
def __init__(
self,
dim: int,
heads: int,
dim_head: int,
rope_type: LTXRopeType,
norm_eps: float = 1e-6,
) -> None:
super().__init__()
self.attn1 = _GemmaAttention(
query_dim=dim,
context_dim=None,
heads=heads,
dim_head=dim_head,
norm_eps=norm_eps,
rope_type=rope_type,
)
self.ff = FeedForward(dim, dim_out=dim)
self.norm_eps = norm_eps
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
norm_hidden_states = torch.nn.functional.rms_norm(
hidden_states, (hidden_states.shape[-1],), eps=self.norm_eps
)
if norm_hidden_states.ndim == 4:
norm_hidden_states = norm_hidden_states.squeeze(1)
attn_output = self.attn1(
norm_hidden_states,
mask=attention_mask,
pe=pe,
)
hidden_states = attn_output + hidden_states
if hidden_states.ndim == 4:
hidden_states = hidden_states.squeeze(1)
norm_hidden_states = torch.nn.functional.rms_norm(
hidden_states, (hidden_states.shape[-1],), eps=self.norm_eps
)
ff_output = self.ff(norm_hidden_states)
hidden_states = ff_output + hidden_states
if hidden_states.ndim == 4:
hidden_states = hidden_states.squeeze(1)
return hidden_states
class _GemmaAttention(nn.Module):
"""Attention implementation aligned with LTX-2 text encoder."""
def __init__(
self,
query_dim: int,
context_dim: int | None,
heads: int,
dim_head: int,
norm_eps: float,
rope_type: LTXRopeType,
) -> None:
super().__init__()
inner_dim = dim_head * heads
context_dim = query_dim if context_dim is None else context_dim
self.heads = heads
self.dim_head = dim_head
self.rope_type = rope_type
self.q_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
self.k_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
self.to_q = nn.Linear(query_dim, inner_dim, bias=True)
self.to_k = nn.Linear(context_dim, inner_dim, bias=True)
self.to_v = nn.Linear(context_dim, inner_dim, bias=True)
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim, bias=True), nn.Identity())
def forward(
self,
x: torch.Tensor,
context: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
k_pe: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
q = self.to_q(x)
context = x if context is None else context
k = self.to_k(context)
v = self.to_v(context)
q = self.q_norm(q)
k = self.k_norm(k)
if pe is not None:
q = apply_ltx_rotary_emb(q, pe, self.rope_type)
k = apply_ltx_rotary_emb(k, pe if k_pe is None else k_pe, self.rope_type)
b, q_len, _ = q.shape
k_len = k.shape[1]
q = q.view(b, q_len, self.heads, self.dim_head).transpose(1, 2)
k = k.view(b, k_len, self.heads, self.dim_head).transpose(1, 2)
v = v.view(b, k_len, self.heads, self.dim_head).transpose(1, 2)
if mask is not None:
if mask.ndim == 2:
mask = mask.unsqueeze(0)
if mask.ndim == 3:
mask = mask.unsqueeze(1)
out = torch.nn.functional.scaled_dot_product_attention(
q,
k,
v,
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
)
out = out.transpose(1, 2).reshape(b, q_len, -1)
return self.to_out(out)
class Embeddings1DConnector(nn.Module):
"""Transformer connector that refines Gemma embeddings for LTX-2."""
_supports_gradient_checkpointing = True
def __init__(self, config: GemmaConnectorConfig) -> None:
super().__init__()
self.num_attention_heads = config.num_attention_heads
self.inner_dim = config.num_attention_heads * config.attention_head_dim
self.positional_embedding_theta = config.positional_embedding_theta
self.positional_embedding_max_pos = config.positional_embedding_max_pos
self.rope_type = config.rope_type
self.double_precision_rope = config.double_precision_rope
self.transformer_1d_blocks = nn.ModuleList(
[
_BasicTransformerBlock1D(
dim=self.inner_dim,
heads=config.num_attention_heads,
dim_head=config.attention_head_dim,
rope_type=config.rope_type,
)
for _ in range(config.num_layers)
]
)
self.num_learnable_registers = config.num_learnable_registers
if self.num_learnable_registers:
self.learnable_registers = nn.Parameter(
torch.rand(
self.num_learnable_registers,
self.inner_dim,
dtype=torch.bfloat16,
)
* 2.0
- 1.0
)
def _replace_padded_with_learnable_registers(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
assert hidden_states.shape[1] % self.num_learnable_registers == 0, (
f"Hidden states sequence length {hidden_states.shape[1]} must be divisible by "
f"num_learnable_registers {self.num_learnable_registers}."
)
num_registers_duplications = (
hidden_states.shape[1] // self.num_learnable_registers
)
learnable_registers = torch.tile(
self.learnable_registers, (num_registers_duplications, 1)
)
attention_mask_binary = (
attention_mask.squeeze(1).squeeze(1).unsqueeze(-1) >= -9000.0
).int()
non_zero_hidden_states = hidden_states[
:, attention_mask_binary.squeeze().bool(), :
]
non_zero_nums = non_zero_hidden_states.shape[1]
pad_length = hidden_states.shape[1] - non_zero_nums
adjusted_hidden_states = torch.nn.functional.pad(
non_zero_hidden_states, pad=(0, 0, 0, pad_length), value=0
)
flipped_mask = torch.flip(attention_mask_binary, dims=[1])
hidden_states = flipped_mask * adjusted_hidden_states + (
1 - flipped_mask
) * learnable_registers
attention_mask = torch.full_like(
attention_mask,
0.0,
dtype=attention_mask.dtype,
device=attention_mask.device,
)
return hidden_states, attention_mask
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
if self.num_learnable_registers:
hidden_states, attention_mask = (
self._replace_padded_with_learnable_registers(
hidden_states, attention_mask
)
)
indices_grid = torch.arange(
hidden_states.shape[1],
dtype=torch.float32,
device=hidden_states.device,
)
indices_grid = indices_grid[None, None, :]
freq_grid_generator = (
generate_ltx_freq_grid_np
if self.double_precision_rope
else generate_ltx_freq_grid_pytorch
)
freqs_cis = precompute_ltx_freqs_cis(
indices_grid=indices_grid,
dim=self.inner_dim,
out_dtype=hidden_states.dtype,
theta=self.positional_embedding_theta,
max_pos=self.positional_embedding_max_pos,
num_attention_heads=self.num_attention_heads,
rope_type=self.rope_type,
freq_grid_generator=freq_grid_generator,
)
for block in self.transformer_1d_blocks:
hidden_states = block(
hidden_states, attention_mask=attention_mask, pe=freqs_cis
)
hidden_states = torch.nn.functional.rms_norm(
hidden_states, (hidden_states.shape[-1],), eps=1e-6
)
return hidden_states, attention_mask
class LTX2GemmaTextEncoderModel(TextEncoder):
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__(config)
arch = config.arch_config
self.feature_extractor_linear = GemmaFeaturesExtractorProjLinear(
in_features=arch.feature_extractor_in_features,
out_features=arch.feature_extractor_out_features,
)
connector_config = GemmaConnectorConfig(
num_attention_heads=arch.connector_num_attention_heads,
attention_head_dim=arch.connector_attention_head_dim,
num_layers=arch.connector_num_layers,
positional_embedding_theta=arch.connector_positional_embedding_theta,
positional_embedding_max_pos=arch.connector_positional_embedding_max_pos,
rope_type=LTXRopeType(arch.connector_rope_type),
double_precision_rope=arch.connector_double_precision_rope,
num_learnable_registers=arch.connector_num_learnable_registers,
)
self.embeddings_connector = Embeddings1DConnector(connector_config)
self.audio_embeddings_connector = Embeddings1DConnector(connector_config)
self.gemma_model_path = arch.gemma_model_path
self.gemma_dtype = arch.gemma_dtype
self.padding_side = arch.padding_side
self._gemma_model: Gemma3ForConditionalGeneration | None = None
def named_parameters(self, prefix: str = "", recurse: bool = True):
for name, param in super().named_parameters(
prefix=prefix, recurse=recurse
):
if name.startswith("gemma_model."):
continue
yield name, param
@property
def gemma_model(self) -> Gemma3ForConditionalGeneration:
if self._gemma_model is None:
gemma_path = self.gemma_model_path
if not gemma_path:
raise ValueError(
"gemma_model_path must be set (expected text_encoder/gemma)."
)
dtype = getattr(torch, self.gemma_dtype, torch.bfloat16)
self._gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
gemma_path,
local_files_only=True,
torch_dtype=dtype,
)
# Configure model-level attention implementation when using TORCH_SDPA.
# Note: torch.backends.cuda.enable_*_sdp() settings should be configured
# at application/pipeline initialization level, not here, to avoid
# unexpected side effects across the application.
if os.getenv("FASTVIDEO_ATTENTION_BACKEND") == "TORCH_SDPA":
if hasattr(self._gemma_model.config, "attn_implementation"):
self._gemma_model.config.attn_implementation = "sdpa"
if hasattr(self._gemma_model.config, "_attn_implementation"):
self._gemma_model.config._attn_implementation = "sdpa"
device = next(self.feature_extractor_linear.parameters()).device
self._gemma_model.to(device=device)
self._gemma_model.eval()
return self._gemma_model
def _run_feature_extractor(
self,
hidden_states: tuple[torch.Tensor, ...],
attention_mask: torch.Tensor,
padding_side: str,
) -> torch.Tensor:
encoded_text_features = torch.stack(hidden_states, dim=-1)
if os.getenv("LTX2_FASTVIDEO_GEMMA_LOG", ""):
for idx, layer in enumerate(hidden_states):
_debug_gemma_log_line(
f"fastvideo:gemma_hidden_state_{idx}"
f":sum={layer.float().sum().item():.6f}"
)
_debug_gemma_log_line(
"fastvideo:gemma_hidden_states_stack"
f":sum={encoded_text_features.float().sum().item():.6f}"
)
encoded_text_features_dtype = encoded_text_features.dtype
sequence_lengths = attention_mask.sum(dim=-1)
normed_text_features = _norm_and_concat_padded_batch(
encoded_text_features, sequence_lengths, padding_side=padding_side
)
return self.feature_extractor_linear(
normed_text_features.to(encoded_text_features_dtype)
)
def _convert_to_additive_mask(
self, attention_mask: torch.Tensor, dtype: torch.dtype
) -> torch.Tensor:
return (attention_mask - 1).to(dtype).reshape(
(attention_mask.shape[0], 1, -1, attention_mask.shape[-1])
) * torch.finfo(dtype).max
def _run_connectors(
self,
encoded_input: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
connector_attention_mask = self._convert_to_additive_mask(
attention_mask, encoded_input.dtype
)
encoded, encoded_connector_attention_mask = self.embeddings_connector(
encoded_input, connector_attention_mask
)
attention_mask = (encoded_connector_attention_mask < 0.000001).to(
torch.int64
)
attention_mask = attention_mask.reshape(
[encoded.shape[0], encoded.shape[1], 1]
)
encoded = encoded * attention_mask
encoded_for_audio, _ = self.audio_embeddings_connector(
encoded_input, connector_attention_mask
)
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
if input_ids is None:
raise ValueError("input_ids is required for Gemma text encoding.")
if attention_mask is None:
attention_mask = torch.ones_like(input_ids)
model = self.gemma_model
input_ids = input_ids.to(device=model.device)
attention_mask = attention_mask.to(device=model.device)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
encoded_inputs = self._run_feature_extractor(
outputs.hidden_states,
attention_mask,
padding_side=self.padding_side,
)
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
_debug_log_line(
"fastvideo:gemma_feature"
f":sum={encoded_inputs.float().sum().item():.6f} "
f"shape={tuple(encoded_inputs.shape)}"
)
video_encoding, audio_encoding, attention_mask = self._run_connectors(
encoded_inputs, attention_mask
)
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
_debug_log_line(
"fastvideo:gemma_video_encoding"
f":sum={video_encoding.float().sum().item():.6f} "
f"shape={tuple(video_encoding.shape)}"
)
_debug_log_line(
"fastvideo:gemma_audio_encoding"
f":sum={audio_encoding.float().sum().item():.6f} "
f"shape={tuple(audio_encoding.shape)}"
)
hidden_states = (audio_encoding, ) if output_hidden_states else None
return BaseEncoderOutput(
last_hidden_state=video_encoding,
hidden_states=hidden_states,
attention_mask=attention_mask,
)
def load_weights(
self, weights: Iterable[tuple[str, torch.Tensor]]
) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if name == "aggregate_embed.weight":
name = "feature_extractor_linear.aggregate_embed.weight"
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
def _norm_and_concat_padded_batch(
encoded_text: torch.Tensor,
sequence_lengths: torch.Tensor,
padding_side: str = "right",
) -> torch.Tensor:
b, t, d, l = encoded_text.shape
device = encoded_text.device
token_indices = torch.arange(t, device=device)[None, :]
if padding_side == "right":
mask = token_indices < sequence_lengths[:, None]
elif padding_side == "left":
start_indices = t - sequence_lengths[:, None]
mask = token_indices >= start_indices
else:
raise ValueError(
f"padding_side must be 'left' or 'right', got {padding_side}"
)
mask = mask.reshape(b, t, 1, 1)
eps = 1e-6
masked = encoded_text.masked_fill(~mask, 0.0)
denom = (sequence_lengths * d).view(b, 1, 1, 1)
mean = masked.sum(dim=(1, 2), keepdim=True) / (denom + eps)
x_min = encoded_text.masked_fill(~mask, float("inf")).amin(
dim=(1, 2), keepdim=True
)
x_max = encoded_text.masked_fill(~mask, float("-inf")).amax(
dim=(1, 2), keepdim=True
)
range_ = x_max - x_min
normed = 8 * (encoded_text - mean) / (range_ + eps)
normed = normed.reshape(b, t, -1)
mask_flattened = mask.reshape(b, t, 1).expand(-1, -1, d * l)
normed = normed.masked_fill(~mask_flattened, 0.0)
return normed
+419
View File
@@ -0,0 +1,419 @@
# SPDX-License-Identifier: Apache-2.0
"""
SigLIP Vision Encoder for FastVideo.
SigLIP (Sigmoid Loss for Language-Image Pre-training) is similar to CLIP
but uses sigmoid loss instead of contrastive loss, and doesn't use a CLS token.
"""
from collections.abc import Iterable
import torch
import torch.nn as nn
from fastvideo.attention import LocalAttention
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.siglip import SiglipVisionArchConfig, SiglipVisionConfig
from fastvideo.distributed import divide, get_tp_world_size
from fastvideo.layers.activation import get_act_fn
from fastvideo.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import ImageEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
logger = init_logger(__name__)
class SiglipVisionEmbeddings(nn.Module):
"""
SigLIP vision embeddings - similar to CLIP but without class embedding.
"""
def __init__(self, arch_config: SiglipVisionArchConfig):
super().__init__()
self.arch_config = arch_config
self.embed_dim = arch_config.hidden_size
self.image_size = arch_config.image_size
self.patch_size = arch_config.patch_size
# SigLIP uses valid padding, so non-divisible sizes work (edge pixels are ignored)
self.patch_embedding = nn.Conv2d(
in_channels=arch_config.num_channels,
out_channels=self.embed_dim,
kernel_size=self.patch_size,
stride=self.patch_size,
padding="valid", # SigLIP uses valid padding
)
# Integer division - with valid padding, edge pixels are ignored
self.num_patches = (self.image_size // self.patch_size) ** 2
self.num_positions = self.num_patches
self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)
self.register_buffer(
"position_ids",
torch.arange(self.num_positions).expand((1, -1)),
persistent=False,
)
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
target_dtype = self.patch_embedding.weight.dtype
patch_embeds = self.patch_embedding(
pixel_values.to(dtype=target_dtype)
) # shape = [*, embed_dim, grid, grid]
embeddings = patch_embeds.flatten(2).transpose(1, 2)
embeddings = embeddings + self.position_embedding(self.position_ids)
return embeddings
class SiglipAttention(nn.Module):
"""Multi-headed attention for SigLIP."""
def __init__(
self,
arch_config: SiglipVisionArchConfig,
enable_scale: bool = True,
is_causal: bool = False,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
self.arch_config = arch_config
self.embed_dim = arch_config.hidden_size
self.num_heads = arch_config.num_attention_heads
self.head_dim = self.embed_dim // self.num_heads
if self.head_dim * self.num_heads != self.embed_dim:
raise ValueError(
f"embed_dim must be divisible by num_heads "
f"(got `embed_dim`: {self.embed_dim} and `num_heads`: {self.num_heads})."
)
self.scale = self.head_dim**-0.5 if enable_scale else None
self.dropout = arch_config.attention_dropout
self.qkv_proj = QKVParallelLinear(
hidden_size=self.embed_dim,
head_size=self.head_dim,
total_num_heads=self.num_heads,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.out_proj = RowParallelLinear(
input_size=self.embed_dim,
output_size=self.embed_dim,
quant_config=quant_config,
prefix=f"{prefix}.out_proj",
)
self.tp_size = get_tp_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = LocalAttention(
self.num_heads_per_partition,
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=is_causal,
supported_attention_backends=arch_config._supported_attention_backends,
)
def forward(self, hidden_states: torch.Tensor):
"""Input shape: Batch x Time x Channel"""
qkv_states, _ = self.qkv_proj(hidden_states)
query_states, key_states, value_states = qkv_states.chunk(3, dim=-1)
query_states = query_states.reshape(
query_states.shape[0], query_states.shape[1],
self.num_heads_per_partition, self.head_dim
)
key_states = key_states.reshape(
key_states.shape[0], key_states.shape[1],
self.num_heads_per_partition, self.head_dim
)
value_states = value_states.reshape(
value_states.shape[0], value_states.shape[1],
self.num_heads_per_partition, self.head_dim
)
attn_output = self.attn(query_states, key_states, value_states)
attn_output = attn_output.reshape(
attn_output.shape[0], attn_output.shape[1],
self.num_heads_per_partition * self.head_dim
)
attn_output, _ = self.out_proj(attn_output)
return attn_output, None
class SiglipMLP(nn.Module):
def __init__(
self,
arch_config: SiglipVisionArchConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.arch_config = arch_config
self.activation_fn = get_act_fn(arch_config.hidden_act)
self.fc1 = ColumnParallelLinear(
arch_config.hidden_size,
arch_config.intermediate_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.fc1",
)
self.fc2 = RowParallelLinear(
arch_config.intermediate_size,
arch_config.hidden_size,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.fc2",
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc1(hidden_states)
hidden_states = self.activation_fn(hidden_states)
hidden_states, _ = self.fc2(hidden_states)
return hidden_states
class SiglipEncoderLayer(nn.Module):
def __init__(
self,
arch_config: SiglipVisionArchConfig,
enable_scale: bool = True,
is_causal: bool = False,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.self_attn = SiglipAttention(
arch_config,
enable_scale=enable_scale,
is_causal=is_causal,
quant_config=quant_config,
prefix=f"{prefix}.self_attn",
)
self.layer_norm1 = nn.LayerNorm(
arch_config.hidden_size, eps=arch_config.layer_norm_eps
)
self.mlp = SiglipMLP(
arch_config,
quant_config=quant_config,
prefix=f"{prefix}.mlp",
)
self.layer_norm2 = nn.LayerNorm(
arch_config.hidden_size, eps=arch_config.layer_norm_eps
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
# SigLIP uses post-norm (like original ViT)
residual = hidden_states
hidden_states = self.layer_norm1(hidden_states)
hidden_states, _ = self.self_attn(hidden_states=hidden_states)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.layer_norm2(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states
class SiglipEncoder(nn.Module):
"""SigLIP encoder consisting of transformer layers."""
def __init__(
self,
arch_config: SiglipVisionArchConfig,
enable_scale: bool = True,
is_causal: bool = False,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.arch_config = arch_config
if num_hidden_layers_override is None:
num_hidden_layers = arch_config.num_hidden_layers
else:
num_hidden_layers = num_hidden_layers_override
self.layers = nn.ModuleList([
SiglipEncoderLayer(
arch_config=arch_config,
enable_scale=enable_scale,
is_causal=is_causal,
quant_config=quant_config,
prefix=f"{prefix}.layers.{layer_idx}",
)
for layer_idx in range(num_hidden_layers)
])
def forward(
self,
inputs_embeds: torch.Tensor,
return_all_hidden_states: bool,
) -> torch.Tensor | list[torch.Tensor]:
hidden_states_pool = [inputs_embeds]
hidden_states = inputs_embeds
for encoder_layer in self.layers:
hidden_states = encoder_layer(hidden_states)
if return_all_hidden_states:
hidden_states_pool.append(hidden_states)
if return_all_hidden_states:
return hidden_states_pool
return [hidden_states]
class SiglipVisionTransformer(nn.Module):
def __init__(
self,
config: SiglipVisionConfig,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
require_post_norm: bool | None = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
arch_config = config.arch_config
embed_dim = arch_config.hidden_size
self.embeddings = SiglipVisionEmbeddings(arch_config)
self.encoder = SiglipEncoder(
arch_config=arch_config,
enable_scale=config.enable_scale,
is_causal=config.is_causal,
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
prefix=f"{prefix}.encoder",
)
num_hidden_layers = arch_config.num_hidden_layers
if len(self.encoder.layers) > arch_config.num_hidden_layers:
raise ValueError(
f"The original encoder only has {num_hidden_layers} "
f"layers, but you requested {len(self.encoder.layers)} layers."
)
# Post layer norm (applied to output)
if require_post_norm is None:
require_post_norm = len(self.encoder.layers) == num_hidden_layers
if require_post_norm:
self.post_layernorm = nn.LayerNorm(embed_dim, eps=arch_config.layer_norm_eps)
else:
self.post_layernorm = None
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: list[int] | None = None,
) -> torch.Tensor:
hidden_states = self.embeddings(pixel_values)
return_all_hidden_states = feature_sample_layers is not None
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
return_all_hidden_states=return_all_hidden_states,
)
if not return_all_hidden_states:
encoder_outputs = encoder_outputs[0]
# Apply post-layernorm
if self.post_layernorm is not None:
if isinstance(encoder_outputs, list):
encoder_outputs[-1] = self.post_layernorm(encoder_outputs[-1])
else:
encoder_outputs = self.post_layernorm(encoder_outputs)
return encoder_outputs
class SiglipVisionModel(ImageEncoder):
"""SigLIP Vision Model for FastVideo."""
config_class = SiglipVisionConfig
main_input_name = "pixel_values"
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
def __init__(self, config: SiglipVisionConfig) -> None:
super().__init__(config)
self.config = config
self.vision_model = SiglipVisionTransformer(
config=config,
quant_config=config.quant_config,
num_hidden_layers_override=config.num_hidden_layers_override,
require_post_norm=config.require_post_norm,
prefix=f"{config.prefix}.vision_model",
)
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: list[int] | None = None,
**kwargs,
) -> BaseEncoderOutput:
last_hidden_state = self.vision_model(pixel_values, feature_sample_layers)
return BaseEncoderOutput(last_hidden_state=last_hidden_state)
@property
def device(self):
return next(self.parameters()).device
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
layer_count = len(self.vision_model.encoder.layers)
for name, loaded_weight in weights:
# Skip projection layers if any
if name.startswith("visual_projection"):
continue
# Skip head if any
if "head" in name:
continue
# Post layernorm handling
if (name.startswith("vision_model.post_layernorm")
and self.vision_model.post_layernorm is None):
continue
# Omit layers when num_hidden_layers_override is set
if name.startswith("vision_model.encoder.layers"):
layer_idx = int(name.split(".")[3])
if layer_idx >= layer_count:
continue
# Handle QKV projection weight mapping
for (param_name, weight_name, shard_id) in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
if name in params_dict:
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
if name in params_dict:
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
+190 -18
View File
@@ -80,11 +80,15 @@ 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"),
"tokenizer_2": (TokenizerLoader, "transformers"),
"image_processor": (ImageProcessorLoader, "transformers"),
"feature_extractor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
}
@@ -242,6 +246,47 @@ class TextEncoderLoader(ComponentLoader):
model_config.pop("model_type", None)
model_config.pop("tokenizer_class", None)
model_config.pop("torch_dtype", None)
repo_root = os.path.dirname(model_path)
index_path = os.path.join(repo_root, "model_index.json")
gemma_path = ""
gemma_path_from_candidate = False
if os.path.isfile(index_path):
try:
with open(index_path, encoding="utf-8") as f:
model_index = json.load(f)
gemma_path = model_index.get("gemma_model_path", "")
except json.JSONDecodeError:
gemma_path = ""
if not gemma_path:
candidate = os.path.normpath(os.path.join(model_path, "gemma"))
if os.path.isdir(candidate):
gemma_path = candidate
gemma_path_from_candidate = True
model_config["gemma_model_path"] = gemma_path
if gemma_path and not gemma_path_from_candidate:
if not os.path.isabs(gemma_path):
model_config["gemma_model_path"] = os.path.normpath(
os.path.join(repo_root, gemma_path)
)
transformer_config_path = os.path.join(
repo_root, "transformer", "config.json"
)
if os.path.isfile(transformer_config_path):
try:
with open(transformer_config_path, encoding="utf-8") as f:
transformer_config = json.load(f)
if (
"connector_double_precision_rope" not in model_config
or not model_config["connector_double_precision_rope"]
):
if transformer_config.get("double_precision_rope") is True:
model_config["connector_double_precision_rope"] = True
if "connector_rope_type" not in model_config:
rope_type = transformer_config.get("rope_type")
if rope_type is not None:
model_config["connector_rope_type"] = rope_type
except json.JSONDecodeError:
pass
logger.info("HF Model config: %s", model_config)
# @TODO(Wei): Better way to handle this?
@@ -346,6 +391,8 @@ class TextEncoderLoader(ComponentLoader):
from fastvideo.platforms import current_platform
logger.info("Loading text encoder with cpu_offload: %s", use_cpu_offload)
if use_cpu_offload:
pin_cpu_memory = fastvideo_args.pin_cpu_memory and is_pin_memory_available()
# Disable FSDP for MPS as it's not compatible
@@ -489,8 +536,20 @@ class TokenizerLoader(ComponentLoader):
# in v0, this was same string as encoder_name "ClipTextModel"
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
# other method of config?
padding_size="right",
)
padding_side = None
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
try:
arch_config = fastvideo_args.pipeline_config.text_encoder_configs[
0
].arch_config
padding_side = getattr(arch_config, "padding_side", None)
except Exception:
padding_side = None
if padding_side:
tokenizer.padding_side = padding_side
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
tokenizer.pad_token = tokenizer.eos_token
logger.info("Loaded tokenizer: %s", tokenizer.__class__.__name__)
return tokenizer
@@ -501,15 +560,11 @@ 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)
config.pop("_name_or_path", None)
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."
)
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
fastvideo_args.model_paths["vae"] = model_path
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
from fastvideo.platforms import current_platform
if fastvideo_args.vae_cpu_offload:
@@ -543,8 +598,29 @@ class VAELoader(ComponentLoader):
vae.load_state_dict(sd, strict=False)
return vae.eval()
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# LTX-2 uses CausalVideoAutoencoder with nested "vae" config
if class_name == "CausalVideoAutoencoder" and "vae" in config:
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(config).to(target_device)
if hasattr(vae, "set_tiling_config"):
vae_config = fastvideo_args.pipeline_config.vae_config
vae.set_tiling_config(
spatial_tile_size_in_pixels=getattr(
vae_config, "ltx2_spatial_tile_size_in_pixels", 512),
spatial_tile_overlap_in_pixels=getattr(
vae_config, "ltx2_spatial_tile_overlap_in_pixels", 64),
temporal_tile_size_in_frames=getattr(
vae_config, "ltx2_temporal_tile_size_in_frames", 64),
temporal_tile_overlap_in_frames=getattr(
vae_config,
"ltx2_temporal_tile_overlap_in_frames", 24),
)
else:
config.pop("_class_name", None)
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -553,23 +629,108 @@ class VAELoader(ComponentLoader):
raise ValueError(f"No safetensors files found in {model_path}")
# Common case: a single `.safetensors` checkpoint file.
# Some models may be sharded into multiple files; in that case we merge.
if len(safetensors_list) == 1:
loaded = safetensors_load_file(safetensors_list[0])
else:
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
# LTX-2 CausalVideoAutoencoder needs per_channel_statistics remapping
if class_name == "CausalVideoAutoencoder" and "vae" in config:
per_channel_prefixes = (
"per_channel_statistics.",
"vae.per_channel_statistics.",
)
remapped = {}
for key, tensor in loaded.items():
remapped[key] = tensor
for prefix in per_channel_prefixes:
if key.startswith(prefix):
suffix = key[len(prefix):]
remapped.setdefault(
f"encoder.per_channel_statistics.{suffix}",
tensor,
)
remapped.setdefault(
f"decoder.per_channel_statistics.{suffix}",
tensor,
)
break
loaded = remapped
vae.load_state_dict(loaded, strict=False)
return vae.eval()
class AudioDecoderLoader(ComponentLoader):
"""Loader for LTX-2 audio decoder (audio_vae component)."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name", None) or "LTX2AudioDecoder"
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
target_device = get_local_torch_device()
precision = getattr(
fastvideo_args.pipeline_config, "audio_decoder_precision", "bf16"
)
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
audio_decoder = model_cls(config).to(target_device)
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
decoder_state = {}
for name, tensor in loaded.items():
if name.startswith("decoder."):
decoder_state[name.replace("decoder.", "")] = tensor
elif name.startswith("per_channel_statistics."):
decoder_state[name] = tensor
target_module = getattr(audio_decoder, "model", audio_decoder)
target_module.load_state_dict(decoder_state, strict=False)
return audio_decoder.eval()
class VocoderLoader(ComponentLoader):
"""Loader for LTX-2 vocoder."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name", None) or "LTX2Vocoder"
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
target_device = get_local_torch_device()
precision = getattr(
fastvideo_args.pipeline_config, "vocoder_precision", "bf16"
)
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
vocoder = model_cls(config).to(target_device)
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors")
)
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
target_module = getattr(vocoder, "model", vocoder)
target_module.load_state_dict(loaded, strict=False)
return vocoder.eval()
class TransformerLoader(ComponentLoader):
"""Loader for transformer."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
config.pop("_name_or_path", None)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
if cls_name is None:
@@ -670,7 +831,7 @@ class TransformerLoader(ComponentLoader):
)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
logger.info("Loaded model with %.2fB parameters, with cpu_offload: %s", total_params / 1e9, fastvideo_args.dit_cpu_offload)
assert next(model.parameters()).dtype == default_dtype, (
"Model dtype does not match default dtype"
@@ -679,7 +840,18 @@ class TransformerLoader(ComponentLoader):
model = model.eval()
if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload:
enable_layerwise_offload(model)
# Check if model has nn.ModuleList for layerwise offload compatibility
has_module_list = any(
isinstance(m, nn.ModuleList) for m in model.children()
)
if has_module_list:
enable_layerwise_offload(model)
else:
logger.warning(
"Layerwise offload requested but model %s does not have "
"nn.ModuleList structure. Skipping layerwise offload.",
cls_name
)
return model
@@ -786,4 +958,4 @@ class PipelineComponentLoader:
)
# Load the module
return loader.load(component_model_path, fastvideo_args)
return loader.load(component_model_path, fastvideo_args)
+4 -3
View File
@@ -78,9 +78,10 @@ def hf_to_custom_state_dict(
for source_param_name, full_tensor in hf_param_sd: # type: ignore
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is None:
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
+17 -1
View File
@@ -26,6 +26,10 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"HunyuanVideo15Transformer3DModel":
("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
"HYWorldTransformer3DModel":
("dits", "hyworld", "HYWorldTransformer3DModel"),
"CausalHunyuanVideo15Transformer3DModel":
("dits", "causal_hunyuanvideo15", "CausalHunyuanVideo15Transformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
@@ -33,6 +37,7 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -54,20 +59,30 @@ _TEXT_ENCODER_MODELS = {
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
"Qwen2_5_VLForConditionalGeneration":
("encoders", "reason1", "Reason1TextEncoder"),
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
"CLIPVisionModel": ("encoders", "clip", "CLIPVisionModel"),
"SiglipVisionModel": ("encoders", "siglip", "SiglipVisionModel"),
}
_VAE_MODELS = {
"AutoencoderKLHunyuanVideo":
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLHYWorld": ("vaes", "hyworldvae", "AutoencoderKLHYWorld"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
}
_AUDIO_MODELS = {
"LTX2AudioEncoder": ("audio", "ltx2_audio_vae", "LTX2AudioEncoder"),
"LTX2AudioDecoder": ("audio", "ltx2_audio_vae", "LTX2AudioDecoder"),
"LTX2Vocoder": ("audio", "ltx2_audio_vae", "LTX2Vocoder"),
}
_SCHEDULERS = {
@@ -91,6 +106,7 @@ _FAST_VIDEO_MODELS = {
**_TEXT_ENCODER_MODELS,
**_IMAGE_ENCODER_MODELS,
**_VAE_MODELS,
**_AUDIO_MODELS,
**_SCHEDULERS,
}
@@ -155,8 +155,8 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
self.sigmas = sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
self.sigma_min = sigma_min if sigma_min is not None else self.sigmas[-1].item()
self.sigma_max = sigma_max if sigma_max is not None else self.sigmas[0].item()
BaseScheduler.__init__(self)
@@ -289,6 +289,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
sigmas: list[float] | None = None,
mu: float | None = None,
timesteps: list[float] | None = None,
extra_one_step: bool = False,
) -> None:
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -350,7 +351,10 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
if timesteps_array is None:
t_max = self._sigma_to_t(self.sigma_max)
t_min = self._sigma_to_t(self.sigma_min)
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
if extra_one_step:
timesteps_array = np.linspace(t_max, t_min, num_inference_steps + 1)[:-1]
else:
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
sigmas_array = timesteps_array / self.config.num_train_timesteps
else:
sigmas_array = np.array(sigmas).astype(np.float32)
@@ -644,7 +648,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
self,
clean_latent: torch.Tensor,
noise: torch.Tensor,
timestep: torch.IntTensor,
timestep: torch.Tensor,
) -> torch.Tensor:
"""
+3
View File
@@ -5,6 +5,9 @@ from typing import Any
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move it elsewhere
def auto_attributes(init_func):
+1 -1
View File
@@ -79,7 +79,7 @@ class ParallelTiledVAE(ABC):
def decode(self, z: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = z.shape
tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_stride_width // self.spatial_compression_ratio
tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio
tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio
num_sample_frames = (num_frames -
1) * self.temporal_compression_ratio + 1
+1 -1
View File
@@ -663,7 +663,7 @@ class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = False
self.use_tiling = True
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from diffusers and HY-WorldPlay
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from fastvideo.models.vaes.hunyuan15vae import AutoencoderKLHunyuanVideo15
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
class AutoencoderKLHYWorld(AutoencoderKLHunyuanVideo15):
# TODO(mingjia): add temporal caching support for HYWorld VAE
def __init__(
self,
config: Hunyuan15VAEConfig,
) -> None:
AutoencoderKLHunyuanVideo15.__init__(self, config)
File diff suppressed because it is too large Load Diff
+228
View File
@@ -0,0 +1,228 @@
import math
import torch
try:
from torch.distributed.tensor import DTensor
except ImportError:
# handle old pytorch versions
Dtensor = None
# This code is modified from the GitHub repository of KellerJordan:
# https://github.com/KellerJordan/Muon/blob/master/muon.py
def zeropower_via_newtonschulz5(G, steps=5):
"""
Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a
quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose
of minimizing steps, it turns out to be empirically effective to keep increasing the slope at
zero even beyond the point where the iteration no longer converges all the way to one everywhere
on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T
where S' is diagonal with S_{ii}' ~ Uniform(0.5, 1.5), which turns out not to hurt model
performance at all relative to UV^T, where USV^T = G is the SVD.
"""
if isinstance(G, DTensor):
device_mesh = G.device_mesh
G = G.full_tensor()
else:
device_mesh = None
assert len(G.shape) >= 2
a, b, c = (3.4445, -4.7750, 2.0315)
X = G
if G.size(-2) > G.size(-1):
X = X.mT
# Ensure spectral norm is at most 1
X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7)
# Perform the NS iterations
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A # quintic computation strategy adapted from suggestion by @jxbz, @leloykun, and @YouJiacheng
X = a * X + B @ X
if G.size(-2) > G.size(-1):
X = X.mT
if device_mesh is not None:
return DTensor.from_local(X, device_mesh)
else:
return X
class Muon(torch.optim.Optimizer):
"""
Muon - MomentUm Orthogonalized by Newton-schulz
Arguments:
muon_params: The parameters to be optimized by Muon.
lr: The learning rate. The updates will have spectral norm of `lr`. (0.02 is a good default)
momentum: The momentum used by the internal SGD. (0.95 is a good default)
nesterov: Whether to use Nesterov-style momentum in the internal SGD. (recommended)
ns_steps: The number of Newton-Schulz iterations to run. (6 is probably always enough)
adamw_params: The parameters to be optimized by AdamW. Any parameters in `muon_params` which are
{0, 1}-D or are detected as being the embed or lm_head will be optimized by AdamW as well.
adamw_lr: The learning rate for the internal AdamW.
adamw_betas: The betas for the internal AdamW.
adamw_eps: The epsilon for the internal AdamW.
adamw_wd: The weight decay for the internal AdamW.
"""
def __init__(
self,
lr=1e-3,
wd=0.1,
muon_params=None,
momentum=0.95,
nesterov=True,
ns_steps=5,
adamw_params=None,
adamw_betas=(0.95, 0.95),
adamw_eps=1e-8,
):
defaults = dict(
lr=lr,
wd=wd,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
)
params = list(muon_params)
adamw_params = list(adamw_params) if adamw_params is not None else []
params.extend(adamw_params)
super().__init__(params, defaults)
# Sort parameters into those for which we will use Muon, and those for which we will not
for p in muon_params:
# Use Muon for every parameter in muon_params which is >= 2D and doesn't look like an embedding or head layer
assert p.ndim >= 2, p.ndim
self.state[p]["use_muon"] = True
for p in adamw_params:
# Do not use Muon for parameters in adamw_params
self.state[p]["use_muon"] = False
def adjust_lr_for_muon(self, lr, param_shape):
A, B = param_shape[:2]
# We adjust the learning rate and weight decay based on the size of the parameter matrix
# as describted in the paper
adjusted_ratio = 0.2 * math.sqrt(max(A, B))
adjusted_lr = lr * adjusted_ratio
return adjusted_lr
def step(self, closure=None):
"""Perform a single optimization step.
Args:
closure (Callable, optional): A closure that reevaluates the model
and returns the loss.
"""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
############################
# Muon #
############################
params = [p for p in group["params"] if self.state[p]["use_muon"]]
lr = group["lr"]
wd = group["wd"]
momentum = group["momentum"]
# generate weight updates in distributed fashion
for p in params:
# sanity check
g = p.grad
if g is None:
continue
if g.ndim > 2:
g = g.view(g.size(0), -1)
assert g is not None
# calc update
state = self.state[p]
if "momentum_buffer" not in state:
state["momentum_buffer"] = torch.zeros_like(g)
buf = state["momentum_buffer"]
buf.mul_(momentum).add_(g)
g = g.add(buf, alpha=momentum) if group["nesterov"] else buf
g = g.bfloat16()
u = zeropower_via_newtonschulz5(g, steps=group["ns_steps"])
# scale update
adjusted_lr = self.adjust_lr_for_muon(lr, p.shape)
# apply weight decay
p.data.mul_(1 - lr * wd)
# apply update
p.data.add_(u.view(p.shape), alpha=-adjusted_lr)
############################
# AdamW backup #
############################
params = [
p for p in group["params"] if not self.state[p]["use_muon"]
]
lr = group['lr']
beta1, beta2 = group["adamw_betas"]
eps = group["adamw_eps"]
weight_decay = group["wd"]
for p in params:
g = p.grad
if g is None:
continue
state = self.state[p]
if "step" not in state:
state["step"] = 0
state["moment1"] = torch.zeros_like(g)
state["moment2"] = torch.zeros_like(g)
state["step"] += 1
step = state["step"]
buf1 = state["moment1"]
buf2 = state["moment2"]
buf1.lerp_(g, 1 - beta1)
buf2.lerp_(g.square(), 1 - beta2)
g = buf1 / (eps + buf2.sqrt())
bias_correction1 = 1 - beta1**step
bias_correction2 = 1 - beta2**step
scale = bias_correction1 / bias_correction2**0.5
p.data.mul_(1 - lr * weight_decay)
p.data.add_(g, alpha=-lr / scale)
return loss
# help function to create the Muon optimizer
def get_muon_optimizer(model,
lr=1e-3,
weight_decay=0.1,
momentum=0.95,
adamw_betas=(0.95, 0.95),
adamw_eps=1e-8):
muon_params = [
p for name, p in model.named_parameters()
if p.requires_grad and p.ndim >= 2
]
adamw_params = [
p for name, p in model.named_parameters()
if p.requires_grad and not (p.ndim >= 2)
]
return Muon(
lr=lr,
wd=weight_decay,
muon_params=muon_params,
momentum=momentum,
adamw_params=adamw_params,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
)
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan causal DMD pipeline implementation.
This module wires the causal DMD denoising stage into the modular pipeline.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
Hy15CausalDMDDenosingStage,
InputValidationStage,
Hy15ImageEncodingStage,
LatentPreparationStage,
TextEncodingStage)
# isort: on
logger = init_logger(__name__)
class Hy15CausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
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"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=Hy15CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = Hy15CausalDMDPipeline
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
Hy15ImageEncodingStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15ImageToVideoPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HunyuanVideo15ImageToVideoPipeline
@@ -0,0 +1 @@
@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
"""
HYWorld video diffusion pipeline implementation.
This module contains an implementation of the HYWorld video diffusion pipeline
using the modular pipeline architecture with HYWorld-specific denoising stage
for chunk-based video generation with context frame selection.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (
ConditioningStage, DecodingStage, HYWorldDenoisingStage,
InputValidationStage, LatentPreparationStage, TextEncodingStage,
TimestepPreparationStage, HYWorldImageEncodingStage)
logger = init_logger(__name__)
class HYWorldPipeline(ComposedPipelineBase):
"""
HYWorld video diffusion pipeline.
This pipeline implements chunk-based video generation with context frame
selection for 3D-aware generation using HYWorldDenoisingStage.
Note: HYWorld only uses a single LLM-based text encoder, unlike SDXL-style
dual encoder setups. The text_encoder_2/tokenizer_2 are not used.
"""
# Include image_encoder and feature_extractor for I2V support with SigLIP
# Note: guider (ClassifierFreeGuidance) is not needed - FastVideo handles CFG differently
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler",
"text_encoder_2", "tokenizer_2", "image_encoder", "feature_extractor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with HYWorld-specific denoising stage."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
]))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=HYWorldImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("feature_extractor"),
vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=HYWorldDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
pipeline=self))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HYWorldPipeline
@@ -0,0 +1,150 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 text-to-video pipeline.
"""
import os
from typing import Any
from transformers import AutoTokenizer
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import PipelineComponentLoader
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (DecodingStage, InputValidationStage,
LTX2AudioDecodingStage,
LTX2DenoisingStage,
LTX2LatentPreparationStage,
TextEncodingStage)
logger = init_logger(__name__)
class LTX2Pipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder",
"tokenizer",
"transformer",
"vae",
"audio_vae",
"vocoder",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(
stage_name="input_validation_stage",
stage=InputValidationStage(),
)
self.add_stage(
stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=LTX2LatentPreparationStage(
transformer=self.get_module("transformer"), ),
)
self.add_stage(
stage_name="denoising_stage",
stage=LTX2DenoisingStage(
transformer=self.get_module("transformer"), ),
)
self.add_stage(
stage_name="audio_decoding_stage",
stage=LTX2AudioDecodingStage(
audio_decoder=self.get_module("audio_vae"),
vocoder=self.get_module("vocoder"),
),
)
self.add_stage(
stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")),
)
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
tokenizer = self.get_module("tokenizer")
if tokenizer is not None:
tokenizer.padding_side = "left"
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
def load_modules(
self,
fastvideo_args: FastVideoArgs,
loaded_modules: dict[str, Any] | None = None,
) -> dict[str, Any]:
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
model_index.pop("workload_type", None)
if len(model_index) <= 1:
raise ValueError(
"model_index.json must contain at least one pipeline module")
required_modules = self.required_config_modules
modules: dict[str, Any] = {}
for module_name, module_spec in model_index.items():
if not isinstance(module_spec, list) or len(module_spec) < 1:
continue
transformers_or_diffusers = module_spec[0]
if transformers_or_diffusers is None:
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
if module_name not in required_modules:
continue
if loaded_modules is not None and module_name in loaded_modules:
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
if module_name == "tokenizer" and not os.path.isdir(
component_model_path):
gemma_path = os.path.join(self.model_path, "text_encoder",
"gemma")
if os.path.isdir(gemma_path):
component_model_path = gemma_path
else:
raise ValueError(
"Tokenizer directory missing and Gemma weights were not found."
)
module = PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=fastvideo_args,
)
logger.info("Loaded module %s from %s", module_name,
component_model_path)
modules[module_name] = module
if "tokenizer" in required_modules and "tokenizer" not in modules:
gemma_path = os.path.join(self.model_path, "text_encoder", "gemma")
if os.path.isdir(gemma_path):
modules["tokenizer"] = AutoTokenizer.from_pretrained(
gemma_path, local_files_only=True)
for module_name in required_modules:
if module_name not in modules or modules[module_name] is None:
raise ValueError(
f"Required module {module_name} was not loaded properly")
return modules
EntryClass = LTX2Pipeline
@@ -84,6 +84,7 @@ class ComposedPipelineBase(ABC):
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
# fastvideo_args.dit_cpu_offload = False
with self.profiler_controller.region("profiler_region_model_loading"):
self.modules = self.load_modules(fastvideo_args, loaded_modules)
@@ -287,6 +288,7 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
model_index.pop("_name_or_path", None)
model_index.pop("workload_type", None)
if "boundary_ratio" in model_index and model_index[
"boundary_ratio"] is not None:
@@ -132,6 +132,9 @@ class ForwardBatch:
keyboard_cond: torch.Tensor | None = None # Shape: (B, T, K)
grid_sizes: torch.Tensor | None = None # Shape: (3,) [F,H,W]
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
# Latent dimensions
height_latents: list[int] | int | None = None
width_latents: list[int] | int | None = None
@@ -229,6 +232,7 @@ class TrainingBatch:
image_latents: torch.Tensor | None = None
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
video_latent: torch.Tensor | None = None
# ODE trajectory supervision
trajectory_latents: torch.Tensor | None = None
@@ -237,6 +241,9 @@ class TrainingBatch:
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
timesteps: torch.Tensor | None = None
use_gt_trajectory: bool = False
trajectory_timesteps: torch.Tensor | None = None
start_timestep_index: int | None = None
sigmas: torch.Tensor | None = None
noise: torch.Tensor | None = None
+4
View File
@@ -28,6 +28,9 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
"HYWorldPipeline": "hyworld",
"Hy15CausalDMDPipeline": "hunyuan15",
"HunyuanVideo15ImageToVideoPipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"Cosmos2_5Pipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
@@ -35,6 +38,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"LongCatPipeline": "longcat",
"LongCatImageToVideoPipeline": "longcat",
"LongCatVideoContinuationPipeline": "longcat",
"LTX2Pipeline": "ltx2",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
@@ -303,7 +303,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
@@ -422,4 +422,4 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
num_processed_samples = 0
num_processed_samples = 0
@@ -28,16 +28,12 @@ from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages import (
DecodingStage, DenoisingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage, Hy15ImageEncodingStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
logger = init_logger(__name__)
@@ -47,7 +43,8 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
@@ -61,19 +58,19 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
assert fastvideo_args.pipeline_config.flow_shift == 5
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
denoising_strength=1.0)
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")],
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
@@ -82,6 +79,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
@@ -95,11 +95,13 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
args):
"""Preprocess text-only data and generate trajectory information."""
num_encoders = len(self.prompt_encoding_stage.text_encoders)
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
@@ -130,12 +132,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
@@ -144,61 +143,48 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
negative_prompt_embeds_list = []
negative_prompt_masks_list = []
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = prompt_embeds_list
batch.prompt_attention_mask = prompt_masks_list
batch.negative_prompt_embeds = negative_prompt_embeds_list
batch.negative_attention_mask = negative_prompt_masks_list
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.num_frames = args.num_frames
batch.fps = args.train_fps
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.num_inference_steps = 48
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.image_encoding_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch, fastvideo_args)
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
trajectory_latents.append(result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
extra_features = {
@@ -209,10 +195,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
if j in [5, 7]:
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
@@ -227,7 +214,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
text_embedding = prompt_embeds_list[0].float().cpu().numpy()
text_mask = prompt_masks_list[0].cpu().numpy()
text_embedding_2 = prompt_embeds_list[1].float().cpu(
).numpy()
text_mask_2 = prompt_masks_list[1].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
@@ -253,6 +244,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"trajectory_latents"],
trajectory_timesteps=sample_extra_features[
"trajectory_timesteps"],
text_embedding_2=text_embedding_2,
text_mask=text_mask,
text_mask_2=text_mask_2,
)
batch_data.append(record)
@@ -58,7 +58,7 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
@@ -181,4 +181,4 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text
EntryClass = PreprocessPipeline_Text
+18 -20
View File
@@ -1,9 +1,7 @@
import argparse
import os
from typing import Any
from fastvideo import PipelineConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.distributed import (
get_world_size, maybe_init_distributed_environment_and_model_parallel)
from fastvideo.fastvideo_args import FastVideoArgs
@@ -16,38 +14,38 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
def main(args) -> None:
args.model_path = maybe_download_model(args.model_path)
# args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
print(pipeline_config.__class__.__name__)
kwargs: dict[str, Any] = {}
if args.preprocess_task == "text_only":
kwargs = {
"text_encoder_cpu_offload": False,
}
else:
# Full config for video/image processing
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
pipeline_config.update_config_from_dict(kwargs)
# kwargs: dict[str, Any] = {}
# if args.preprocess_task == "text_only":
# kwargs = {
# "text_encoder_cpu_offload": False,
# }
# else:
# # Full config for video/image processing
# kwargs = {
# "vae_precision": "fp32",
# "vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
# }
# pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pipeline_config=pipeline_config,
)
if args.preprocess_task == "t2v":
@@ -134,4 +132,4 @@ if __name__ == "__main__":
)
args = parser.parse_args()
main(args)
main(args)
@@ -24,4 +24,4 @@ if __name__ == "__main__":
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
fastvideo_args = FastVideoArgs.from_cli_args(args)
main(fastvideo_args)
main(fastvideo_args)
+14 -1
View File
@@ -8,6 +8,7 @@ complete diffusion pipelines.
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.hy15_causal_denoising import Hy15CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
@@ -17,13 +18,19 @@ from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage)
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage,
HYWorldImageEncodingStage)
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.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage)
from fastvideo.pipelines.stages.text_encoding import (Cosmos25TextEncodingStage,
@@ -44,18 +51,24 @@ __all__ = [
"LatentPreparationStage",
"CosmosLatentPreparationStage",
"Cosmos25LatentPreparationStage",
"LTX2LatentPreparationStage",
"LTX2AudioDecodingStage",
"Hy15CausalDMDDenosingStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"MatrixGameCausalDenoisingStage",
"HYWorldDenoisingStage",
"CosmosDenoisingStage",
"Cosmos25DenoisingStage",
"LTX2DenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"MatrixGameImageEncodingStage",
"Hy15ImageEncodingStage",
"HYWorldImageEncodingStage",
"RefImageEncodingStage",
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
@@ -45,12 +45,12 @@ class CausalDMDDenosingStage(DenoisingStage):
self.transformer_2 = transformer_2
self.vae = vae
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
self.num_transformer_blocks = self.transformer.config.num_layers
self.num_frames_per_block = self.transformer.config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.sliding_window_num_frames
try:
self.local_attn_size = getattr(self.transformer.model,
self.local_attn_size = getattr(self.transformer.config,
"local_attn_size",
-1) # type: ignore
except Exception:
@@ -412,8 +412,8 @@ class CausalDMDDenosingStage(DenoisingStage):
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = self.transformer.attention_head_dim
num_attention_heads = self.transformer.config.num_attention_heads
attention_head_dim = self.transformer.config.attention_head_dim
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
+18 -10
View File
@@ -26,7 +26,7 @@ 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.platforms import AttentionBackendEnum
from fastvideo.utils import dict_to_3d_list, masks_like
from fastvideo.utils import dict_to_3d_list, masks_like, PRECISION_TO_TYPE
try:
from fastvideo.attention.backends.sliding_tile_attn import (
@@ -210,7 +210,10 @@ class DenoisingStage(PipelineStage):
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
assert self.vae is not None, "VAE is not provided for TI2V task"
z = self.vae.encode(batch.pil_image).mean.float()
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
self.vae = self.vae.to(get_local_torch_device())
z = self.vae.encode(batch.pil_image.to(vae_dtype)).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
@@ -223,6 +226,9 @@ class DenoisingStage(PipelineStage):
else:
z = z * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to('cpu')
latent_model_input = latent_model_input.squeeze(0)
_, mask2 = masks_like([latent_model_input], zero=True)
@@ -232,18 +238,18 @@ class DenoisingStage(PipelineStage):
latent_model_input = latent_model_input.to(get_local_torch_device())
latents = latent_model_input
F = batch.num_frames
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
seq_len = ((F - 1) // temporal_scale +
1) * (batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size[1] *
patch_size[2])
temporal_scale = fastvideo_args.pipeline_config.vae_config.temporal_compression_ratio
spatial_scale = fastvideo_args.pipeline_config.vae_config.spatial_compression_ratio
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
seq_len = ((F - 1) // temporal_scale + 1) * (
batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size * patch_size)
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = [latents]
logger.info("timesteps: %s", timesteps)
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -408,6 +414,7 @@ class DenoisingStage(PipelineStage):
**action_kwargs,
)
assert batch.do_classifier_free_guidance, "do_classifier_free_guidance is not supported"
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
@@ -462,6 +469,7 @@ class DenoisingStage(PipelineStage):
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_timesteps.append(torch.zeros_like(t))
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
@@ -0,0 +1,361 @@
import math
import torch # type: ignore
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.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.utils import PRECISION_TO_TYPE
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
logger = init_logger(__name__)
class Hy15CausalDMDDenosingStage(CausalDMDDenosingStage):
"""
Denoising stage for causal diffusion.
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
if isinstance(self.transformer.config.patch_size, tuple):
patch_ratio = self.transformer.config.patch_size[
1] * self.transformer.config.patch_size[2]
elif isinstance(self.transformer.config.patch_size, int):
patch_ratio = self.transformer.config.patch_size**2
else:
raise ValueError(
f"Unsupported patch size type: {type(self.transformer.config.patch_size)}"
)
self.frame_seq_length = latent_seq_length // patch_ratio
# TODO(will): make this a parameter once we add i2v support
independent_first_frame = self.transformer.independent_first_frame if hasattr(
self.transformer, 'independent_first_frame') else False
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
self.scheduler.set_timesteps(num_inference_steps=1000,
extra_one_step=True,
device=get_local_torch_device())
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
logger.info("[causal_denoising] timesteps: %s", timesteps)
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
# Image kwargs (kept empty unless caller provides compatible args)
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# STA
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Latents and prompts
assert batch.latents is not None, "latents must be provided"
latents = batch.latents # [B, C, T, H, W]
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
pos_start_base = 0
num_blocks = math.ceil(t / self.num_frames_per_block)
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
# Initialize txt kv cache
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
txt_kv_cache = self.transformer(
txt_inference=True,
vision_inference=False,
encoder_hidden_states=prompt_embeds,
encoder_hidden_states_image=image_embeds,
encoder_attention_mask=batch.prompt_attention_mask,
timestep=torch.zeros([latents.shape[0]], device=latents.device),
cache_txt=True,
)
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert self.vae is not None, "VAE is not provided for causal video gen task"
self.vae = self.vae.to(get_local_torch_device())
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
first_frame_latent = self.vae.encode(
batch.pil_image.to(vae_dtype)).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent = first_frame_latent * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
t_zero = torch.zeros([latents.shape[0], 1],
device=latents.device,
dtype=torch.long)
if batch.video_latent is not None:
video_latent_chunk = batch.video_latent[:, :, start_index:
start_index + 1, :, :]
first_frame_input = torch.cat([
first_frame_latent,
video_latent_chunk,
torch.zeros_like(first_frame_latent),
],
dim=1)
else:
first_frame_input = first_frame_latent.clone()
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.transformer(
txt_inference=False,
vision_inference=True,
hidden_states=first_frame_input.to(target_dtype),
timestep=t_zero,
kv_cache=kv_cache1,
txt_kv_cache=txt_kv_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
rope_start_idx=start_index,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
vision_input_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"vision_inference": True,
"txt_inference": False,
},
)
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for current_num_frames in block_sizes:
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# use BTCHW for DMD conversion routines
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
# Copy for pred conversion
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
if batch.video_latent is not None:
video_latent_chunk = batch.video_latent[:, :,
start_index:
start_index +
current_num_frames, :, :]
latent_model_input = torch.cat([
latent_model_input,
video_latent_chunk,
torch.zeros_like(current_latents),
],
dim=1)
elif batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(target_dtype)
],
dim=2)
# Prepare inputs
t_expand = t_cur.repeat(latent_model_input.shape[0])
# Attention metadata if needed
if (vsa_available and self.attn_backend
== VideoSparseAttentionBackend):
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=i, # type: ignore
raw_latent_shape=(current_num_frames, h,
w), # type: ignore
patch_size=fastvideo_args.pipeline_config.
dit_config.patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(), # type: ignore
) # type: ignore
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw, kv_cache1 = self.transformer(
hidden_states=latent_model_input,
timestep=t_expanded_noise,
kv_cache=kv_cache1,
txt_kv_cache=txt_kv_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
rope_start_idx=start_index,
**vision_input_kwargs,
)
pred_noise_btchw = pred_noise_btchw.permute(
0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1],
dtype=torch.long,
device=pred_video_btchw.device)
noise = torch.randn(
video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(
0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(
0, 2, 1, 3, 4)
if progress_bar is not None:
progress_bar.update()
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Re-run with context timestep to update KV cache using clean context
context_noise = getattr(fastvideo_args.pipeline_config,
"context_noise", 0)
t_context = torch.ones([latents.shape[0], current_num_frames],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
if batch.video_latent is not None:
video_latent_chunk = batch.video_latent[:, :, start_index:
start_index +
current_num_frames, :, :]
context_bcthw = torch.cat([
context_bcthw,
video_latent_chunk,
torch.zeros_like(current_latents),
],
dim=1)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
_, kv_cache1 = self.transformer(
hidden_states=context_bcthw,
timestep=t_context,
kv_cache=kv_cache1,
txt_kv_cache=txt_kv_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
rope_start_idx=start_index,
**vision_input_kwargs,
)
start_index += current_num_frames
batch.latents = latents
return batch
@@ -0,0 +1,449 @@
# SPDX-License-Identifier: Apache-2.0
"""
HYWorld denoising stage for chunk-based video generation with context frame selection.
This stage implements the bi_rollout denoising logic from HYWorld, which processes
video generation in chunks with camera-aware context frame selection for temporal consistency.
"""
import torch
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.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.utils import dict_to_3d_list
from fastvideo.models.dits.hyworld.retrieval_context import (
generate_points_in_sphere, select_aligned_memory_frames)
logger = init_logger(__name__)
class HYWorldDenoisingStage(DenoisingStage):
"""
Denoising stage for HYWorld-style chunk-based video generation.
This stage implements bi_rollout denoising with:
- Chunk-based processing (generates video in chunks, e.g., 4 frames at a time) - Context frame selection based on camera view alignment
- 3D-aware generation using view matrices and camera intrinsics
- Support for action conditioning
- Dual timestep handling (context frames use different timestep than current frames) - Context frame selection based on camera view alignment
"""
def __init__(
self,
transformer,
scheduler,
pipeline=None,
transformer_2=None,
vae=None,
) -> None:
super().__init__(transformer, scheduler, pipeline, transformer_2, vae)
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run the chunk-based denoising loop with context frame selection.
Args:
batch: The current batch information. Must contain:
- viewmats: torch.Tensor | None - Camera view matrices (B, T, 4, 4)
- Ks: torch.Tensor | None - Camera intrinsics (B, T, 3, 3)
- action: torch.Tensor | None - Action conditioning (B, T)
- chunk_latent_frames: int - Number of frames per chunk (default: 4)
These can be passed via batch.extra dict or as direct attributes.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
# Extract HYWorld-specific parameters from batch.extra or batch attributes
viewmats = getattr(batch, "viewmats", None) or batch.extra.get(
"viewmats", None)
Ks = getattr(batch, "Ks", None) or batch.extra.get("Ks", None)
action = getattr(batch, "action", None) or batch.extra.get(
"action", None)
chunk_latent_frames = (getattr(batch, "chunk_latent_frames", None)
or batch.extra.get("chunk_latent_frames", 4))
stabilization_level = 15
points_local = (getattr(batch, "points_local", None)
or batch.extra.get("points_local", None))
if viewmats is None or Ks is None or action is None:
raise ValueError(
"viewmats, Ks, and action are required for HYWorld denoising. "
"Please provide them in batch.extra['viewmats'], batch.extra['Ks'], "
"and batch.extra['action']")
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta,
},
)
# Setup precision and autocast settings
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = (len(timesteps) -
num_inference_steps * self.scheduler.order)
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_image":
image_embeds,
"mask_strategy":
dict_to_3d_list(None, t_max=50, l_max=60, h_max=24),
},
)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert not torch.isnan(
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
latent_model_input = latents.to(target_dtype)
assert latent_model_input.shape[0] == 1, "only support batch size 1"
device = get_local_torch_device()
# Generate local points if not provided
if points_local is None:
points_local = generate_points_in_sphere(50000, 8.0).to(device)
else:
points_local = points_local.to(device)
# Use conditional latents directly (prepared by HYWorldImageEncodingStage)
# batch.image_latent is already [1, 33, T, H, W] with first frame encoded, rest zeros
cond_latents = batch.image_latent
# Calculate chunk configuration
latent_frames = latents.shape[2]
chunk_num = latent_frames // chunk_latent_frames
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# Main chunk processing loop
for chunk_i in range(chunk_num):
if chunk_i > 0:
# Select context frames based on camera alignment
current_frame_idx = chunk_i * chunk_latent_frames
selected_frame_indices = []
for chunk_start_idx in range(
current_frame_idx,
current_frame_idx + chunk_latent_frames,
4, # Process every 4 frames
):
selected_history_frame_id = select_aligned_memory_frames(
viewmats[0].cpu().detach().numpy(),
chunk_start_idx,
memory_frames=20,
temporal_context_size=12,
pred_latent_size=4,
points_local=points_local,
device=device,
)
selected_frame_indices.extend(selected_history_frame_id)
selected_frame_indices = sorted(
list(set(selected_frame_indices)))
# Remove current chunk frames from context
to_remove = list(
range(current_frame_idx,
current_frame_idx + chunk_latent_frames))
selected_frame_indices = [
x for x in selected_frame_indices if x not in to_remove
]
# Extract context frames
context_latents = latents[:, :, selected_frame_indices]
context_w2c = viewmats[:, selected_frame_indices]
context_Ks = Ks[:, selected_frame_indices]
context_action = action[:, selected_frame_indices]
self.scheduler.set_timesteps(num_inference_steps, device=device)
# Define chunk boundaries
start_idx = chunk_i * chunk_latent_frames
end_idx = chunk_i * chunk_latent_frames + chunk_latent_frames
# Denoising loop for this chunk
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if chunk_i == 0:
# First chunk: standard processing
timestep_input = torch.full(
(chunk_latent_frames, ),
t.item(),
device=device,
dtype=timesteps.dtype,
)
latent_model_input = latents[:, :, :chunk_latent_frames]
cond_latents_input = cond_latents[:, :, :
chunk_latent_frames]
else:
# Subsequent chunks: use context frames with different timesteps
t_now = torch.full(
(chunk_latent_frames, ),
t.item(),
device=device,
dtype=timesteps.dtype,
)
t_ctx = torch.full(
(len(selected_frame_indices), ),
stabilization_level - 1,
device=device,
dtype=timesteps.dtype,
)
timestep_input = torch.cat([t_ctx, t_now], dim=0)
latents_model_now = latents[:, :, start_idx:end_idx]
latent_model_input = torch.cat(
[context_latents, latents_model_now], dim=2)
cond_latents_input = cond_latents[:, :, :
latent_model_input.
shape[2]]
# Prepare viewmats, Ks, action for current chunk
viewmats_input = viewmats[:, start_idx:end_idx]
Ks_input = Ks[:, start_idx:end_idx]
action_input = action[:, start_idx:end_idx]
if chunk_i > 0:
viewmats_input = torch.cat(
[context_w2c, viewmats_input], dim=1)
Ks_input = torch.cat([context_Ks, Ks_input], dim=1)
action_input = torch.cat([context_action, action_input],
dim=1)
# Prepare latent input (concatenate with cond_latents if needed)
latents_concat = torch.concat(
[latent_model_input, cond_latents_input], dim=1)
# Note: Unlike some other pipelines, HYWorld runs CFG sequentially (two passes)
# rather than batching pos/neg together, following the original implementation
latents_concat = self.scheduler.scale_model_input(
latents_concat, t)
# Keep batch size 1 for sequential CFG
t_expand_txt = t.unsqueeze(0)
t_expand = timestep_input
viewmats_input = viewmats_input.to(device)
Ks_input = Ks_input.to(device)
action_input = action_input.reshape(-1).to(device)
with torch.autocast(
device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled,
):
current_model = self.transformer
batch.is_cfg_negative = False
# Prepare transformer kwargs with HYWorld-specific inputs
# Note: batch size 1 for sequential CFG (matching original HY-WorldPlay)
transformer_kwargs = {
**image_kwargs,
"timestep": t_expand,
"timestep_txt": t_expand_txt,
"viewmats": viewmats_input.to(target_dtype),
"Ks": Ks_input.to(target_dtype),
"action": action_input.to(target_dtype),
}
# Set encoder_attention_mask for positive/negative conditioning
pos_transformer_kwargs = {
**transformer_kwargs, "encoder_attention_mask":
batch.prompt_attention_mask
}
neg_transformer_kwargs = {
**transformer_kwargs, "encoder_attention_mask":
batch.negative_attention_mask
}
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
noise_pred = current_model(
latents_concat,
prompt_embeds,
**pos_transformer_kwargs,
)
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
noise_pred_uncond = current_model(
latents_concat,
neg_prompt_embeds,
**neg_transformer_kwargs,
)
noise_pred_text = noise_pred
noise_pred = noise_pred_uncond + batch.guidance_scale * (
noise_pred_text - noise_pred_uncond)
# Apply guidance rescale if needed
if batch.guidance_rescale > 0.0:
noise_pred = self.rescale_noise_cfg(
noise_pred,
noise_pred_text,
guidance_rescale=batch.guidance_rescale,
)
# Step scheduler - update only the current chunk's latents
latent_model_input = self.scheduler.step(
noise_pred,
t,
latent_model_input,
**extra_step_kwargs,
return_dict=False)[0]
# Update only the current chunk's latents
latents[:, :, start_idx:
end_idx] = latent_model_input[:, :,
-chunk_latent_frames:]
# Save trajectory latents if needed
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
trajectory_latents.append(latents.clone())
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
# Handle trajectory output
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
else:
trajectory_tensor = None
trajectory_timesteps_tensor = None
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()
# Update batch with final latents
batch.latents = latents
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify HYWorld denoising stage inputs."""
result = VerificationResult()
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.min_dims(1)])
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)
# Check for HYWorld-specific inputs
viewmats = getattr(batch, "viewmats", None) or batch.extra.get(
"viewmats", None)
Ks = getattr(batch, "Ks", None) or batch.extra.get("Ks", None)
action = getattr(batch, "action", None) or batch.extra.get(
"action", None)
if viewmats is None:
result.add_failure(
"viewmats",
"viewmats must be provided in batch.extra['viewmats'] or as batch.viewmats",
)
else:
result.add_check("viewmats", viewmats, V.is_tensor)
if Ks is None:
result.add_failure(
"Ks", "Ks must be provided in batch.extra['Ks'] or as batch.Ks")
else:
result.add_check("Ks", Ks, V.is_tensor)
if action is None:
result.add_failure(
"action",
"action must be provided in batch.extra['action'] or as batch.action",
)
else:
result.add_check("action", action, V.is_tensor)
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
if batch.do_classifier_free_guidance:
result.add_check(
"negative_prompt_embeds",
batch.negative_prompt_embeds,
V.list_not_empty,
)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify HYWorld denoising stage outputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
+195 -4
View File
@@ -115,10 +115,10 @@ class Hy15ImageEncodingStage(ImageEncodingStage):
"""
Encode the prompt into image encoder hidden states.
"""
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
# if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
raw_latent_shape = list(batch.raw_latent_shape)
raw_latent_shape[1] = 1
@@ -127,6 +127,197 @@ class Hy15ImageEncodingStage(ImageEncodingStage):
return batch
class HYWorldImageEncodingStage(ImageEncodingStage):
"""
Stage for encoding image prompts into embeddings for HYWorld models.
Uses SigLIP (or other vision encoder) to encode reference images for I2V tasks.
Also encodes reference image with VAE for conditional latent.
"""
def __init__(self, image_encoder=None, image_processor=None, vae=None):
super().__init__(image_encoder=image_encoder,
image_processor=image_processor)
self.vae = vae
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify image encoding stage inputs."""
return VerificationResult()
@torch.no_grad()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states and VAE latents.
For I2V:
- encodes the reference image using SigLIP → image_embeds
- encodes the reference image using VAE → image_latent (expanded to full temporal dim)
For T2V: creates zero embeddings
The image_latent is expanded to match the full temporal dimension of the video latent,
following the original HunyuanVideo-1.5 implementation where:
- First frame contains the encoded reference image
- All other frames are zeros
- Mask channel is 1 for first frame, 0 for rest
"""
device = get_local_torch_device()
# Default vision embed dimensions for HunyuanVideo1.5/HYWorld
num_vision_tokens = 729 # (384/14)^2 for SigLIP
vision_dim = 1152 # SigLIP hidden size
# Get temporal dimension from raw_latent_shape (set by LatentPreparationStage)
raw_latent_shape = list(batch.raw_latent_shape)
latent_channels = raw_latent_shape[1]
latent_temporal = raw_latent_shape[2] # T dimension
latent_height = raw_latent_shape[3]
latent_width = raw_latent_shape[4]
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
if batch.pil_image is None:
# T2V case: create zero embeddings for image_embeds
batch.image_embeds = [
torch.zeros(1, num_vision_tokens, vision_dim, device=device)
]
# T2V: create zero latents for image_latent with full temporal dimension
# Shape: [B, latent_channels + 1 (mask channel), T, H, W]
batch.image_latent = torch.zeros(1,
latent_channels + 1,
latent_temporal,
latent_height,
latent_width,
device=device)
else:
image = batch.pil_image
# 1. Encode with SigLIP for image_embeds
if self.image_encoder is not None:
self.image_encoder = self.image_encoder.to(device)
# Get model dtype for proper precision matching (HY-WorldPlay uses fp16)
model_dtype = next(self.image_encoder.parameters()).dtype
# Preprocess image for SigLIP
# Convert to numpy and resize to target resolution (matching HY-WorldPlay)
import numpy as np
if not isinstance(image, np.ndarray):
image_np = np.array(image)
else:
image_np = image
# Resize to target resolution BEFORE SigLIP preprocessing
from fastvideo.models.dits.hyworld.data_utils import resize_and_center_crop
image_np = resize_and_center_crop(image_np,
target_width=batch.width,
target_height=batch.height)
image_inputs = self.image_processor.preprocess(
images=image_np, return_tensors="pt").to(
device=device, dtype=model_dtype) # Match model dtype!
pixel_values = image_inputs['pixel_values']
with set_forward_context(current_timestep=0,
attn_metadata=None):
outputs = self.image_encoder(pixel_values=pixel_values)
image_embeds = outputs.last_hidden_state
batch.image_embeds = [image_embeds]
if fastvideo_args.image_encoder_cpu_offload:
self.image_encoder.to('cpu')
else:
batch.image_embeds = [
torch.zeros(1, num_vision_tokens, vision_dim, device=device)
]
# 2. Encode with VAE for image_latent (conditional latent for I2V)
if self.vae is not None:
from torchvision import transforms
from PIL import Image as PILImage
import numpy as np
# Preprocess image for VAE
if isinstance(image, np.ndarray):
image = PILImage.fromarray(image)
# Get target size from batch
origin_size = image.size
target_height, target_width = batch.height, batch.width
original_width, original_height = origin_size
scale_factor = max(target_width / original_width,
target_height / original_height)
resize_width = int(round(original_width * scale_factor))
resize_height = int(round(original_height * scale_factor))
ref_image_transform = transforms.Compose([
transforms.Resize(
(resize_height, resize_width),
interpolation=transforms.InterpolationMode.LANCZOS),
transforms.CenterCrop((target_height, target_width)),
transforms.ToTensor(),
transforms.Normalize([0.5], [0.5])
])
ref_images_pixel_values = ref_image_transform(image)
ref_images_pixel_values = (ref_images_pixel_values.unsqueeze(
0).unsqueeze(2).to(device))
# Encode with VAE
self.vae = self.vae.to(device)
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
cond_latents = self.vae.encode(
ref_images_pixel_values).mode()
cond_latents.mul_(self.vae.config.scaling_factor)
# cond_latents shape: [1, 32, 1, H//compression, W//compression]
# Expand to full temporal dimension: [1, 32, T, H, W]
# First frame contains the encoded image, rest are zeros
expanded_latent = cond_latents.repeat(1, 1, latent_temporal, 1,
1)
expanded_latent[:, :,
1:, :, :] = 0.0 # Zero out all frames except first
# Create mask: [1, 1, T, H, W]
# First frame mask = 1 (conditional), rest = 0
mask = torch.zeros(1,
1,
latent_temporal,
latent_height,
latent_width,
device=device,
dtype=expanded_latent.dtype)
mask[:, :, 0, :, :] = 1.0 # First frame is conditional
# Concatenate latent and mask: [1, 33, T, H, W]
batch.image_latent = torch.cat([expanded_latent, mask], dim=1)
if fastvideo_args.vae_cpu_offload:
self.vae.to('cpu')
else:
# No VAE available, create zero latents with full temporal dimension
batch.image_latent = torch.zeros(1,
latent_channels + 1,
latent_temporal,
latent_height,
latent_width,
device=device)
# Initialize video latent placeholder
raw_latent_shape[1] = 1
batch.video_latent = torch.zeros(tuple(raw_latent_shape), device=device)
return batch
class MatrixGameImageEncodingStage(ImageEncodingStage):
CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073]
CLIP_STD = [0.26862954, 0.26130258, 0.27577711]
@@ -120,9 +120,9 @@ class InputValidationStage(PipelineStage):
else:
# Standard Wan logic
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
dh, dw = patch_size * vae_stride, patch_size * vae_stride
max_area = 480 * 848
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
"""
Audio decoding stage for LTX-2 pipelines.
"""
from __future__ import annotations
import os
import torch
from fastvideo.models.dits.ltx2 import DEFAULT_LTX2_VOCODER_OUTPUT_SAMPLE_RATE
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
class LTX2AudioDecodingStage(PipelineStage):
"""Decode LTX-2 audio latents into a waveform."""
def __init__(self, audio_decoder, vocoder) -> None:
super().__init__()
self.audio_decoder = audio_decoder
self.vocoder = vocoder
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
audio_latents = batch.extra.get("ltx2_audio_latents")
if audio_latents is None:
return batch
device = get_local_torch_device()
self.audio_decoder = self.audio_decoder.to(device)
self.vocoder = self.vocoder.to(device)
audio_latents = audio_latents.to(device)
disable_autocast = os.getenv("LTX2_DISABLE_AUDIO_AUTOCAST", "1") == "1"
with torch.no_grad(), torch.autocast(
device_type="cuda",
dtype=audio_latents.dtype,
enabled=not disable_autocast,
):
decoded_spec = self.audio_decoder(audio_latents)
audio_wave = self.vocoder(decoded_spec).squeeze(0).float()
# Move to CPU for pickling across process boundary
batch.extra["audio"] = audio_wave.cpu()
batch.extra[
"audio_sample_rate"] = DEFAULT_LTX2_VOCODER_OUTPUT_SAMPLE_RATE
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("audio_latents", batch.extra.get("ltx2_audio_latents"),
V.none_or_tensor)
return result
@@ -0,0 +1,308 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 denoising stage using the native sigma schedule.
"""
from __future__ import annotations
import math
import os
import torch
from tqdm.auto import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.logger import init_logger
from fastvideo.models.dits.ltx2 import (
AudioLatentShape, DEFAULT_LTX2_AUDIO_CHANNELS,
DEFAULT_LTX2_AUDIO_DOWNSAMPLE, DEFAULT_LTX2_AUDIO_HOP_LENGTH,
DEFAULT_LTX2_AUDIO_MEL_BINS, DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
VideoLatentShape)
from fastvideo.utils import PRECISION_TO_TYPE
BASE_SHIFT_ANCHOR = 1024
MAX_SHIFT_ANCHOR = 4096
# Official distilled sigma schedule (8 denoising steps)
# From LTX-2/packages/ltx-pipelines/src/ltx_pipelines/utils/constants.py
DISTILLED_SIGMA_VALUES = [
1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0
]
logger = init_logger(__name__)
def _ltx2_sigmas(
steps: int,
latent: torch.Tensor | None,
device: torch.device,
max_shift: float = 2.05,
base_shift: float = 0.95,
stretch: bool = True,
terminal: float = 0.1,
) -> torch.Tensor:
tokens = math.prod(
latent.shape[2:]) if latent is not None else MAX_SHIFT_ANCHOR
sigmas = torch.linspace(1.0,
0.0,
steps + 1,
device=device,
dtype=torch.float32)
mm = (max_shift - base_shift) / (MAX_SHIFT_ANCHOR - BASE_SHIFT_ANCHOR)
b = base_shift - mm * BASE_SHIFT_ANCHOR
sigma_shift = tokens * mm + b
numerator = math.exp(sigma_shift)
sigmas = torch.where(
sigmas != 0,
numerator / (numerator + (1 / sigmas - 1)),
torch.zeros_like(sigmas),
)
if stretch:
non_zero_mask = sigmas != 0
non_zero_sigmas = sigmas[non_zero_mask]
one_minus_z = 1.0 - non_zero_sigmas
scale_factor = one_minus_z[-1] / (1.0 - terminal)
stretched = 1.0 - (one_minus_z / scale_factor)
sigmas = sigmas.clone()
sigmas[non_zero_mask] = stretched
return sigmas
class LTX2DenoisingStage(PipelineStage):
"""Run the LTX-2 denoising loop over the sigma schedule."""
def __init__(self, transformer) -> None:
super().__init__()
self.transformer = transformer
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
if batch.latents is None:
raise ValueError("Latents must be provided before denoising.")
latents = batch.latents
prompt_embeds = batch.prompt_embeds[0]
prompt_mask = None
neg_prompt_embeds = None
neg_prompt_mask = None
# Only load negative prompts if CFG is actually enabled
if batch.do_classifier_free_guidance:
assert batch.negative_prompt_embeds is not None, (
"CFG is enabled but negative_prompt_embeds is None")
neg_prompt_embeds = batch.negative_prompt_embeds[0]
# Ensure text conditioning is on the same device as latents.
if prompt_embeds.device != latents.device:
prompt_embeds = prompt_embeds.to(latents.device)
if prompt_mask is not None and prompt_mask.device != latents.device:
prompt_mask = prompt_mask.to(latents.device)
if neg_prompt_embeds is not None and neg_prompt_embeds.device != latents.device:
neg_prompt_embeds = neg_prompt_embeds.to(latents.device)
if neg_prompt_mask is not None and neg_prompt_mask.device != latents.device:
neg_prompt_mask = neg_prompt_mask.to(latents.device)
target_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
disable_autocast = os.getenv("LTX2_DISABLE_AUTOCAST", "1") == "1"
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast and (
not disable_autocast)
# Use official distilled sigma schedule for 8 steps (distilled models)
use_distilled_sigmas = os.getenv("LTX2_USE_DISTILLED_SIGMAS",
"1") == "1"
if use_distilled_sigmas and batch.num_inference_steps == 8:
sigmas = torch.tensor(
DISTILLED_SIGMA_VALUES,
device=latents.device,
dtype=torch.float32,
)
logger.info("[LTX2] Using official distilled sigma schedule")
else:
sigmas = _ltx2_sigmas(
steps=batch.num_inference_steps,
latent=None,
device=latents.device,
)
if hasattr(self.transformer, "patchifier"):
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
token_count = self.transformer.patchifier.get_token_count(
video_shape)
else:
token_count = 1
timestep_template = torch.ones(
(latents.shape[0], token_count),
device=latents.device,
dtype=torch.float32,
)
audio_prompt_embeds = batch.extra.get("ltx2_audio_prompt_embeds")
audio_neg_embeds = batch.extra.get("ltx2_audio_negative_embeds")
audio_context_p = audio_prompt_embeds[0] if audio_prompt_embeds else None
audio_context_n = audio_neg_embeds[0] if audio_neg_embeds else None
audio_latents = None
audio_timestep_template = None
if audio_context_p is not None:
fps_value = batch.fps
if isinstance(fps_value, list):
fps_value = fps_value[0] if fps_value else None
if fps_value is None:
fps_value = 1.0
duration = float(batch.num_frames) / float(fps_value)
audio_shape = AudioLatentShape.from_duration(
batch=latents.shape[0],
duration=duration,
channels=DEFAULT_LTX2_AUDIO_CHANNELS,
mel_bins=DEFAULT_LTX2_AUDIO_MEL_BINS,
sample_rate=DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
hop_length=DEFAULT_LTX2_AUDIO_HOP_LENGTH,
audio_latent_downsample_factor=DEFAULT_LTX2_AUDIO_DOWNSAMPLE,
)
audio_generator = None
if fastvideo_args.ltx2_initial_latent_path and batch.seed is not None:
audio_generator = torch.Generator(
device=latents.device).manual_seed(batch.seed)
elif batch.generator is not None:
if isinstance(batch.generator, list):
audio_generator = batch.generator[0]
else:
audio_generator = batch.generator
if audio_generator is not None and audio_generator.device.type != latents.device.type:
if batch.seed is None:
audio_generator = torch.Generator(device=latents.device)
else:
audio_generator = torch.Generator(
device=latents.device).manual_seed(batch.seed)
audio_patch_shape = (
audio_shape.batch,
audio_shape.frames,
audio_shape.channels * audio_shape.mel_bins,
)
audio_latents_patch = torch.randn(
audio_patch_shape,
generator=audio_generator,
device=latents.device,
dtype=latents.dtype,
)
if hasattr(self.transformer, "audio_patchifier"):
audio_latents = self.transformer.audio_patchifier.unpatchify(
audio_latents_patch, audio_shape)
else:
audio_latents = audio_latents_patch.view(
audio_shape.batch,
audio_shape.frames,
audio_shape.channels,
audio_shape.mel_bins,
).permute(0, 2, 1, 3).contiguous()
audio_timestep_template = torch.ones(
(latents.shape[0], audio_shape.frames),
device=latents.device,
dtype=torch.float32,
)
logger.info(
"[LTX2] Denoising start: steps=%d dtype=%s guidance=%s "
"sigmas_shape=%s latents_shape=%s",
batch.num_inference_steps,
target_dtype,
batch.guidance_scale,
tuple(sigmas.shape),
tuple(latents.shape),
)
for step_index in tqdm(range(len(sigmas) - 1)):
sigma = sigmas[step_index]
sigma_next = sigmas[step_index + 1]
timestep = timestep_template * sigma
audio_timestep = (audio_timestep_template * sigma
if audio_timestep_template is not None else None)
with torch.autocast(
device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled,
), set_forward_context(
current_timestep=sigma.item(),
attn_metadata=None,
forward_batch=batch,
):
pos_outputs = self.transformer(
hidden_states=latents.to(target_dtype),
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_mask,
timestep=timestep,
audio_hidden_states=audio_latents,
audio_encoder_hidden_states=audio_context_p,
audio_timestep=audio_timestep,
)
if isinstance(pos_outputs, tuple):
pos_denoised, pos_audio = pos_outputs
else:
pos_denoised = pos_outputs
pos_audio = None
# Only run negative pass if CFG is enabled
if batch.do_classifier_free_guidance:
neg_outputs = self.transformer(
hidden_states=latents.to(target_dtype),
encoder_hidden_states=neg_prompt_embeds,
encoder_attention_mask=neg_prompt_mask,
timestep=timestep,
audio_hidden_states=audio_latents,
audio_encoder_hidden_states=audio_context_n,
audio_timestep=audio_timestep,
)
if isinstance(neg_outputs, tuple):
neg_denoised, neg_audio = neg_outputs
else:
neg_denoised = neg_outputs
neg_audio = None
pos_denoised = pos_denoised + (batch.guidance_scale - 1) * (
pos_denoised - neg_denoised)
if pos_audio is not None and neg_audio is not None:
pos_audio = pos_audio + (batch.guidance_scale -
1) * (pos_audio - neg_audio)
sigma_value = sigma.to(torch.float32) if isinstance(
sigma, torch.Tensor) else torch.tensor(
float(sigma),
device=latents.device,
dtype=torch.float32,
)
dt = sigma_next - sigma
velocity = ((latents.float() - pos_denoised.float()) /
sigma_value).to(latents.dtype)
latents = (latents.float() + velocity.float() * dt).to(
latents.dtype)
if pos_audio is not None and audio_latents is not None:
audio_velocity = ((audio_latents.float() - pos_audio.float()) /
sigma_value).to(audio_latents.dtype)
audio_latents = (audio_latents.float() +
audio_velocity.float() * dt).to(
audio_latents.dtype)
batch.latents = latents
batch.extra["ltx2_audio_latents"] = audio_latents
logger.info("[LTX2] Denoising done.")
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
return result
@@ -0,0 +1,189 @@
# SPDX-License-Identifier: Apache-2.0
"""
Latent preparation stage for LTX-2 pipelines.
"""
from pathlib import Path
import torch
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
class LTX2LatentPreparationStage(PipelineStage):
"""Prepare initial LTX-2 latents without relying on a diffusers scheduler."""
def __init__(self, transformer) -> None:
super().__init__()
self.transformer = transformer
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
latent_num_frames = self._adjust_video_length(batch, fastvideo_args)
if not batch.prompt_embeds:
batch_size = 1
elif isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
else:
batch_size = batch.prompt_embeds[0].shape[0]
batch_size *= batch.num_videos_per_prompt
if not batch.prompt_embeds:
transformer_dtype = next(self.transformer.parameters()).dtype
device = get_local_torch_device()
dummy_prompt = torch.zeros(
batch_size,
0,
self.transformer.hidden_size,
device=device,
dtype=transformer_dtype,
)
batch.prompt_embeds = [dummy_prompt]
batch.negative_prompt_embeds = []
batch.do_classifier_free_guidance = False
dtype = batch.prompt_embeds[0].dtype
device = get_local_torch_device()
generator = batch.generator
latents = batch.latents
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
height = batch.height
width = batch.width
latent_path = fastvideo_args.ltx2_initial_latent_path
if height is None or width is None:
raise ValueError("Height and width must be provided")
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
if height % spatial_ratio != 0 or width % spatial_ratio != 0:
raise ValueError(
f"Height and width must be divisible by {spatial_ratio} "
f"but are {height} and {width}.")
shape = (
batch_size,
self.transformer.num_channels_latents,
num_frames,
height // spatial_ratio,
width // spatial_ratio,
)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, "
f"but requested an effective batch size of {batch_size}.")
if latents is None:
if latent_path:
loaded_latents = self._load_initial_latent(
latent_path, device, dtype)
if loaded_latents is not None:
latents = loaded_latents
else:
latents = randn_tensor(
shape,
generator=generator,
device=device,
dtype=dtype,
)
self._save_initial_latent(latent_path, latents)
else:
latents = randn_tensor(
shape,
generator=generator,
device=device,
dtype=dtype,
)
else:
latents = latents.to(device)
batch.latents = latents
batch.raw_latent_shape = shape
return batch
def _adjust_video_length(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> int | None:
if not fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames:
return None
temporal_scale_factor = (fastvideo_args.pipeline_config.vae_config.
arch_config.temporal_compression_ratio)
video_length = batch.num_frames
return int((video_length - 1) // temporal_scale_factor + 1)
def _load_initial_latent(
self,
latent_path: str,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor | None:
path = Path(latent_path)
if not path.exists():
return None
payload = torch.load(path, map_location=device)
if isinstance(payload, dict):
if "video_latent" in payload:
latent = payload["video_latent"]
elif "latent" in payload:
latent = payload["latent"]
else:
latent = None
else:
latent = payload
if not torch.is_tensor(latent):
raise TypeError(f"Expected tensor for initial latent in {path}")
logger.info("[LTX2] Loaded initial latent from %s", path)
return latent.to(device=device, dtype=dtype)
def _save_initial_latent(self, latent_path: str,
latents: torch.Tensor) -> None:
path = Path(latent_path)
path.parent.mkdir(parents=True, exist_ok=True)
if path.exists():
return
torch.save({"video_latent": latents.detach().cpu()}, path)
logger.info("[LTX2] Saved initial latent to %s", path)
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check(
"prompt_or_embeds",
None,
lambda _: V.string_or_list_strings(batch.prompt) or not batch.
prompt_embeds or V.list_not_empty(batch.prompt_embeds),
)
if batch.prompt_embeds:
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_of_tensors)
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
V.positive_int)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("num_frames", batch.num_frames, V.positive_int)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("latents", batch.latents, V.none_or_tensor)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple)
return result
+15 -3
View File
@@ -11,14 +11,11 @@ from typing import Any
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
class TextEncodingStage(PipelineStage):
"""
@@ -39,6 +36,7 @@ class TextEncodingStage(PipelineStage):
super().__init__()
self.tokenizers = tokenizers
self.text_encoders = text_encoders
self._last_audio_embeds: list[torch.Tensor] | None = None
@torch.no_grad()
def forward(
@@ -70,6 +68,8 @@ class TextEncodingStage(PipelineStage):
encoder_index=all_indices,
return_attention_mask=True,
)
if self._last_audio_embeds is not None:
batch.extra["ltx2_audio_prompt_embeds"] = self._last_audio_embeds
for pe in prompt_embeds_list:
batch.prompt_embeds.append(pe)
@@ -86,6 +86,9 @@ class TextEncodingStage(PipelineStage):
encoder_index=all_indices,
return_attention_mask=True,
)
if self._last_audio_embeds is not None:
batch.extra[
"ltx2_audio_negative_embeds"] = self._last_audio_embeds
assert batch.negative_prompt_embeds is not None
for ne in neg_embeds_list:
@@ -184,10 +187,13 @@ class TextEncodingStage(PipelineStage):
embeds_list: list[torch.Tensor] = []
attn_masks_list: list[torch.Tensor] = []
audio_embeds_list: list[torch.Tensor] = []
preprocess_funcs = fastvideo_args.pipeline_config.preprocess_text_funcs
postprocess_funcs = fastvideo_args.pipeline_config.postprocess_text_funcs
encoder_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
is_ltx2 = getattr(fastvideo_args.pipeline_config.dit_config, "prefix",
"") == "ltx2"
if return_type not in ("list", "dict", "stack"):
raise ValueError(
@@ -259,6 +265,11 @@ class TextEncodingStage(PipelineStage):
except Exception:
prompt_embeds, attention_mask = postprocess_func(
outputs, attention_mask)
if is_ltx2 and getattr(outputs, "hidden_states", None):
audio_embed = outputs.hidden_states[0]
if dtype is not None:
audio_embed = audio_embed.to(dtype=dtype)
audio_embeds_list.append(audio_embed)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
@@ -266,6 +277,7 @@ class TextEncodingStage(PipelineStage):
if return_attention_mask:
attn_masks_list.append(attention_mask)
self._last_audio_embeds = audio_embeds_list if is_ltx2 else None
return self.return_embeds(embeds_list, attn_masks_list, return_type,
return_attention_mask, indices)
@@ -0,0 +1,153 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for SigLIP vision encoder using HYWorld model."""
import gc
import json
import os
import pytest
import torch
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import SiglipVisionModel as HFSiglipVisionModel
from fastvideo.configs.models.encoders import SiglipVisionConfig
from fastvideo.configs.models.encoders.siglip import SiglipVisionArchConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29505"
# HYWorld model path - SigLIP image encoder
MODEL_ID = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v"
IMAGE_ENCODER_SUBFOLDER = "image_encoder"
@pytest.mark.usefixtures("distributed_setup")
def test_siglip_encoder():
"""
Test compatibility between FastVideo SigLIP encoder and HuggingFace implementation.
The test verifies that both implementations:
- Load models with the same weights and parameters
- Produce nearly identical outputs for the same input
"""
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
logger.info("Loading SigLIP models from %s/%s", MODEL_ID, IMAGE_ENCODER_SUBFOLDER)
# Load HuggingFace implementation
hf_model = HFSiglipVisionModel.from_pretrained(
MODEL_ID, subfolder=IMAGE_ENCODER_SUBFOLDER
).to(torch.float16).to(device).eval()
# Load the config from Hugging Face and extract vision_config
config_path = hf_hub_download(
repo_id=MODEL_ID,
filename=f"{IMAGE_ENCODER_SUBFOLDER}/config.json"
)
with open(config_path) as f:
full_config = json.load(f)
# Get vision config from the full config
vision_config_dict = full_config.get("vision_config", full_config)
# Create FastVideo config with vision-specific settings
arch_config = SiglipVisionArchConfig(
hidden_size=vision_config_dict.get("hidden_size", 1152),
image_size=vision_config_dict.get("image_size", 384),
intermediate_size=vision_config_dict.get("intermediate_size", 4304),
num_attention_heads=vision_config_dict.get("num_attention_heads", 16),
num_hidden_layers=vision_config_dict.get("num_hidden_layers", 27),
patch_size=vision_config_dict.get("patch_size", 14),
)
config = SiglipVisionConfig(arch_config=arch_config)
# Create FastVideo model
from fastvideo.models.encoders.siglip import SiglipVisionModel
fv_model = SiglipVisionModel(config).to(torch.float16).to(device)
# Load weights from safetensors via Hugging Face
weights_path = hf_hub_download(
repo_id=MODEL_ID,
filename=f"{IMAGE_ENCODER_SUBFOLDER}/model.safetensors"
)
state_dict = load_file(weights_path)
# Filter to only vision_model weights (keep the vision_model. prefix)
vision_weights = [
(name, weight)
for name, weight in state_dict.items()
if name.startswith("vision_model.")
]
fv_model.load_weights(vision_weights)
fv_model.eval()
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
params1 = dict(hf_model.named_parameters())
params2 = dict(fv_model.named_parameters())
logger.info("HuggingFace model has %d parameters", len(params1))
logger.info("FastVideo model has %d parameters", len(params2))
# Compare non-stacked parameters
# Note: HF uses param names directly, FV adds "vision_model." prefix
for name1, param1 in sorted(params1.items()):
# Map HF param name to FV param name
name2 = "vision_model." + name1
skip = False
for param_name, weight_name, shard_id in fv_model.config.arch_config.stacked_params_mapping:
if weight_name in name1:
skip = True
break
# stacked params (qkv) are more troublesome to compare
if skip:
continue
if name2 in params2:
param2 = params2[name2]
param2 = param2.to_local().to(device) if isinstance(param2, DTensor) else param2.to(device)
assert_close(param1, param2, atol=1e-4, rtol=1e-4)
gc.collect()
torch.cuda.empty_cache()
# Test with sample images
batch_size = 2
image_size = fv_model.config.arch_config.image_size
# Create random pixel values
pixel_values = torch.randn(batch_size, 3, image_size, image_size).to(device).to(torch.float16)
logger.info("Testing SigLIP encoder with random pixel values of shape %s", pixel_values.shape)
with torch.no_grad():
# Get embeddings from HuggingFace implementation
hf_outputs = hf_model(pixel_values=pixel_values)
# Get embeddings from FastVideo implementation
with set_forward_context(current_timestep=0, attn_metadata=None):
fv_outputs = fv_model(pixel_values=pixel_values)
# Compare last hidden states
hf_hidden_state = hf_outputs.last_hidden_state
fv_hidden_state = fv_outputs.last_hidden_state
logger.info("HF hidden state shape: %s", hf_hidden_state.shape)
logger.info("FV hidden state shape: %s", fv_hidden_state.shape)
assert hf_hidden_state.shape == fv_hidden_state.shape, \
f"Hidden state shapes don't match: {hf_hidden_state.shape} vs {fv_hidden_state.shape}"
# Compare outputs with tolerance for numerical differences
assert_close(hf_hidden_state, fv_hidden_state, atol=1e-2, rtol=1e-3)
logger.info("SigLIP encoder test passed - outputs match within tolerance")
gc.collect()
torch.cuda.empty_cache()
+24 -10
View File
@@ -67,17 +67,23 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="H100:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(gpu="H100:1",
image=image,
timeout=1200,
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol})
def run_encoder_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol})
def run_vae_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol})
def run_transformer_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
@app.function(
gpu="L40S:4",
@@ -89,13 +95,21 @@ def run_transformer_tests():
def run_ssim_tests():
run_test("export HF_HOME='/root/data/.cache' && export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="L40S:4",
image=image,
timeout=900,
secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})],
volumes={"/root/data": model_vol})
def run_training_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
@app.function(gpu="L40S:2",
image=image,
timeout=900,
secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})],
volumes={"/root/data": model_vol})
def run_training_lora_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
@app.function(gpu="H100:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests_VSA():
@@ -80,9 +80,48 @@ WAN_I2V_PARAMS = {
"text-encoder-precision": ("fp32",)
}
# LTX-2 distilled one-stage params (no refine/upscale)
# Official defaults: height=512, width=768, num_frames=121, fps=24, seed=10
# Using num_frames=41 for faster CI (still valid: 41 = 8×5 + 1)
LTX2_T2V_PARAMS = {
"num_gpus": 2,
"model_path": "FastVideo/LTX2-Distilled-Diffusers",
"height": 512,
"width": 768,
"num_frames": 41, # Shorter for CI; official default is 121
"num_inference_steps": 8, # Distilled uses 8 steps
"guidance_scale": 1.0, # No CFG for distilled
"embedded_cfg_scale": 6,
"seed": 1024,
"sp_size": 2,
"tp_size": 1,
"fps": 24,
"neg_prompt": (
"blurry, out of focus, overexposed, underexposed, low contrast, washed out colors, "
"excessive noise, grainy texture, poor lighting, flickering, motion blur, distorted "
"proportions, unnatural skin tones, deformed facial features, asymmetrical face, "
"missing facial features, extra limbs, disfigured hands, wrong hand count, artifacts "
"around text, inconsistent perspective, camera shake, incorrect depth of field, "
"background too sharp, background clutter, distracting reflections, harsh shadows, "
"inconsistent lighting direction, color banding, cartoonish rendering, 3D CGI look, "
"unrealistic materials, uncanny valley effect, incorrect ethnicity, wrong gender, "
"exaggerated expressions, wrong gaze direction, mismatched lip sync, silent or muted "
"audio, distorted voice, robotic voice, echo, background noise, off-sync audio, "
"incorrect dialogue, added dialogue, repetitive speech, jittery movement, awkward "
"pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, "
"flat lighting, inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts."
),
"ltx2_vae_tiling": True,
"ltx2_vae_spatial_tile_size_in_pixels": 512,
"ltx2_vae_spatial_tile_overlap_in_pixels": 64,
"ltx2_vae_temporal_tile_size_in_frames": 64,
"ltx2_vae_temporal_tile_overlap_in_frames": 24,
}
MODEL_TO_PARAMS = {
"FastHunyuan-diffusers": HUNYUAN_PARAMS,
"Wan2.1-T2V-1.3B-Diffusers": WAN_T2V_PARAMS,
# "ltx2_diffusers": LTX2_T2V_PARAMS,
}
I2V_MODEL_TO_PARAMS = {
@@ -229,18 +268,26 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"flow_shift": BASE_PARAMS["flow_shift"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
"use_fsdp_inference": True,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
}
if "flow_shift" in BASE_PARAMS:
init_kwargs["flow_shift"] = BASE_PARAMS["flow_shift"]
if BASE_PARAMS.get("vae_sp"):
init_kwargs["vae_sp"] = True
init_kwargs["vae_tiling"] = True
if "text-encoder-precision" in BASE_PARAMS:
init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
# LTX2-specific VAE tiling parameters
if BASE_PARAMS.get("ltx2_vae_tiling"):
init_kwargs["ltx2_vae_tiling"] = True
init_kwargs["ltx2_vae_spatial_tile_size_in_pixels"] = BASE_PARAMS.get("ltx2_vae_spatial_tile_size_in_pixels", 512)
init_kwargs["ltx2_vae_spatial_tile_overlap_in_pixels"] = BASE_PARAMS.get("ltx2_vae_spatial_tile_overlap_in_pixels", 64)
init_kwargs["ltx2_vae_temporal_tile_size_in_frames"] = BASE_PARAMS.get("ltx2_vae_temporal_tile_size_in_frames", 64)
init_kwargs["ltx2_vae_temporal_tile_overlap_in_frames"] = BASE_PARAMS.get("ltx2_vae_temporal_tile_overlap_in_frames", 24)
generation_kwargs = {
"num_inference_steps": num_inference_steps,
@@ -121,9 +121,9 @@ def test_distributed_training():
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 3,
'grad_norm': 0.2,
'step_time': 2.5,
'avg_step_time': 5,
'grad_norm': 0.5,
'step_time': 5,
'train_loss': 0.04
}
@@ -0,0 +1,156 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pytest
import torch
from fastvideo.configs.models.dits import HYWorldConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.distributed.parallel_state import (
get_sp_parallel_rank,
get_sp_world_size,
)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
MODEL_PATH = maybe_download_model("FastVideo/HY-WorldPlay-Bidirectional-Diffusers")
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
REFERENCE_LATENT = -197132.85557549074 # Pre-computed reference value
@pytest.mark.usefixtures("distributed_setup")
def test_hyworld_transformer():
transformer_path = TRANSFORMER_PATH
sp_rank = get_sp_parallel_rank()
sp_world_size = get_sp_world_size()
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=False,
use_fsdp_inference=False,
pipeline_config=PipelineConfig(
dit_config=HYWorldConfig(), dit_precision=precision_str
),
)
args.device = device
loader = TransformerLoader()
model = loader.load(transformer_path, args).to(device, dtype=precision)
model.eval()
total_params = sum(p.numel() for p in model.parameters())
weight_sum = sum(p.to(torch.float64).sum().item() for p in model.parameters())
weight_mean = weight_sum / total_params
logger.info("Total parameters: %s", total_params)
logger.info("Weight sum: %s", weight_sum)
logger.info("Weight mean: %s", weight_mean)
torch.manual_seed(42)
# Create inputs for the model
batch_size = 1
seq_len = 30
seq_len_2 = 70
num_frames = 16
latent_height = 120
latent_width = 30
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(
batch_size, 65, num_frames, latent_height, latent_width,
device=device, dtype=precision
)
if sp_world_size > 1:
chunk_per_rank = hidden_states.shape[2] // sp_world_size
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) * chunk_per_rank]
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(
batch_size, seq_len + 1, 3584, device=device, dtype=precision
)
# Create attention mask for encoder_hidden_states
encoder_attention_mask = torch.ones(
batch_size, seq_len + 1, device=device, dtype=torch.bool
)
encoder_hidden_states[:, 15:] = 0
encoder_attention_mask[:, 15:] = False
encoder_hidden_states_2 = torch.randn(
batch_size, seq_len_2 + 1, 1472, device=device, dtype=precision
)
encoder_attention_mask_2 = torch.ones(
batch_size, seq_len_2 + 1, device=device, dtype=torch.bool
)
encoder_hidden_states_2[:, 39:] = 0
encoder_attention_mask_2[:, 39:] = False
# Image embeddings
image_embeds = torch.zeros(
batch_size, 729, 1152, dtype=precision, device=device
)
# Action tensor [B*T] - discrete action per frame, first frame of each batch is 0
action = torch.randint(0, 10, (batch_size * num_frames,), device=device)
action[::num_frames] = 0 # First frame of each batch is 0
action = action.to(dtype=precision)
# Camera view matrices [B, T, 4, 4] - 4x4 extrinsic/view matrices
viewmats = torch.eye(4, dtype=precision, device=device).unsqueeze(0).unsqueeze(0).expand(
batch_size, num_frames, -1, -1
).contiguous()
# Camera intrinsics [B, T, 3, 3] - 3x3 intrinsic matrices
Ks = torch.eye(3, dtype=precision, device=device).unsqueeze(0).unsqueeze(0).expand(
batch_size, num_frames, -1, -1
).contiguous()
# Timestep [B*T] - one timestep value per frame
timestep = torch.full((batch_size * num_frames,), 500, device=device, dtype=precision)
# Timestep for text [B] - one timestep value per batch
timestep_txt = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(data_type="dummy")
with torch.no_grad():
with torch.amp.autocast("cuda", dtype=precision):
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output = model(
hidden_states=hidden_states,
encoder_hidden_states=[encoder_hidden_states, encoder_hidden_states_2],
encoder_attention_mask=[encoder_attention_mask, encoder_attention_mask_2],
encoder_hidden_states_image=[image_embeds],
timestep=timestep,
timestep_txt=timestep_txt,
action=action,
viewmats=viewmats,
Ks=Ks,
)
latent = output.double().sum().item()
diff = abs(REFERENCE_LATENT - latent)
relative_diff = diff / abs(REFERENCE_LATENT)
logger.info(f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}")
logger.info(f"Absolute diff: {diff}, Relative diff: {relative_diff * 100:.4f}%")
# Allow 0.5% relative difference
assert relative_diff < 0.005, f"Output latents differ significantly: relative diff = {relative_diff * 100:.4f}% (max allowed: 0.5%)"
+76 -48
View File
@@ -28,8 +28,6 @@ from fastvideo.distributed import (cleanup_dist_env_and_memory,
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
@@ -42,6 +40,7 @@ from fastvideo.training.training_utils import (
shift_timestep)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
from fastvideo.optim.muon import get_muon_optimizer
vsa_available = is_vsa_available()
@@ -85,7 +84,6 @@ class DistillationPipeline(TrainingPipeline):
if training_args.real_score_model_path:
logger.info("Loading real score transformer from: %s",
training_args.real_score_model_path)
training_args.override_transformer_cls_name = "WanTransformer3DModel"
# TODO(will): can use deepcopy instead if the model is the same
self.real_score_transformer = self.load_module_from_path(
training_args.real_score_model_path, "transformer",
@@ -111,7 +109,6 @@ class DistillationPipeline(TrainingPipeline):
if training_args.fake_score_model_path:
logger.info("Loading fake score transformer from: %s",
training_args.fake_score_model_path)
training_args.override_transformer_cls_name = "WanTransformer3DModel"
self.fake_score_transformer = self.load_module_from_path(
training_args.fake_score_model_path, "transformer",
training_args)
@@ -146,8 +143,8 @@ class DistillationPipeline(TrainingPipeline):
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
# self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
# shift=self.timestep_shift)
if self.training_args.boundary_ratio is not None:
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
@@ -192,16 +189,24 @@ class DistillationPipeline(TrainingPipeline):
if fake_score_lr == 0.0:
fake_score_lr = training_args.learning_rate
betas_str = training_args.fake_score_betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
betas_str = training_args.fake_score_betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.fake_score_optimizer = get_muon_optimizer(
self.fake_score_transformer,
lr=fake_score_lr,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
self.fake_score_lr_scheduler = get_scheduler(
training_args.fake_score_lr_scheduler,
@@ -218,13 +223,23 @@ class DistillationPipeline(TrainingPipeline):
fake_score_params_2 = list(
filter(lambda p: p.requires_grad,
self.fake_score_transformer_2.parameters()))
self.fake_score_optimizer_2 = torch.optim.AdamW(
fake_score_params_2,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
self.fake_score_optimizer_2 = torch.optim.AdamW(
fake_score_params_2,
lr=fake_score_lr,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.fake_score_optimizer_2 = get_muon_optimizer(
self.fake_score_transformer_2,
lr=fake_score_lr,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
self.fake_score_lr_scheduler_2 = get_scheduler(
training_args.fake_score_lr_scheduler,
optimizer=self.fake_score_optimizer_2,
@@ -272,21 +287,6 @@ class DistillationPipeline(TrainingPipeline):
self.generator_ema: EMA_FSDP | None = None
self.generator_ema_2: EMA_FSDP | None = None
if (self.training_args.ema_decay
is not None) and (self.training_args.ema_decay > 0.0):
self.generator_ema = EMA_FSDP(self.transformer,
decay=self.training_args.ema_decay)
logger.info("Initialized generator EMA with decay=%s",
self.training_args.ema_decay)
# Initialize EMA for transformer_2 if it exists
if self.transformer_2 is not None:
self.generator_ema_2 = EMA_FSDP(
self.transformer_2, decay=self.training_args.ema_decay)
logger.info("Initialized generator EMA_2 with decay=%s",
self.training_args.ema_decay)
else:
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
def load_module_from_path(self, model_path: str, module_type: str,
training_args: "TrainingArgs"):
@@ -572,6 +572,7 @@ class DistillationPipeline(TrainingPipeline):
training_batch.input_kwargs = {
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_hidden_states_image": training_batch.image_embeds,
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep,
"return_dict": False,
@@ -723,10 +724,18 @@ class DistillationPipeline(TrainingPipeline):
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).detach().unflatten(0,
(1, generator_pred_video.shape[1]))
noisy_latent_copy = noisy_latent.clone()
if training_batch.video_latent is not None:
noisy_latent_copy = torch.cat([
noisy_latent_copy,
training_batch.video_latent,
torch.zeros_like(noisy_latent_copy),
],
dim=2)
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
noisy_latent_copy, timestep, training_batch.conditional_dict,
training_batch)
current_fake_score_transformer = self._get_fake_score_transformer(
timestep)
@@ -742,7 +751,7 @@ class DistillationPipeline(TrainingPipeline):
# real_score_transformer cond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
noisy_latent_copy, timestep, training_batch.conditional_dict,
training_batch)
current_real_score_transformer = self._get_real_score_transformer(
timestep)
@@ -758,7 +767,7 @@ class DistillationPipeline(TrainingPipeline):
# real_score_transformer uncond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.unconditional_dict,
noisy_latent_copy, timestep, training_batch.unconditional_dict,
training_batch)
# Use same transformer as conditional forward for consistency
real_score_pred_noise_uncond = current_real_score_transformer(
@@ -779,9 +788,16 @@ class DistillationPipeline(TrainingPipeline):
original_latent - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
original_latent.float(),
(original_latent.float() - grad.float()).detach())
if self.training_args.use_context_forcing and training_batch.trajectory_latents is not None:
context_forcing_length = training_batch.trajectory_latents.shape[1]
dmd_loss = 0.5 * F.mse_loss(
original_latent.float()[:, context_forcing_length:],
(original_latent.float()[:, context_forcing_length:] -
grad.float()[:, context_forcing_length:]).detach())
else:
dmd_loss = 0.5 * F.mse_loss(
original_latent.float(),
(original_latent.float() - grad.float()).detach())
training_batch.dmd_latent_vis_dict.update({
"training_batch_dmd_fwd_clean_latent":
@@ -831,6 +847,13 @@ class DistillationPipeline(TrainingPipeline):
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
fake_score_timestep).unflatten(0,
(1, generator_pred_video.shape[1]))
if training_batch.video_latent is not None:
noisy_generator_pred_video = torch.cat([
noisy_generator_pred_video,
training_batch.video_latent,
torch.zeros_like(noisy_generator_pred_video),
],
dim=2)
with set_forward_context(current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
@@ -877,7 +900,7 @@ class DistillationPipeline(TrainingPipeline):
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
super()._prepare_dit_inputs(training_batch)
# super()._prepare_dit_inputs(training_batch)
conditional_dict = {
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
@@ -898,7 +921,7 @@ class DistillationPipeline(TrainingPipeline):
self.video_latent_shape = training_batch.latents.shape
self.video_latent_shape_sp = training_batch.latents.shape
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
@@ -1239,6 +1262,7 @@ class DistillationPipeline(TrainingPipeline):
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch, steps)
batch.prompt_attention_mask = []
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
@@ -1249,8 +1273,11 @@ class DistillationPipeline(TrainingPipeline):
)
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
batch_negative, training_args)
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
if len(result_batch.prompt_embeds) == 1:
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
else:
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds, result_batch.prompt_attention_mask
logger.info(
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
@@ -1354,6 +1381,7 @@ class DistillationPipeline(TrainingPipeline):
if hasattr(self, 'transformer_2') and self.transformer_2 is not None:
self.transformer_2.train()
gc.collect()
torch.cuda.empty_cache()
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
File diff suppressed because it is too large Load Diff
+222 -53
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import sys
from copy import deepcopy
from typing import Any, cast
@@ -13,14 +14,15 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
# from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
# WanCausalDMDPipeline)
from fastvideo.pipelines.basic.hunyuan15.hunyuan15_causal_dmd_pipeline import Hy15CausalDMDPipeline
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.decoding import DecodingStage
logger = init_logger(__name__)
@@ -35,16 +37,18 @@ class ODEInitTrainingPipeline(TrainingPipeline):
- minimizing MSE to the stored next latent at timestep t_next
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
_required_config_modules = [
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Match the preprocess/generation scheduler for consistent stepping
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
training=True)
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# # Match the preprocess/generation scheduler for consistent stepping
# self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
# shift=fastvideo_args.pipeline_config.flow_shift,
# sigma_min=0.0,
# extra_one_step=True)
# self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
# training=True)
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
@@ -56,18 +60,36 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.text_encoder = self.get_module("text_encoder")
self.text_encoder.requires_grad_(False)
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer")
],
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
# self.noise_scheduler = SelfForcingFlowMatchScheduler(
# shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
extra_one_step=True,
device=get_local_torch_device())
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
dtype=torch.long,
device=get_local_torch_device())
self.dmd_denoising_steps = torch.tensor(
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
torch.tensor([0],
@@ -78,7 +100,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
logger.info("warped self.dmd_denoising_steps: %s",
self.dmd_denoising_steps)
else:
raise ValueError("warp_denoising_step must be true")
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32))).cuda()
self.dmd_denoising_steps = timesteps[self.dmd_denoising_steps]
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
get_local_torch_device())
@@ -98,7 +123,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
# Warm start validation with current transformer
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
self.validation_pipeline = Hy15CausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
@@ -109,7 +134,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
dit_cpu_offload=False)
def _get_next_batch(
self,
@@ -117,15 +142,40 @@ class ODEInitTrainingPipeline(TrainingPipeline):
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
device = get_local_torch_device()
encoder_hidden_states = batch['text_embedding'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_hidden_states_2 = batch['text_embedding_2'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_attention_mask = batch['text_mask'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_attention_mask_2 = batch['text_mask_2'].to(
device, dtype=torch.bfloat16).squeeze(0)
encoder_hidden_states_image = [
torch.zeros(1,
729,
1152,
device=get_local_torch_device(),
dtype=torch.bfloat16)
]
infos = batch['info_list']
if encoder_hidden_states.dim() < 3:
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
infos[0]["caption"],
self.training_args,
encoder_index=[0],
return_attention_mask=True,
)
encoder_hidden_states = prompt_embeds_list[0].to(
device, dtype=torch.bfloat16)
encoder_attention_mask = prompt_masks_list[0].to(
device, dtype=torch.bfloat16)
# Trajectory tensors may include a leading singleton batch dim per row
trajectory_latents = batch['trajectory_latents']
if trajectory_latents.dim() == 7:
@@ -154,11 +204,13 @@ class ODEInitTrainingPipeline(TrainingPipeline):
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
# Move to device
device = get_local_torch_device()
training_batch.encoder_hidden_states = encoder_hidden_states.to(
device, dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
device, dtype=torch.bfloat16)
training_batch.encoder_hidden_states = [
encoder_hidden_states, encoder_hidden_states_2
]
training_batch.encoder_attention_mask = [
encoder_attention_mask, encoder_attention_mask_2
]
training_batch.encoder_hidden_states_image = encoder_hidden_states_image
training_batch.infos = infos
return training_batch, trajectory_latents.to(
@@ -197,6 +249,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
dtype=torch.long).repeat(1, num_frame)
return timestep
else:
num_pad_frames = 0
if num_frame % num_frame_per_block != 0:
# Pad num_frame to be divisible by num_frame_per_block
num_pad_frames = num_frame_per_block - (num_frame %
num_frame_per_block)
num_frame += num_pad_frames
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, num_frame],
device=self.device,
@@ -207,12 +265,15 @@ class ODEInitTrainingPipeline(TrainingPipeline):
num_frame_per_block)
timestep[:, :, 1:] = timestep[:, :, 0:1]
timestep = timestep.reshape(timestep.shape[0], -1)
if num_pad_frames > 0:
timestep = timestep[:, num_pad_frames:]
return timestep
def _step_predict_next_latent(
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor
encoder_attention_mask: torch.Tensor,
encoder_hidden_states_image: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
torch.Tensor]]:
latent_vis_dict: dict[str, torch.Tensor] = {}
@@ -225,7 +286,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36], dtype=torch.long).cpu()
[0, 12, 24, 36, 50], dtype=torch.long).cpu()
# [0, 1, 2, 3], dtype=torch.long).cpu()
logger.info("self._cached_closest_idx_per_dmd: %s",
self._cached_closest_idx_per_dmd)
@@ -241,6 +302,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
traj_latents,
dim=1,
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
# relevant_traj_latents = traj_latents
logger.info("relevant_traj_latents: %s", relevant_traj_latents.shape)
# assert relevant_traj_latents.shape[0] == 1
@@ -251,51 +313,149 @@ class ODEInitTrainingPipeline(TrainingPipeline):
num_frames,
3,
uniform_timestep=False)
logger.info("indexes: %s", indexes.shape)
logger.info("indexes: %s", indexes)
# noisy_input = relevant_traj_latents[indexes]
noisy_input = torch.gather(
latents = torch.gather(
relevant_traj_latents,
dim=1,
index=indexes.reshape(B, 1, num_frames, 1, 1,
1).expand(-1, -1, -1, num_channels, height,
width).to(self.device)).squeeze(1)
noisy_input = torch.cat([
latents,
torch.zeros_like(latents),
torch.zeros_like(latents[:, :, 0:1])
],
dim=2)
timestep = self.dmd_denoising_steps[indexes]
logger.info("selected timestep for rank %s: %s",
self.global_rank,
timestep,
local_main_process_only=False)
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.permute(
latent_vis_dict["noisy_input"] = latents.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
4).detach().clone().cpu()
input_kwargs = {
logger.info("timestep: %s", timestep)
txt_input_kwargs = {
"txt_inference":
True,
"vision_inference":
False,
"encoder_hidden_states":
encoder_hidden_states,
"encoder_hidden_states_image":
encoder_hidden_states_image,
"encoder_attention_mask":
encoder_attention_mask,
"timestep":
torch.zeros([latents.shape[0]],
device=latents.device,
dtype=torch.bfloat16),
"cache_txt":
True,
}
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
txt_kv_cache = self.transformer(**txt_input_kwargs)
vision_input_kwargs = {
"txt_inference": False,
"vision_inference": True,
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep.to(device, dtype=torch.bfloat16),
"return_dict": False,
"txt_kv_cache": txt_kv_cache,
}
# Predict noise and step the scheduler to obtain next latent
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
noise_pred = self.transformer(**vision_input_kwargs).permute(
0, 2, 1, 3, 4)
from fastvideo.models.utils import pred_noise_to_pred_video
pred_video = pred_noise_to_pred_video(
pred_noise=noise_pred.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
noise_input_latent=latents.flatten(0, 1),
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
scheduler=self.modules["scheduler"]).unflatten(
0, noise_pred.shape[:2])
scheduler=self.noise_scheduler).unflatten(0, noise_pred.shape[:2])
latent_vis_dict["pred_video"] = pred_video.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
return pred_video, target_latent, timestep, latent_vis_dict
# def _step_predict_next_latent(
# self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
# encoder_hidden_states: torch.Tensor,
# encoder_attention_mask: torch.Tensor,
# encoder_hidden_states_image: torch.Tensor
# ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
# torch.Tensor]]:
# latent_vis_dict: dict[str, torch.Tensor] = {}
# device = get_local_torch_device()
# target_latent = traj_latents[:, -1]
# del traj_latents
# # Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
# B, num_frames, num_channels, height, width = target_latent.shape
# indexes = self._get_timestep( # [B, num_frames]
# 0,
# 1000,
# B,
# num_frames,
# 3,
# uniform_timestep=False)
# timestep = self.noise_scheduler.timesteps[indexes.cpu()].to(device)
# latents = self.noise_scheduler.add_noise(target_latent.flatten(0, 1), torch.randn_like(target_latent.flatten(0, 1)), timestep.flatten(0, 1)).unflatten(0, (B, num_frames))
# noisy_input = torch.cat([latents, torch.zeros_like(latents), torch.zeros_like(latents[:, :, 0:1])], dim=2)
# # Prepare inputs for transformer
# latent_vis_dict["noisy_input"] = latents.permute(
# 0, 2, 1, 3, 4).detach().clone().cpu()
# latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
# 4).detach().clone().cpu()
# logger.info("timestep: %s", timestep)
# txt_input_kwargs = {
# "txt_inference": True,
# "vision_inference": False,
# "encoder_hidden_states": encoder_hidden_states,
# "encoder_hidden_states_image": encoder_hidden_states_image,
# "encoder_attention_mask": encoder_attention_mask,
# "timestep": torch.zeros([latents.shape[0]], device=latents.device, dtype=torch.bfloat16),
# "cache_txt": True,
# }
# with set_forward_context(current_timestep=timestep,
# attn_metadata=None,
# forward_batch=None):
# txt_kv_cache = self.transformer(**txt_input_kwargs)
# vision_input_kwargs = {
# "txt_inference": False,
# "vision_inference": True,
# "hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
# "timestep": timestep.to(device, dtype=torch.bfloat16),
# "txt_kv_cache": txt_kv_cache,
# }
# # Predict noise and step the scheduler to obtain next latent
# with set_forward_context(current_timestep=timestep,
# attn_metadata=None,
# forward_batch=None):
# noise_pred = self.transformer(**vision_input_kwargs).permute(0, 2, 1, 3, 4)
# from fastvideo.models.utils import pred_noise_to_pred_video
# pred_video = pred_noise_to_pred_video(
# pred_noise=noise_pred.flatten(0, 1),
# noise_input_latent=latents.flatten(0, 1),
# timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
# scheduler=self.noise_scheduler).unflatten(
# 0, noise_pred.shape[:2])
# latent_vis_dict["pred_video"] = pred_video.permute(
# 0, 2, 1, 3, 4).detach().clone().cpu()
# return pred_video, target_latent, timestep, latent_vis_dict
def train_one_step(self, training_batch): # type: ignore[override]
self.transformer.train()
self.optimizer.zero_grad()
@@ -307,8 +467,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
for _ in range(args.gradient_accumulation_steps):
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
training_batch)
# traj_latents = traj_latents[:, :, :21]
text_embeds = training_batch.encoder_hidden_states
text_attention_mask = training_batch.encoder_attention_mask
image_embeds = training_batch.encoder_hidden_states_image
assert traj_latents.shape[0] == 1
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
@@ -318,7 +480,8 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
traj_latents, traj_timesteps, text_embeds, text_attention_mask,
image_embeds)
training_batch.latent_vis_dict.update(latent_vis_dict)
@@ -356,6 +519,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
except Exception:
grad_value = 0.0
training_batch.grad_norm = grad_value
if training_batch.current_timestep % 10 == 0:
gc.collect()
torch.cuda.empty_cache()
return training_batch
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
@@ -367,14 +534,13 @@ class ODEInitTrainingPipeline(TrainingPipeline):
assert latent_key in latents_vis_dict and latents_vis_dict[
latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
pixel_latent = self.decoding_stage.decode(latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
video_artifact = self.tracker.video(
video, fps=16, format="mp4") # change to 16 for Wan2.1
video, fps=24, format="mp4") # change to 16 for Wan2.1
if video_artifact is not None:
tracker_loss_dict[latent_key] = video_artifact
# Clean up references
@@ -383,6 +549,9 @@ class ODEInitTrainingPipeline(TrainingPipeline):
if self.global_rank == 0 and tracker_loss_dict:
self.tracker.log_artifacts(tracker_loss_dict, step)
gc.collect()
torch.cuda.empty_cache()
def main(args) -> None:
logger.info("Starting ODE-init training pipeline...")

Some files were not shown because too many files have changed in this diff Show More