Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
38f9c41e46 | ||
|
|
b1f88ad5d5 | ||
|
|
90a598bd9e | ||
|
|
90d86d5a79 | ||
|
|
7d373cd2c4 | ||
|
|
0806218156 | ||
|
|
8c6056fbe2 | ||
|
|
e4705349d0 | ||
|
|
50e63840c6 | ||
|
|
f1ec0cde18 | ||
|
|
1b503554d1 | ||
|
|
351ceb7c59 | ||
|
|
10875e0d7b | ||
|
|
1eaae8a10b | ||
|
|
59e00f6164 | ||
|
|
745cc05b10 | ||
|
|
c5dc244871 | ||
|
|
dbf3917bf4 | ||
|
|
050f189c95 | ||
|
|
029216029f |
@@ -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[@]}"
|
||||
@@ -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__":
|
||||
|
||||
@@ -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()
|
||||
@@ -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[@]}"
|
||||
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -2,5 +2,16 @@ from fastvideo.configs.models.base import ModelConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig)
|
||||
|
||||
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
__all__ = [
|
||||
"ModelConfig",
|
||||
"VAEConfig",
|
||||
"DiTConfig",
|
||||
"EncoderConfig",
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.configs.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 audio VAE and vocoder configuration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioArchConfig(ArchConfig):
|
||||
architectures: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioEncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioEncoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2AudioDecoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2AudioDecoder"]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VocoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
|
||||
architectures=["LTX2Vocoder"]))
|
||||
@@ -3,11 +3,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
|
||||
|
||||
@@ -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"
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 VAE configuration.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VAEArchConfig(VAEArchConfig):
|
||||
# Mirrors LTX-2 safetensors metadata config under "vae"
|
||||
_class_name: str = "CausalVideoAutoencoder"
|
||||
dims: int = 3
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 128
|
||||
encoder_blocks: list = field(default_factory=list)
|
||||
decoder_blocks: list = field(default_factory=list)
|
||||
patch_size: int = 4
|
||||
norm_layer: str = "pixel_norm"
|
||||
latent_log_var: str = "uniform"
|
||||
encoder_spatial_padding_mode: str = "zeros"
|
||||
decoder_spatial_padding_mode: str = "reflect"
|
||||
causal_decoder: bool = False
|
||||
timestep_conditioning: bool = True
|
||||
use_quant_conv: bool = False
|
||||
scaling_factor: float = 1.0
|
||||
normalize_latent_channels: bool = False
|
||||
|
||||
# Match FastVideo naming for compression ratios (LTX-2 default)
|
||||
temporal_compression_ratio: int = 8
|
||||
spatial_compression_ratio: int = 32
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2VAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
|
||||
|
||||
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
|
||||
ltx2_spatial_tile_size_in_pixels: int = 512
|
||||
ltx2_spatial_tile_overlap_in_pixels: int = 64
|
||||
ltx2_temporal_tile_size_in_frames: int = 64
|
||||
ltx2_temporal_tile_overlap_in_frames: int = 24
|
||||
@@ -4,6 +4,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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 = ""
|
||||
@@ -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 = ""
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
@@ -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()),
|
||||
])
|
||||
])
|
||||
@@ -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"]
|
||||
@@ -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,
|
||||
|
||||
@@ -18,6 +18,8 @@ import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -389,6 +391,11 @@ class VideoGenerator:
|
||||
if batch.save_video:
|
||||
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
audio = output_batch.extra.get("audio")
|
||||
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
|
||||
if (audio is not None and audio_sample_rate is not None and
|
||||
not self._mux_audio(output_path, audio, audio_sample_rate)):
|
||||
logger.warning("Audio mux failed; saved video without audio.")
|
||||
|
||||
if batch.return_frames:
|
||||
return frames
|
||||
@@ -396,6 +403,7 @@ class VideoGenerator:
|
||||
return {
|
||||
"samples": samples,
|
||||
"frames": frames,
|
||||
"audio": output_batch.extra.get("audio"),
|
||||
"prompts": prompt,
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
@@ -405,6 +413,98 @@ class VideoGenerator:
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _mux_audio(
|
||||
video_path: str,
|
||||
audio: torch.Tensor | np.ndarray,
|
||||
sample_rate: int,
|
||||
) -> bool:
|
||||
"""Mux audio into video using PyAV."""
|
||||
try:
|
||||
import av
|
||||
except ImportError:
|
||||
logger.warning("PyAV not installed; cannot mux audio. "
|
||||
"Install with: pip install av")
|
||||
return False
|
||||
|
||||
if torch.is_tensor(audio):
|
||||
audio_np = audio.detach().cpu().float().numpy()
|
||||
else:
|
||||
audio_np = np.asarray(audio, dtype=np.float32)
|
||||
|
||||
if audio_np.ndim == 1:
|
||||
audio_np = audio_np[:, None]
|
||||
elif audio_np.ndim == 2:
|
||||
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
|
||||
audio_np = audio_np.T
|
||||
else:
|
||||
logger.warning("Unexpected audio shape %s; skipping mux.",
|
||||
audio_np.shape)
|
||||
return False
|
||||
|
||||
audio_np = np.clip(audio_np, -1.0, 1.0)
|
||||
audio_int16 = (audio_np * 32767.0).astype(np.int16)
|
||||
num_channels = audio_int16.shape[1]
|
||||
layout = "stereo" if num_channels == 2 else "mono"
|
||||
|
||||
try:
|
||||
import wave
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
out_path = os.path.join(tmpdir, "muxed.mp4")
|
||||
wav_path = os.path.join(tmpdir, "audio.wav")
|
||||
|
||||
# Write audio to WAV file
|
||||
with wave.open(wav_path, "wb") as wav_file:
|
||||
wav_file.setnchannels(num_channels)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(sample_rate)
|
||||
wav_file.writeframes(audio_int16.tobytes())
|
||||
|
||||
# Open input video and audio
|
||||
input_video = av.open(video_path)
|
||||
input_audio = av.open(wav_path)
|
||||
|
||||
# Create output with both streams
|
||||
output = av.open(out_path, mode="w")
|
||||
|
||||
# Add video stream (copy codec from input)
|
||||
in_video_stream = input_video.streams.video[0]
|
||||
out_video_stream = output.add_stream(
|
||||
codec_name=in_video_stream.codec_context.name,
|
||||
rate=in_video_stream.average_rate,
|
||||
)
|
||||
out_video_stream.width = in_video_stream.width
|
||||
out_video_stream.height = in_video_stream.height
|
||||
out_video_stream.pix_fmt = in_video_stream.pix_fmt
|
||||
|
||||
# Add audio stream (AAC)
|
||||
out_audio_stream = output.add_stream("aac", rate=sample_rate)
|
||||
out_audio_stream.layout = layout
|
||||
|
||||
# Remux video (decode and re-encode to be safe)
|
||||
for frame in input_video.decode(video=0):
|
||||
for packet in out_video_stream.encode(frame):
|
||||
output.mux(packet)
|
||||
for packet in out_video_stream.encode():
|
||||
output.mux(packet)
|
||||
|
||||
# Encode audio
|
||||
for frame in input_audio.decode(audio=0):
|
||||
frame.pts = None # Let encoder assign PTS
|
||||
for packet in out_audio_stream.encode(frame):
|
||||
output.mux(packet)
|
||||
for packet in out_audio_stream.encode():
|
||||
output.mux(packet)
|
||||
|
||||
input_video.close()
|
||||
input_audio.close()
|
||||
output.close()
|
||||
shutil.move(out_path, video_path)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning("Audio mux failed: %s", e)
|
||||
return False
|
||||
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None) -> None:
|
||||
|
||||
@@ -166,6 +166,14 @@ class FastVideoArgs:
|
||||
# Prompt text file for batch processing
|
||||
prompt_txt: str | None = None
|
||||
|
||||
# LTX-2 VAE tiling overrides
|
||||
ltx2_vae_tiling: bool | None = None
|
||||
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
|
||||
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
|
||||
ltx2_vae_temporal_tile_size_in_frames: int | None = None
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
|
||||
ltx2_initial_latent_path: str | None = None
|
||||
|
||||
# model paths for correct deallocation
|
||||
model_paths: dict[str, str] = field(default_factory=dict)
|
||||
model_loaded: dict[str, bool] = field(default_factory=lambda: {
|
||||
@@ -203,8 +211,44 @@ class FastVideoArgs:
|
||||
logger.error("Failed to load V-MoBA config from %s: %s",
|
||||
self.moba_config_path, e)
|
||||
raise
|
||||
self._apply_ltx2_vae_overrides()
|
||||
self.check_fastvideo_args()
|
||||
|
||||
def _apply_ltx2_vae_overrides(self) -> None:
|
||||
if self.pipeline_config is None:
|
||||
return
|
||||
vae_config = self.pipeline_config.vae_config
|
||||
has_any = any(value is not None for value in (
|
||||
self.ltx2_vae_spatial_tile_size_in_pixels,
|
||||
self.ltx2_vae_spatial_tile_overlap_in_pixels,
|
||||
self.ltx2_vae_temporal_tile_size_in_frames,
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames,
|
||||
))
|
||||
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
|
||||
"vae_tiling"):
|
||||
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
|
||||
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
|
||||
self.pipeline_config.vae_tiling = True
|
||||
|
||||
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
|
||||
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
|
||||
vae_config.ltx2_spatial_tile_size_in_pixels = (
|
||||
self.ltx2_vae_spatial_tile_size_in_pixels)
|
||||
if hasattr(
|
||||
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
|
||||
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
|
||||
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
|
||||
self.ltx2_vae_spatial_tile_overlap_in_pixels)
|
||||
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
|
||||
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
|
||||
vae_config.ltx2_temporal_tile_size_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_size_in_frames)
|
||||
if hasattr(
|
||||
vae_config, "ltx2_temporal_tile_overlap_in_frames"
|
||||
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
|
||||
vae_config.ltx2_temporal_tile_overlap_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames)
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
# Model and path configuration
|
||||
@@ -325,6 +369,44 @@ class FastVideoArgs:
|
||||
"Path to a text file containing prompts (one per line) for batch processing",
|
||||
)
|
||||
|
||||
# LTX-2 VAE tiling overrides
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-tiling",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.ltx2_vae_tiling,
|
||||
help="Enable LTX-2 VAE tiling overrides.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-spatial-tile-size-in-pixels",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
|
||||
help="LTX-2 VAE spatial tile size in pixels.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-spatial-tile-overlap-in-pixels",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
|
||||
help="LTX-2 VAE spatial tile overlap in pixels.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-temporal-tile-size-in-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
|
||||
help="LTX-2 VAE temporal tile size in frames.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-vae-temporal-tile-overlap-in-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
|
||||
help="LTX-2 VAE temporal tile overlap in frames.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-initial-latent-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_initial_latent_path,
|
||||
help="Path to load/save a precomputed LTX-2 initial latent.",
|
||||
)
|
||||
|
||||
# LoRA parameters (inference-time adapter loading)
|
||||
parser.add_argument(
|
||||
"--lora-path",
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
"""
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -67,17 +67,23 @@ def run_test(pytest_command: str):
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="H100:1",
|
||||
image=image,
|
||||
timeout=1200,
|
||||
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_encoder_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/encoders -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1200, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_vae_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/vaes -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_transformer_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
run_test("export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
|
||||
@app.function(
|
||||
gpu="L40S:4",
|
||||
@@ -89,13 +95,21 @@ def run_transformer_tests():
|
||||
def run_ssim_tests():
|
||||
run_test("export HF_HOME='/root/data/.cache' && export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:4",
|
||||
image=image,
|
||||
timeout=900,
|
||||
secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_training_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
|
||||
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
@app.function(gpu="L40S:2",
|
||||
image=image,
|
||||
timeout=900,
|
||||
secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_training_lora_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
|
||||
run_test("export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP")
|
||||
|
||||
@app.function(gpu="H100:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests_VSA():
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -80,9 +80,48 @@ WAN_I2V_PARAMS = {
|
||||
"text-encoder-precision": ("fp32",)
|
||||
}
|
||||
|
||||
# LTX-2 distilled one-stage params (no refine/upscale)
|
||||
# Official defaults: height=512, width=768, num_frames=121, fps=24, seed=10
|
||||
# Using num_frames=41 for faster CI (still valid: 41 = 8×5 + 1)
|
||||
LTX2_T2V_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
"model_path": "FastVideo/LTX2-Distilled-Diffusers",
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 41, # Shorter for CI; official default is 121
|
||||
"num_inference_steps": 8, # Distilled uses 8 steps
|
||||
"guidance_scale": 1.0, # No CFG for distilled
|
||||
"embedded_cfg_scale": 6,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 1,
|
||||
"fps": 24,
|
||||
"neg_prompt": (
|
||||
"blurry, out of focus, overexposed, underexposed, low contrast, washed out colors, "
|
||||
"excessive noise, grainy texture, poor lighting, flickering, motion blur, distorted "
|
||||
"proportions, unnatural skin tones, deformed facial features, asymmetrical face, "
|
||||
"missing facial features, extra limbs, disfigured hands, wrong hand count, artifacts "
|
||||
"around text, inconsistent perspective, camera shake, incorrect depth of field, "
|
||||
"background too sharp, background clutter, distracting reflections, harsh shadows, "
|
||||
"inconsistent lighting direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, wrong gender, "
|
||||
"exaggerated expressions, wrong gaze direction, mismatched lip sync, silent or muted "
|
||||
"audio, distorted voice, robotic voice, echo, background noise, off-sync audio, "
|
||||
"incorrect dialogue, added dialogue, repetitive speech, jittery movement, awkward "
|
||||
"pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, "
|
||||
"flat lighting, inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts."
|
||||
),
|
||||
"ltx2_vae_tiling": True,
|
||||
"ltx2_vae_spatial_tile_size_in_pixels": 512,
|
||||
"ltx2_vae_spatial_tile_overlap_in_pixels": 64,
|
||||
"ltx2_vae_temporal_tile_size_in_frames": 64,
|
||||
"ltx2_vae_temporal_tile_overlap_in_frames": 24,
|
||||
}
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"FastHunyuan-diffusers": HUNYUAN_PARAMS,
|
||||
"Wan2.1-T2V-1.3B-Diffusers": WAN_T2V_PARAMS,
|
||||
# "ltx2_diffusers": LTX2_T2V_PARAMS,
|
||||
}
|
||||
|
||||
I2V_MODEL_TO_PARAMS = {
|
||||
@@ -229,18 +268,26 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"flow_shift": BASE_PARAMS["flow_shift"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"use_fsdp_inference": True,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
}
|
||||
if "flow_shift" in BASE_PARAMS:
|
||||
init_kwargs["flow_shift"] = BASE_PARAMS["flow_shift"]
|
||||
if BASE_PARAMS.get("vae_sp"):
|
||||
init_kwargs["vae_sp"] = True
|
||||
init_kwargs["vae_tiling"] = True
|
||||
if "text-encoder-precision" in BASE_PARAMS:
|
||||
init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
|
||||
# LTX2-specific VAE tiling parameters
|
||||
if BASE_PARAMS.get("ltx2_vae_tiling"):
|
||||
init_kwargs["ltx2_vae_tiling"] = True
|
||||
init_kwargs["ltx2_vae_spatial_tile_size_in_pixels"] = BASE_PARAMS.get("ltx2_vae_spatial_tile_size_in_pixels", 512)
|
||||
init_kwargs["ltx2_vae_spatial_tile_overlap_in_pixels"] = BASE_PARAMS.get("ltx2_vae_spatial_tile_overlap_in_pixels", 64)
|
||||
init_kwargs["ltx2_vae_temporal_tile_size_in_frames"] = BASE_PARAMS.get("ltx2_vae_temporal_tile_size_in_frames", 64)
|
||||
init_kwargs["ltx2_vae_temporal_tile_overlap_in_frames"] = BASE_PARAMS.get("ltx2_vae_temporal_tile_overlap_in_frames", 24)
|
||||
|
||||
generation_kwargs = {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
|
||||
@@ -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%)"
|
||||
@@ -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
@@ -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
Reference in New Issue
Block a user