Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
38f9c41e46 | ||
|
|
b1f88ad5d5 | ||
|
|
90a598bd9e | ||
|
|
90d86d5a79 | ||
|
|
7d373cd2c4 | ||
|
|
0806218156 | ||
|
|
8c6056fbe2 | ||
|
|
e4705349d0 | ||
|
|
50e63840c6 | ||
|
|
f1ec0cde18 |
@@ -0,0 +1,169 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=dmd_3333
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_3333_output/dmd_3333_%j.out
|
||||
#SBATCH --error=dmd_3333_output/dmd_3333_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate wei-fv-distill
|
||||
export HOME="/mnt/weka/home/hao.zhang/wei"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=32
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Hunyuan1.5:
|
||||
GENERATOR_MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
# REAL_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled"
|
||||
# FAKE_SCORE_MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
# DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/vidprom_16k_text_embed"
|
||||
DATA_DIR_2="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/fv-ode-preprocessing-16k-hy15-121"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name SFhy1.5_t2v_distill_self_forcing_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_3333_1e-5_1e-5_cfg6_corrected_scheduler"
|
||||
--wandb_run_name "self_forcing_3333_context_forcing"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 848
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--simulate_generator_forward
|
||||
--log_visualization
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
# --resume-from-checkpoint "/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_1333/checkpoint-300"
|
||||
--init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_worldplay_df_init_3333/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors
|
||||
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wei/hy15_distilled_ode_init_3333/checkpoint-2400/transformer/diffusion_pytorch_model.safetensors
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_TOTAL_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_TOTAL_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
# --data_path_2 "$DATA_DIR_2"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "8"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
--text-encoder-cpu-offload
|
||||
# --vae_cpu_offload True
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 100
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 200
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--warp_denoising_step
|
||||
--dmd_denoising_steps '1000,875,750,625,500,375,250,125'
|
||||
# --dmd_denoising_steps '1000,760,520,280'
|
||||
# --dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.5
|
||||
--fake_score_learning_rate 5e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
)
|
||||
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
# --use-context-forcing True
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/hy15_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -2,14 +2,14 @@ from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
OUTPUT_PATH = "video_samples_hy15_t2v_distilled"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
@@ -18,15 +18,35 @@ def main():
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
# init_weights_from_safetensors="/mnt/weka/home/hao.zhang/wei/SFhy1.5_distill_1333_5e-5_2e-6_cfg3.5/checkpoint-500/ema/generator_ema.safetensors"
|
||||
)
|
||||
|
||||
# json_path = "/mnt/weka/home/hao.zhang/wei/FastVideo/data/mixkit_i2v_full_720p.json"
|
||||
# with open(json_path, 'r') as f:
|
||||
# data_list = json.load(f)["data"]
|
||||
|
||||
# # Now you can index into data_list however you like
|
||||
# # For example: data_list[0], data_list[1:3], etc.
|
||||
# for data in data_list:
|
||||
# prompt = data["prompt"]
|
||||
# image_path = data["image_path"]
|
||||
# generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, image_path=image_path, num_frames=121, fps=24)
|
||||
# return
|
||||
|
||||
# prompt = (
|
||||
# "In a brightly lit studio, a photographer wearing a denim jacket focuses intently, capturing shots with a professional camera. Facing him, a model stands gracefully, adjusting her long, flowing hair with delicate movements. The scene is characterized by strong contrasts; the model's soft pink attire and gentle gestures complement the rugged, precise demeanor of the photographer. Positioned against a minimalist backdrop, the pair work seamlessly, with the camera\u2019s lens pointed directly at the model, capturing her elegance. The soft, diffused lighting casts a gentle glow on both subjects, creating an airy and ethereal atmosphere perfect for a high-fashion photo shoot."
|
||||
# )
|
||||
|
||||
# video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=121, fps=24, image_path="data/1.png")
|
||||
# return
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
@@ -35,7 +55,7 @@ def main():
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,36 +1,38 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_MODE=offline
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
DATA_DIR="data/ode-preprocessing-hy15-test/"
|
||||
VALIDATION_DATASET_FILE="data/validation_64.json"
|
||||
NUM_GPUS=1
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--output_dir "ode_init_hy15_test"
|
||||
--wandb_run_name "vidprom_bz128_1e-5"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
# --warp_denoising_step
|
||||
# --log_visualization
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_latent_t 19
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--warp_denoising_step
|
||||
--num_frames 73
|
||||
--dmd_denoising_steps "1000,750,500,250"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--num_gpus 1
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
@@ -51,18 +53,17 @@ dataset_args=(
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
# --log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
@@ -78,6 +79,7 @@ miscellaneous_args=(
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=hy15_ode_3333
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=hy15_ode_3333_output/hy15_ode_3333_%j.out
|
||||
#SBATCH --error=hy15_ode_3333_output/hy15_ode_3333_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export NCCL_DEBUG_SUBSYS=INIT,NET
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate wei-fv-distill
|
||||
export HOME="/mnt/weka/home/hao.zhang/wei"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=64
|
||||
|
||||
MODEL_PATH="weizhou03/SFHunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/wl/sharefs/wl/release/fv-ode-preprocessing-16k-hy15-121"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "hy15_ode_init"
|
||||
--output_dir "/mnt/weka/home/hao.zhang/wei/hy15_ode_init_1333_new"
|
||||
--wandb_run_name "hy15_ode_init_1333_new"
|
||||
--warp_denoising_step
|
||||
--dmd_denoising_steps '1000,760,520,280,0'
|
||||
--log_visualization
|
||||
--visualization_steps 100
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 31
|
||||
--num_height 480
|
||||
--num_width 848
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_TOTAL_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_TOTAL_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--text-encoder-cpu-offload
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 400
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
# --log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 10
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
# --vae_cpu_offload True
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 6
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -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
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -127,7 +127,10 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
|
||||
def ode_text_only_record_creator(
|
||||
video_name: str, text_embedding: np.ndarray, caption: str,
|
||||
trajectory_latents: np.ndarray,
|
||||
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
|
||||
trajectory_timesteps: np.ndarray,
|
||||
text_mask: np.ndarray | None = None,
|
||||
text_embedding_2: np.ndarray | None = None,
|
||||
text_mask_2: np.ndarray | None = None) -> dict[str, Any]:
|
||||
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
|
||||
|
||||
Args:
|
||||
@@ -165,6 +168,25 @@ def ode_text_only_record_creator(
|
||||
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
|
||||
})
|
||||
|
||||
if text_embedding_2 is not None:
|
||||
record.update({
|
||||
"text_embedding_2_bytes": text_embedding_2.tobytes(),
|
||||
"text_embedding_2_shape": list(text_embedding_2.shape),
|
||||
"text_embedding_2_dtype": str(text_embedding_2.dtype),
|
||||
})
|
||||
if text_mask is not None:
|
||||
record.update({
|
||||
"text_mask_bytes": text_mask.tobytes(),
|
||||
"text_mask_shape": list(text_mask.shape),
|
||||
"text_mask_dtype": str(text_mask.dtype),
|
||||
})
|
||||
if text_mask_2 is not None:
|
||||
record.update({
|
||||
"text_mask_2_bytes": text_mask_2.tobytes(),
|
||||
"text_mask_2_shape": list(text_mask_2.shape),
|
||||
"text_mask_2_dtype": str(text_mask_2.dtype),
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
|
||||
@@ -187,4 +209,4 @@ def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"caption": caption,
|
||||
}
|
||||
return record
|
||||
return record
|
||||
@@ -90,6 +90,15 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_embedding_2_bytes", pa.binary()),
|
||||
pa.field("text_embedding_2_shape", pa.list_(pa.int64())),
|
||||
pa.field("text_embedding_2_dtype", pa.string()),
|
||||
pa.field("text_mask_bytes", pa.binary()),
|
||||
pa.field("text_mask_shape", pa.list_(pa.int64())),
|
||||
pa.field("text_mask_dtype", pa.string()),
|
||||
pa.field("text_mask_2_bytes", pa.binary()),
|
||||
pa.field("text_mask_2_shape", pa.list_(pa.int64())),
|
||||
pa.field("text_mask_2_dtype", pa.string()),
|
||||
# --- ODE Trajectory ---
|
||||
pa.field("trajectory_latents_bytes", pa.binary()),
|
||||
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
|
||||
@@ -115,4 +124,4 @@ pyarrow_schema_text_only = pa.schema([
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("caption", pa.string()),
|
||||
])
|
||||
])
|
||||
@@ -17,6 +17,7 @@ from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -655,7 +656,7 @@ class TextDataset(torch.utils.data.IterableDataset,
|
||||
self.seed = seed
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer_path = os.path.join(maybe_download_model(args.model_path), "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
@@ -758,4 +759,4 @@ class TextDataset(torch.utils.data.IterableDataset,
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
@@ -154,8 +154,14 @@ def collate_rows_from_parquet_schema(rows,
|
||||
) if rng else random.random()) < cfg_rate:
|
||||
data = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
if row[f"{tensor_name}_dtype"] == "float32":
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
elif row[f"{tensor_name}_dtype"] == "int64":
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.int64).reshape(shape).copy()
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {row[f"{tensor_name}_dtype"]}")
|
||||
tensor = torch.from_numpy(data)
|
||||
# if len(data.shape) == 3:
|
||||
# B, L, D = tensor.shape
|
||||
@@ -168,7 +174,7 @@ def collate_rows_from_parquet_schema(rows,
|
||||
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
|
||||
|
||||
# Stack tensors with special handling for text embeddings
|
||||
if tensor_name == 'text_embedding':
|
||||
if tensor_name == 'null':
|
||||
# Handle text embeddings with padding
|
||||
padded_tensors = []
|
||||
attention_masks = []
|
||||
|
||||
@@ -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"
|
||||
@@ -943,6 +946,8 @@ class TrainingArgs(FastVideoArgs):
|
||||
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
use_context_forcing: bool = False
|
||||
use_ode_init: bool = False
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -998,6 +1003,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to parquet files")
|
||||
parser.add_argument("--data-path-2",
|
||||
type=str,
|
||||
required=False,
|
||||
help="Path to parquet files")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
@@ -1091,6 +1100,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--validation-steps",
|
||||
type=float,
|
||||
help="Number of validation steps")
|
||||
parser.add_argument("--visualization-steps",
|
||||
type=float,
|
||||
help="Number of visualization steps")
|
||||
parser.add_argument("--log-validation",
|
||||
action=StoreBoolean,
|
||||
help="Whether to log validation results")
|
||||
@@ -1139,6 +1151,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--gradient-accumulation-steps",
|
||||
type=int,
|
||||
help="Number of steps to accumulate gradients")
|
||||
parser.add_argument("--optimizer-type",
|
||||
type=str,
|
||||
choices=["adamw", "muon"],
|
||||
help="Optimizer type")
|
||||
parser.add_argument("--learning-rate",
|
||||
type=float,
|
||||
required=True,
|
||||
@@ -1365,6 +1381,12 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
parser.add_argument("--use-context-forcing",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use context forcing")
|
||||
parser.add_argument("--use-ode-init",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use ODE init")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@@ -0,0 +1,808 @@
|
||||
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Any, Dict, Optional, List
|
||||
import math
|
||||
|
||||
import torch
|
||||
# import torch._dynamo
|
||||
# torch._dynamo.config.cache_size_limit = 128
|
||||
# try:
|
||||
# torch._dynamo.config.recompile_limit = 128
|
||||
# except AttributeError:
|
||||
# pass
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
|
||||
from torch.nn.attention.flex_attention import BlockMask
|
||||
flex_attention = torch.compile(
|
||||
flex_attention, dynamic=False, mode="default")
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.configs.models.dits import HunyuanVideo15Config
|
||||
from fastvideo.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
# TODO(will-PY-refactor): RMSNorm ....
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed, _apply_rotary_emb
|
||||
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
|
||||
unpatchify)
|
||||
from fastvideo.models.dits.base import CachableDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
from fastvideo.models.dits.hunyuanvideo15 import (
|
||||
HunyuanRMSNorm,
|
||||
HunyuanVideo15TimeEmbedding,
|
||||
HunyuanVideo15ByT5TextProjection,
|
||||
HunyuanVideo15ImageProjection,
|
||||
SingleTokenRefiner,
|
||||
FinalLayer)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
A multimodal DiT block with separate modulation for text and image/video,
|
||||
using distributed attention and linear layers.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
dtype: torch.dtype | None = None,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.head_dim = hidden_size // num_attention_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
# Image modulation components
|
||||
self.img_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mod",
|
||||
)
|
||||
|
||||
# Fused operations for image stream
|
||||
self.img_attn_norm = LayerNormScaleShift(hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.img_mlp_residual = ScaleResidual()
|
||||
|
||||
# Image attention components
|
||||
self.img_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_qkv")
|
||||
|
||||
self.img_attn_q_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
|
||||
self.img_attn_k_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
self.img_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_proj")
|
||||
|
||||
self.img_mlp = MLP(hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mlp")
|
||||
|
||||
# Text modulation components
|
||||
self.txt_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.txt_mod",
|
||||
)
|
||||
|
||||
# Fused operations for text stream
|
||||
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.txt_mlp_residual = ScaleResidual()
|
||||
|
||||
# Text attention components
|
||||
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
# QK norm layers for text
|
||||
self.txt_attn_q_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
|
||||
self.txt_attn_k_norm = HunyuanRMSNorm(self.head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
self.txt_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
|
||||
|
||||
self.max_attention_size = 21 * 1590 if local_attn_size == -1 else local_attn_size * 1590
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=self.num_attention_heads,
|
||||
head_size=self.head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn"
|
||||
)
|
||||
|
||||
def forward_txt(
|
||||
self,
|
||||
txt: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
cache_txt: bool = False,
|
||||
):
|
||||
txt_mod_outputs = self.txt_mod(vec)
|
||||
(
|
||||
txt_attn_shift,
|
||||
txt_attn_scale,
|
||||
txt_attn_gate,
|
||||
txt_mlp_shift,
|
||||
txt_mlp_scale,
|
||||
txt_mlp_gate,
|
||||
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
|
||||
|
||||
# Prepare text for attention using fused operation
|
||||
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
|
||||
|
||||
# Get QKV for text
|
||||
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
|
||||
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
|
||||
self.num_attention_heads, -1)
|
||||
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
|
||||
2]
|
||||
# Apply QK-Norm if needed
|
||||
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
|
||||
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
|
||||
|
||||
t_kv = {}
|
||||
if cache_txt:
|
||||
t_kv["k_txt"] = txt_k
|
||||
t_kv["v_txt"] = txt_v
|
||||
|
||||
txt_attn = self.attn(txt_q, txt_k, txt_v)
|
||||
# Process text attention output
|
||||
txt_attn_out, _ = self.txt_attn_proj(
|
||||
txt_attn.reshape(batch_size, text_seq_len, -1))
|
||||
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
|
||||
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
|
||||
|
||||
# Process text MLP
|
||||
txt_mlp_out = self.txt_mlp(txt_mlp_input)
|
||||
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
|
||||
|
||||
return txt, t_kv
|
||||
|
||||
def forward_vision(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple,
|
||||
block_mask: BlockMask,
|
||||
kv_cache: dict | None = None,
|
||||
txt_kv_cache: list | None = None,
|
||||
current_start: int = 0,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# Process modulation vectors
|
||||
if vec.dim() == 3:
|
||||
img_mod_outputs = self.img_mod(vec).unflatten(dim=-1, sizes=(6, -1))
|
||||
(
|
||||
img_attn_shift,
|
||||
img_attn_scale,
|
||||
img_attn_gate,
|
||||
img_mlp_shift,
|
||||
img_mlp_scale,
|
||||
img_mlp_gate,
|
||||
) = torch.chunk(img_mod_outputs, 6, dim=2)
|
||||
else:
|
||||
img_mod_outputs = self.img_mod(vec)
|
||||
(
|
||||
img_attn_shift,
|
||||
img_attn_scale,
|
||||
img_attn_gate,
|
||||
img_mlp_shift,
|
||||
img_mlp_scale,
|
||||
img_mlp_gate,
|
||||
) = torch.chunk(img_mod_outputs, 6, dim=-1)
|
||||
|
||||
# Prepare image for attention using fused operation
|
||||
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
|
||||
# Get QKV for image
|
||||
img_qkv, _ = self.img_attn_qkv(img_attn_input)
|
||||
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
|
||||
self.num_attention_heads, -1)
|
||||
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
|
||||
2]
|
||||
|
||||
# Apply QK-Norm if needed
|
||||
img_q = self.img_attn_q_norm(img_q).to(img_v)
|
||||
img_k = self.img_attn_k_norm(img_k).to(img_v)
|
||||
|
||||
# Apply rotary embeddings
|
||||
cos, sin = freqs_cis
|
||||
img_q = _apply_rotary_emb(img_q, cos, sin, is_neox_style=False)
|
||||
img_k = _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
|
||||
|
||||
# Apply flex_attention
|
||||
# Does not support SP padding for now
|
||||
if kv_cache is None:
|
||||
q = img_q
|
||||
k = torch.cat([img_k, txt_kv_cache["k_txt"]], dim=1)
|
||||
v = torch.cat([img_v, txt_kv_cache["v_txt"]], dim=1)
|
||||
# Padding for flex attention
|
||||
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
|
||||
padded_kv_length = math.ceil(k.shape[1] / 128) * 128 - k.shape[1]
|
||||
padded_roped_query = torch.cat(
|
||||
[q,
|
||||
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
|
||||
device=q.device, dtype=v.dtype)],
|
||||
dim=1
|
||||
)
|
||||
|
||||
padded_roped_key = torch.cat(
|
||||
[k, torch.zeros([k.shape[0], padded_kv_length, k.shape[2], k.shape[3]],
|
||||
device=k.device, dtype=v.dtype)],
|
||||
dim=1
|
||||
)
|
||||
|
||||
padded_v = torch.cat(
|
||||
[v, torch.zeros([v.shape[0], padded_kv_length, v.shape[2], v.shape[3]],
|
||||
device=v.device, dtype=v.dtype)],
|
||||
dim=1
|
||||
)
|
||||
|
||||
img_attn = flex_attention(
|
||||
query=padded_roped_query.transpose(2, 1),
|
||||
key=padded_roped_key.transpose(2, 1),
|
||||
value=padded_v.transpose(2, 1),
|
||||
block_mask=block_mask
|
||||
)[:, :, :-padded_length].transpose(2, 1)
|
||||
|
||||
assert img_attn.shape[1] == image_seq_len
|
||||
updated_kv_cache = None
|
||||
else:
|
||||
current_end = current_start + img_q.shape[1]
|
||||
num_new_tokens = img_q.shape[1]
|
||||
sink_tokens = self.sink_size * 1590
|
||||
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
|
||||
kv_cache_size = self.max_attention_size
|
||||
|
||||
# Clone cache to avoid in-place modification during gradient checkpointing
|
||||
k_cache = kv_cache["k"].clone()
|
||||
v_cache = kv_cache["v"].clone()
|
||||
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
# Calculate the number of new tokens added in this step
|
||||
# Shift existing cache content left to discard oldest tokens
|
||||
# Clone the source slice to avoid overlapping memory error
|
||||
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
|
||||
num_rolled_tokens = kv_cache_size - num_new_tokens - sink_tokens
|
||||
k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# Insert the new keys/values at the end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
assert local_end_index == self.max_attention_size
|
||||
else:
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
assert local_start_index >= 0
|
||||
q = img_q
|
||||
k = torch.cat([k_cache[:, :local_start_index], img_k, txt_kv_cache["k_txt"]], dim=1)
|
||||
v = torch.cat([v_cache[:, :local_start_index], img_v, txt_kv_cache["v_txt"]], dim=1)
|
||||
img_attn = self.attn(q, k, v)
|
||||
|
||||
k_cache[:, local_start_index:local_end_index] = img_k
|
||||
v_cache[:, local_start_index:local_end_index] = img_v
|
||||
|
||||
updated_kv_cache = {
|
||||
"k": k_cache,
|
||||
"v": v_cache,
|
||||
"global_end_index": torch.tensor([current_end], dtype=torch.long, device=k_cache.device),
|
||||
"local_end_index": torch.tensor([local_end_index], dtype=torch.long, device=k_cache.device)
|
||||
}
|
||||
|
||||
img_attn_out, _ = self.img_attn_proj(
|
||||
img_attn.view(batch_size, image_seq_len, -1))
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
|
||||
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
|
||||
|
||||
# Process image MLP
|
||||
img_mlp_out = self.img_mlp(img_mlp_input)
|
||||
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
|
||||
|
||||
return img, updated_kv_cache
|
||||
|
||||
def forward(
|
||||
self,
|
||||
txt_inference=False,
|
||||
vision_inference=False,
|
||||
**kwargs
|
||||
):
|
||||
if txt_inference:
|
||||
return self.forward_txt(**kwargs)
|
||||
elif vision_inference:
|
||||
return self.forward_vision(**kwargs)
|
||||
else:
|
||||
raise ValueError("txt_inference and vision_inference cannot be both False")
|
||||
|
||||
|
||||
class CausalHunyuanVideo15Transformer3DModel(CachableDiT):
|
||||
r"""
|
||||
A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5).
|
||||
"""
|
||||
|
||||
# shard single stream, double stream blocks, and refiner_blocks
|
||||
_fsdp_shard_conditions = HunyuanVideo15Config()._fsdp_shard_conditions
|
||||
_compile_conditions = HunyuanVideo15Config()._compile_conditions
|
||||
_supported_attention_backends = HunyuanVideo15Config(
|
||||
)._supported_attention_backends
|
||||
param_names_mapping = HunyuanVideo15Config().param_names_mapping
|
||||
reverse_param_names_mapping = HunyuanVideo15Config(
|
||||
).reverse_param_names_mapping
|
||||
lora_param_names_mapping = HunyuanVideo15Config().lora_param_names_mapping
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: HunyuanVideo15Config,
|
||||
hf_config: dict[str, Any],
|
||||
) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.out_channels = config.out_channels or config.in_channels
|
||||
self.patch_size = (config.patch_size_t, config.patch_size, config.patch_size)
|
||||
|
||||
# 1. Latent and condition embedders
|
||||
self.img_in = PatchEmbed(self.patch_size,
|
||||
config.in_channels,
|
||||
self.hidden_size,
|
||||
prefix=f"{config.prefix}.img_in")
|
||||
self.image_embedder = HunyuanVideo15ImageProjection(config.image_embed_dim, self.hidden_size)
|
||||
|
||||
self.txt_in = SingleTokenRefiner(config.text_embed_dim,
|
||||
self.hidden_size,
|
||||
config.num_attention_heads,
|
||||
depth=config.num_refiner_layers,
|
||||
dtype=None,
|
||||
prefix=f"{config.prefix}.txt_in")
|
||||
|
||||
self.txt_in_2 = HunyuanVideo15ByT5TextProjection(config.text_embed_2_dim, 2048, self.hidden_size)
|
||||
|
||||
self.time_in = HunyuanVideo15TimeEmbedding(self.hidden_size, use_meanflow=config.use_meanflow)
|
||||
|
||||
self.cond_type_embed = nn.Embedding(3, self.hidden_size)
|
||||
|
||||
# 3. Dual stream transformer blocks
|
||||
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
MMDoubleStreamBlock(
|
||||
hidden_size=self.hidden_size,
|
||||
num_attention_heads=config.num_attention_heads,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
local_attn_size=config.local_attn_size,
|
||||
sink_size=config.sink_size,
|
||||
dtype=None,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{config.prefix}.double_blocks.{i}"
|
||||
)
|
||||
for i in range(config.num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 5. Output projection
|
||||
self.final_layer = FinalLayer(self.hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
prefix=f"{config.prefix}.final_layer")
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.num_frame_per_block = config.num_frames_per_block
|
||||
self.local_attn_size = config.local_attn_size
|
||||
self.block_mask = None
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def get_text_and_mask(
|
||||
self,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_2: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
encoder_attention_mask_2: torch.Tensor,
|
||||
encoder_hidden_states_image: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
):
|
||||
batch_size, txt_seq_len = encoder_hidden_states.shape[0], encoder_hidden_states.shape[1]
|
||||
# qwen text embedding
|
||||
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep, encoder_attention_mask)
|
||||
|
||||
encoder_hidden_states_cond_emb = self.cond_type_embed(
|
||||
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
|
||||
)
|
||||
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
|
||||
|
||||
# byt5 text embedding
|
||||
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
|
||||
|
||||
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
|
||||
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
|
||||
)
|
||||
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
|
||||
|
||||
# image embed
|
||||
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
|
||||
is_t2v = torch.all(encoder_hidden_states_image == 0)
|
||||
if is_t2v:
|
||||
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
|
||||
encoder_attention_mask_3 = torch.zeros(
|
||||
(batch_size, encoder_hidden_states_3.shape[1]),
|
||||
dtype=encoder_attention_mask.dtype,
|
||||
device=encoder_attention_mask.device,
|
||||
)
|
||||
else:
|
||||
encoder_attention_mask_3 = torch.ones(
|
||||
(batch_size, encoder_hidden_states_3.shape[1]),
|
||||
dtype=encoder_attention_mask.dtype,
|
||||
device=encoder_attention_mask.device,
|
||||
)
|
||||
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
|
||||
2
|
||||
* torch.ones_like(
|
||||
encoder_hidden_states_3[:, :, 0],
|
||||
dtype=torch.long,
|
||||
)
|
||||
)
|
||||
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
|
||||
|
||||
# reorder and combine text tokens: combine valid tokens first, then padding
|
||||
encoder_attention_mask = encoder_attention_mask.bool()
|
||||
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
|
||||
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
|
||||
new_encoder_hidden_states = []
|
||||
new_encoder_attention_mask = []
|
||||
|
||||
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
encoder_hidden_states_2,
|
||||
encoder_attention_mask_2,
|
||||
encoder_hidden_states_3,
|
||||
encoder_attention_mask_3,
|
||||
):
|
||||
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
|
||||
new_encoder_hidden_states.append(
|
||||
torch.cat(
|
||||
[
|
||||
image[image_mask], # valid image
|
||||
text_2[text_mask_2], # valid byt5
|
||||
text[text_mask], # valid mllm
|
||||
image[~image_mask], # invalid image (zeroed)
|
||||
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
|
||||
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
# Apply same reordering to attention masks
|
||||
new_encoder_attention_mask.append(
|
||||
torch.cat(
|
||||
[
|
||||
image_mask[image_mask],
|
||||
text_mask_2[text_mask_2],
|
||||
text_mask[text_mask],
|
||||
image_mask[~image_mask],
|
||||
text_mask_2[~text_mask_2],
|
||||
text_mask[~text_mask],
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
|
||||
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
|
||||
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
|
||||
assert encoder_hidden_states.shape[0] == 1
|
||||
return encoder_hidden_states, encoder_attention_mask
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1,
|
||||
text_seq_len: int = 0
|
||||
) -> BlockMask:
|
||||
"""
|
||||
we will divide the token sequence into the following format
|
||||
[1 latent frame] [1 latent frame] ... [1 latent frame]
|
||||
We use flexattention to construct the attention mask
|
||||
"""
|
||||
total_length = num_frames * frame_seqlen
|
||||
total_kv_length = total_length + text_seq_len
|
||||
|
||||
total_length_tensor = torch.tensor(total_length, device=device)
|
||||
total_kv_length_tensor = torch.tensor(total_kv_length, device=device)
|
||||
|
||||
# we do right padding to get to a multiple of 128
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
kv_padded_length = math.ceil(total_kv_length / 128) * 128 - total_kv_length
|
||||
|
||||
ends = torch.zeros(total_length + padded_length,
|
||||
device=device, dtype=torch.long)
|
||||
|
||||
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
|
||||
frame_indices = torch.arange(
|
||||
# start=frame_seqlen,
|
||||
start=0,
|
||||
end=total_length,
|
||||
step=frame_seqlen * num_frame_per_block,
|
||||
device=device
|
||||
)
|
||||
# frame_indices = torch.cat([torch.tensor([0], device=device), frame_indices])
|
||||
|
||||
for i, tmp in enumerate(frame_indices):
|
||||
# if i == 0:
|
||||
# ends[tmp:tmp + frame_seqlen] = tmp + frame_seqlen
|
||||
# else:
|
||||
ends[tmp:tmp + frame_seqlen * num_frame_per_block] = tmp + \
|
||||
frame_seqlen * num_frame_per_block
|
||||
|
||||
def attention_mask(b, h, q_idx, kv_idx):
|
||||
if local_attn_size == -1:
|
||||
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx) | ((kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor))
|
||||
else:
|
||||
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx) | ((kv_idx >= total_length_tensor) & (kv_idx < total_kv_length_tensor))
|
||||
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
|
||||
|
||||
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
|
||||
KV_LEN=total_kv_length + kv_padded_length, _compile=False, device=device)
|
||||
|
||||
# if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
# print(
|
||||
# f" cache a block wise causal mask with block size of {num_frame_per_block} frames")
|
||||
# print(block_mask)
|
||||
|
||||
# import imageio
|
||||
# import numpy as np
|
||||
# from torch.nn.attention.flex_attention import create_mask
|
||||
|
||||
# mask = create_mask(attention_mask, B=None, H=None, Q_LEN=total_length +
|
||||
# padded_length, KV_LEN=total_length + padded_length, device=device)
|
||||
# import cv2
|
||||
# mask = cv2.resize(mask[0, 0].cpu().float().numpy(), (1024, 1024))
|
||||
# imageio.imwrite("mask_%d.jpg" % (0), np.uint8(255. * mask))
|
||||
|
||||
return block_mask
|
||||
|
||||
def forward_txt(
|
||||
self,
|
||||
encoder_hidden_states: List[torch.Tensor],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: Optional[List[torch.Tensor]] = None,
|
||||
encoder_attention_mask: Optional[List[torch.Tensor]] = None,
|
||||
guidance: Optional[torch.Tensor] = None,
|
||||
timestep_r: Optional[torch.LongTensor] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
cache_txt: bool = False,
|
||||
**kwargs
|
||||
):
|
||||
# Check that the timestep is only consisted of 0s
|
||||
assert torch.all(timestep == 0), "Timestep for txt must be only consisted of 0s"
|
||||
|
||||
if cache_txt:
|
||||
_kv_cache_new = []
|
||||
transformer_num_layers = len(self.double_blocks)
|
||||
for _ in range(transformer_num_layers):
|
||||
_kv_cache_new.append(
|
||||
{"k_vision": None, "v_vision": None, "k_txt": None, "v_txt": None}
|
||||
)
|
||||
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
|
||||
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
|
||||
|
||||
# 2. Conditional embeddings
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
|
||||
else:
|
||||
temb = self.time_in(timestep, timestep_r=timestep_r)
|
||||
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
|
||||
|
||||
encoder_hidden_states, encoder_attention_mask = self.get_text_and_mask(
|
||||
encoder_hidden_states,
|
||||
encoder_hidden_states_2,
|
||||
encoder_attention_mask,
|
||||
encoder_attention_mask_2,
|
||||
encoder_hidden_states_image,
|
||||
timestep
|
||||
)
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states[encoder_attention_mask.bool().to(encoder_hidden_states.device)].unsqueeze(0)
|
||||
|
||||
# 4. Transformer blocks
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
encoder_hidden_states, t_kv = block(
|
||||
txt_inference=True,
|
||||
vision_inference=False,
|
||||
txt=encoder_hidden_states,
|
||||
vec=temb,
|
||||
cache_txt=cache_txt,
|
||||
)
|
||||
|
||||
if cache_txt:
|
||||
_kv_cache_new[index]["k_txt"] = t_kv["k_txt"]
|
||||
_kv_cache_new[index]["v_txt"] = t_kv["v_txt"]
|
||||
|
||||
if cache_txt:
|
||||
return _kv_cache_new
|
||||
|
||||
def forward_vision(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
guidance: Optional[torch.Tensor] = None,
|
||||
timestep_r: Optional[torch.LongTensor] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
kv_cache: dict | None = None,
|
||||
txt_kv_cache: list | None = None,
|
||||
current_start: int = 0,
|
||||
rope_start_idx: int = 0,
|
||||
):
|
||||
assert txt_kv_cache is not None, "txt_kv_cache must be provided"
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# 1. RoPE
|
||||
# Get rotary embeddings
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||
(post_patch_num_frames, post_patch_height, post_patch_width), self.hidden_size,
|
||||
self.num_attention_heads, self.config.rope_axes_dim, self.config.rope_theta, start_frame=rope_start_idx)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# 2. Conditional embeddings
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
|
||||
else:
|
||||
temb = self.time_in(timestep, timestep_r=timestep_r)
|
||||
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
|
||||
|
||||
hidden_states = self.img_in(hidden_states)
|
||||
|
||||
# Prepare block-wise causal attention mask
|
||||
if kv_cache is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
text_seq_len=txt_kv_cache[0]["k_txt"].shape[1]
|
||||
)
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block_index, block in enumerate(self.double_blocks):
|
||||
hidden_states, new_cache = self._gradient_checkpointing_func(
|
||||
block,
|
||||
txt_inference=False,
|
||||
vision_inference=True,
|
||||
img=hidden_states,
|
||||
vec=temb,
|
||||
freqs_cis=freqs_cis,
|
||||
block_mask=self.block_mask,
|
||||
kv_cache=kv_cache[block_index] if kv_cache is not None else None,
|
||||
txt_kv_cache=txt_kv_cache[block_index],
|
||||
current_start=current_start
|
||||
)
|
||||
if new_cache is not None and kv_cache is not None:
|
||||
for k in new_cache.keys():
|
||||
kv_cache[block_index][k] = new_cache[k].clone()
|
||||
|
||||
else:
|
||||
for block_index, block in enumerate(self.double_blocks):
|
||||
hidden_states, new_cache = block(
|
||||
txt_inference=False,
|
||||
vision_inference=True,
|
||||
img=hidden_states,
|
||||
vec=temb,
|
||||
freqs_cis=freqs_cis,
|
||||
block_mask=self.block_mask,
|
||||
kv_cache=kv_cache[block_index] if kv_cache is not None else None,
|
||||
txt_kv_cache=txt_kv_cache[block_index],
|
||||
current_start=current_start
|
||||
)
|
||||
if new_cache is not None and kv_cache is not None:
|
||||
for k in new_cache.keys():
|
||||
kv_cache[block_index][k] = new_cache[k].clone()
|
||||
|
||||
# Final layer processing
|
||||
hidden_states = self.final_layer(hidden_states, temb)
|
||||
# Unpatchify to get original shape
|
||||
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
|
||||
|
||||
return hidden_states, kv_cache
|
||||
|
||||
def forward(
|
||||
self,
|
||||
txt_inference=False,
|
||||
vision_inference=False,
|
||||
**kwargs,
|
||||
):
|
||||
if txt_inference:
|
||||
return self.forward_txt(**kwargs)
|
||||
elif vision_inference:
|
||||
return self.forward_vision(**kwargs)
|
||||
else:
|
||||
raise ValueError("txt_inference and vision_inference cannot be both False")
|
||||
@@ -127,8 +127,9 @@ class HunyuanVideo15TimeEmbedding(nn.Module):
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
timestep_r: Optional[torch.Tensor] = None,
|
||||
timestep_seq_len: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
timesteps_emb = self.timestep_embedder(timestep)
|
||||
timesteps_emb = self.timestep_embedder(timestep, timestep_seq_len)
|
||||
|
||||
if timestep_r is not None:
|
||||
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
|
||||
@@ -473,6 +474,7 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
|
||||
guidance: Optional[torch.Tensor] = None,
|
||||
timestep_r: Optional[torch.LongTensor] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
**kwargs
|
||||
):
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
|
||||
@@ -494,7 +496,12 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# 2. Conditional embeddings
|
||||
temb = self.time_in(timestep, timestep_r=timestep_r)
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
temb = self.time_in(timestep.flatten(), timestep_r=timestep_r, timestep_seq_len=ts_seq_len)
|
||||
else:
|
||||
temb = self.time_in(timestep, timestep_r=timestep_r)
|
||||
# temb is [bs, seq_len, inner_dim] if ts_seq_len is not None, otherwise [bs, inner_dim]
|
||||
|
||||
hidden_states = self.img_in(hidden_states)
|
||||
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
|
||||
@@ -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
|
||||
@@ -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:
|
||||
@@ -829,7 +831,7 @@ class TransformerLoader(ComponentLoader):
|
||||
)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
logger.info("Loaded model with %.2fB parameters, with cpu_offload: %s", total_params / 1e9, fastvideo_args.dit_cpu_offload)
|
||||
|
||||
assert next(model.parameters()).dtype == default_dtype, (
|
||||
"Model dtype does not match default dtype"
|
||||
|
||||
@@ -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:
|
||||
|
||||
"""
|
||||
|
||||
@@ -5,6 +5,9 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# TODO(PY): move it elsewhere
|
||||
def auto_attributes(init_func):
|
||||
|
||||
@@ -663,7 +663,7 @@ class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
|
||||
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
|
||||
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
|
||||
# intermediate tiles together, the memory requirement can be lowered.
|
||||
self.use_tiling = False
|
||||
self.use_tiling = True
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
|
||||
@@ -0,0 +1,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
|
||||
@@ -84,6 +84,7 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
# fastvideo_args.dit_cpu_offload = False
|
||||
with self.profiler_controller.region("profiler_region_model_loading"):
|
||||
self.modules = self.load_modules(fastvideo_args, loaded_modules)
|
||||
|
||||
@@ -287,6 +288,7 @@ class ComposedPipelineBase(ABC):
|
||||
# remove keys that are not pipeline modules
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
model_index.pop("_name_or_path", None)
|
||||
model_index.pop("workload_type", None)
|
||||
if "boundary_ratio" in model_index and model_index[
|
||||
"boundary_ratio"] is not None:
|
||||
|
||||
@@ -232,6 +232,7 @@ class TrainingBatch:
|
||||
image_latents: torch.Tensor | None = None
|
||||
infos: list[dict[str, Any]] | None = None
|
||||
mask_lat_size: torch.Tensor | None = None
|
||||
video_latent: torch.Tensor | None = None
|
||||
|
||||
# ODE trajectory supervision
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
@@ -240,6 +241,9 @@ class TrainingBatch:
|
||||
# Transformer inputs
|
||||
noisy_model_input: torch.Tensor | None = None
|
||||
timesteps: torch.Tensor | None = None
|
||||
use_gt_trajectory: bool = False
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
start_timestep_index: int | None = None
|
||||
sigmas: torch.Tensor | None = None
|
||||
noise: torch.Tensor | None = None
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -303,7 +303,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
with torch.no_grad():
|
||||
# Filter out invalid samples (those with all zeros)
|
||||
valid_indices = []
|
||||
for i, pixel_values in enumerate(data["pixel_values"]):
|
||||
@@ -422,4 +422,4 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
if num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
num_processed_samples = 0
|
||||
num_processed_samples = 0
|
||||
@@ -28,16 +28,12 @@ from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage, DenoisingStage, InputValidationStage, LatentPreparationStage,
|
||||
TextEncodingStage, TimestepPreparationStage, Hy15ImageEncodingStage)
|
||||
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -47,7 +43,8 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
"""ODE Trajectory preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
|
||||
"transformer", "scheduler"
|
||||
]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
@@ -61,19 +58,19 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
assert fastvideo_args.pipeline_config.flow_shift == 5
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
|
||||
denoising_strength=1.0)
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2")
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2")
|
||||
],
|
||||
))
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
@@ -82,6 +79,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None)))
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=Hy15ImageEncodingStage(image_encoder=None,
|
||||
image_processor=None))
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
@@ -95,11 +95,13 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
args):
|
||||
"""Preprocess text-only data and generate trajectory information."""
|
||||
|
||||
num_encoders = len(self.prompt_encoding_stage.text_encoders)
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
with torch.no_grad():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
@@ -130,12 +132,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
encoder_index=list(range(num_encoders)),
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
sampling_params = SamplingParam.from_pretrained(args.model_path)
|
||||
|
||||
@@ -144,61 +143,48 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
sampling_params.negative_prompt,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
encoder_index=list(range(num_encoders)),
|
||||
return_attention_mask=True,
|
||||
)
|
||||
negative_prompt_embed = negative_prompt_embeds_list[0][0]
|
||||
negative_prompt_attention_mask = negative_prompt_masks_list[
|
||||
0][0]
|
||||
else:
|
||||
negative_prompt_embed = None
|
||||
negative_prompt_attention_mask = None
|
||||
negative_prompt_embeds_list = []
|
||||
negative_prompt_masks_list = []
|
||||
|
||||
trajectory_latents = []
|
||||
trajectory_timesteps = []
|
||||
trajectory_decoded = []
|
||||
|
||||
for i, (prompt_embed, prompt_attention_mask) in enumerate(
|
||||
zip(prompt_embeds, prompt_attention_masks,
|
||||
strict=False)):
|
||||
prompt_embed = prompt_embed.unsqueeze(0)
|
||||
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
|
||||
# Collect the trajectory data (text-to-video generation)
|
||||
batch = ForwardBatch(**shallow_asdict(sampling_params), )
|
||||
batch.prompt_embeds = prompt_embeds_list
|
||||
batch.prompt_attention_mask = prompt_masks_list
|
||||
batch.negative_prompt_embeds = negative_prompt_embeds_list
|
||||
batch.negative_attention_mask = negative_prompt_masks_list
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
batch.return_trajectory_decoded = False
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
batch.num_frames = args.num_frames
|
||||
batch.fps = args.train_fps
|
||||
|
||||
# Collect the trajectory data (text-to-video generation)
|
||||
batch = ForwardBatch(**shallow_asdict(sampling_params), )
|
||||
batch.prompt_embeds = [prompt_embed]
|
||||
batch.prompt_attention_mask = [prompt_attention_mask]
|
||||
batch.negative_prompt_embeds = [negative_prompt_embed]
|
||||
batch.negative_attention_mask = [
|
||||
negative_prompt_attention_mask
|
||||
]
|
||||
batch.num_inference_steps = 48
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
batch.return_trajectory_decoded = False
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
batch.fps = args.train_fps
|
||||
batch.guidance_scale = 6.0
|
||||
batch.do_classifier_free_guidance = True
|
||||
result_batch = self.input_validation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.timestep_preparation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.latent_preparation_stage(
|
||||
result_batch, fastvideo_args)
|
||||
result_batch = self.image_encoding_stage(
|
||||
result_batch, fastvideo_args)
|
||||
result_batch = self.denoising_stage(result_batch,
|
||||
fastvideo_args)
|
||||
result_batch = self.decoding_stage(result_batch, fastvideo_args)
|
||||
|
||||
result_batch = self.input_validation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.timestep_preparation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.latent_preparation_stage(
|
||||
result_batch, fastvideo_args)
|
||||
result_batch = self.denoising_stage(result_batch,
|
||||
fastvideo_args)
|
||||
result_batch = self.decoding_stage(result_batch,
|
||||
fastvideo_args)
|
||||
|
||||
trajectory_latents.append(
|
||||
result_batch.trajectory_latents.cpu())
|
||||
trajectory_timesteps.append(
|
||||
result_batch.trajectory_timesteps.cpu())
|
||||
trajectory_decoded.append(result_batch.trajectory_decoded)
|
||||
trajectory_latents.append(result_batch.trajectory_latents.cpu())
|
||||
trajectory_timesteps.append(
|
||||
result_batch.trajectory_timesteps.cpu())
|
||||
trajectory_decoded.append(result_batch.trajectory_decoded)
|
||||
|
||||
# Prepare extra features for text-only processing
|
||||
extra_features = {
|
||||
@@ -209,10 +195,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
if batch.return_trajectory_decoded:
|
||||
for i, decoded_frames in enumerate(trajectory_decoded):
|
||||
for j, decoded_frame in enumerate(decoded_frames):
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
if j in [5, 7]:
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data: list[dict[str, Any]] = []
|
||||
@@ -227,7 +214,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
text_embedding = prompt_embeds_list[0].float().cpu().numpy()
|
||||
text_mask = prompt_masks_list[0].cpu().numpy()
|
||||
text_embedding_2 = prompt_embeds_list[1].float().cpu(
|
||||
).numpy()
|
||||
text_mask_2 = prompt_masks_list[1].cpu().numpy()
|
||||
|
||||
# Get extra features for this sample
|
||||
sample_extra_features = {}
|
||||
@@ -253,6 +244,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
"trajectory_latents"],
|
||||
trajectory_timesteps=sample_extra_features[
|
||||
"trajectory_timesteps"],
|
||||
text_embedding_2=text_embedding_2,
|
||||
text_mask=text_mask,
|
||||
text_mask_2=text_mask_2,
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
|
||||
@@ -58,7 +58,7 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
with torch.no_grad():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
@@ -181,4 +181,4 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
|
||||
self.preprocess_text_only(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_Text
|
||||
EntryClass = PreprocessPipeline_Text
|
||||
@@ -1,9 +1,7 @@
|
||||
import argparse
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.distributed import (
|
||||
get_world_size, maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -16,38 +14,38 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
|
||||
PreprocessPipeline_Text)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
# args.model_path = maybe_download_model(args.model_path)
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
print(pipeline_config.__class__.__name__)
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if args.preprocess_task == "text_only":
|
||||
kwargs = {
|
||||
"text_encoder_cpu_offload": False,
|
||||
}
|
||||
else:
|
||||
# Full config for video/image processing
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
|
||||
}
|
||||
pipeline_config.update_config_from_dict(kwargs)
|
||||
# kwargs: dict[str, Any] = {}
|
||||
# if args.preprocess_task == "text_only":
|
||||
# kwargs = {
|
||||
# "text_encoder_cpu_offload": False,
|
||||
# }
|
||||
# else:
|
||||
# # Full config for video/image processing
|
||||
# kwargs = {
|
||||
# "vae_precision": "fp32",
|
||||
# "vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
|
||||
# }
|
||||
# pipeline_config.update_config_from_dict(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=args.model_path,
|
||||
num_gpus=get_world_size(),
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
if args.preprocess_task == "t2v":
|
||||
@@ -134,4 +132,4 @@ if __name__ == "__main__":
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
@@ -24,4 +24,4 @@ if __name__ == "__main__":
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
fastvideo_args = FastVideoArgs.from_cli_args(args)
|
||||
main(fastvideo_args)
|
||||
main(fastvideo_args)
|
||||
@@ -8,6 +8,7 @@ complete diffusion pipelines.
|
||||
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
|
||||
from fastvideo.pipelines.stages.hy15_causal_denoising import Hy15CausalDMDDenosingStage
|
||||
from fastvideo.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
|
||||
@@ -52,6 +53,7 @@ __all__ = [
|
||||
"Cosmos25LatentPreparationStage",
|
||||
"LTX2LatentPreparationStage",
|
||||
"LTX2AudioDecodingStage",
|
||||
"Hy15CausalDMDDenosingStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
|
||||
@@ -45,12 +45,12 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
self.transformer_2 = transformer_2
|
||||
self.vae = vae
|
||||
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
||||
self.num_transformer_blocks = len(self.transformer.blocks)
|
||||
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
|
||||
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
|
||||
self.num_transformer_blocks = self.transformer.config.num_layers
|
||||
self.num_frames_per_block = self.transformer.config.num_frames_per_block
|
||||
self.sliding_window_num_frames = self.transformer.config.sliding_window_num_frames
|
||||
|
||||
try:
|
||||
self.local_attn_size = getattr(self.transformer.model,
|
||||
self.local_attn_size = getattr(self.transformer.config,
|
||||
"local_attn_size",
|
||||
-1) # type: ignore
|
||||
except Exception:
|
||||
@@ -412,8 +412,8 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
num_attention_heads = self.transformer.config.num_attention_heads
|
||||
attention_head_dim = self.transformer.config.attention_head_dim
|
||||
if self.local_attn_size != -1:
|
||||
kv_cache_size = self.local_attn_size * self.frame_seq_length
|
||||
else:
|
||||
|
||||
@@ -26,7 +26,7 @@ from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like, PRECISION_TO_TYPE
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
@@ -210,7 +210,10 @@ class DenoisingStage(PipelineStage):
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
assert self.vae is not None, "VAE is not provided for TI2V task"
|
||||
z = self.vae.encode(batch.pil_image).mean.float()
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
z = self.vae.encode(batch.pil_image.to(vae_dtype)).mean.float()
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
@@ -223,6 +226,9 @@ class DenoisingStage(PipelineStage):
|
||||
else:
|
||||
z = z * self.vae.scaling_factor
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae = self.vae.to('cpu')
|
||||
|
||||
latent_model_input = latent_model_input.squeeze(0)
|
||||
_, mask2 = masks_like([latent_model_input], zero=True)
|
||||
|
||||
@@ -232,18 +238,18 @@ class DenoisingStage(PipelineStage):
|
||||
latent_model_input = latent_model_input.to(get_local_torch_device())
|
||||
latents = latent_model_input
|
||||
F = batch.num_frames
|
||||
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
|
||||
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
seq_len = ((F - 1) // temporal_scale +
|
||||
1) * (batch.height // spatial_scale) * (
|
||||
batch.width // spatial_scale) // (patch_size[1] *
|
||||
patch_size[2])
|
||||
temporal_scale = fastvideo_args.pipeline_config.vae_config.temporal_compression_ratio
|
||||
spatial_scale = fastvideo_args.pipeline_config.vae_config.spatial_compression_ratio
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
|
||||
seq_len = ((F - 1) // temporal_scale + 1) * (
|
||||
batch.height // spatial_scale) * (
|
||||
batch.width // spatial_scale) // (patch_size * patch_size)
|
||||
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = [latents]
|
||||
|
||||
logger.info("timesteps: %s", timesteps)
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -408,6 +414,7 @@ class DenoisingStage(PipelineStage):
|
||||
**action_kwargs,
|
||||
)
|
||||
|
||||
assert batch.do_classifier_free_guidance, "do_classifier_free_guidance is not supported"
|
||||
if batch.do_classifier_free_guidance:
|
||||
batch.is_cfg_negative = True
|
||||
with set_forward_context(
|
||||
@@ -462,6 +469,7 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
trajectory_tensor: torch.Tensor | None = None
|
||||
if trajectory_latents:
|
||||
trajectory_timesteps.append(torch.zeros_like(t))
|
||||
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
|
||||
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
|
||||
dim=0)
|
||||
|
||||
@@ -0,0 +1,361 @@
|
||||
import math
|
||||
import torch # type: ignore
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
VideoSparseAttentionBackend = None # type: ignore
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class Hy15CausalDMDDenosingStage(CausalDMDDenosingStage):
|
||||
"""
|
||||
Denoising stage for causal diffusion.
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
|
||||
if isinstance(self.transformer.config.patch_size, tuple):
|
||||
patch_ratio = self.transformer.config.patch_size[
|
||||
1] * self.transformer.config.patch_size[2]
|
||||
elif isinstance(self.transformer.config.patch_size, int):
|
||||
patch_ratio = self.transformer.config.patch_size**2
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported patch size type: {type(self.transformer.config.patch_size)}"
|
||||
)
|
||||
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
# TODO(will): make this a parameter once we add i2v support
|
||||
independent_first_frame = self.transformer.independent_first_frame if hasattr(
|
||||
self.transformer, 'independent_first_frame') else False
|
||||
# Timesteps for DMD
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
self.scheduler.set_timesteps(num_inference_steps=1000,
|
||||
extra_one_step=True,
|
||||
device=get_local_torch_device())
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
logger.info("[causal_denoising] timesteps: %s", timesteps)
|
||||
|
||||
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
else:
|
||||
boundary_timestep = None
|
||||
high_noise_timesteps = None
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert not torch.isnan(
|
||||
image_embeds[0]).any(), "image_embeds contains nan"
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
# STA
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
# Latents and prompts
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents # [B, C, T, H, W]
|
||||
b, c, t, h, w = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
# Initialize or reset caches
|
||||
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
|
||||
pos_start_base = 0
|
||||
num_blocks = math.ceil(t / self.num_frames_per_block)
|
||||
block_sizes = [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
# Initialize txt kv cache
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch):
|
||||
txt_kv_cache = self.transformer(
|
||||
txt_inference=True,
|
||||
vision_inference=False,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
encoder_hidden_states_image=image_embeds,
|
||||
encoder_attention_mask=batch.prompt_attention_mask,
|
||||
timestep=torch.zeros([latents.shape[0]], device=latents.device),
|
||||
cache_txt=True,
|
||||
)
|
||||
|
||||
first_frame_latent = None
|
||||
if batch.pil_image is not None:
|
||||
# Causal video gen directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert self.vae is not None, "VAE is not provided for causal video gen task"
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
first_frame_latent = self.vae.encode(
|
||||
batch.pil_image.to(vae_dtype)).mean.float()
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
first_frame_latent -= self.vae.shift_factor.to(
|
||||
first_frame_latent.device, first_frame_latent.dtype)
|
||||
else:
|
||||
first_frame_latent -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
|
||||
first_frame_latent.device, first_frame_latent.dtype)
|
||||
else:
|
||||
first_frame_latent = first_frame_latent * self.vae.scaling_factor
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae = self.vae.to("cpu")
|
||||
|
||||
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
|
||||
t_zero = torch.zeros([latents.shape[0], 1],
|
||||
device=latents.device,
|
||||
dtype=torch.long)
|
||||
if batch.video_latent is not None:
|
||||
video_latent_chunk = batch.video_latent[:, :, start_index:
|
||||
start_index + 1, :, :]
|
||||
first_frame_input = torch.cat([
|
||||
first_frame_latent,
|
||||
video_latent_chunk,
|
||||
torch.zeros_like(first_frame_latent),
|
||||
],
|
||||
dim=1)
|
||||
else:
|
||||
first_frame_input = first_frame_latent.clone()
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch):
|
||||
self.transformer(
|
||||
txt_inference=False,
|
||||
vision_inference=True,
|
||||
hidden_states=first_frame_input.to(target_dtype),
|
||||
timestep=t_zero,
|
||||
kv_cache=kv_cache1,
|
||||
txt_kv_cache=txt_kv_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
rope_start_idx=start_index,
|
||||
)
|
||||
|
||||
start_index += 1
|
||||
block_sizes.pop(0)
|
||||
latents[:, :, :1, :, :] = first_frame_latent
|
||||
|
||||
vision_input_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"vision_inference": True,
|
||||
"txt_inference": False,
|
||||
},
|
||||
)
|
||||
|
||||
# DMD loop in causal blocks
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
for current_num_frames in block_sizes:
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
# use BTCHW for DMD conversion routines
|
||||
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
|
||||
video_raw_latent_shape = noise_latents_btchw.shape
|
||||
|
||||
for i, t_cur in enumerate(timesteps):
|
||||
# Copy for pred conversion
|
||||
noise_latents = noise_latents_btchw.clone()
|
||||
latent_model_input = current_latents.to(target_dtype)
|
||||
|
||||
if batch.video_latent is not None:
|
||||
video_latent_chunk = batch.video_latent[:, :,
|
||||
start_index:
|
||||
start_index +
|
||||
current_num_frames, :, :]
|
||||
latent_model_input = torch.cat([
|
||||
latent_model_input,
|
||||
video_latent_chunk,
|
||||
torch.zeros_like(current_latents),
|
||||
],
|
||||
dim=1)
|
||||
elif batch.image_latent is not None and independent_first_frame and start_index == 0:
|
||||
latent_model_input = torch.cat([
|
||||
latent_model_input,
|
||||
batch.image_latent.to(target_dtype)
|
||||
],
|
||||
dim=2)
|
||||
|
||||
# Prepare inputs
|
||||
t_expand = t_cur.repeat(latent_model_input.shape[0])
|
||||
|
||||
# Attention metadata if needed
|
||||
if (vsa_available and self.attn_backend
|
||||
== VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=i, # type: ignore
|
||||
raw_latent_shape=(current_num_frames, h,
|
||||
w), # type: ignore
|
||||
patch_size=fastvideo_args.pipeline_config.
|
||||
dit_config.patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.
|
||||
VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(), # type: ignore
|
||||
) # type: ignore
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(latent_model_input.shape[0], current_num_frames),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
pred_noise_btchw, kv_cache1 = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=t_expanded_noise,
|
||||
kv_cache=kv_cache1,
|
||||
txt_kv_cache=txt_kv_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
rope_start_idx=start_index,
|
||||
**vision_input_kwargs,
|
||||
)
|
||||
pred_noise_btchw = pred_noise_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
# Convert pred noise to pred video with FM Euler scheduler utilities
|
||||
pred_video_btchw = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
scheduler=self.scheduler).unflatten(
|
||||
0, pred_noise_btchw.shape[:2])
|
||||
|
||||
if i < len(timesteps) - 1:
|
||||
next_timestep = timesteps[i + 1] * torch.ones(
|
||||
[1],
|
||||
dtype=torch.long,
|
||||
device=pred_video_btchw.device)
|
||||
noise = torch.randn(
|
||||
video_raw_latent_shape,
|
||||
dtype=pred_video_btchw.dtype,
|
||||
generator=(batch.generator[0] if isinstance(
|
||||
batch.generator, list) else
|
||||
batch.generator)).to(self.device)
|
||||
noise_btchw = noise
|
||||
noise_latents_btchw = self.scheduler.add_noise(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1),
|
||||
next_timestep).unflatten(
|
||||
0, pred_video_btchw.shape[:2])
|
||||
current_latents = noise_latents_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
else:
|
||||
current_latents = pred_video_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
# Write back and advance
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = current_latents
|
||||
|
||||
# Re-run with context timestep to update KV cache using clean context
|
||||
context_noise = getattr(fastvideo_args.pipeline_config,
|
||||
"context_noise", 0)
|
||||
t_context = torch.ones([latents.shape[0], current_num_frames],
|
||||
device=latents.device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = current_latents.to(target_dtype)
|
||||
if batch.video_latent is not None:
|
||||
video_latent_chunk = batch.video_latent[:, :, start_index:
|
||||
start_index +
|
||||
current_num_frames, :, :]
|
||||
context_bcthw = torch.cat([
|
||||
context_bcthw,
|
||||
video_latent_chunk,
|
||||
torch.zeros_like(current_latents),
|
||||
],
|
||||
dim=1)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
|
||||
_, kv_cache1 = self.transformer(
|
||||
hidden_states=context_bcthw,
|
||||
timestep=t_context,
|
||||
kv_cache=kv_cache1,
|
||||
txt_kv_cache=txt_kv_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
rope_start_idx=start_index,
|
||||
**vision_input_kwargs,
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -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
|
||||
|
||||
@@ -120,9 +120,9 @@ class InputValidationStage(PipelineStage):
|
||||
else:
|
||||
# Standard Wan logic
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
|
||||
max_area = 480 * 832
|
||||
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
dh, dw = patch_size * vae_stride, patch_size * vae_stride
|
||||
max_area = 480 * 848
|
||||
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||
|
||||
scale = max(ow / iw, oh / ih)
|
||||
|
||||
@@ -28,8 +28,6 @@ from fastvideo.distributed import (cleanup_dist_env_and_memory,
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
TrainingBatch)
|
||||
@@ -42,6 +40,7 @@ from fastvideo.training.training_utils import (
|
||||
shift_timestep)
|
||||
from fastvideo.utils import (is_vsa_available, maybe_download_model,
|
||||
set_random_seed, verify_model_config_and_directory)
|
||||
from fastvideo.optim.muon import get_muon_optimizer
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
@@ -85,7 +84,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if training_args.real_score_model_path:
|
||||
logger.info("Loading real score transformer from: %s",
|
||||
training_args.real_score_model_path)
|
||||
training_args.override_transformer_cls_name = "WanTransformer3DModel"
|
||||
# TODO(will): can use deepcopy instead if the model is the same
|
||||
self.real_score_transformer = self.load_module_from_path(
|
||||
training_args.real_score_model_path, "transformer",
|
||||
@@ -111,7 +109,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if training_args.fake_score_model_path:
|
||||
logger.info("Loading fake score transformer from: %s",
|
||||
training_args.fake_score_model_path)
|
||||
training_args.override_transformer_cls_name = "WanTransformer3DModel"
|
||||
self.fake_score_transformer = self.load_module_from_path(
|
||||
training_args.fake_score_model_path, "transformer",
|
||||
training_args)
|
||||
@@ -146,8 +143,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=self.timestep_shift)
|
||||
# self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
# shift=self.timestep_shift)
|
||||
|
||||
if self.training_args.boundary_ratio is not None:
|
||||
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
|
||||
@@ -192,16 +189,24 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if fake_score_lr == 0.0:
|
||||
fake_score_lr = training_args.learning_rate
|
||||
|
||||
betas_str = training_args.fake_score_betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
self.fake_score_optimizer = torch.optim.AdamW(
|
||||
fake_score_params,
|
||||
lr=fake_score_lr,
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
if training_args.optimizer_type == "adamw":
|
||||
betas_str = training_args.fake_score_betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
self.fake_score_optimizer = torch.optim.AdamW(
|
||||
fake_score_params,
|
||||
lr=fake_score_lr,
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
elif training_args.optimizer_type == "muon":
|
||||
self.fake_score_optimizer = get_muon_optimizer(
|
||||
self.fake_score_transformer,
|
||||
lr=fake_score_lr,
|
||||
weight_decay=training_args.weight_decay,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
|
||||
|
||||
self.fake_score_lr_scheduler = get_scheduler(
|
||||
training_args.fake_score_lr_scheduler,
|
||||
@@ -218,13 +223,23 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fake_score_params_2 = list(
|
||||
filter(lambda p: p.requires_grad,
|
||||
self.fake_score_transformer_2.parameters()))
|
||||
self.fake_score_optimizer_2 = torch.optim.AdamW(
|
||||
fake_score_params_2,
|
||||
lr=fake_score_lr,
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
if training_args.optimizer_type == "adamw":
|
||||
self.fake_score_optimizer_2 = torch.optim.AdamW(
|
||||
fake_score_params_2,
|
||||
lr=fake_score_lr,
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
elif training_args.optimizer_type == "muon":
|
||||
self.fake_score_optimizer_2 = get_muon_optimizer(
|
||||
self.fake_score_transformer_2,
|
||||
lr=fake_score_lr,
|
||||
weight_decay=training_args.weight_decay,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported optimizer type: {training_args.optimizer_type}")
|
||||
|
||||
self.fake_score_lr_scheduler_2 = get_scheduler(
|
||||
training_args.fake_score_lr_scheduler,
|
||||
optimizer=self.fake_score_optimizer_2,
|
||||
@@ -272,21 +287,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
self.generator_ema_2: EMA_FSDP | None = None
|
||||
if (self.training_args.ema_decay
|
||||
is not None) and (self.training_args.ema_decay > 0.0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer,
|
||||
decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA with decay=%s",
|
||||
self.training_args.ema_decay)
|
||||
|
||||
# Initialize EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None:
|
||||
self.generator_ema_2 = EMA_FSDP(
|
||||
self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA_2 with decay=%s",
|
||||
self.training_args.ema_decay)
|
||||
else:
|
||||
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
|
||||
|
||||
def load_module_from_path(self, model_path: str, module_type: str,
|
||||
training_args: "TrainingArgs"):
|
||||
@@ -572,6 +572,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_hidden_states_image": training_batch.image_embeds,
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"timestep": timestep,
|
||||
"return_dict": False,
|
||||
@@ -723,10 +724,18 @@ class DistillationPipeline(TrainingPipeline):
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).detach().unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
noisy_latent_copy = noisy_latent.clone()
|
||||
|
||||
if training_batch.video_latent is not None:
|
||||
noisy_latent_copy = torch.cat([
|
||||
noisy_latent_copy,
|
||||
training_batch.video_latent,
|
||||
torch.zeros_like(noisy_latent_copy),
|
||||
],
|
||||
dim=2)
|
||||
# fake_score_transformer forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
noisy_latent_copy, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
current_fake_score_transformer = self._get_fake_score_transformer(
|
||||
timestep)
|
||||
@@ -742,7 +751,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
# real_score_transformer cond forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
noisy_latent_copy, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
current_real_score_transformer = self._get_real_score_transformer(
|
||||
timestep)
|
||||
@@ -758,7 +767,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
# real_score_transformer uncond forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.unconditional_dict,
|
||||
noisy_latent_copy, timestep, training_batch.unconditional_dict,
|
||||
training_batch)
|
||||
# Use same transformer as conditional forward for consistency
|
||||
real_score_pred_noise_uncond = current_real_score_transformer(
|
||||
@@ -779,9 +788,16 @@ class DistillationPipeline(TrainingPipeline):
|
||||
original_latent - real_score_pred_video).mean()
|
||||
grad = torch.nan_to_num(grad)
|
||||
|
||||
dmd_loss = 0.5 * F.mse_loss(
|
||||
original_latent.float(),
|
||||
(original_latent.float() - grad.float()).detach())
|
||||
if self.training_args.use_context_forcing and training_batch.trajectory_latents is not None:
|
||||
context_forcing_length = training_batch.trajectory_latents.shape[1]
|
||||
dmd_loss = 0.5 * F.mse_loss(
|
||||
original_latent.float()[:, context_forcing_length:],
|
||||
(original_latent.float()[:, context_forcing_length:] -
|
||||
grad.float()[:, context_forcing_length:]).detach())
|
||||
else:
|
||||
dmd_loss = 0.5 * F.mse_loss(
|
||||
original_latent.float(),
|
||||
(original_latent.float() - grad.float()).detach())
|
||||
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"training_batch_dmd_fwd_clean_latent":
|
||||
@@ -831,6 +847,13 @@ class DistillationPipeline(TrainingPipeline):
|
||||
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
|
||||
fake_score_timestep).unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
if training_batch.video_latent is not None:
|
||||
noisy_generator_pred_video = torch.cat([
|
||||
noisy_generator_pred_video,
|
||||
training_batch.video_latent,
|
||||
torch.zeros_like(noisy_generator_pred_video),
|
||||
],
|
||||
dim=2)
|
||||
|
||||
with set_forward_context(current_timestep=training_batch.timesteps,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
@@ -877,7 +900,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
super()._prepare_dit_inputs(training_batch)
|
||||
# super()._prepare_dit_inputs(training_batch)
|
||||
conditional_dict = {
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
@@ -898,7 +921,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.video_latent_shape = training_batch.latents.shape
|
||||
|
||||
self.video_latent_shape_sp = training_batch.latents.shape
|
||||
|
||||
|
||||
return training_batch
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
@@ -1239,6 +1262,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(
|
||||
sampling_param, training_args, validation_batch, steps)
|
||||
batch.prompt_attention_mask = []
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
@@ -1249,8 +1273,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
|
||||
batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
if len(result_batch.prompt_embeds) == 1:
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
else:
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds, result_batch.prompt_attention_mask
|
||||
|
||||
logger.info(
|
||||
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
@@ -1354,6 +1381,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if hasattr(self, 'transformer_2') and self.transformer_2 is not None:
|
||||
self.transformer_2.train()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import gc
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any, cast
|
||||
@@ -13,14 +14,15 @@ from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
# from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
# WanCausalDMDPipeline)
|
||||
from fastvideo.pipelines.basic.hunyuan15.hunyuan15_causal_dmd_pipeline import Hy15CausalDMDPipeline
|
||||
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases)
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -35,16 +37,18 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
- minimizing MSE to the stored next latent at timestep t_next
|
||||
"""
|
||||
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Match the preprocess/generation scheduler for consistent stepping
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# # Match the preprocess/generation scheduler for consistent stepping
|
||||
# self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
# shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
# sigma_min=0.0,
|
||||
# extra_one_step=True)
|
||||
# self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
|
||||
# training=True)
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
|
||||
@@ -56,18 +60,36 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.text_encoder = self.get_module("text_encoder")
|
||||
self.text_encoder.requires_grad_(False)
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder")
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer")
|
||||
],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
|
||||
# self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
# shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
|
||||
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
extra_one_step=True,
|
||||
device=get_local_torch_device())
|
||||
|
||||
logger.info("dmd_denoising_steps: %s",
|
||||
self.training_args.pipeline_config.dmd_denoising_steps)
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
self.dmd_denoising_steps = torch.tensor(
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
@@ -78,7 +100,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
logger.info("warped self.dmd_denoising_steps: %s",
|
||||
self.dmd_denoising_steps)
|
||||
else:
|
||||
raise ValueError("warp_denoising_step must be true")
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32))).cuda()
|
||||
self.dmd_denoising_steps = timesteps[self.dmd_denoising_steps]
|
||||
|
||||
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
|
||||
get_local_torch_device())
|
||||
@@ -98,7 +123,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
# Warm start validation with current transformer
|
||||
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
|
||||
self.validation_pipeline = Hy15CausalDMDPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
@@ -109,7 +134,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True)
|
||||
dit_cpu_offload=False)
|
||||
|
||||
def _get_next_batch(
|
||||
self,
|
||||
@@ -117,15 +142,40 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
# Required fields from parquet (ODE trajectory schema)
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
device = get_local_torch_device()
|
||||
encoder_hidden_states = batch['text_embedding'].to(
|
||||
device, dtype=torch.bfloat16).squeeze(0)
|
||||
encoder_hidden_states_2 = batch['text_embedding_2'].to(
|
||||
device, dtype=torch.bfloat16).squeeze(0)
|
||||
encoder_attention_mask = batch['text_mask'].to(
|
||||
device, dtype=torch.bfloat16).squeeze(0)
|
||||
encoder_attention_mask_2 = batch['text_mask_2'].to(
|
||||
device, dtype=torch.bfloat16).squeeze(0)
|
||||
encoder_hidden_states_image = [
|
||||
torch.zeros(1,
|
||||
729,
|
||||
1152,
|
||||
device=get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
]
|
||||
infos = batch['info_list']
|
||||
|
||||
if encoder_hidden_states.dim() < 3:
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
infos[0]["caption"],
|
||||
self.training_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
encoder_hidden_states = prompt_embeds_list[0].to(
|
||||
device, dtype=torch.bfloat16)
|
||||
encoder_attention_mask = prompt_masks_list[0].to(
|
||||
device, dtype=torch.bfloat16)
|
||||
|
||||
# Trajectory tensors may include a leading singleton batch dim per row
|
||||
trajectory_latents = batch['trajectory_latents']
|
||||
if trajectory_latents.dim() == 7:
|
||||
@@ -154,11 +204,13 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
|
||||
|
||||
# Move to device
|
||||
device = get_local_torch_device()
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
device, dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
device, dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = [
|
||||
encoder_hidden_states, encoder_hidden_states_2
|
||||
]
|
||||
training_batch.encoder_attention_mask = [
|
||||
encoder_attention_mask, encoder_attention_mask_2
|
||||
]
|
||||
training_batch.encoder_hidden_states_image = encoder_hidden_states_image
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch, trajectory_latents.to(
|
||||
@@ -197,6 +249,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
dtype=torch.long).repeat(1, num_frame)
|
||||
return timestep
|
||||
else:
|
||||
num_pad_frames = 0
|
||||
if num_frame % num_frame_per_block != 0:
|
||||
# Pad num_frame to be divisible by num_frame_per_block
|
||||
num_pad_frames = num_frame_per_block - (num_frame %
|
||||
num_frame_per_block)
|
||||
num_frame += num_pad_frames
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, num_frame],
|
||||
device=self.device,
|
||||
@@ -207,12 +265,15 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
num_frame_per_block)
|
||||
timestep[:, :, 1:] = timestep[:, :, 0:1]
|
||||
timestep = timestep.reshape(timestep.shape[0], -1)
|
||||
if num_pad_frames > 0:
|
||||
timestep = timestep[:, num_pad_frames:]
|
||||
return timestep
|
||||
|
||||
def _step_predict_next_latent(
|
||||
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
encoder_hidden_states_image: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
|
||||
torch.Tensor]]:
|
||||
latent_vis_dict: dict[str, torch.Tensor] = {}
|
||||
@@ -225,7 +286,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
|
||||
if self._cached_closest_idx_per_dmd is None:
|
||||
self._cached_closest_idx_per_dmd = torch.tensor(
|
||||
[0, 12, 24, 36], dtype=torch.long).cpu()
|
||||
[0, 12, 24, 36, 50], dtype=torch.long).cpu()
|
||||
# [0, 1, 2, 3], dtype=torch.long).cpu()
|
||||
logger.info("self._cached_closest_idx_per_dmd: %s",
|
||||
self._cached_closest_idx_per_dmd)
|
||||
@@ -241,6 +302,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
traj_latents,
|
||||
dim=1,
|
||||
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
|
||||
# relevant_traj_latents = traj_latents
|
||||
logger.info("relevant_traj_latents: %s", relevant_traj_latents.shape)
|
||||
# assert relevant_traj_latents.shape[0] == 1
|
||||
|
||||
@@ -251,51 +313,149 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
num_frames,
|
||||
3,
|
||||
uniform_timestep=False)
|
||||
logger.info("indexes: %s", indexes.shape)
|
||||
logger.info("indexes: %s", indexes)
|
||||
# noisy_input = relevant_traj_latents[indexes]
|
||||
noisy_input = torch.gather(
|
||||
latents = torch.gather(
|
||||
relevant_traj_latents,
|
||||
dim=1,
|
||||
index=indexes.reshape(B, 1, num_frames, 1, 1,
|
||||
1).expand(-1, -1, -1, num_channels, height,
|
||||
width).to(self.device)).squeeze(1)
|
||||
noisy_input = torch.cat([
|
||||
latents,
|
||||
torch.zeros_like(latents),
|
||||
torch.zeros_like(latents[:, :, 0:1])
|
||||
],
|
||||
dim=2)
|
||||
timestep = self.dmd_denoising_steps[indexes]
|
||||
logger.info("selected timestep for rank %s: %s",
|
||||
self.global_rank,
|
||||
timestep,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
latent_vis_dict["noisy_input"] = noisy_input.permute(
|
||||
latent_vis_dict["noisy_input"] = latents.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
|
||||
4).detach().clone().cpu()
|
||||
|
||||
input_kwargs = {
|
||||
logger.info("timestep: %s", timestep)
|
||||
txt_input_kwargs = {
|
||||
"txt_inference":
|
||||
True,
|
||||
"vision_inference":
|
||||
False,
|
||||
"encoder_hidden_states":
|
||||
encoder_hidden_states,
|
||||
"encoder_hidden_states_image":
|
||||
encoder_hidden_states_image,
|
||||
"encoder_attention_mask":
|
||||
encoder_attention_mask,
|
||||
"timestep":
|
||||
torch.zeros([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.bfloat16),
|
||||
"cache_txt":
|
||||
True,
|
||||
}
|
||||
with set_forward_context(current_timestep=timestep,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
txt_kv_cache = self.transformer(**txt_input_kwargs)
|
||||
|
||||
vision_input_kwargs = {
|
||||
"txt_inference": False,
|
||||
"vision_inference": True,
|
||||
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timestep.to(device, dtype=torch.bfloat16),
|
||||
"return_dict": False,
|
||||
"txt_kv_cache": txt_kv_cache,
|
||||
}
|
||||
# Predict noise and step the scheduler to obtain next latent
|
||||
with set_forward_context(current_timestep=timestep,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
noise_pred = self.transformer(**vision_input_kwargs).permute(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=noise_pred.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
noise_input_latent=latents.flatten(0, 1),
|
||||
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
|
||||
scheduler=self.modules["scheduler"]).unflatten(
|
||||
0, noise_pred.shape[:2])
|
||||
scheduler=self.noise_scheduler).unflatten(0, noise_pred.shape[:2])
|
||||
latent_vis_dict["pred_video"] = pred_video.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
|
||||
return pred_video, target_latent, timestep, latent_vis_dict
|
||||
|
||||
# def _step_predict_next_latent(
|
||||
# self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
|
||||
# encoder_hidden_states: torch.Tensor,
|
||||
# encoder_attention_mask: torch.Tensor,
|
||||
# encoder_hidden_states_image: torch.Tensor
|
||||
# ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
|
||||
# torch.Tensor]]:
|
||||
# latent_vis_dict: dict[str, torch.Tensor] = {}
|
||||
# device = get_local_torch_device()
|
||||
# target_latent = traj_latents[:, -1]
|
||||
# del traj_latents
|
||||
|
||||
# # Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
# B, num_frames, num_channels, height, width = target_latent.shape
|
||||
|
||||
# indexes = self._get_timestep( # [B, num_frames]
|
||||
# 0,
|
||||
# 1000,
|
||||
# B,
|
||||
# num_frames,
|
||||
# 3,
|
||||
# uniform_timestep=False)
|
||||
# timestep = self.noise_scheduler.timesteps[indexes.cpu()].to(device)
|
||||
|
||||
# latents = self.noise_scheduler.add_noise(target_latent.flatten(0, 1), torch.randn_like(target_latent.flatten(0, 1)), timestep.flatten(0, 1)).unflatten(0, (B, num_frames))
|
||||
# noisy_input = torch.cat([latents, torch.zeros_like(latents), torch.zeros_like(latents[:, :, 0:1])], dim=2)
|
||||
|
||||
# # Prepare inputs for transformer
|
||||
# latent_vis_dict["noisy_input"] = latents.permute(
|
||||
# 0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
# latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
|
||||
# 4).detach().clone().cpu()
|
||||
|
||||
# logger.info("timestep: %s", timestep)
|
||||
# txt_input_kwargs = {
|
||||
# "txt_inference": True,
|
||||
# "vision_inference": False,
|
||||
# "encoder_hidden_states": encoder_hidden_states,
|
||||
# "encoder_hidden_states_image": encoder_hidden_states_image,
|
||||
# "encoder_attention_mask": encoder_attention_mask,
|
||||
# "timestep": torch.zeros([latents.shape[0]], device=latents.device, dtype=torch.bfloat16),
|
||||
# "cache_txt": True,
|
||||
# }
|
||||
# with set_forward_context(current_timestep=timestep,
|
||||
# attn_metadata=None,
|
||||
# forward_batch=None):
|
||||
# txt_kv_cache = self.transformer(**txt_input_kwargs)
|
||||
|
||||
# vision_input_kwargs = {
|
||||
# "txt_inference": False,
|
||||
# "vision_inference": True,
|
||||
# "hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
|
||||
# "timestep": timestep.to(device, dtype=torch.bfloat16),
|
||||
# "txt_kv_cache": txt_kv_cache,
|
||||
# }
|
||||
# # Predict noise and step the scheduler to obtain next latent
|
||||
# with set_forward_context(current_timestep=timestep,
|
||||
# attn_metadata=None,
|
||||
# forward_batch=None):
|
||||
# noise_pred = self.transformer(**vision_input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
# pred_video = pred_noise_to_pred_video(
|
||||
# pred_noise=noise_pred.flatten(0, 1),
|
||||
# noise_input_latent=latents.flatten(0, 1),
|
||||
# timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
|
||||
# scheduler=self.noise_scheduler).unflatten(
|
||||
# 0, noise_pred.shape[:2])
|
||||
# latent_vis_dict["pred_video"] = pred_video.permute(
|
||||
# 0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
|
||||
# return pred_video, target_latent, timestep, latent_vis_dict
|
||||
|
||||
def train_one_step(self, training_batch): # type: ignore[override]
|
||||
self.transformer.train()
|
||||
self.optimizer.zero_grad()
|
||||
@@ -307,8 +467,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
for _ in range(args.gradient_accumulation_steps):
|
||||
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
|
||||
training_batch)
|
||||
# traj_latents = traj_latents[:, :, :21]
|
||||
text_embeds = training_batch.encoder_hidden_states
|
||||
text_attention_mask = training_batch.encoder_attention_mask
|
||||
image_embeds = training_batch.encoder_hidden_states_image
|
||||
assert traj_latents.shape[0] == 1
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
@@ -318,7 +480,8 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
|
||||
# Forward to predict next latent by stepping scheduler with predicted noise
|
||||
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
|
||||
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
|
||||
traj_latents, traj_timesteps, text_embeds, text_attention_mask,
|
||||
image_embeds)
|
||||
|
||||
training_batch.latent_vis_dict.update(latent_vis_dict)
|
||||
|
||||
@@ -356,6 +519,10 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
except Exception:
|
||||
grad_value = 0.0
|
||||
training_batch.grad_norm = grad_value
|
||||
|
||||
if training_batch.current_timestep % 10 == 0:
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
return training_batch
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
@@ -367,14 +534,13 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
assert latent_key in latents_vis_dict and latents_vis_dict[
|
||||
latent_key] is not None
|
||||
latent = latents_vis_dict[latent_key]
|
||||
pixel_latent = self.validation_pipeline.decoding_stage.decode(
|
||||
latent, training_args)
|
||||
pixel_latent = self.decoding_stage.decode(latent, training_args)
|
||||
|
||||
video = pixel_latent.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
video_artifact = self.tracker.video(
|
||||
video, fps=16, format="mp4") # change to 16 for Wan2.1
|
||||
video, fps=24, format="mp4") # change to 16 for Wan2.1
|
||||
if video_artifact is not None:
|
||||
tracker_loss_dict[latent_key] = video_artifact
|
||||
# Clean up references
|
||||
@@ -383,6 +549,9 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
if self.global_rank == 0 and tracker_loss_dict:
|
||||
self.tracker.log_artifacts(tracker_loss_dict, step)
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting ODE-init training pipeline...")
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import copy
|
||||
import os
|
||||
import gc
|
||||
import time
|
||||
from collections import deque
|
||||
from typing import Any
|
||||
@@ -8,6 +9,7 @@ from typing import Any
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
@@ -19,8 +21,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,
|
||||
@@ -84,35 +86,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 +110,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(
|
||||
@@ -216,7 +200,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
batch_size, num_generated_frames, *self.video_latent_shape[2:]
|
||||
]
|
||||
|
||||
noise = torch.randn(noise_shape, device=self.device, dtype=dtype)
|
||||
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 = 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",
|
||||
@@ -252,7 +240,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,9 +274,16 @@ 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)
|
||||
|
||||
@@ -297,8 +292,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
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 +330,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 +370,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 +390,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 +405,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 +420,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 +445,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)
|
||||
|
||||
@@ -523,9 +538,9 @@ 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
|
||||
|
||||
@@ -538,28 +553,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):
|
||||
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):
|
||||
@@ -979,6 +1002,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)
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
from dataclasses import asdict
|
||||
import math
|
||||
import os
|
||||
import gc
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import deque
|
||||
@@ -25,7 +26,7 @@ from fastvideo.attention.backends.video_sparse_attn import (
|
||||
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
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.distributed import (cleanup_dist_env_and_memory,
|
||||
get_local_torch_device, get_sp_group,
|
||||
@@ -47,6 +48,7 @@ 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()
|
||||
@@ -89,6 +91,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
def set_schemas(self) -> None:
|
||||
self.train_dataset_schema = pyarrow_schema_t2v
|
||||
self.train_dataset_schema_2 = pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing training pipeline...")
|
||||
@@ -108,7 +111,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 +130,24 @@ 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 +167,23 @@ 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 +207,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
|
||||
@@ -575,13 +609,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)
|
||||
device="cpu").manual_seed(self.seed + self.global_rank)
|
||||
logger.info("Initialized random seeds with seed: %s",
|
||||
self.seed + self.global_rank)
|
||||
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
@@ -648,22 +684,25 @@ 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)
|
||||
context_len = int(training_batch.encoder_hidden_states.shape[1])
|
||||
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])
|
||||
|
||||
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:
|
||||
pass
|
||||
|
||||
self.tracker.log(metrics, step)
|
||||
if step % self.training_args.training_state_checkpointing_steps == 0:
|
||||
@@ -676,12 +715,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(
|
||||
|
||||
@@ -256,8 +256,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")
|
||||
@@ -289,8 +287,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",
|
||||
@@ -416,6 +412,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)
|
||||
@@ -453,46 +510,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,
|
||||
@@ -643,18 +699,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)
|
||||
@@ -849,7 +925,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)
|
||||
@@ -1153,6 +1229,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:
|
||||
|
||||
Reference in New Issue
Block a user