Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3c71218987 | ||
|
|
6b4cc8772d | ||
|
|
0df9934f4d | ||
|
|
8140cb4460 | ||
|
|
e42f688466 | ||
|
|
4062953769 | ||
|
|
8cebfdad67 | ||
|
|
d5d061a1f7 | ||
|
|
07db3796d2 | ||
|
|
363d81f3bb | ||
|
|
2e0c5fa979 | ||
|
|
448490b838 | ||
|
|
5546fa9aad | ||
|
|
9718cd974e | ||
|
|
8a623d5aa7 | ||
|
|
d481af1d36 | ||
|
|
eda4148681 | ||
|
|
9cf2a1561a |
@@ -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[@]}"
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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,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":
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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()),
|
||||
])
|
||||
])
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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!"
|
||||
Executable
+35
@@ -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"
|
||||
@@ -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!"
|
||||
@@ -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!"
|
||||
@@ -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!"
|
||||
@@ -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!"
|
||||
@@ -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!"
|
||||
@@ -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!"
|
||||
@@ -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!"
|
||||
Reference in New Issue
Block a user