Compare commits

...
Author SHA1 Message Date
Wei Zhou 3c71218987 ckpt 2026-04-01 02:22:08 +00:00
JerryZhou54 6b4cc8772d launched synthetic video generation for hy1.5 2026-02-08 23:46:16 +00:00
JerryZhou54 0df9934f4d ckpt 2026-02-07 21:43:09 +00:00
JerryZhou54 8140cb4460 small fix 2026-01-28 19:10:28 +00:00
JerryZhou54 e42f688466 fix lint 2026-01-28 15:39:18 +00:00
JerryZhou54 4062953769 Modify distill scripts 2026-01-28 15:32:15 +00:00
JerryZhou54 8cebfdad67 Modify fastvideo/training 2026-01-28 15:32:15 +00:00
JerryZhou54 d5d061a1f7 revert hy15 preprocess related changes 2026-01-28 15:32:13 +00:00
JerryZhou54 07db3796d2 refactor hy15 sf distill 2026-01-28 15:30:21 +00:00
JerryZhou54 363d81f3bb refactor hy15 causal denoising 2026-01-28 15:29:51 +00:00
JerryZhou54 2e0c5fa979 fix lint 2026-01-28 15:29:48 +00:00
JerryZhou54 448490b838 compatible with wan sf 2026-01-28 15:29:21 +00:00
JerryZhou54 5546fa9aad Add context forcing, ode_init to sf training 2026-01-28 15:29:21 +00:00
JerryZhou54 9718cd974e small change 2026-01-28 15:29:21 +00:00
JerryZhou54 8a623d5aa7 ckpt 2026-01-28 15:29:18 +00:00
JerryZhou54 d481af1d36 Ode init running for hy15 2026-01-28 15:28:26 +00:00
JerryZhou54 eda4148681 Ode runnable for hy15 2026-01-28 15:27:23 +00:00
JerryZhou54 9cf2a1561a Add support for ode_init inference for hy15 & support multiple timesteps for hy15 2026-01-28 15:27:22 +00:00
79 changed files with 7848 additions and 410 deletions
@@ -0,0 +1,172 @@
#!/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_tf_init_3333/checkpoint-2800/transformer/diffusion_pytorch_model.safetensors
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_worldplay_df_init_1333/checkpoint-11600/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=(
--optimizer-type "muon"
--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-ode-init True
# --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[@]}"
@@ -0,0 +1,172 @@
#!/bin/bash
#SBATCH --job-name=distill_dmd_t2v_1.3b
#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=distill_dmd_t2v_1.3b_output/distill_dmd_t2v_1.3b_%j.out
#SBATCH --error=distill_dmd_t2v_1.3b_output/distill_dmd_t2v_1.3b_%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
export HF_HOME=/mnt/data/huggingface
export HF_DATASETS_CACHE="/tmp/hf_datasets_${SLURM_JOB_ID}_${SLURM_PROCID}"
mkdir -p "${HF_DATASETS_CACHE}"
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
source ~/miniconda3/bin/activate
conda activate fastvideo
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="/mnt/data/vidprom_filtered_extended_umt5_text_embed"
# DATA_DIR="/mnt/data/mixkit_wan_1.3b_processed_t2v_distributed"
DATA_DIR_2="/mnt/data/mixkit_wan_1.3b_processed_t2v_distributed"
VALIDATION_DATASET_FILE="/mnt/home/zhouw.jerry2017/FastVideo/examples/distill/SFWan2.1-T2V/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
--wandb_run_name causal_forcing_vidprom_extended
--output_dir /mnt/data/wan_cf_df_distill_1.3b
--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 832
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
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 # TODO: check if you can remove this in this script
--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 100
--validation_sampling_steps "4"
--validation_guidance_scale "1.0"
)
optimizer_args=(
--optimizer-type "adamw"
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--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
# --init_weights_from_safetensors "/mnt/data/wan_tf_ode_init_3333/checkpoint-1600/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/data/wan_tf_init_3333/checkpoint-800/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/data/Causal-Forcing-wan1.3b/causal_ode.safetensors"
# --init_weights_from_safetensors "/mnt/data/wan_ar_diffusion_3333_mixkit_shift5/checkpoint-7200/transformer/diffusion_pytorch_model.safetensors"
--init_weights_from_safetensors "/mnt/data/wan_ar_diffusion_3333_mixkit_shift5/checkpoint-4800/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/data/wan_tf_ode_init_3333_mixkit_shift5/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/data/wan_ar_df_diffusion_3333_mixkit_shift5/checkpoint-4800/transformer/diffusion_pytorch_model.safetensors"
)
dmd_args=(
--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.0
--fake_score_learning_rate 2e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
# --use-decoupled-dmd True
# --use-gan-loss True
)
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)
)
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/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
+1 -1
View File
@@ -15,7 +15,7 @@ def main():
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
text_encoder_cpu_offload=False,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
+14 -8
View File
@@ -10,12 +10,13 @@ def main():
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers",
# "IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
dit_cpu_offload=False, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -23,14 +24,19 @@ def main():
# image_encoder_cpu_offload=False,
)
prompt = "一位年轻女性穿着一件粉色的连衣裙,裙子上有白色的装饰和粉色的纽扣。她的头发是紫色的,头上戴着一个红色的大蝴蝶结,显得非常可爱和精致。她还戴着一个红色的领结,整体造型充满了少女感和活力。她的表情温柔,双手轻轻交叉放在身前,姿态优雅。背景是简单的灰色,没有任何多余的装饰,使得人物更加突出。她的妆容清淡自然,突显了她的清新气质。整体画面给人一种甜美、梦幻的感觉,仿佛置身于童话世界中。"
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# prompt = "一位年轻女性穿着一件粉色的连衣裙,裙子上有白色的装饰和粉色的纽扣。她的头发是紫色的,头上戴着一个红色的大蝴蝶结,显得非常可爱和精致。她还戴着一个红色的领结,整体造型充满了少女感和活力。她的表情温柔,双手轻轻交叉放在身前,姿态优雅。背景是简单的灰色,没有任何多余的装饰,使得人物更加突出。她的妆容清淡自然,突显了她的清新气质。整体画面给人一种甜美、梦幻的感觉,仿佛置身于童话世界中。"
# negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical."
# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code."
image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
video = generator.generate_video(prompt, negative_prompt=negative_prompt, image_path=image_path, video_path=control_video_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
# image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
# control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
import json, os
with open("/mnt/weka/home/hao.zhang/wei/FastVideo/data/mixkit_i2v.jsonl", "r") as f:
prompt_image_pairs = json.load(f)
for d in prompt_image_pairs:
prompt = d["prompt"]
image_path = os.path.join("/mnt/weka/home/hao.zhang/wei/FastVideo", d["image_path"])
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
if __name__ == "__main__":
main()
+14 -2
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v_720p_77"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -12,17 +12,29 @@ def main():
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
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,
)
import json, os
with open("/mnt/weka/home/hao.zhang/wei/FastVideo/data/mixkit_i2v_720p.jsonl", "r") as f:
prompt_image_pairs = json.load(f)
for d in prompt_image_pairs:
prompt = d["prompt"]
image_path = os.path.join("/mnt/weka/home/hao.zhang/wei/FastVideo", d["image_path"])
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_PATH, save_video=True, num_frames=77)
return
# I2V is triggered just by passing in an image_path argument
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
# prompt = "In the midst of the joyous New Year's Eve celebration, the cheerful group of friends, their spirits lifted by the festivities, decides to immortalize the moment with a vibrant snapshot"
# image_path = "images/mixkit-a-cheerful-group-of-friends-celebrate-new-years-eve-and-51525.png"
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path)
return
# Generate another video with a different prompt, without reloading the
# model!
@@ -0,0 +1,124 @@
#!/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='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v"
DATA_DIR="your_data_dir"
VALIDATION_DATASET_FILE="your_validation_data_dir"
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "hy15_ode_init"
--output_dir "your_output_dir"
--wandb_run_name "your_wandb_run_name"
--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 21
--num_height 480
--num_width 848
--num_frames 81
--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
)
# 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[@]}"
@@ -0,0 +1,133 @@
#!/bin/bash
#SBATCH --job-name=hy15_tf_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_tf_3333_output/hy15_tf_3333_%j.out
#SBATCH --error=hy15_tf_3333_output/hy15_tf_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_worldplay_tf_init_3333"
--wandb_run_name "hy15_worldplay_tf_init_3333"
--warp_denoising_step
--dmd_denoising_steps '1000,760,520,280,0'
--log_visualization
--visualization_steps 100
--max_train_steps 30000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 848
--num_frames 81
--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=(
--optimizer-type "muon"
--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/hy15_ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,137 @@
#!/bin/bash
#SBATCH --job-name=wan_tf_3333
#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=wan_tf_3333_output/wan_tf_3333_%j.out
#SBATCH --error=wan_tf_3333_output/wan_tf_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
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
source ~/miniconda3/bin/activate
conda activate fastvideo
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/data/mixkit_wan_1.3b_processed_t2v_distributed"
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 "wan_tf_init"
--output_dir "/mnt/data/wan_ar_df_diffusion_3333_mixkit_shift5"
--wandb_run_name "wan_df_init_3333"
--warp_denoising_step
--dmd_denoising_steps '1000,750,500,250,0'
--log_visualization
--visualization_steps 100
--max_train_steps 30000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
--override-transformer-cls-name "CausalWanTransformer3DModel"
)
# 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=(
--optimizer-type "adamw"
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 800
--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/ar_diffusion_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,138 @@
#!/bin/bash
#SBATCH --job-name=wan_tf_3333
#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=wan_tf_3333_output/wan_tf_3333_%j.out
#SBATCH --error=wan_tf_3333_output/wan_tf_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
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
source ~/miniconda3/bin/activate
conda activate fastvideo
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/data/mixkit_wan_1.3b_processed_t2v_distributed"
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 "wan_tf_init"
--output_dir "/mnt/data/wan_ar_diffusion_3333_mixkit_shift5"
--wandb_run_name "wan_tf_init_3333"
--warp_denoising_step
--dmd_denoising_steps '1000,750,500,250,0'
--log_visualization
--visualization_steps 100
--max_train_steps 30000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
--override-transformer-cls-name "CausalWanTransformer3DModel"
)
# 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=(
--optimizer-type "adamw"
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 800
--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"
--use-tf True
)
# 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/ar_diffusion_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,138 @@
#!/bin/bash
#SBATCH --job-name=wan_tf_ode_init_3333
#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=wan_tf_ode_init_3333_output/wan_tf_ode_init_3333_%j.out
#SBATCH --error=wan_tf_ode_init_3333_output/wan_tf_ode_init_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
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
source ~/miniconda3/bin/activate
conda activate fastvideo
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=64
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/data/fv-tf-ode-preprocessing-6k-mixkit-wan-1.3b/"
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 "wan_tf_ode_init"
--output_dir "/mnt/data/wan_tf_ode_init_3333_mixkit_shift5"
--wandb_run_name "wan_tf_ode_init_3333_mixkit_shift5"
--warp_denoising_step
--dmd_denoising_steps '1000,750,500,250'
--log_visualization
--visualization_steps 100
--max_train_steps 30000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
--enable_gradient_checkpointing_type "full"
--init_weights_from_safetensors "/mnt/data/wan_ar_diffusion_3333_mixkit_shift5/checkpoint-2400/transformer/diffusion_pytorch_model.safetensors"
)
# 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=(
--optimizer-type "adamw"
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--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[@]}"
@@ -0,0 +1 @@
The raw mixkit dataset is downloaded and preprocessed following instructions from https://github.com/tianweiy/CausVid/tree/master
@@ -0,0 +1,25 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="/mnt/data/mixkit/merge.txt"
OUTPUT_DIR="/mnt/data/mixkit_wan_1.3b_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--drop_short_ratio 0.0 \
--preprocess_task "t2v"
@@ -0,0 +1,66 @@
#!/bin/bash
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=preprocess_mixkit_wan_1.3b_t2v/preprocess_wan_data_t2v.log
#SBATCH --error=preprocess_mixkit_wan_1.3b_t2v/preprocess_wan_data_t2v.error
#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
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
source ~/miniconda3/bin/activate
conda activate fastvideo
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="/mnt/data/mixkit/merge.txt"
OUTPUT_DIR="/mnt/data/mixkit_wan_1.3b_processed_t2v/"
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $GPU_NUM \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--drop_short_ratio 0.0 \
--preprocess_task "t2v"
@@ -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
+3 -3
View File
@@ -87,10 +87,10 @@ class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: int | None = None
num_hidden_layers_override: int | None = 31
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = True
enable_scale: bool = False
is_causal: bool = False
prefix: str = "clip"
+17 -1
View File
@@ -127,13 +127,29 @@ class Hunyuan15T2V480PConfig(PipelineConfig):
vae_tiling: bool = True
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Hunyuan15DistilledI2V480PConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 7
@dataclass
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 9
@dataclass
class SelfForcingHunyuan15T2V480PConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 5
is_causal: bool = True
# dmd_denoising_steps: list[int] | None = field(
# default_factory=lambda: [1000, 750, 500, 250])
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 875, 750, 625, 500, 375, 250, 125])
warp_denoising_step: bool = True
+5 -1
View File
@@ -8,9 +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 (
@@ -35,6 +35,10 @@ 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":
+3 -1
View File
@@ -41,7 +41,7 @@ class WanT2V480PConfig(PipelineConfig):
vae_sp: bool = False
# Denoising stage
flow_shift: float | None = 3.0
flow_shift: float | None = 5.0
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
@@ -176,6 +176,8 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
flow_shift: float | None = 5.0
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, 680, 360, 180])
warp_denoising_step: bool = True
+10 -2
View File
@@ -17,8 +17,8 @@ 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 = ""
@@ -28,6 +28,14 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
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):
height: int = 720
+5 -1
View File
@@ -5,8 +5,8 @@ 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
@@ -44,6 +44,10 @@ logger = init_logger(__name__)
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"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":
+2 -1
View File
@@ -15,7 +15,8 @@ class WanT2V_1_3B_SamplingParam(SamplingParam):
# Denoising stage
guidance_scale: float = 3.0
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
num_inference_steps: int = 50
teacache_params: WanTeaCacheParams = field(
+2 -1
View File
@@ -10,7 +10,7 @@ from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
from fastvideo.dataset.validation_dataset import ValidationDataset
def getdataset(args) -> VideoCaptionMergedDataset:
def getdataset(args, start_idx: int = 0) -> VideoCaptionMergedDataset:
if args.do_temporal_sample:
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
else:
@@ -36,6 +36,7 @@ def getdataset(args) -> VideoCaptionMergedDataset:
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
start_idx=start_idx,
seed=args.seed)
+20 -1
View File
@@ -165,6 +165,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 +206,4 @@ def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
"text_embedding_dtype": str(text_embedding.dtype),
"caption": caption,
}
return record
return record
+10 -2
View File
@@ -90,6 +90,15 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# pa.field("text_embedding_2_bytes", pa.binary()),
# pa.field("text_embedding_2_shape", pa.list_(pa.int64())),
# pa.field("text_embedding_2_dtype", pa.string()),
# pa.field("text_mask_bytes", pa.binary()),
# pa.field("text_mask_shape", pa.list_(pa.int64())),
# pa.field("text_mask_dtype", pa.string()),
# pa.field("text_mask_2_bytes", pa.binary()),
# pa.field("text_mask_2_shape", pa.list_(pa.int64())),
# pa.field("text_mask_2_dtype", pa.string()),
# --- ODE Trajectory ---
pa.field("trajectory_latents_bytes", pa.binary()),
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
@@ -117,7 +126,6 @@ pyarrow_schema_text_only = pa.schema([
pa.field("caption", pa.string()),
])
pyarrow_schema_matrixgame = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
@@ -156,4 +164,4 @@ pyarrow_schema_matrixgame = pa.schema([
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
])
+19 -23
View File
@@ -17,6 +17,7 @@ from PIL import Image
from transformers import AutoTokenizer
from fastvideo.logger import init_logger
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -40,7 +41,6 @@ class PreprocessBatch:
num_frames: int | None = None
sample_frame_index: list[int] | None = None
sample_num_frames: int | None = None
action_path: str | None = None
# Processed data
pixel_values: torch.Tensor | None = None
@@ -470,18 +470,14 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
self.seed = seed
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
if os.path.exists(tokenizer_path):
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
else:
tokenizer = None
tokenizer_path = os.path.join(maybe_download_model(args.model_path), "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
# Initialize processing stages
self._init_stages(args, transform, transform_topcrop, tokenizer)
# Process metadata
self.processed_batches = self._process_metadata()
self.processed_batches = self._process_metadata(start_idx=start_idx)
def _init_stages(self, args, transform, transform_topcrop,
tokenizer) -> None:
@@ -508,7 +504,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
else:
self.text_encoding_stage = None
def _load_raw_data(self) -> list[dict]:
def _load_raw_data(self, start_idx: int = 0) -> list[dict]:
"""Load raw data from JSON files."""
# Read folder-annotation pairs
with open(self.data_merge_path) as f:
@@ -528,14 +524,12 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
if "action_path" in item and item["action_path"]:
item["action_path"] = opj(folder, item["action_path"])
return data_items
return data_items[start_idx:]
def _process_metadata(self) -> list[PreprocessBatch]:
def _process_metadata(self, start_idx: int = 0) -> list[PreprocessBatch]:
"""Process the raw metadata through all filtering stages."""
raw_data = self._load_raw_data()
raw_data = self._load_raw_data(start_idx=start_idx)
processed_batches = []
# Initialize counters
@@ -551,8 +545,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"),
action_path=item.get("action_path"))
duration=item.get("duration"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
@@ -616,10 +609,15 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
batch = self.image_transform_stage.process(batch)
if self.text_encoding_stage is not None:
batch = self.text_encoding_stage.process(batch)
else:
raise ValueError("Text encoding stage is not initialized")
# Build result dictionary
result = {
"pixel_values": batch.pixel_values,
# "text": batch.text,
# "input_ids": batch.input_ids,
# "cond_mask": batch.cond_mask,
"path": batch.path,
}
@@ -627,14 +625,12 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
result["text"] = batch.text
result["input_ids"] = batch.input_ids
result["cond_mask"] = batch.cond_mask
else:
raise ValueError("Text encoding stage is not initialized")
# Add video-specific fields
if batch.is_video:
result.update({"fps": batch.fps, "duration": batch.duration})
# Add action_path
if batch.action_path:
result["action_path"] = batch.action_path
return result
@@ -672,7 +668,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)
@@ -775,4 +771,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"]
+11 -1
View File
@@ -156,6 +156,14 @@ def collate_rows_from_parquet_schema(rows,
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
@@ -175,6 +183,8 @@ def collate_rows_from_parquet_schema(rows,
for tensor in tensor_list:
if tensor.numel() > 0:
if tensor.ndim == 3:
tensor = tensor.squeeze(0)
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
@@ -222,4 +232,4 @@ def collate_rows_from_parquet_schema(rows,
if info_list and 'caption' in info_list[0]:
batch_data['caption_text'] = [info['caption'] for info in info_list]
return batch_data
return batch_data
+32 -4
View File
@@ -2,6 +2,7 @@
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
import os
import pathlib
import json
import datasets
from torch.utils.data import IterableDataset
@@ -32,10 +33,7 @@ class ValidationDataset(IterableDataset):
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".json":
data = datasets.load_dataset("json",
data_files=self.filename.as_posix(),
split="train",
field="data")
data = self._load_json_without_hf_locking(self.filename)
elif self.filename.suffix == ".parquet":
data = datasets.load_dataset("parquet",
data_files=self.filename.as_posix(),
@@ -162,3 +160,33 @@ class ValidationDataset(IterableDataset):
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
@staticmethod
def _load_json_without_hf_locking(filename: pathlib.Path) -> list[dict]:
"""Load JSON validation files without using datasets.load_dataset lock files.
This avoids distributed lock contention when many ranks initialize
validation simultaneously on shared filesystems.
"""
with filename.open("r", encoding="utf-8") as f:
payload = json.load(f)
if isinstance(payload, dict):
if "data" not in payload:
raise ValueError(
f"Validation JSON {filename.as_posix()} must contain a 'data' field when top-level is an object."
)
records = payload["data"]
elif isinstance(payload, list):
records = payload
else:
raise ValueError(
f"Validation JSON {filename.as_posix()} must be either a list of samples or an object containing a 'data' list."
)
if not isinstance(records, list):
raise ValueError(
f"Validation JSON records in {filename.as_posix()} must be a list."
)
return records
+40
View File
@@ -830,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
@@ -859,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 = ""
@@ -874,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"
@@ -934,6 +937,9 @@ class TrainingArgs(FastVideoArgs):
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
use_tf: bool = False
use_decoupled_dmd: bool = False # use decoupled DMD
use_gan_loss: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
@@ -943,6 +949,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":
@@ -998,6 +1006,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,
@@ -1091,6 +1103,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")
@@ -1139,6 +1154,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,
@@ -1329,6 +1348,21 @@ class TrainingArgs(FastVideoArgs):
help=
"Whether to warp denoising step according to the scheduler time shift"
)
parser.add_argument(
"--use-tf",
action=StoreBoolean,
help="Whether to use Teacher-forcing for finetuning"
)
parser.add_argument(
"--use-decoupled-dmd",
action=StoreBoolean,
help="Whether to use decoupled DMD"
)
parser.add_argument(
"--use-gan-loss",
action=StoreBoolean,
help="Whether to use GAN loss"
)
# Self-forcing specific arguments
parser.add_argument(
@@ -1365,6 +1399,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
@@ -0,0 +1,948 @@
# 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
# Apply flex_attention
# Does not support SP padding for now
if kv_cache is None:
is_tf = (img_q.shape[1] == cos.shape[0] * 2)
if is_tf:
q_chunk = torch.chunk(img_q, 2, dim=1)
k_chunk = torch.chunk(img_k, 2, dim=1)
roped_query = []
roped_key = []
for ii in range(2):
rq = _apply_rotary_emb(q_chunk[ii], cos, sin, is_neox_style=False)
rk = _apply_rotary_emb(k_chunk[ii], cos, sin, is_neox_style=False)
roped_query.append(rq)
roped_key.append(rk)
img_q = torch.cat(roped_query, dim=1)
img_k = torch.cat(roped_key, dim=1)
else:
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)
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:
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)
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_teacher_forcing_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=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
"""
# debug
DEBUG = False
if DEBUG:
num_frames = 9
frame_seqlen = 256
total_length = num_frames * frame_seqlen * 2
total_kv_length = total_length + text_seq_len
# 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
total_length_tensor = torch.tensor(total_length, device=device)
total_kv_length_tensor = torch.tensor(total_kv_length, device=device)
clean_ends = num_frames * frame_seqlen
# for clean context frames, we can construct their flex attention mask based on a [start, end] interval
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
# for noisy frames, we need two intervals to construct the flex attention mask [context_start, context_end] [noisy_start, noisy_end]
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_noise_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_noise_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
attention_block_size = frame_seqlen * num_frame_per_block
frame_indices = torch.arange(
start=0,
end=num_frames * frame_seqlen,
step=attention_block_size,
device=device, dtype=torch.long
)
# attention for clean context frames
for start in frame_indices:
context_ends[start:start + attention_block_size] = start + attention_block_size
noisy_image_start_list = torch.arange(
num_frames * frame_seqlen, total_length,
step=attention_block_size,
device=device, dtype=torch.long
)
noisy_image_end_list = noisy_image_start_list + attention_block_size
# attention for noisy frames
for block_index, (start, end) in enumerate(zip(noisy_image_start_list, noisy_image_end_list)):
# attend to noisy tokens within the same block
noise_noise_starts[start:end] = start
noise_noise_ends[start:end] = end
# attend to context tokens in previous blocks
# noise_context_starts[start:end] = 0
noise_context_ends[start:end] = block_index * attention_block_size
def attention_mask(b, h, q_idx, kv_idx):
# first design the mask for clean frames
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
# then design the mask for noisy frames
# noisy frames will attend to all clean preceeding clean frames + itself
C1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
C2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
noise_mask = (q_idx >= clean_ends) & (C1 | C2)
eye_mask = q_idx == kv_idx
text_mask = (kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor)
return eye_mask | clean_mask | noise_mask | text_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 DEBUG:
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_kv_length + kv_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
@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,
clean_hidden_states: Optional[torch.Tensor] = None,
aug_timestep: Optional[torch.Tensor] = None,
):
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)
if clean_hidden_states is not None:
clean_hidden_states = self.img_in(clean_hidden_states)
hidden_states = torch.cat([clean_hidden_states, hidden_states], dim=1)
if aug_timestep is None:
aug_timestep = torch.zeros_like(timestep)
if aug_timestep.dim() == 2:
ts_seq_len = aug_timestep.shape[1]
aug_temb = self.time_in(aug_timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
aug_temb = self.time_in(aug_timestep, timestep_r=timestep_r)
temb = torch.cat([aug_temb, temb], dim=1)
# Prepare block-wise causal attention mask
if kv_cache is None:
if clean_hidden_states is not None:
self.block_mask = self._prepare_teacher_forcing_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,
text_seq_len=txt_kv_cache[0]["k_txt"].shape[1]
)
else:
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()
if clean_hidden_states is not None:
hidden_states = hidden_states[:, hidden_states.shape[1] // 2:]
# 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")
+149 -11
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.platforms import AttentionBackendEnum, current_platform
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
@@ -88,10 +88,25 @@ class CausalWanSelfAttention(nn.Module):
cache_start = current_start
cos, sin = freqs_cis
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
if kv_cache is None:
is_tf = (q.shape[1] == cos.shape[0] * 2)
if is_tf:
q_chunk = torch.chunk(q, 2, dim=1)
k_chunk = torch.chunk(k, 2, dim=1)
roped_query = []
roped_key = []
for ii in range(2):
rq = _apply_rotary_emb(q_chunk[ii], cos, sin, is_neox_style=False).type_as(v)
rk = _apply_rotary_emb(k_chunk[ii], cos, sin, is_neox_style=False).type_as(v)
roped_query.append(rq)
roped_key.append(rk)
roped_query = torch.cat(roped_query, dim=1)
roped_key = torch.cat(roped_key, dim=1)
else:
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
# Padding for flex attention
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_roped_query = torch.cat(
@@ -120,6 +135,8 @@ class CausalWanSelfAttention(nn.Module):
block_mask=block_mask
)[:, :, :-padded_length].transpose(2, 1)
else:
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
frame_seqlen = q.shape[1]
current_end = current_start + roped_query.shape[1]
sink_tokens = self.sink_size * frame_seqlen
@@ -286,6 +303,8 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -294,10 +313,13 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -376,6 +398,94 @@ class CausalWanTransformer3DModel(BaseDiT):
self.__post_init__()
@staticmethod
def _prepare_teacher_forcing_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1
) -> 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
"""
# debug
DEBUG = False
if DEBUG:
num_frames = 9
frame_seqlen = 256
total_length = num_frames * frame_seqlen * 2
# we do right padding to get to a multiple of 128
padded_length = math.ceil(total_length / 128) * 128 - total_length
clean_ends = num_frames * frame_seqlen
# for clean context frames, we can construct their flex attention mask based on a [start, end] interval
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
# for noisy frames, we need two intervals to construct the flex attention mask [context_start, context_end] [noisy_start, noisy_end]
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_noise_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_noise_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
attention_block_size = frame_seqlen * num_frame_per_block
frame_indices = torch.arange(
start=0,
end=num_frames * frame_seqlen,
step=attention_block_size,
device=device, dtype=torch.long
)
# attention for clean context frames
for start in frame_indices:
context_ends[start:start + attention_block_size] = start + attention_block_size
noisy_image_start_list = torch.arange(
num_frames * frame_seqlen, total_length,
step=attention_block_size,
device=device, dtype=torch.long
)
noisy_image_end_list = noisy_image_start_list + attention_block_size
# attention for noisy frames
for block_index, (start, end) in enumerate(zip(noisy_image_start_list, noisy_image_end_list)):
# attend to noisy tokens within the same block
noise_noise_starts[start:end] = start
noise_noise_ends[start:end] = end
# attend to context tokens in previous blocks
# noise_context_starts[start:end] = 0
noise_context_ends[start:end] = block_index * attention_block_size
def attention_mask(b, h, q_idx, kv_idx):
# first design the mask for clean frames
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
# then design the mask for noisy frames
# noisy frames will attend to all clean preceeding clean frames + itself
C1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
C2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
noise_mask = (q_idx >= clean_ends) & (C1 | C2)
eye_mask = q_idx == kv_idx
return eye_mask | clean_mask | noise_mask
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
KV_LEN=total_length + padded_length, _compile=False, device=device)
if DEBUG:
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
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
@@ -551,6 +661,7 @@ class CausalWanTransformer3DModel(BaseDiT):
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
start_frame: int = 0,
clean_hidden_states: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
@@ -562,6 +673,10 @@ class CausalWanTransformer3DModel(BaseDiT):
else:
encoder_hidden_states_image = None
if timestep.ndim == 1:
# [B] -> [B, F]
timestep = timestep.unsqueeze(1).repeat(1, hidden_states.shape[2])
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
@@ -588,13 +703,21 @@ class CausalWanTransformer3DModel(BaseDiT):
# Construct blockwise causal attn mask
if self.block_mask 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
)
if clean_hidden_states is not None:
self.block_mask = self._prepare_teacher_forcing_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,
)
else:
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
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
@@ -607,6 +730,17 @@ class CausalWanTransformer3DModel(BaseDiT):
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
if clean_hidden_states is not None:
clean_hidden_states = self.patch_embedding(clean_hidden_states)
clean_hidden_states = clean_hidden_states.flatten(2).transpose(1, 2)
hidden_states = torch.cat([clean_hidden_states, hidden_states], dim=1)
aug_timestep = torch.zeros_like(timestep)
temb_clean = self.condition_embedder.time_embedder(aug_timestep.flatten(), None)
timestep_proj_clean = self.condition_embedder.time_modulation(temb_clean)
timestep_proj_clean = timestep_proj_clean.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=aug_timestep.shape)
timestep_proj = torch.cat([timestep_proj_clean, timestep_proj], dim=1)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
@@ -630,6 +764,9 @@ class CausalWanTransformer3DModel(BaseDiT):
timestep_proj, freqs_cis,
block_mask=self.block_mask)
if clean_hidden_states is not None:
hidden_states = hidden_states[:, hidden_states.shape[1] // 2:]
# 5. Output norm, projection & unpatchify
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
@@ -644,12 +781,13 @@ class CausalWanTransformer3DModel(BaseDiT):
def forward(
self,
*args,
clean_hidden_states: torch.Tensor | None = None,
**kwargs
):
if kwargs.get('kv_cache', None) is not None:
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
return self._forward_train(*args, clean_hidden_states=clean_hidden_states, **kwargs)
def unpatchify(self, x, grid_sizes):
+18 -3
View File
@@ -127,8 +127,9 @@ class HunyuanVideo15TimeEmbedding(nn.Module):
self,
timestep: torch.Tensor,
timestep_r: Optional[torch.Tensor] = None,
timestep_seq_len: int | None = None,
) -> torch.Tensor:
timesteps_emb = self.timestep_embedder(timestep)
timesteps_emb = self.timestep_embedder(timestep, timestep_seq_len)
if timestep_r is not None:
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
@@ -473,6 +474,7 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
**kwargs
):
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
@@ -494,7 +496,12 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
else:
temb = self.time_in(timestep, timestep_r=timestep_r)
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
@@ -698,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)
@@ -850,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
+173 -5
View File
@@ -8,6 +8,8 @@ import numpy as np
import torch
import torch.nn as nn
from einops import repeat
import fastvideo.envs as envs
from fastvideo.attention import (DistributedAttention, DistributedAttention_VSA,
LocalAttention)
@@ -152,6 +154,32 @@ class WanSelfAttention(nn.Module):
"""
pass
class WanGanCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens=None, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
crossattn_cache (List[dict], *optional*): Contains the cached key and value tensors for context embedding.
"""
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
qq = self.norm_q(self.to_q(context)[0]).view(b, 1, -1, d)
kk = self.norm_k(self.to_k(x)[0]).view(b, -1, n, d)
vv = self.to_v(x)[0].view(b, -1, n, d)
# compute attention
x = self.attn(qq, kk, vv)
# output
x = x.flatten(2)
x, _ = self.to_out(x)
return x
class WanT2VCrossAttention(WanSelfAttention):
@@ -245,6 +273,98 @@ class WanI2VCrossAttention(WanSelfAttention):
x, _ = self.to_out(x)
return x
class GanAttentionBlock(nn.Module):
def __init__(self,
dim=1536,
ffn_dim=8192,
num_heads=12,
window_size=(-1, -1),
qk_norm=True,
cross_attn_norm=True,
eps=1e-6):
super().__init__()
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.window_size = window_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
# layers
# self.norm1 = WanLayerNorm(dim, eps)
# self.self_attn = WanSelfAttention(dim, num_heads, window_size, qk_norm,
# eps)
self.norm3 = FP32LayerNorm(
dim, eps,
elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.ffn = nn.Sequential(
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
nn.Linear(ffn_dim, dim))
self.cross_attn = WanGanCrossAttention(dim, num_heads,
(-1, -1),
qk_norm,
eps)
# modulation
# self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states,
encoder_hidden_states,
# seq_lens,
# grid_sizes,
# freqs,
# context,
# context_lens,
):
r"""
Args:
x(Tensor): Shape [B, L, C]
e(Tensor): Shape [B, 6, C]
seq_lens(Tensor): Shape [B], length of each sequence in batch
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
# assert e.dtype == torch.float32
# with amp.autocast(dtype=torch.float32):
# e = (self.modulation + e).chunk(6, dim=1)
# assert e[0].dtype == torch.float32
# # self-attention
# y = self.self_attn(
# self.norm1(x) * (1 + e[1]) + e[0], seq_lens, grid_sizes,
# freqs)
# # with amp.autocast(dtype=torch.float32):
# x = x + y * e[2]
# cross-attention & ffn function
def cross_attn_ffn(x, context):
token = context + self.cross_attn(self.norm3(x), context)
y = self.ffn(self.norm2(token)) + token # * (1 + e[4]) + e[3])
# with amp.autocast(dtype=torch.float32):
# x = x + y * e[5]
return y
x = cross_attn_ffn(hidden_states, encoder_hidden_states)
return x
class RegisterTokens(nn.Module):
def __init__(self, num_registers: int, dim: int):
super().__init__()
self.register_tokens = nn.Parameter(torch.randn(num_registers, dim) * 0.02)
self.rms_norm = RMSNorm(dim, eps=1e-6)
def forward(self):
return self.rms_norm(self.register_tokens)
def reset_parameters(self):
nn.init.normal_(self.register_tokens, std=0.02)
class WanTransformerBlock(nn.Module):
@@ -630,6 +750,26 @@ class WanTransformer3DModel(CachableDiT):
self.cnt = 0
self.__post_init__()
def adding_cls_branch(self, atten_dim=1536, num_class=4, time_embed_dim=0) -> None:
self._cls_pred_branch = nn.Sequential(
# Input: [B, 384, 21, 60, 104]
nn.LayerNorm(atten_dim * 3 + time_embed_dim),
nn.Linear(atten_dim * 3 + time_embed_dim, 1536),
nn.SiLU(),
nn.Linear(atten_dim, num_class)
)
self._cls_pred_branch.requires_grad_(True)
num_registers = 3
self._register_tokens = RegisterTokens(num_registers=num_registers, dim=atten_dim)
self._register_tokens.requires_grad_(True)
gan_ca_blocks = []
for _ in range(num_registers):
block = GanAttentionBlock()
gan_ca_blocks.append(block)
self._gan_ca_blocks = nn.ModuleList(gan_ca_blocks)
self._gan_ca_blocks.requires_grad_(True)
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
@@ -637,6 +777,8 @@ class WanTransformer3DModel(CachableDiT):
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
classify_mode=False,
concat_time_embeddings=False,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
@@ -727,6 +869,16 @@ class WanTransformer3DModel(CachableDiT):
assert encoder_hidden_states.dtype == orig_dtype
final_x = None
if classify_mode:
assert self._register_tokens is not None
assert self._gan_ca_blocks is not None
assert self._cls_pred_branch is not None
final_x = []
registers = repeat(self._register_tokens(), "n d -> b n d", b=hidden_states.shape[0])
gan_idx = 0
# 4. Transformer blocks
# if caching is enabled, we might be able to skip the forward pass
should_skip_forward = self.should_skip_forward_for_cached_states(
@@ -740,19 +892,33 @@ class WanTransformer3DModel(CachableDiT):
if enable_teacache:
original_hidden_states = hidden_states.clone()
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
for ii, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
else:
for block in self.blocks:
else:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask)
timestep_proj, freqs_cis, attention_mask)
if classify_mode and ii in [13, 21, 29]:
gan_token = registers[:, gan_idx: gan_idx + 1]
final_x.append(self._gan_ca_blocks[gan_idx](hidden_states, gan_token))
gan_idx += 1
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
if classify_mode:
final_x = torch.cat(final_x, dim=1)
if concat_time_embeddings:
cls_input = torch.cat([final_x.float(), (10 * temb[:, None, :]).float()], dim=1).view(final_x.shape[0], -1)
final_x = self._cls_pred_branch(cls_input)
else:
cls_input = final_x.float().view(final_x.shape[0], -1)
final_x = self._cls_pred_branch(cls_input)
# 5. Output norm, projection & unpatchify
if temb.dim() == 3:
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
@@ -778,6 +944,8 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
if classify_mode:
return output, torch.nan_to_num(final_x)
return output
def maybe_cache_states(self, hidden_states: torch.Tensor,
+17 -7
View File
@@ -391,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
@@ -558,10 +560,9 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
class_name = config.get("_class_name")
assert class_name is not None, (
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
)
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."
fastvideo_args.model_paths["vae"] = model_path
from fastvideo.platforms import current_platform
@@ -729,6 +730,7 @@ class TransformerLoader(ComponentLoader):
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:
@@ -807,6 +809,11 @@ class TransformerLoader(ComponentLoader):
or cls_name == "Cosmos25Transformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25"
)
add_cls_branch = False
if hasattr(fastvideo_args, "_adding_cls_branch") and fastvideo_args._adding_cls_branch:
strict_load = False
add_cls_branch = True
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={"config": dit_config, "hf_config": hf_config},
@@ -826,10 +833,11 @@ class TransformerLoader(ComponentLoader):
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
add_cls_branch=add_cls_branch,
)
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"
@@ -867,9 +875,11 @@ class SchedulerLoader(ComponentLoader):
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
if fastvideo_args.pipeline_config.flow_shift is not None and "shift" in config:
config["shift"] = fastvideo_args.pipeline_config.flow_shift
scheduler = scheduler_cls(**config)
if fastvideo_args.pipeline_config.flow_shift is not None:
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
# if fastvideo_args.pipeline_config.flow_shift is not None:
# scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
if fastvideo_args.pipeline_config.timesteps_scale is not None:
scheduler.set_timesteps_scale(
fastvideo_args.pipeline_config.timesteps_scale
+12 -7
View File
@@ -75,6 +75,7 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
add_cls_branch: bool = False,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -96,6 +97,8 @@ def maybe_load_fsdp_model(
logger.info("Loading model with default_dtype: %s", default_dtype)
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
if add_cls_branch:
model.adding_cls_branch(atten_dim=1536, num_class=1, time_embed_dim=0)
# Check if we should use FSDP
use_fsdp = training_mode or fsdp_inference
@@ -340,7 +343,7 @@ def load_model_from_full_model_state_dict(
unused_keys)
# List of allowed parameter name patterns
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l", "_cls_pred_branch", "_register_tokens", "_gan_ca_blocks"] # Can be extended as needed
for new_param_name in unused_keys:
if not any(pattern in new_param_name
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
@@ -353,14 +356,16 @@ def load_model_from_full_model_state_dict(
meta_sharded_param = meta_sd.get(new_param_name)
if not hasattr(meta_sharded_param, "device_mesh"):
# Initialize with zeros
sharded_tensor = torch.zeros_like(meta_sharded_param,
device=device,
dtype=param_dtype)
# sharded_tensor = torch.zeros_like(meta_sharded_param,
# device=device,
# dtype=param_dtype)
sharded_tensor = meta_sharded_param.to(device=device, dtype=param_dtype)
else:
# Initialize with zeros and distribute
full_tensor = torch.zeros_like(meta_sharded_param,
device=device,
dtype=param_dtype)
# full_tensor = torch.zeros_like(meta_sharded_param,
# device=device,
# dtype=param_dtype)
full_tensor = meta_sharded_param.to(device=device, dtype=param_dtype)
sharded_tensor = distribute_tensor(
full_tensor,
meta_sharded_param.device_mesh,
+4 -3
View File
@@ -78,9 +78,10 @@ def hf_to_custom_state_dict(
for source_param_name, full_tensor in hf_param_sd: # type: ignore
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is None:
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
+2
View File
@@ -28,6 +28,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("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"),
@@ -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:
"""
@@ -69,6 +69,9 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
# use this scheduler for ODE trajectory
timestep = timestep.unsqueeze(0)
model_output = model_output.permute(0, 2, 1, 3, 4)
sample = sample.permute(0, 2, 1, 3, 4)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
timestep = timestep.to(model_output.device)
@@ -81,7 +84,8 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
self.inverse_timesteps or self.reverse_sigmas) else 0
else:
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
prev_sample = sample + model_output * (sigma_ - sigma)
prev_sample = sample.flatten(0, 1) + model_output.flatten(0, 1) * (sigma_ - sigma)
prev_sample = prev_sample.unflatten(0, model_output.shape[:2]).permute(0, 2, 1, 3, 4)
if isinstance(prev_sample, torch.Tensor | float) and not return_dict:
return (prev_sample, )
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
+1 -1
View File
@@ -663,7 +663,7 @@ class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = False
self.use_tiling = True
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
+1
View File
@@ -1219,6 +1219,7 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
out = torch.clamp(out, min=-1.0, max=1.0)
self.clear_cache()
else:
raise NotImplementedError("Feature cache is not supported for WanVAE")
out = ParallelTiledVAE.decode(self, z)
return out
+1 -1
View File
@@ -277,7 +277,7 @@ def load_video(
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images, original_fps if return_fps else pil_images
return (pil_images, original_fps) if return_fps else pil_images
def get_default_height_width(
+228
View File
@@ -0,0 +1,228 @@
import math
import torch
try:
from torch.distributed.tensor import DTensor
except ImportError:
# handle old pytorch versions
Dtensor = None
# This code is modified from the GitHub repository of KellerJordan:
# https://github.com/KellerJordan/Muon/blob/master/muon.py
def zeropower_via_newtonschulz5(G, steps=5):
"""
Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a
quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose
of minimizing steps, it turns out to be empirically effective to keep increasing the slope at
zero even beyond the point where the iteration no longer converges all the way to one everywhere
on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T
where S' is diagonal with S_{ii}' ~ Uniform(0.5, 1.5), which turns out not to hurt model
performance at all relative to UV^T, where USV^T = G is the SVD.
"""
if isinstance(G, DTensor):
device_mesh = G.device_mesh
G = G.full_tensor()
else:
device_mesh = None
assert len(G.shape) >= 2
a, b, c = (3.4445, -4.7750, 2.0315)
X = G
if G.size(-2) > G.size(-1):
X = X.mT
# Ensure spectral norm is at most 1
X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7)
# Perform the NS iterations
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A # quintic computation strategy adapted from suggestion by @jxbz, @leloykun, and @YouJiacheng
X = a * X + B @ X
if G.size(-2) > G.size(-1):
X = X.mT
if device_mesh is not None:
return DTensor.from_local(X, device_mesh)
else:
return X
class Muon(torch.optim.Optimizer):
"""
Muon - MomentUm Orthogonalized by Newton-schulz
Arguments:
muon_params: The parameters to be optimized by Muon.
lr: The learning rate. The updates will have spectral norm of `lr`. (0.02 is a good default)
momentum: The momentum used by the internal SGD. (0.95 is a good default)
nesterov: Whether to use Nesterov-style momentum in the internal SGD. (recommended)
ns_steps: The number of Newton-Schulz iterations to run. (6 is probably always enough)
adamw_params: The parameters to be optimized by AdamW. Any parameters in `muon_params` which are
{0, 1}-D or are detected as being the embed or lm_head will be optimized by AdamW as well.
adamw_lr: The learning rate for the internal AdamW.
adamw_betas: The betas for the internal AdamW.
adamw_eps: The epsilon for the internal AdamW.
adamw_wd: The weight decay for the internal AdamW.
"""
def __init__(
self,
lr=1e-3,
wd=0.1,
muon_params=None,
momentum=0.95,
nesterov=True,
ns_steps=5,
adamw_params=None,
adamw_betas=(0.95, 0.95),
adamw_eps=1e-8,
):
defaults = dict(
lr=lr,
wd=wd,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
)
params = list(muon_params)
adamw_params = list(adamw_params) if adamw_params is not None else []
params.extend(adamw_params)
super().__init__(params, defaults)
# Sort parameters into those for which we will use Muon, and those for which we will not
for p in muon_params:
# Use Muon for every parameter in muon_params which is >= 2D and doesn't look like an embedding or head layer
assert p.ndim >= 2, p.ndim
self.state[p]["use_muon"] = True
for p in adamw_params:
# Do not use Muon for parameters in adamw_params
self.state[p]["use_muon"] = False
def adjust_lr_for_muon(self, lr, param_shape):
A, B = param_shape[:2]
# We adjust the learning rate and weight decay based on the size of the parameter matrix
# as describted in the paper
adjusted_ratio = 0.2 * math.sqrt(max(A, B))
adjusted_lr = lr * adjusted_ratio
return adjusted_lr
def step(self, closure=None):
"""Perform a single optimization step.
Args:
closure (Callable, optional): A closure that reevaluates the model
and returns the loss.
"""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
############################
# Muon #
############################
params = [p for p in group["params"] if self.state[p]["use_muon"]]
lr = group["lr"]
wd = group["wd"]
momentum = group["momentum"]
# generate weight updates in distributed fashion
for p in params:
# sanity check
g = p.grad
if g is None:
continue
if g.ndim > 2:
g = g.view(g.size(0), -1)
assert g is not None
# calc update
state = self.state[p]
if "momentum_buffer" not in state:
state["momentum_buffer"] = torch.zeros_like(g)
buf = state["momentum_buffer"]
buf.mul_(momentum).add_(g)
g = g.add(buf, alpha=momentum) if group["nesterov"] else buf
g = g.bfloat16()
u = zeropower_via_newtonschulz5(g, steps=group["ns_steps"])
# scale update
adjusted_lr = self.adjust_lr_for_muon(lr, p.shape)
# apply weight decay
p.data.mul_(1 - lr * wd)
# apply update
p.data.add_(u.view(p.shape), alpha=-adjusted_lr)
############################
# AdamW backup #
############################
params = [
p for p in group["params"] if not self.state[p]["use_muon"]
]
lr = group['lr']
beta1, beta2 = group["adamw_betas"]
eps = group["adamw_eps"]
weight_decay = group["wd"]
for p in params:
g = p.grad
if g is None:
continue
state = self.state[p]
if "step" not in state:
state["step"] = 0
state["moment1"] = torch.zeros_like(g)
state["moment2"] = torch.zeros_like(g)
state["step"] += 1
step = state["step"]
buf1 = state["moment1"]
buf2 = state["moment2"]
buf1.lerp_(g, 1 - beta1)
buf2.lerp_(g.square(), 1 - beta2)
g = buf1 / (eps + buf2.sqrt())
bias_correction1 = 1 - beta1**step
bias_correction2 = 1 - beta2**step
scale = bias_correction1 / bias_correction2**0.5
p.data.mul_(1 - lr * weight_decay)
p.data.add_(g, alpha=-lr / scale)
return loss
# help function to create the Muon optimizer
def get_muon_optimizer(model,
lr=1e-3,
weight_decay=0.1,
momentum=0.95,
adamw_betas=(0.95, 0.95),
adamw_eps=1e-8):
muon_params = [
p for name, p in model.named_parameters()
if p.requires_grad and p.ndim >= 2
]
adamw_params = [
p for name, p in model.named_parameters()
if p.requires_grad and not (p.ndim >= 2)
]
return Muon(
lr=lr,
wd=weight_decay,
muon_params=muon_params,
momentum=momentum,
adamw_params=adamw_params,
adamw_betas=adamw_betas,
adamw_eps=adamw_eps,
)
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan causal DMD pipeline implementation.
This module wires the causal DMD denoising stage into the modular pipeline.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
Hy15CausalDMDDenosingStage,
InputValidationStage,
Hy15ImageEncodingStage,
LatentPreparationStage,
TextEncodingStage)
# isort: on
logger = init_logger(__name__)
class Hy15CausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=Hy15CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = Hy15CausalDMDPipeline
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
Hy15ImageEncodingStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15ImageToVideoPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HunyuanVideo15ImageToVideoPipeline
@@ -287,6 +287,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:
@@ -123,6 +123,7 @@ class ForwardBatch:
# Latent tensors
latents: torch.Tensor | None = None
clean_latents: torch.Tensor | None = None
raw_latent_shape: tuple[int, ...] | None = None
noise_pred: torch.Tensor | None = None
image_latent: torch.Tensor | None = None
@@ -232,6 +233,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
@@ -240,6 +242,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
+2
View File
@@ -29,6 +29,8 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
"HYWorldPipeline": "hyworld",
"Hy15CausalDMDPipeline": "hunyuan15",
"HunyuanVideo15ImageToVideoPipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"Cosmos2_5Pipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
@@ -0,0 +1,319 @@
# SPDX-License-Identifier: Apache-2.0
"""
ODE Trajectory Data Preprocessing pipeline implementation.
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
using the modular pipeline architecture.
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
"""
import os
from collections.abc import Iterator
from typing import Any
import pyarrow as pa
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import (
ode_text_only_record_creator)
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.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, Hy15ImageEncodingStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
logger = init_logger(__name__)
class Hy15ODEPreprocessPipeline(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int
def get_pyarrow_schema(self) -> pa.Schema:
"""Return the PyArrow schema for ODE Trajectory pipeline."""
return pyarrow_schema_ode_trajectory_text_only
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.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="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", 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"),
scheduler=self.get_module("scheduler"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
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.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
# Add fps and duration if available in data
if "fps" in data:
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
if "duration" in data:
valid_data["duration"] = [
data["duration"][i] for i in valid_indices
]
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
sampling_params = SamplingParam.from_pretrained(args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
else:
negative_prompt_embeds_list = []
negative_prompt_masks_list = []
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
# 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
assert batch.guidance_scale == 6, "guidance_scale must be 6 for hy15_ode_trajectory"
assert batch.do_classifier_free_guidance, "do_classifier_free_guidance must be True for hy15_ode_trajectory"
# 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
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)
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 = {
"trajectory_latents": trajectory_latents,
"trajectory_timesteps": trajectory_timesteps
}
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
if j in [50]:
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]] = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
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 = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).numpy()
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (text-only ODE schema)
record: dict[str, Any] = ode_text_only_record_creator(
video_name=video_name,
text_embedding=text_embedding,
caption=valid_data["text"][idx],
trajectory_latents=sample_extra_features[
"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)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
self.get_pyarrow_schema())
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self.num_processed_samples = 0
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush(write_remainder=True)
if written:
logger.info("Final flush wrote %s samples", written)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for video preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing videos",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.video_data: dict[str, Any] = {} # Store video metadata and paths
self.latent_data: dict[str, Any] = {} # Store latent tensors
self.preprocess_text_and_trajectory(fastvideo_args, args)
EntryClass = Hy15ODEPreprocessPipeline
@@ -283,8 +283,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
table = pq.read_table(os.path.join(root, file))
start_idx += table.num_rows
logger.info("Start index: %s", start_idx)
# Loading dataset
train_dataset = getdataset(args)
train_dataset = getdataset(args, start_idx=start_idx)
train_dataloader = DataLoader(
train_dataset,
@@ -301,6 +303,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
for batch_idx, data in enumerate(pbar):
if data is None:
logger.warning("Batch is None")
continue
with torch.inference_mode():
@@ -313,6 +316,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
num_processed_samples += len(valid_indices)
if not valid_indices:
logger.warning("No valid indices")
continue
# Create new batch with only valid samples
@@ -24,8 +24,7 @@ from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import (
ode_text_only_record_creator)
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
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 (
@@ -180,13 +179,14 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.guidance_scale = 3.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.timestep_preparation_stage(
# result_batch, fastvideo_args)
result_batch.timesteps = self.get_module("scheduler").timesteps
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
@@ -209,10 +209,12 @@ 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 [49, 50]:
local_rank = int(os.getenv("RANK", 0))
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{local_rank}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
@@ -0,0 +1,329 @@
# SPDX-License-Identifier: Apache-2.0
"""
ODE Trajectory Data Preprocessing pipeline implementation.
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
using the modular pipeline architecture.
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
"""
import os
from collections.abc import Iterator
from typing import Any
from unittest import result
import pyarrow as pa
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import (
ode_text_only_record_creator)
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory_text_only, pyarrow_schema_t2v
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.utils import save_decoded_latents_as_video, shallow_asdict
from fastvideo.distributed import get_local_torch_device
logger = init_logger(__name__)
class PreprocessPipeline_TF_ODE(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int
def get_pyarrow_schema(self) -> pa.Schema:
"""Return the PyArrow schema for ODE Trajectory pipeline."""
return pyarrow_schema_ode_trajectory_text_only
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")],
))
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", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
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")))
def _process_batch(self, batch):
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
latent = batch['vae_latent'].float()
infos = batch['info_list'][0]
if (hasattr(self.modules["vae"], "shift_factor")
and self.modules["vae"].shift_factor is not None):
if isinstance(self.modules["vae"].shift_factor, torch.Tensor):
latent -= self.modules["vae"].shift_factor.to(
latent.device, latent.dtype)
else:
latent -= self.modules["vae"].shift_factor
else:
raise ValueError("vae.shift_factor is not set")
if isinstance(self.modules["vae"].scaling_factor, torch.Tensor):
latent = latent * self.modules["vae"].scaling_factor.to(
latent.device, latent.dtype)
else:
latent = latent * self.modules["vae"].scaling_factor
# Move to device
device = get_local_torch_device()
encoder_hidden_states = encoder_hidden_states.to(
device, dtype=torch.bfloat16)
encoder_attention_mask = encoder_attention_mask.to(
device, dtype=torch.bfloat16)
return encoder_hidden_states, encoder_attention_mask, latent.to(device, dtype=torch.bfloat16), infos
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data and generate trajectory information."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
sampling_params = SamplingParam.from_pretrained(args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
prompt_embeds, prompt_attention_masks, clean_latents, infos = self._process_batch(data)
assert prompt_embeds.shape[0] == 1, "Only one prompt embeds are supported"
self.num_processed_samples += prompt_embeds.shape[0]
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = [prompt_embeds]
batch.prompt_attention_mask = [prompt_attention_masks]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.clean_latents = clean_latents
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(
# result_batch, fastvideo_args)
result_batch.timesteps = self.get_module("scheduler").timesteps
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
if batch.return_trajectory_decoded:
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_decoded.append(result_batch.trajectory_decoded)
result_batch.trajectory_latents = result_batch.trajectory_latents[:, [0, 12, 24, 36, -2, -1]]
result_batch.trajectory_timesteps = result_batch.trajectory_timesteps[[0, 12, 24, 36, -2, -1]]
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
# Prepare extra features for text-only processing
extra_features = {
"trajectory_latents": trajectory_latents,
"trajectory_timesteps": trajectory_timesteps
}
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
if j in [50, 51]:
local_rank = int(os.getenv("RANK", 0))
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{local_rank}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate([infos["file_name"]]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds.float().cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).numpy()
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (text-only ODE schema)
record: dict[str, Any] = ode_text_only_record_creator(
video_name=video_name,
text_embedding=text_embedding,
caption=infos["caption"],
trajectory_latents=sample_extra_features[
"trajectory_latents"],
trajectory_timesteps=sample_extra_features[
"trajectory_timesteps"],
)
batch_data.append(record)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
self.get_pyarrow_schema())
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self.num_processed_samples = 0
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush(write_remainder=True)
if written:
logger.info("Final flush wrote %s samples", written)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
self.train_dataset, self.preprocess_dataloader = build_parquet_map_style_dataloader(
args.data_merge_path,
args.preprocess_video_batch_size,
# parquet_schema=pyarrow_schema_ode_trajectory_text_only,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.dataloader_num_workers,
cfg_rate=0.0,
drop_last=True,
text_padding_length=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=args.seed)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for video preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing videos",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.video_data: dict[str, Any] = {} # Store video metadata and paths
self.latent_data: dict[str, Any] = {} # Store latent tensors
self.preprocess_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_TF_ODE
+27 -14
View File
@@ -12,6 +12,10 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_tf_ode import (
PreprocessPipeline_TF_ODE)
from fastvideo.pipelines.preprocess.hy15.hy15_ode_preprocess_pipeline import (
Hy15ODEPreprocessPipeline)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
@@ -24,7 +28,7 @@ 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"
@@ -32,17 +36,17 @@ def main(args) -> None:
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
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)
# 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,
@@ -51,6 +55,7 @@ def main(args) -> None:
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
pipeline_config=pipeline_config,
init_weights_from_safetensors=args.init_weights_from_safetensors,
)
if args.preprocess_task == "t2v":
PreprocessPipeline = PreprocessPipeline_T2V
@@ -62,12 +67,19 @@ def main(args) -> None:
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
elif args.preprocess_task == "tf_ode":
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
PreprocessPipeline = PreprocessPipeline_TF_ODE
elif args.preprocess_task == "hy15_ode_trajectory":
assert args.flow_shift is not None, "flow_shift is required for hy15_ode_trajectory"
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
PreprocessPipeline = Hy15ODEPreprocessPipeline
elif args.preprocess_task == "matrixgame":
PreprocessPipeline = PreprocessPipeline_MatrixGame
else:
raise ValueError(
f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only, matrixgame")
f"Valid options: t2v, i2v, ode_trajectory, hy15_ode_trajectory, text_only, matrixgame, tf_ode")
logger.info("Preprocess task: %s using %s", args.preprocess_task,
PreprocessPipeline.__name__)
@@ -115,7 +127,7 @@ if __name__ == "__main__":
"--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only", "ode_trajectory", "matrixgame"],
choices=["t2v", "i2v", "text_only", "ode_trajectory", "hy15_ode_trajectory", "matrixgame", "tf_ode"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
@@ -138,6 +150,7 @@ if __name__ == "__main__":
help=
"The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument("--init_weights_from_safetensors", type=str, default=None)
args = parser.parse_args()
main(args)
+2
View File
@@ -8,6 +8,7 @@ complete diffusion pipelines.
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.hy15_causal_denoising import Hy15CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
@@ -52,6 +53,7 @@ __all__ = [
"Cosmos25LatentPreparationStage",
"LTX2LatentPreparationStage",
"LTX2AudioDecodingStage",
"Hy15CausalDMDDenosingStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
+44 -8
View File
@@ -14,7 +14,7 @@ try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
except:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
@@ -22,7 +22,7 @@ try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
except:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
@@ -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:
@@ -112,6 +112,13 @@ class CausalDMDDenosingStage(DenoisingStage):
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
if batch.do_classifier_free_guidance:
raise ValueError("Classifier free guidance is not supported for causal DMD denoising")
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"
# Initialize or reset caches
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
@@ -218,6 +225,11 @@ class CausalDMDDenosingStage(DenoisingStage):
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# self.scheduler.set_timesteps(num_inference_steps=48,
# denoising_strength=1.0)
# timesteps = self.scheduler.timesteps.to(get_local_torch_device())
logger.info(f"timesteps: {timesteps}")
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
@@ -296,6 +308,30 @@ class CausalDMDDenosingStage(DenoisingStage):
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
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):
pred_noise_uncond_btchw = current_model(
latent_model_input,
neg_prompt_embeds,
t_expanded_noise,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
pred_noise_btchw = pred_noise_uncond_btchw + batch.guidance_scale * (
pred_noise_btchw - pred_noise_uncond_btchw)
# Convert pred noise to pred video with FM Euler scheduler utilities
if boundary_timestep is not None and t_cur >= boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
@@ -412,8 +448,8 @@ class CausalDMDDenosingStage(DenoisingStage):
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = self.transformer.attention_head_dim
num_attention_heads = self.transformer.config.num_attention_heads
attention_head_dim = self.transformer.config.attention_head_dim
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
+18 -13
View File
@@ -54,22 +54,23 @@ class DecodingStage(PipelineStage):
"""Convert normalized latents into the VAE's expected latent space."""
# Some VAEs handle latent (de)normalization internally.
if bool(getattr(self.vae, "handles_latent_denorm", False)):
raise NotImplementedError("handles_latent_denorm is not supported for WanVAE")
return latents
cfg = getattr(self.vae, "config", None)
# MatrixGame-style: z = z * std + mean
if (cfg is not None and hasattr(cfg, "latents_mean")
and hasattr(cfg, "latents_std")):
latents_mean = torch.tensor(cfg.latents_mean,
device=latents.device,
dtype=latents.dtype).view(
1, -1, 1, 1, 1)
latents_std = torch.tensor(cfg.latents_std,
device=latents.device,
dtype=latents.dtype).view(
1, -1, 1, 1, 1)
return latents * latents_std + latents_mean
# if (cfg is not None and hasattr(cfg, "latents_mean")
# and hasattr(cfg, "latents_std")):
# latents_mean = torch.tensor(cfg.latents_mean,
# device=latents.device,
# dtype=latents.dtype).view(
# 1, -1, 1, 1, 1)
# latents_std = torch.tensor(cfg.latents_std,
# device=latents.device,
# dtype=latents.dtype).view(
# 1, -1, 1, 1, 1)
# return latents * latents_std + latents_mean
# Diffusers-style: scaling_factor (+ optional shift_factor)
if hasattr(self.vae, "scaling_factor"):
@@ -86,6 +87,10 @@ class DecodingStage(PipelineStage):
latents.device, latents.dtype)
else:
latents = latents + self.vae.shift_factor
else:
raise NotImplementedError("shift_factor is not supported for WanVAE")
else:
raise NotImplementedError("handles_latent_denorm is not supported for WanVAE")
return latents
@@ -121,8 +126,8 @@ class DecodingStage(PipelineStage):
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.pipeline_config.vae_tiling:
# self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
+38 -11
View File
@@ -26,27 +26,27 @@ 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 (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
except:
st_attn_available = False
try:
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
from fastvideo.utils import is_vmoba_available
vmoba_attn_available = is_vmoba_available()
except ImportError:
except:
vmoba_attn_available = False
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
except:
vsa_available = False
logger = init_logger(__name__)
@@ -176,6 +176,13 @@ class DenoisingStage(PipelineStage):
},
)
tf_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"clean_hidden_states": batch.clean_latents,
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
@@ -210,7 +217,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 +233,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,9 +245,11 @@ 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
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
if isinstance(patch_size, int):
patch_size = (patch_size, patch_size, patch_size)
seq_len = ((F - 1) // temporal_scale +
1) * (batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size[1] *
@@ -242,8 +257,9 @@ class DenoisingStage(PipelineStage):
# 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):
@@ -311,9 +327,9 @@ class DenoisingStage(PipelineStage):
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
])
timestep = temp_ts.unsqueeze(0)
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
t_expand = timestep.repeat(latent_model_input.shape[0], 1).to(get_local_torch_device())
else:
t_expand = t.repeat(latent_model_input.shape[0])
t_expand = t.repeat(latent_model_input.shape[0]).to(get_local_torch_device())
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
@@ -406,8 +422,11 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
**tf_kwargs,
)
assert batch.do_classifier_free_guidance, "do_classifier_free_guidance must be True"
assert current_guidance_scale == 6.0, "guidance_scale must be 6.0"
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
@@ -423,6 +442,7 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
**tf_kwargs,
)
noise_pred_text = noise_pred
@@ -462,9 +482,16 @@ class DenoisingStage(PipelineStage):
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_timesteps.append(torch.zeros_like(t))
if batch.clean_latents is not None:
trajectory_latents.append(batch.clean_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)
# assert trajectory_tensor.shape[1] == 1, "trajectory_tensor should have 1 frame"
# assert trajectory_timesteps_tensor.shape[0] == 1, "trajectory_timesteps_tensor should have 1 frame"
else:
trajectory_tensor = None
trajectory_timesteps_tensor = None
@@ -0,0 +1,356 @@
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
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:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except:
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)
# 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
# block_sizes[0] = 1
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
+4 -4
View File
@@ -115,10 +115,10 @@ class Hy15ImageEncodingStage(ImageEncodingStage):
"""
Encode the prompt into image encoder hidden states.
"""
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
# if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
raw_latent_shape = list(batch.raw_latent_shape)
raw_latent_shape[1] = 1
@@ -118,11 +118,12 @@ class InputValidationStage(PipelineStage):
oh, ow = batch.height, batch.width
img = img.resize((ow, oh), Image.LANCZOS)
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
if isinstance(patch_size, int):
patch_size = (patch_size, patch_size, patch_size)
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
max_area = 480 * 848
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
@@ -19,7 +19,7 @@ try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
except:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
@@ -27,7 +27,7 @@ try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
except:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
+321
View File
@@ -0,0 +1,321 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any, cast
import numpy as np
import torch
import torch.nn.functional as F
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
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.pipeline_batch_info import TrainingBatch
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
logger = init_logger(__name__)
class ARDiffusionTrainingPipeline(TrainingPipeline):
"""
Training pipeline for AR diffusion using precomputed denoising trajectories.
Supervision: predict the next latent in the stored trajectory by
- feeding current latent at timestep t into the transformer to predict noise
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Match the preprocess/generation scheduler for consistent stepping
assert fastvideo_args.pipeline_config.flow_shift == 5.0, "flow_shift must be 5.0"
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_t2v
def initialize_training_pipeline(self, training_args: TrainingArgs):
super().initialize_training_pipeline(training_args)
self.noise_scheduler = self.get_module("scheduler")
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "timestep_shift must be 5.0"
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)
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.vae))
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.manual_idx = 0
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
# Warm start validation with current transformer
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
def _get_next_batch(
self,
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
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']
latent = batch['vae_latent'].float()
infos = batch['info_list']
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latent -= self.vae.shift_factor.to(
latent.device, latent.dtype)
else:
latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent = latent * self.vae.scaling_factor.to(
latent.device, latent.dtype)
else:
latent = latent * self.vae.scaling_factor
# [B, C, T, H, W] -> [B, T, C, H, W]
latent = latent.permute(0, 2, 1, 3, 4)
# 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.infos = infos
return training_batch, latent[:, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
)
def _get_timestep(self,
min_timestep: int,
max_timestep: int,
batch_size: int,
num_frame: int,
num_frame_per_block: int,
uniform_timestep: bool = False) -> torch.Tensor:
if uniform_timestep:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, 1],
device=self.device,
dtype=torch.long).repeat(1, num_frame)
return timestep
else:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, num_frame],
device=self.device,
dtype=torch.long)
# logger.info(f"individual timestep: {timestep}")
# make the noise level the same within every block
timestep = timestep.reshape(timestep.shape[0], -1,
num_frame_per_block)
timestep[:, :, 1:] = timestep[:, :, 0:1]
timestep = timestep.reshape(timestep.shape[0], -1)
return timestep
def _step_predict_next_latent(
self, latent: torch.Tensor,
encoder_hidden_states: 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()
B, num_frames, num_channels, height, width = 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)
noisy_input = self.noise_scheduler.add_noise(
latent.flatten(0, 1),
torch.randn_like(latent.flatten(0, 1)),
timestep.flatten(0, 1)).unflatten(0, (B, num_frames))
clean_input = latent.clone()
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
latent_vis_dict["x0"] = latent.permute(0, 2, 1, 3,
4).detach().clone().cpu()
logger.info("timestep: %s", timestep)
input_kwargs = {
"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,
}
if self.training_args.use_tf:
raise NotImplementedError("Teacher-forcing is not implemented yet")
input_kwargs["clean_hidden_states"] = clean_input.permute(0, 2, 1, 3, 4)
# 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)
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),
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
scheduler=self.modules["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, latent, timestep, latent_vis_dict
def train_one_step(self, training_batch): # type: ignore[override]
self.transformer.train()
self.optimizer.zero_grad()
training_batch.total_loss = 0.0
args = cast(TrainingArgs, self.training_args)
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
for _ in range(args.gradient_accumulation_steps):
training_batch, latent = self._get_next_batch(
training_batch)
text_embeds = training_batch.encoder_hidden_states
text_attention_mask = training_batch.encoder_attention_mask
assert latent.shape[0] == 1
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
latent, text_embeds)
training_batch.latent_vis_dict.update(latent_vis_dict)
mask = t != 0
# Compute loss
loss = F.mse_loss(noise_pred[mask],
target_latent[mask],
reduction="mean")
loss = loss / args.gradient_accumulation_steps
with set_forward_context(current_timestep=t,
attn_metadata=None,
forward_batch=None):
loss.backward()
avg_loss = loss.detach().clone()
training_batch.total_loss += avg_loss.item()
# Clip grad and step optimizers
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in self.transformer.parameters() if p.requires_grad],
args.max_grad_norm if args.max_grad_norm is not None else 0.0)
self.optimizer.step()
self.lr_scheduler.step()
if grad_norm is None:
grad_value = 0.0
else:
try:
if isinstance(grad_norm, torch.Tensor):
grad_value = float(grad_norm.detach().float().item())
else:
grad_value = float(grad_norm)
except Exception:
grad_value = 0.0
training_batch.grad_norm = grad_value
return training_batch
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
tracker_loss_dict: dict[str, Any] = {}
latents_vis_dict = training_batch.latent_vis_dict
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
for latent_key in latent_log_keys:
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.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
if video_artifact is not None:
tracker_loss_dict[latent_key] = video_artifact
# Clean up references
del video, pixel_latent, latent
if self.global_rank == 0 and tracker_loss_dict:
self.tracker.log_artifacts(tracker_loss_dict, step)
def main(args) -> None:
logger.info("Starting AR diffusion training pipeline...")
pipeline = ARDiffusionTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("AR diffusion training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
+394 -56
View File
@@ -28,8 +28,6 @@ from fastvideo.distributed import (cleanup_dist_env_and_memory,
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
@@ -39,9 +37,10 @@ from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
shift_timestep, unshift_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,8 @@ 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"
if training_args.use_gan_loss:
training_args._adding_cls_branch = True
self.fake_score_transformer = self.load_module_from_path(
training_args.fake_score_model_path, "transformer",
training_args)
@@ -146,8 +145,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 +191,25 @@ 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 +226,25 @@ 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 +292,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 +577,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,
@@ -698,8 +704,155 @@ class DistillationPipeline(TrainingPipeline):
"generator_timestep"] = target_timestep.float().detach()
return pred_video
def _dmd_decoupled_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch,
exit_timestep: float) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
original_latent = generator_pred_video
with torch.no_grad():
def get_dm_training_batch():
dm_timestep = torch.randint(0,
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
dm_timestep = shift_timestep(
dm_timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
dm_timestep = dm_timestep.clamp(self.min_timestep, self.max_timestep)
dm_noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_pred_video.dtype)
dm_noisy_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), dm_noise.flatten(0, 1),
dm_timestep).detach().unflatten(0,
(1, generator_pred_video.shape[1]))
return dm_timestep, dm_noisy_latent, dm_noise
def get_ca_training_batch(noise):
from fastvideo.training.training_utils import unshift_timestep
exit_timestep_int = int(unshift_timestep(exit_timestep, self.timestep_shift, self.num_train_timestep))
ca_timestep = torch.randint(0,
exit_timestep_int, [1],
device=self.device,
dtype=torch.long)
ca_timestep = shift_timestep(
ca_timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
ca_timestep = ca_timestep.clamp(self.min_timestep, self.max_timestep)
ca_noisy_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
ca_timestep).detach().unflatten(0,
(1, generator_pred_video.shape[1]))
return ca_timestep, ca_noisy_latent
dm_timestep, dm_noisy_latent, noise = get_dm_training_batch()
ca_timestep, ca_noisy_latent = get_ca_training_batch(noise)
# Calculate DM fake score
training_batch = self._build_distill_input_kwargs(
dm_noisy_latent.clone(), dm_timestep, training_batch.conditional_dict,
training_batch)
current_fake_score_transformer = self._get_fake_score_transformer(
dm_timestep)
dm_fake_score_pred_noise = current_fake_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
dm_faker_score_pred_video = pred_noise_to_pred_video(
pred_noise=dm_fake_score_pred_noise.flatten(0, 1),
noise_input_latent=dm_noisy_latent.flatten(0, 1),
timestep=dm_timestep,
scheduler=self.noise_scheduler).unflatten(
0, dm_fake_score_pred_noise.shape[:2])
# Calculate DM real score
training_batch = self._build_distill_input_kwargs(
dm_noisy_latent.clone(), dm_timestep, training_batch.conditional_dict,
training_batch)
current_real_score_transformer = self._get_real_score_transformer(
dm_timestep)
dm_real_score_pred_noise_cond = current_real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
dm_pred_real_video_cond = pred_noise_to_pred_video(
pred_noise=dm_real_score_pred_noise_cond.flatten(0, 1),
noise_input_latent=dm_noisy_latent.flatten(0, 1),
timestep=dm_timestep,
scheduler=self.noise_scheduler).unflatten(
0, dm_real_score_pred_noise_cond.shape[:2])
# Calculate CA real score cond and uncond
training_batch = self._build_distill_input_kwargs(
ca_noisy_latent.clone(), ca_timestep, training_batch.conditional_dict,
training_batch)
ca_real_score_pred_noise_cond = current_real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
ca_pred_real_video_cond = pred_noise_to_pred_video(
pred_noise=ca_real_score_pred_noise_cond.flatten(0, 1),
noise_input_latent=ca_noisy_latent.flatten(0, 1),
timestep=ca_timestep,
scheduler=self.noise_scheduler).unflatten(
0, ca_real_score_pred_noise_cond.shape[:2])
training_batch = self._build_distill_input_kwargs(
ca_noisy_latent.clone(), ca_timestep, training_batch.unconditional_dict,
training_batch)
ca_real_score_pred_noise_uncond = current_real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
ca_pred_real_video_uncond = pred_noise_to_pred_video(
pred_noise=ca_real_score_pred_noise_uncond.flatten(0, 1),
noise_input_latent=ca_noisy_latent.flatten(0, 1),
timestep=ca_timestep,
scheduler=self.noise_scheduler).unflatten(
0, ca_real_score_pred_noise_uncond.shape[:2])
# Calculate normalized DM loss
dm_grad = dm_faker_score_pred_video - dm_pred_real_video_cond
# dm_grad = (dm_faker_score_pred_video - dm_pred_real_video_cond) / torch.abs(original_latent - dm_pred_real_video_cond).mean()
ca_grad = (self.real_score_guidance_scale - 1) * (ca_pred_real_video_uncond - ca_pred_real_video_cond)
# ca_grad = ca_grad / torch.abs(original_latent - (self.real_score_guidance_scale - 1) * ca_pred_real_video_cond).mean()
grad = dm_grad + ca_grad
grad = torch.nan_to_num(grad)
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":
training_batch.latents,
"generator_pred_video":
original_latent.detach(),
"real_score_pred_video":
dm_pred_real_video_cond.detach(),
"faker_score_pred_video":
dm_faker_score_pred_video.detach(),
"dmd_timestep":
ca_timestep.detach(),
})
return dmd_loss
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
training_batch: TrainingBatch,
exit_timestep: float) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
original_latent = generator_pred_video
with torch.no_grad():
@@ -723,10 +876,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 +903,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 +919,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 +940,23 @@ 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())
if self.training_args.use_gan_loss:
gan_generator_loss = self._gan_generator_forward(
generator_pred_video=generator_pred_video,
training_batch=training_batch,
exit_timestep=exit_timestep)
dmd_loss += gan_generator_loss
training_batch.dmd_latent_vis_dict.update({
"training_batch_dmd_fwd_clean_latent":
@@ -798,6 +973,105 @@ class DistillationPipeline(TrainingPipeline):
return dmd_loss
def _gan_generator_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch,
exit_timestep: float) -> torch.Tensor:
idx = torch.argmin(exit_timestep - self.denoising_step_list)
critic_timestep = exit_timestep.clamp(self.min_timestep, self.max_timestep) * torch.ones([1, 1], device=self.device, dtype=torch.long)
# if idx == len(self.denoising_step_list) - 1:
# start_timestep = 0
# end_timestep = int(unshift_timestep(self.denoising_step_list[idx], self.timestep_shift, self.num_train_timestep))
# else:
# start_timestep = int(unshift_timestep(self.denoising_step_list[idx + 1], self.timestep_shift, self.num_train_timestep))
# end_timestep = int(unshift_timestep(self.denoising_step_list[idx], self.timestep_shift, self.num_train_timestep))
# critic_timestep = torch.randint(start_timestep, end_timestep, [1], device=self.device, dtype=torch.long)
# critic_timestep = shift_timestep(critic_timestep, self.timestep_shift, self.num_train_timestep)
# critic_timestep = critic_timestep.clamp(self.min_timestep, self.max_timestep)
critic_noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_pred_video.dtype)
noisy_fake_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), critic_noise.flatten(0, 1),
critic_timestep).detach().unflatten(0,
(1, generator_pred_video.shape[1]))
training_batch = self._build_distill_input_kwargs(
noisy_fake_latent, critic_timestep, training_batch.conditional_dict,
training_batch)
current_fake_score_transformer = self._get_fake_score_transformer(
critic_timestep)
_, noisy_fake_logit = current_fake_score_transformer(classify_mode=True,
**training_batch.input_kwargs)
gan_generator_loss = F.softplus(1 - noisy_fake_logit.float()).mean()
return gan_generator_loss
def _gan_critic_forward(self, training_batch: TrainingBatch, clean_latent: torch.Tensor) -> torch.Tensor:
with torch.no_grad(), set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
training_batch)
else:
generator_pred_video, exit_timestep = self._generator_forward(training_batch)
critic_timestep = exit_timestep.clamp(self.min_timestep, self.max_timestep) * torch.ones([1, 1], device=self.device, dtype=torch.long)
# idx = torch.argmin(exit_timestep - self.denoising_step_list)
# if idx == len(self.denoising_step_list) - 1:
# start_timestep = 0
# end_timestep = int(unshift_timestep(self.denoising_step_list[idx], self.timestep_shift, self.num_train_timestep))
# else:
# start_timestep = int(unshift_timestep(self.denoising_step_list[idx + 1], self.timestep_shift, self.num_train_timestep))
# end_timestep = int(unshift_timestep(self.denoising_step_list[idx], self.timestep_shift, self.num_train_timestep))
# critic_timestep = torch.randint(start_timestep, end_timestep, [1], device=self.device, dtype=torch.long)
# critic_timestep = shift_timestep(critic_timestep, self.timestep_shift, self.num_train_timestep)
# critic_timestep = critic_timestep.clamp(self.min_timestep, self.max_timestep)
# Calculate fake latent score
critic_noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_pred_video.dtype)
noisy_fake_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), critic_noise.flatten(0, 1),
critic_timestep).detach().unflatten(0,
(1, generator_pred_video.shape[1]))
training_batch = self._build_distill_input_kwargs(
noisy_fake_latent, critic_timestep, training_batch.conditional_dict,
training_batch)
current_fake_score_transformer = self._get_fake_score_transformer(
critic_timestep)
_, noisy_fake_logit = current_fake_score_transformer(classify_mode=True,
**training_batch.input_kwargs)
# Calculate real latent score
# critic_noise = torch.randn(self.video_latent_shape,
# device=self.device,
# dtype=generator_pred_video.dtype)
noisy_real_latent = self.noise_scheduler.add_noise(
clean_latent.flatten(0, 1), critic_noise.flatten(0, 1),
critic_timestep).detach().unflatten(0,
(1, clean_latent.shape[1]))
training_batch = self._build_distill_input_kwargs(
noisy_real_latent, critic_timestep, training_batch.conditional_dict,
training_batch)
current_fake_score_transformer = self._get_fake_score_transformer(
critic_timestep)
_, noisy_real_logit = current_fake_score_transformer(classify_mode=True,
**training_batch.input_kwargs)
gan_critic_loss = F.softplus(noisy_real_logit.float()).mean() - F.softplus(1 - noisy_fake_logit.float()).mean()
# gan_critic_loss = F.softplus(1 - noisy_real_logit.mean()) + F.softplus(1 + noisy_fake_logit.mean())
return gan_critic_loss
def faker_score_forward(
self, training_batch: TrainingBatch
) -> tuple[TrainingBatch, torch.Tensor]:
@@ -805,10 +1079,10 @@ class DistillationPipeline(TrainingPipeline):
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video = self._generator_multi_step_simulation_forward(
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
training_batch)
else:
generator_pred_video = self._generator_forward(training_batch)
generator_pred_video, exit_timestep = self._generator_forward(training_batch)
fake_score_timestep = torch.randint(0,
self.num_train_timestep, [1],
@@ -831,6 +1105,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 +1158,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,
@@ -953,6 +1234,52 @@ class DistillationPipeline(TrainingPipeline):
training_batch.infos = infos
return training_batch
def _get_next_batch_2(
self,
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
batch = next(self.train_loader_iter_2, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
self.train_loader_iter_2 = iter(self.train_dataloader_2)
batch = next(self.train_loader_iter_2)
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
latent = batch['vae_latent'].float()
infos = batch['info_list']
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latent -= self.vae.shift_factor.to(
latent.device, latent.dtype)
else:
latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent = latent * self.vae.scaling_factor.to(
latent.device, latent.dtype)
else:
latent = latent * self.vae.scaling_factor
# [B, C, T, H, W] -> [B, T, C, H, W]
latent = latent.permute(0, 2, 1, 3, 4)
# 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.infos = infos
return training_batch, latent[:, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
)
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
gradient_accumulation_steps = getattr(self.training_args,
'gradient_accumulation_steps', 1)
@@ -980,17 +1307,23 @@ class DistillationPipeline(TrainingPipeline):
current_timestep=batch_gen.timesteps,
attn_metadata=batch_gen.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video = self._generator_multi_step_simulation_forward(
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
batch_gen)
else:
generator_pred_video = self._generator_forward(
generator_pred_video, exit_timestep = self._generator_forward(
batch_gen)
with set_forward_context(current_timestep=batch_gen.timesteps,
attn_metadata=batch_gen.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=batch_gen)
if self.training_args.use_decoupled_dmd:
dmd_loss = self._dmd_decoupled_forward(
generator_pred_video=generator_pred_video,
training_batch=batch_gen,
exit_timestep=exit_timestep)
else:
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=batch_gen)
with set_forward_context(
current_timestep=batch_gen.timesteps,
@@ -1239,6 +1572,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 +1583,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 +1691,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):
@@ -0,0 +1,662 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import sys
from copy import deepcopy
from typing import Any, cast
import numpy as np
import torch
import torch.nn.functional as F
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
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.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__)
class Hy15ODEInitTrainingPipeline(TrainingPipeline):
"""
Training pipeline for ODE-init using precomputed denoising trajectories.
Supervision: predict the next latent in the stored trajectory by
- feeding current latent at timestep t into the transformer to predict noise
- stepping the scheduler with the predicted noise
- minimizing MSE to the stored next latent at timestep t_next
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
]
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
def initialize_training_pipeline(self, training_args: TrainingArgs):
super().initialize_training_pipeline(training_args)
self.noise_scheduler = self.get_module("scheduler")
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
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
extra_one_step=True,
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],
dtype=torch.float32))).cuda()
logger.info("timesteps: %s", timesteps)
self.dmd_denoising_steps = timesteps[1000 -
self.dmd_denoising_steps]
logger.info("warped self.dmd_denoising_steps: %s",
self.dmd_denoising_steps)
else:
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())
logger.info("denoising_step_list: %s", self.dmd_denoising_steps)
logger.info(
"Initialized ODE-init training pipeline with %s denoising steps",
len(self.dmd_denoising_steps))
# Cache for nearest trajectory index per DMD step (computed lazily on first batch)
self._cached_closest_idx_per_dmd = None
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.manual_idx = 0
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
# Warm start validation with current transformer
self.validation_pipeline = Hy15CausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=False)
def _get_next_batch(
self,
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
# Required fields from parquet (ODE trajectory schema)
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:
# [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
trajectory_latents = trajectory_latents[:, 0]
elif trajectory_latents.dim() == 6:
# already [B, S, C, T, H, W]
pass
else:
raise ValueError(
f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
)
trajectory_timesteps = batch['trajectory_timesteps']
if trajectory_timesteps.dim() == 3:
# [B, 1, S] -> [B, S]
trajectory_timesteps = trajectory_timesteps[:, 0]
elif trajectory_timesteps.dim() == 2:
# [B, S]
pass
else:
raise ValueError(
f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
)
# [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
# Move to device
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[:, :, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
), trajectory_timesteps.to(
device)
def _get_timestep(self,
min_timestep: int,
max_timestep: int,
batch_size: int,
num_frame: int,
num_frame_per_block: int,
uniform_timestep: bool = False) -> torch.Tensor:
if uniform_timestep:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, 1],
device=self.device,
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,
dtype=torch.long)
# logger.info(f"individual timestep: {timestep}")
# make the noise level the same within every block
timestep = timestep.reshape(timestep.shape[0], -1,
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_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]
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
B, S, num_frames, num_channels, height, width = traj_latents.shape
# 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, 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)
logger.info(
"corresponding timesteps: %s", self.noise_scheduler.timesteps[
self._cached_closest_idx_per_dmd])
# Select the K indexes from traj_latents using self._cached_closest_idx_per_dmd
# traj_latents: [B, S, C, T, H, W], self._cached_closest_idx_per_dmd: [K]
# Output: [B, K, C, T, H, W]
assert self._cached_closest_idx_per_dmd is not None
relevant_traj_latents = torch.index_select(
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
indexes = self._get_timestep( # [B, num_frames]
0,
len(self.dmd_denoising_steps),
B,
num_frames,
3,
uniform_timestep=False)
# noisy_input = relevant_traj_latents[indexes]
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]
# 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 _df_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)
noise_pred = noise_pred.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 _tf_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)
clean_input = torch.cat([
target_latent,
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,
"clean_hidden_states": clean_input.permute(0, 2, 1, 3, 4),
}
# 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)
noise_pred = noise_pred.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()
training_batch.total_loss = 0.0
args = cast(TrainingArgs, self.training_args)
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
for _ in range(args.gradient_accumulation_steps):
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
training_batch)
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]
_, S = traj_latents.shape[0], traj_latents.shape[1]
if S < 2:
raise ValueError("Trajectory must contain at least 2 steps")
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._tf_step_predict_next_latent(
traj_latents, traj_timesteps, text_embeds, text_attention_mask,
image_embeds)
training_batch.latent_vis_dict.update(latent_vis_dict)
mask = t != 0
# Compute loss
loss = F.mse_loss(noise_pred[mask],
target_latent[mask],
reduction="mean")
loss = loss / args.gradient_accumulation_steps
with set_forward_context(current_timestep=t,
attn_metadata=None,
forward_batch=None):
loss.backward()
avg_loss = loss.detach().clone()
training_batch.total_loss += avg_loss.item()
# Clip grad and step optimizers
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in self.transformer.parameters() if p.requires_grad],
args.max_grad_norm if args.max_grad_norm is not None else 0.0)
self.optimizer.step()
self.lr_scheduler.step()
if grad_norm is None:
grad_value = 0.0
else:
try:
if isinstance(grad_norm, torch.Tensor):
grad_value = float(grad_norm.detach().float().item())
else:
grad_value = float(grad_norm)
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,
training_args: TrainingArgs, step: int):
tracker_loss_dict: dict[str, Any] = {}
latents_vis_dict = training_batch.latent_vis_dict
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
for latent_key in latent_log_keys:
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.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=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
del video, pixel_latent, latent
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 Hy15 ODE-init training pipeline...")
pipeline = Hy15ODEInitTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Hy15 ODE-init training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
File diff suppressed because it is too large Load Diff
+214 -14
View File
@@ -18,6 +18,7 @@ from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
@@ -39,6 +40,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Match the preprocess/generation scheduler for consistent stepping
assert fastvideo_args.pipeline_config.flow_shift == 5.0, "flow_shift must be 5.0"
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
@@ -57,22 +59,22 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
assert self.timestep_shift == 5.0, "timestep_shift must be 5.0"
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)
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.vae))
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],
dtype=torch.float32))).cuda()
logger.info("timesteps: %s", timesteps)
self.dmd_denoising_steps = timesteps[1000 -
self.dmd_denoising_steps]
logger.info("warped self.dmd_denoising_steps: %s",
@@ -83,8 +85,6 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
get_local_torch_device())
logger.info("denoising_step_list: %s", self.dmd_denoising_steps)
logger.info(
"Initialized ODE-init training pipeline with %s denoising steps",
len(self.dmd_denoising_steps))
@@ -111,6 +111,64 @@ class ODEInitTrainingPipeline(TrainingPipeline):
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
def _get_next_gt_batch(
self,
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
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)
logger.info("batch keys: %s", batch.keys())
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
# Trajectory tensors may include a leading singleton batch dim per row
trajectory_latents = batch['trajectory_latents']
if trajectory_latents.dim() == 7:
# [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
trajectory_latents = trajectory_latents[:, 0]
elif trajectory_latents.dim() == 6:
# already [B, S, C, T, H, W]
pass
else:
raise ValueError(
f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
)
trajectory_timesteps = batch['trajectory_timesteps']
if trajectory_timesteps.dim() == 3:
# [B, 1, S] -> [B, S]
trajectory_timesteps = trajectory_timesteps[:, 0]
elif trajectory_timesteps.dim() == 2:
# [B, S]
pass
else:
raise ValueError(
f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
)
# [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
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.infos = infos
return training_batch, trajectory_latents[:, :, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
), trajectory_timesteps.to(
device)
def _get_next_batch(
self,
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
@@ -161,8 +219,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
device, dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
return training_batch, trajectory_latents[:, :, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
), trajectory_timesteps.to(
device)
## TEMP Used for loading the sf .pt files directly
"""
@@ -296,6 +358,144 @@ class ODEInitTrainingPipeline(TrainingPipeline):
return pred_video, target_latent, timestep, latent_vis_dict
def _tf_step_predict_next_latent(
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: 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)
noisy_input = 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))
clean_input = target_latent.clone()
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.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)
input_kwargs = {
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"clean_hidden_states": clean_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep.to(device, dtype=torch.bfloat16),
"return_dict": False,
}
# 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)
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),
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
scheduler=self.modules["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 _tf_ode_step_predict_next_latent(
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: 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()
clean_latent = traj_latents[:, -1]
target_latent = traj_latents[:, -2]
ode_latents_valid = traj_latents[:, :-1]
# self._cached_closest_idx_per_dmd = torch.tensor(
# [0, 16, 32, 41], dtype=torch.long).cpu()
# self.dmd_denoising_steps = traj_timesteps[0, self._cached_closest_idx_per_dmd]
# logger.info(f"corresponding timesteps: {self.dmd_denoising_steps}")
# relevant_traj_latents = torch.index_select(
# ode_latents_valid,
# dim=1,
# index=self._cached_closest_idx_per_dmd.to(device))
relevant_traj_latents = ode_latents_valid
# 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,
len(self.dmd_denoising_steps),
B,
num_frames,
3,
uniform_timestep=True)
timestep = self.dmd_denoising_steps[indexes.cpu()].to(device)
noisy_input = 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)
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
4).detach().clone().cpu()
latent_vis_dict["clean_latent"] = clean_latent.permute(0, 2, 1, 3,
4).detach().clone().cpu()
logger.info("timestep: %s", timestep)
input_kwargs = {
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"clean_hidden_states": clean_latent.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep.to(device, dtype=torch.bfloat16),
"return_dict": False,
}
# 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)
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),
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
scheduler=self.modules["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()
@@ -313,11 +513,11 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
_, S = traj_latents.shape[0], traj_latents.shape[1]
if S < 2:
raise ValueError("Trajectory must contain at least 2 steps")
# if S < 2:
# raise ValueError("Trajectory must contain at least 2 steps")
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
noise_pred, target_latent, t, latent_vis_dict = self._tf_ode_step_predict_next_latent(
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
training_batch.latent_vis_dict.update(latent_vis_dict)
@@ -362,12 +562,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
training_args: TrainingArgs, step: int):
tracker_loss_dict: dict[str, Any] = {}
latents_vis_dict = training_batch.latent_vis_dict
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
latent_log_keys = ['noisy_input', 'x0', 'pred_video', 'clean_latent']
for latent_key in latent_log_keys:
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(
pixel_latent = self.decoding_stage.decode(
latent, training_args)
video = pixel_latent.cpu().float()
@@ -8,7 +8,6 @@ from typing import Any
import numpy as np
import torch
import torch.distributed as dist
from einops import rearrange
from tqdm.auto import tqdm
import fastvideo.envs as envs
@@ -19,8 +18,8 @@ from fastvideo.fastvideo_args import 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.models.utils import pred_noise_to_pred_video
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import SelfForcingFlowMatchScheduler
from fastvideo.pipelines import TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.training.training_utils import (EMA_FSDP,
@@ -68,7 +67,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self.noise_scheduler = SelfForcingFlowMatchScheduler(
num_inference_steps=1000,
shift=5.0,
shift=self.timestep_shift,
sigma_min=0.0,
extra_one_step=True,
training=True)
@@ -84,35 +83,23 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self.last_step_only = getattr(training_args, 'last_step_only', False)
self.context_noise = getattr(training_args, 'context_noise', 0)
self.kv_cache1: list[dict[str, Any]] | None = None
self.crossattn_cache: list[dict[str, Any]] | None = None
logger.info("Self-forcing generator update ratio: %s",
self.dfake_gen_update_ratio)
logger.info("RANK: %s, exiting initialize_training_pipeline",
self.global_rank,
local_main_process_only=False)
def generate_and_sync_list(self, num_blocks: int, num_denoising_steps: int,
def generate_and_sync_list(self, num_blocks: int, start_timestep_index: int,
end_timestep_index: int,
device: torch.device) -> list[int]:
"""Generate and synchronize random exit flags across distributed processes."""
logger.info(
"RANK: %s, enter generate_and_sync_list blocks=%s steps=%s device=%s",
self.global_rank,
num_blocks,
num_denoising_steps,
str(device),
local_main_process_only=False)
rank = dist.get_rank() if dist.is_initialized() else 0
if rank == 0:
# Generate random indices
indices = torch.randint(low=0,
high=num_denoising_steps,
indices = torch.randint(low=start_timestep_index,
high=end_timestep_index,
size=(num_blocks, ),
device=device)
if self.last_step_only:
indices = torch.ones_like(indices) * (num_denoising_steps - 1)
indices = torch.ones_like(indices) * (end_timestep_index - 1)
else:
indices = torch.empty(num_blocks, dtype=torch.long, device=device)
@@ -120,12 +107,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
dist.broadcast(indices,
src=0) # Broadcast the random indices to all ranks
flags = indices.tolist()
logger.info(
"RANK: %s, exit generate_and_sync_list flags_len=%s first=%s",
self.global_rank,
len(flags),
flags[0] if len(flags) > 0 else None,
local_main_process_only=False)
return flags
def generator_loss(
@@ -138,14 +119,21 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
with set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
generator_pred_video = self._generator_multi_step_simulation_forward(
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
training_batch)
with set_forward_context(current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
dmd_loss = self._dmd_forward(
if self.training_args.use_decoupled_dmd:
dmd_loss = self._dmd_decoupled_forward(
generator_pred_video=generator_pred_video,
training_batch=training_batch,
exit_timestep=exit_timestep)
else:
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=training_batch)
training_batch=training_batch,
exit_timestep=exit_timestep)
log_dict = {
"dmdtrain_gradient_norm": torch.tensor(0.0, device=self.device)
@@ -186,7 +174,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
# During training, the number of generated frames should be uniformly sampled from
# [21, self.num_training_frames], but still being a multiple of self.num_frame_per_block
min_num_frames = 20 if self.independent_first_frame else 21
min_num_frames = num_training_frames - 1 if self.independent_first_frame else num_training_frames
max_num_frames = num_training_frames - 1 if self.independent_first_frame else num_training_frames
assert max_num_frames % self.num_frame_per_block == 0
assert min_num_frames % self.num_frame_per_block == 0
@@ -205,23 +193,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
num_generated_frames += 1
min_num_frames += 1
# Create noise with dynamic shape
if initial_latent is not None:
noise_shape = [
batch_size, num_generated_frames - 1,
*self.video_latent_shape[2:]
]
if training_batch.use_gt_trajectory and training_batch.trajectory_latents is not None:
noise = training_batch.trajectory_latents.to(self.device,
dtype=dtype)
else:
noise_shape = [
batch_size, num_generated_frames, *self.video_latent_shape[2:]
]
noise = torch.randn(noise_shape, device=self.device, dtype=dtype)
if self.sp_world_size > 1:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
noise = torch.randn_like(training_batch.latents).to(self.device, dtype=dtype)
batch_size, num_frames, num_channels, height, width = noise.shape
@@ -252,7 +228,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
# Step 1: Initialize KV cache to all zeros
cache_frames = num_generated_frames + num_input_frames
self.kv_cache1, self.crossattn_cache = self._initialize_simulation_caches(
kv_cache1, crossattn_cache = self._initialize_simulation_caches(
batch_size, dtype, self.device, max_num_frames=cache_frames)
# Step 2: Cache context feature
@@ -286,19 +262,33 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
all_num_frames = [self.num_frame_per_block] * num_blocks
if self.independent_first_frame and initial_latent is None:
all_num_frames = [1] + all_num_frames
num_denoising_steps = len(self.denoising_step_list)
if training_batch.use_gt_trajectory and training_batch.trajectory_latents is not None:
start_timestep_index = training_batch.start_timestep_index
end_timestep_index = len(self.denoising_step_list)
else:
start_timestep_index = 0
end_timestep_index = len(self.denoising_step_list)
exit_flags = self.generate_and_sync_list(len(all_num_frames),
num_denoising_steps,
start_timestep_index,
end_timestep_index,
device=noise.device)
start_gradient_frame_index = max(0, num_output_frames - 21)
start_gradient_frame_index = max(
0, num_output_frames - num_training_frames)
for block_index, current_num_frames in enumerate(all_num_frames):
noisy_input = noise[:, current_start_frame -
num_input_frames:current_start_frame +
current_num_frames - num_input_frames]
index = start_timestep_index
current_timestep = self.denoising_step_list[index]
assert index < len(
self.denoising_step_list
), "Index is greater than the number of denoising steps"
# Step 3.1: Spatial denoising loop
for index, current_timestep in enumerate(self.denoising_step_list):
while index < len(self.denoising_step_list):
if self.same_step_across_blocks:
exit_flag = (index == exit_flags[0])
else:
@@ -329,8 +319,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=current_start_frame).permute(
@@ -369,8 +359,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=current_start_frame).permute(
@@ -389,8 +379,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
start_frame=current_start_frame).permute(
@@ -404,6 +394,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
0, pred_flow.shape[:2])
break
index += 1
current_timestep = self.denoising_step_list[index]
# Step 3.2: record the model's output
output[:, current_start_frame:current_start_frame +
current_num_frames] = denoised_pred
@@ -416,6 +409,17 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
context_timestep).unflatten(0, denoised_pred.shape[:2])
with torch.no_grad():
if training_batch.video_latent is not None:
denoised_pred = torch.cat([
denoised_pred,
training_batch.
video_latent[:,
current_start_frame:current_start_frame +
current_num_frames],
torch.zeros_like(denoised_pred),
],
dim=2)
training_batch_temp = self._build_distill_input_kwargs(
denoised_pred, context_timestep,
training_batch.conditional_dict, training_batch)
@@ -430,8 +434,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
timestep=training_batch_temp.input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.
input_kwargs.get('encoder_hidden_states_image'),
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame * self.frame_seq_length,
start_frame=current_start_frame)
@@ -445,10 +449,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
# Slice last 21 frames if we generated more
gradient_mask = None
if pred_image_or_video.shape[1] > 21:
if pred_image_or_video.shape[1] > num_training_frames:
with torch.no_grad():
# Re-encode to get image latent
latent_to_decode = pred_image_or_video[:, :-20, ...]
latent_to_decode = pred_image_or_video[:, :-(
num_training_frames - 1), ...]
# Decode to video
latent_to_decode = latent_to_decode.permute(
0, 2, 1, 3, 4) # [B, C, F, H, W]
@@ -479,8 +484,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
image_latent = image_latent.permute(0, 2, 1, 3,
4) # [B, F, C, H, W]
pred_image_or_video_last_21 = torch.cat(
[image_latent, pred_image_or_video[:, -20:, ...]], dim=1)
pred_image_or_video_last_21 = torch.cat([
image_latent,
pred_image_or_video[:, -(num_training_frames - 1):, ...]
],
dim=1)
else:
pred_image_or_video_last_21 = pred_image_or_video
@@ -510,6 +518,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self.denoising_step_list[exit_flags[0]],
dtype=torch.float32,
device=self.device)
exit_timestep = self.denoising_step_list[exit_flags[0]]
# Store gradient mask information for debugging
if gradient_mask is not None:
@@ -523,11 +532,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
min_num_frames, dtype=torch.float32, device=self.device)
# Clean up caches
assert self.kv_cache1 is not None
assert self.crossattn_cache is not None
self._reset_simulation_caches(self.kv_cache1, self.crossattn_cache)
assert kv_cache1 is not None
assert crossattn_cache is not None
self._reset_simulation_caches(kv_cache1, crossattn_cache)
return final_output if gradient_mask is not None else pred_image_or_video
return (final_output, exit_timestep) if gradient_mask is not None else (pred_image_or_video, exit_timestep)
def _initialize_simulation_caches(
self,
@@ -538,28 +547,36 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
max_num_frames: int | None = None,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
num_transformer_blocks = len(self.transformer.blocks)
num_transformer_blocks = self.transformer.config.num_layers
latent_shape = self.video_latent_shape_sp
_, num_frames, _, height, width = latent_shape
_, p_h, p_w = self.transformer.patch_size
post_patch_height = height // p_h
post_patch_width = width // p_w
if isinstance(self.transformer.config.patch_size, tuple) or isinstance(self.transformer.config.patch_size, list):
ph, pw = self.transformer.config.patch_size[
1], self.transformer.config.patch_size[2]
elif isinstance(self.transformer.config.patch_size, int):
ph, pw = self.transformer.config.patch_size, self.transformer.config.patch_size
else:
raise ValueError(
f"Unsupported patch size type: {type(self.transformer.config.patch_size)}"
)
post_patch_height = height // ph
post_patch_width = width // pw
frame_seq_length = post_patch_height * post_patch_width
self.frame_seq_length = frame_seq_length
# Get model configuration parameters - handle FSDP wrapping
num_attention_heads = getattr(self.transformer, 'num_attention_heads',
None)
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
None)
text_len = getattr(self.transformer, 'text_len', None)
num_attention_heads = self.transformer.config.num_attention_heads
attention_head_dim = self.transformer.config.attention_head_dim
text_len = getattr(self.transformer.config, 'text_len', None)
if max_num_frames is None:
max_num_frames = num_frames
num_max_frames = max(max_num_frames, num_frames)
kv_cache_size = num_max_frames * frame_seq_length
local_attn_size = getattr(self.transformer.config, 'local_attn_size',
-1)
kv_cache_size = num_max_frames * frame_seq_length if local_attn_size == -1 else local_attn_size * frame_seq_length
kv_cache = []
for _ in range(num_transformer_blocks):
@@ -796,6 +813,44 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
if self.training_args.use_gan_loss:
# Freeze the base fake-score models and only train GAN head params.
# We must handle both experts because timestep routing can select either.
fake_score_models = [self.fake_score_transformer]
if self.fake_score_transformer_2 is not None:
fake_score_models.append(self.fake_score_transformer_2)
for fs_model in fake_score_models:
fs_model.requires_grad_(False)
for name, param in fs_model.named_parameters():
if "_cls_pred_branch" in name or "_gan_ca_blocks" in name or "_register_tokens" in name:
param.requires_grad_(True)
# Prevent accidental gradient carry-over from the previous critic update.
self.fake_score_optimizer.zero_grad(set_to_none=True)
if self.fake_score_transformer_2 is not None:
self.fake_score_optimizer_2.zero_grad(set_to_none=True)
batch_gan_critic = copy.deepcopy(training_batch)
batch_gan_critic, real_latent = self._get_next_batch_2(batch_gan_critic)
with set_forward_context(current_timestep=batch_gan_critic.timesteps,
attn_metadata=batch_gan_critic.attn_metadata):
critic_gan_loss = self._gan_critic_forward(batch_gan_critic, real_latent)
(critic_gan_loss / gradient_accumulation_steps).backward()
total_critic_loss += critic_gan_loss.detach().item()
if self.train_fake_score_transformer_2 and self.fake_score_transformer_2 is not None:
self._clip_model_grad_norm_(batch_gan_critic,
self.fake_score_transformer_2)
self.fake_score_optimizer_2.step()
self.fake_score_lr_scheduler_2.step()
else:
self._clip_model_grad_norm_(batch_gan_critic,
self.fake_score_transformer)
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
for fs_model in fake_score_models:
fs_model.requires_grad_(True)
avg_critic_loss = torch.tensor(total_critic_loss /
gradient_accumulation_steps,
device=self.device)
@@ -979,6 +1034,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
logger.info("Starting training from scratch")
self.train_loader_iter = iter(self.train_dataloader)
if getattr(self, "train_dataloader_2", None) is not None:
self.train_loader_iter_2 = iter(self.train_dataloader_2)
step_times: deque[float] = deque(maxlen=100)
+94 -51
View File
@@ -13,19 +13,18 @@ import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
# from fastvideo.attention.backends.video_sparse_attn import (
# VideoSparseAttentionMetadataBuilder)
# from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory_text_only, pyarrow_schema_t2v, pyarrow_schema_text_only
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_local_torch_device, get_sp_group,
@@ -47,9 +46,14 @@ from fastvideo.training.training_utils import (
shard_latents_across_sp)
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
from fastvideo.optim.muon import get_muon_optimizer
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
try:
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
except:
vsa_available = False
vmoba_available = False
logger = init_logger(__name__)
@@ -88,7 +92,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self) -> None:
self.train_dataset_schema = pyarrow_schema_t2v
self.train_dataset_schema = pyarrow_schema_text_only
self.train_dataset_schema_2 = pyarrow_schema_t2v
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
@@ -108,7 +113,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
set_random_seed(self.seed + self.global_rank)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
@@ -127,16 +132,25 @@ class TrainingPipeline(LoRAPipeline, ABC):
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
# Parse betas from string format "beta1,beta2"
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.optimizer = get_muon_optimizer(
self.transformer,
lr=training_args.learning_rate,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(
f"Unsupported optimizer type: {training_args.optimizer_type}")
self.init_steps = 0
logger.info("optimizer: %s", self.optimizer)
@@ -156,13 +170,25 @@ class TrainingPipeline(LoRAPipeline, ABC):
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.optimizer_type == "adamw":
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
elif training_args.optimizer_type == "muon":
self.optimizer_2 = get_muon_optimizer(
self.transformer_2,
lr=training_args.learning_rate,
weight_decay=training_args.weight_decay,
)
else:
raise ValueError(
f"Unsupported optimizer type: {training_args.optimizer_type}"
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
@@ -186,6 +212,19 @@ class TrainingPipeline(LoRAPipeline, ABC):
text_len, # type: ignore[attr-defined]
seed=self.seed)
if getattr(training_args, 'data_path_2', None) is not None:
self.train_dataset_2, self.train_dataloader_2 = build_parquet_map_style_dataloader(
training_args.data_path_2,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema_2,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=self.seed)
self.noise_scheduler = noise_scheduler
if self.training_args.boundary_ratio is not None:
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
@@ -359,7 +398,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
patch_size = self.training_args.pipeline_config.dit_config.patch_size
current_vsa_sparsity = training_batch.current_vsa_sparsity
assert latents_shape is not None
assert training_batch.timesteps is not None
# assert training_batch.timesteps is not None
if envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
if not vsa_available:
raise ImportError(
@@ -590,15 +629,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed)
device=current_platform.device_name).manual_seed(self.seed +
self.global_rank)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
self.seed + self.global_rank)
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -663,26 +702,28 @@ class TrainingPipeline(LoRAPipeline, ABC):
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
try:
metrics["batch_size"] = int(
training_batch.raw_latent_shape[0])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (
training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] // patch_h) * (
training_batch.raw_latent_shape[4] // patch_w)
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
arch_config = self.training_args.pipeline_config.dit_config.arch_config
arch_config = self.training_args.pipeline_config.dit_config.arch_config
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
except Exception:
pass
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
@@ -695,12 +736,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_visualization and step % self.training_args.visualization_steps == 0:
self.visualize_intermediate_latents(training_batch,
self.training_args, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
with self.profiler_controller.region(
"profiler_region_training_validation"):
if self.training_args.log_visualization:
self.visualize_intermediate_latents(
training_batch, self.training_args, step)
self._log_validation(self.transformer, self.training_args,
step)
gpu_memory_usage = current_platform.get_torch_device(
+199 -50
View File
@@ -257,8 +257,6 @@ def save_distillation_checkpoint(
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -290,8 +288,6 @@ def save_distillation_checkpoint(
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
if generator_ema_2 is not None:
generator_2_states["ema"] = generator_ema_2.state_dict()
generator_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
@@ -417,6 +413,67 @@ def save_distillation_checkpoint(
rank,
local_main_process_only=False)
# Persist EMA separately to avoid shape mismatches across ranks.
# Supports:
# - mode="rank0_full": save consolidated EMA only on rank 0
# - mode="local_shard": save per-rank EMA shard for each rank
try:
if generator_ema is not None and getattr(generator_ema, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema,
generator_transformer, rank,
save_dir, "generator_ema")
elif generator_ema is not None and getattr(generator_ema, "mode",
None) == "local_shard":
# Save per-rank shard
ema_dir_shard = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard, exist_ok=True)
ema_shard_path = os.path.join(ema_dir_shard,
f"generator_ema_rank{rank}.pt")
torch.save(generator_ema.state_dict(), ema_shard_path)
logger.info(
"rank: %s, saved generator EMA shard (local_shard) to %s",
rank,
ema_shard_path,
local_main_process_only=False)
# Also consolidate EMA to a single full-state file on rank 0 by applying EMA to the model and gathering
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema, generator_transformer, rank, save_dir,
"generator_ema")
except Exception as e:
logger.warning("rank: %s, failed saving EMA separately: %s", rank,
str(e))
try:
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema_2,
generator_transformer_2, rank,
save_dir, "generator_ema_2")
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "local_shard":
# Save per-rank shard for EMA_2
ema_dir_shard_2 = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard_2, exist_ok=True)
ema2_shard_path = os.path.join(ema_dir_shard_2,
f"generator_ema_2_rank{rank}.pt")
torch.save(generator_ema_2.state_dict(), ema2_shard_path)
logger.info(
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
rank,
ema2_shard_path,
local_main_process_only=False)
# Also consolidate EMA_2 to a single full-state file on rank 0
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema_2, generator_transformer_2, rank, save_dir,
"generator_ema_2")
except Exception as e:
logger.warning("rank: %s, failed saving EMA_2 separately: %s", rank,
str(e))
# Save generator model weights (consolidated) for inference
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
device=None)
@@ -454,46 +511,45 @@ def save_distillation_checkpoint(
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(
generator_transformer_2, device=None)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(generator_transformer_2,
device=None)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(inference_save_dir_2,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2,
"config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2, "config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
def load_checkpoint(transformer,
@@ -644,18 +700,38 @@ def load_distillation_checkpoint(
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
# Load EMA separately if saved in rank0_full mode
if generator_ema is not None:
try:
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully",
rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
if getattr(generator_ema, "mode", None) == "rank0_full":
ema_path = os.path.join(checkpoint_path, "ema",
"generator_ema.pt")
if rank == 0 and os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA (rank0_full) loaded from %s",
rank, ema_path)
elif rank == 0:
logger.info(
"rank: %s, generator EMA file not found at %s; skipping",
rank, ema_path)
elif getattr(generator_ema, "mode", None) == "local_shard":
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
f"generator_ema_rank{rank}.pt")
if os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA shard (local_shard) loaded from %s",
rank, ema_path)
else:
logger.info(
"rank: %s, generator EMA shard file not found at %s; skipping",
rank, ema_path)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank,
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
str(e))
# Load generator_2 distributed checkpoint (MoE support)
@@ -850,7 +926,7 @@ def load_distillation_checkpoint(
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * 0.476986
return latents * vae.config.scaling_factor
elif model_type == "wan":
latents_mean = torch.tensor(vae.latents_mean)
latents_std = 1.0 / torch.tensor(vae.latents_std)
@@ -1169,6 +1245,71 @@ def custom_to_hf_state_dict(
return new_state_dict
def _save_full_ema_safetensors_from_state(
state_dict: dict[str, Any],
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
output_path: str,
) -> None:
"""
Convert a training-format state_dict to HF format and save as safetensors.
"""
diffusers_state_dict = custom_to_hf_state_dict(state_dict,
reverse_param_names_mapping)
save_file(diffusers_state_dict, output_path)
def _save_rank0_full_ema_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
if rank != 0:
return
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
ema_state = ema.state_dict()
_save_full_ema_safetensors_from_state(ema_state,
module.reverse_param_names_mapping,
output_path)
logger.info("rank: %s, saved %s as consolidated EMA safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
def _consolidate_local_shard_ema_and_save_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
try:
# Temporarily apply EMA to the live (sharded) module and gather full CPU state on rank 0
with ema.apply_to_model(module):
cpu_state_full = gather_state_dict_on_cpu_rank0(module, device=None)
if rank == 0:
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
_save_full_ema_safetensors_from_state(
cpu_state_full, module.reverse_param_names_mapping, output_path)
logger.info(
"rank: %s, saved consolidated %s EMA (from local_shard) as safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
except Exception as ce:
logger.warning(
"rank: %s, failed consolidating %s EMA (local_shard): %s", rank,
base_name, str(ce))
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
@@ -1177,6 +1318,14 @@ def shift_timestep(timestep: torch.Tensor, shift: float,
denominator = 1 + (shift - 1) * t
return num_train_timestep * (shift * t / denominator)
def unshift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
# inverse of shift_timestep
if shift == 1:
return timestep
t = timestep / num_train_timestep
denominator = shift - (shift - 1) * t
return num_train_timestep * (t / denominator)
# coding=utf-8
# Copyright 2025 The HuggingFace Inc. team.
+4 -1
View File
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except:
vsa_available = False
logger = init_logger(__name__)
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..7}; do
# Calculate the starting file number for this node
start_file=$((node_id * 4 + 1))
echo "Launching node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 3)).txt"
echo "sbatch --job-name=ode-${node_id} --output=preprocess_output/preprocess-node-${node_id}.out --error=preprocess_output/preprocess-node-${node_id}.err slurms/syn.slurm $start_file $node_id"
sbatch --job-name=ode-${node_id} \
--output=preprocess_output/preprocess-node-${node_id}.out \
--error=preprocess_output/preprocess-node-${node_id}.err \
/home/hal-weiz/FastVideo/hy15/syn.slurm $start_file $node_id
done
echo "All 8 nodes launched successfully!"
+35
View File
@@ -0,0 +1,35 @@
#!/bin/bash
# Define source and destination directories
SOURCE_DIR="/home/hal-weiz/hy15_tf_init_3333/prompts_16k"
DEST_DIR="/home/hal-weiz/hy15_tf_init_3333/prompts_16k_32_files"
# Create destination directory if it doesn't exist
mkdir -p "$DEST_DIR"
echo "Merging files from $SOURCE_DIR to $DEST_DIR..."
# Loop to create 32 merged files from 64 source files
for i in {1..32}; do
# Calculate source file indices
# i=1 -> uses 1 and 2
# i=2 -> uses 3 and 4
# ...
# i=32 -> uses 63 and 64
idx1=$(( (i-1)*2 + 1 ))
idx2=$(( (i-1)*2 + 2 ))
file1="${SOURCE_DIR}/v2m_${idx1}.txt"
file2="${SOURCE_DIR}/v2m_${idx2}.txt"
outfile="${DEST_DIR}/v2m_${i}.txt"
if [[ -f "$file1" && -f "$file2" ]]; then
# Concatenate the two files
cat "$file1" "$file2" > "$outfile"
echo "Created v2m_${i}.txt from v2m_${idx1}.txt and v2m_${idx2}.txt"
else
echo "Warning: Missing source files for v2m_${i}.txt (expected v2m_${idx1}.txt and v2m_${idx2}.txt)"
fi
done
echo "Done. Created 32 merged files in $DEST_DIR"
+87
View File
@@ -0,0 +1,87 @@
#!/bin/bash
#SBATCH --job-name=hy15_ode_trajectory
#SBATCH --nodes=1
#SBATCH --gres=gpu:4
#SBATCH --ntasks-per-node=4
#SBATCH --cpus-per-task=16
#SBATCH --output=hy15_ode_trajectory/%j.out
#SBATCH --error=hy15_ode_trajectory/%j.err
# conda init
source ~/.venv/bin/activate
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
set -euo pipefail
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
export MODEL_BASE=hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3)
GPU_NUM=1
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
# Run 8 parallel preprocessing jobs on this node (in 2 batches of 4)
for i in {1..4}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
DATA_MERGE_PATH="/home/hal-weiz/hy15_tf_init_3333/prompts_16k_32_files/v2m_${file_num}.txt"
# Create unique output directory based on node and GPU
OUTPUT_DIR="/home/hal-weiz/fv-ode-preprocessing-16k-hy15-121/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
# Distribute 16 CPUs across 4 GPUs (4 CPUs per GPU)
start_cpu=$(( (i-1) * 4 ))
end_cpu=$(( start_cpu + 3 ))
echo "Starting GPU $gpu processing file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 848 \
--num_frames 121 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 24 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "hy15_ode_trajectory" &
done
# Wait for all jobs on this node to complete
wait
echo "All processing blocks completed!"
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..7}; do
# Calculate the starting file number for this node
start_file=$((node_id * 8 + 1))
echo "Launching node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
echo "sbatch --job-name=ode-${node_id} --output=preprocess_output/preprocess-node-${node_id}.out --error=preprocess_output/preprocess-node-${node_id}.err slurms/syn.slurm $start_file $node_id"
sbatch --job-name=ode-${node_id} \
--output=preprocess_output/preprocess-node-${node_id}.out \
--error=preprocess_output/preprocess-node-${node_id}.err \
/mnt/home/zhouw.jerry2017/FastVideo/wan_ode_preprocess/syn.slurm $start_file $node_id
done
echo "All 8 nodes launched successfully!"
+99
View File
@@ -0,0 +1,99 @@
#!/bin/bash
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=8
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=16
#SBATCH --mem=1440G
#SBATCH --exclusive
#SBATCH --time=72:00:00
# conda init
# source ~/conda/miniconda/bin/activate
source ~/miniconda3/bin/activate
PYTHON_VIRTUAL_ENVIRONMENT=fastvideo
conda activate $PYTHON_VIRTUAL_ENVIRONMENT
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
num_gpus=1
export MODEL_BASE=wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3 4 5 6 7)
GPU_NUM=1
MODEL_TYPE="wan"
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
# Run 8 parallel preprocessing jobs on this node
for i in {1..8}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
# DATA_MERGE_PATH="/mnt/data/fv-ode-preprocessing-6k-mixkit-wan-1.3b/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
DATA_MERGE_PATH="/mnt/data/mixkit_wan_1.3b_processed_t2v_distributed/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
# INIT_WEIGHTS_FROM_SAFETENSORS="/mnt/data/wan_tf_init_3333_mixkit_shift5/checkpoint-2400/transformer/diffusion_pytorch_model.safetensors"
INIT_WEIGHTS_FROM_SAFETENSORS="/mnt/data/wan_ar_diffusion_3333_mixkit_shift5/checkpoint-2400/transformer/diffusion_pytorch_model.safetensors"
# Create unique output directory based on node and GPU
OUTPUT_DIR="/mnt/data/fv-tf-ode-preprocessing-6k-mixkit-wan-1.3b/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
end_cpu=$(( start_cpu+1 ))
echo "Starting GPU $gpu processing file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--init_weights_from_safetensors $INIT_WEIGHTS_FROM_SAFETENSORS \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "tf_ode" &
done
# Wait for all jobs on this node to complete
wait
echo "All processing blocks completed!"
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..7}; do
# Calculate the starting file number for this node
start_file=$((node_id * 8 + 1))
echo "Launching node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
echo "sbatch --job-name=ode-${node_id} --output=preprocess_output/preprocess-node-${node_id}.out --error=preprocess_output/preprocess-node-${node_id}.err slurms/syn.slurm $start_file $node_id"
sbatch --job-name=ode-${node_id} \
--output=preprocess_output/preprocess-node-${node_id}.out \
--error=preprocess_output/preprocess-node-${node_id}.err \
/mnt/home/zhouw.jerry2017/FastVideo/wan_preprocess/syn.slurm $start_file $node_id
done
echo "All 8 nodes launched successfully!"
+95
View File
@@ -0,0 +1,95 @@
#!/bin/bash
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=8
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=16
#SBATCH --mem=960G
#SBATCH --exclusive
#SBATCH --time=72:00:00
# conda init
# source ~/conda/miniconda/bin/activate
source ~/miniconda3/bin/activate
PYTHON_VIRTUAL_ENVIRONMENT=fastvideo
conda activate $PYTHON_VIRTUAL_ENVIRONMENT
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
num_gpus=1
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3 4 5 6 7)
GPU_NUM=1
MODEL_TYPE="wan"
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
# Run 8 parallel preprocessing jobs on this node
for i in {1..8}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
DATA_MERGE_PATH="/mnt/home/zhouw.jerry2017/mixkit_prompts/mixkit_prompts_${file_num}.txt"
# Create unique output directory based on node and GPU
OUTPUT_DIR="/mnt/data/fv-ode-preprocessing-6k-mixkit-wan-1.3b/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
end_cpu=$(( start_cpu+1 ))
echo "Starting GPU $gpu processing file mixkit_prompts_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory" &
done
# Wait for all jobs on this node to complete
wait
echo "All processing blocks completed!"
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..7}; do
# Calculate the starting file number for this node
start_file=$((node_id * 8 + 1))
echo "Launching node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
echo "sbatch --job-name=ode-${node_id} --output=preprocess_output/preprocess-node-${node_id}.out --error=preprocess_output/preprocess-node-${node_id}.err slurms/syn.slurm $start_file $node_id"
sbatch --job-name=ode-${node_id} \
--output=preprocess_output/preprocess-node-${node_id}.out \
--error=preprocess_output/preprocess-node-${node_id}.err \
/mnt/home/zhouw.jerry2017/FastVideo/wan_text_preprocess/syn.slurm $start_file $node_id
done
echo "All 8 nodes launched successfully!"
+95
View File
@@ -0,0 +1,95 @@
#!/bin/bash
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=8
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=16
#SBATCH --mem=960G
#SBATCH --exclusive
#SBATCH --time=72:00:00
# conda init
# source ~/conda/miniconda/bin/activate
source ~/miniconda3/bin/activate
PYTHON_VIRTUAL_ENVIRONMENT=fastvideo
conda activate $PYTHON_VIRTUAL_ENVIRONMENT
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
num_gpus=1
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
export HF_HOME=/mnt/data/huggingface
export CUDA_HOME=/usr/local/cuda
export PATH="$CUDA_HOME/bin:$PATH"
export LD_LIBRARY_PATH="$LD_LIBRARY_PATH:/home/shared-bin/local/usr/lib/aarch64-linux-gnu:$CUDA_HOME/lib64"
export CUDACXX="$CUDA_HOME/bin/nvcc"
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3 4 5 6 7)
GPU_NUM=1
MODEL_TYPE="wan"
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
# Run 8 parallel preprocessing jobs on this node
for i in {1..8}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
DATA_MERGE_PATH="/mnt/home/zhouw.jerry2017/vidprom_shards/vidprom_${file_num}.txt"
# Create unique output directory based on node and GPU
OUTPUT_DIR="/mnt/data/vidprom_filtered_extended_umt5_text_embed/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
end_cpu=$(( start_cpu+1 ))
echo "Starting GPU $gpu processing file mixkit_prompts_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "text_only" &
done
# Wait for all jobs on this node to complete
wait
echo "All processing blocks completed!"