Compare commits

...
Author SHA1 Message Date
SolitaryThinker 94d45df2bc prof 2025-10-08 08:56:04 +00:00
JerryZhou54 08a9958de2 Fix small runtime issues 2025-10-07 09:43:13 +00:00
JerryZhou54 4b76388f21 Add matthew's timestep change & add real score guidance scale 2 2025-10-07 09:43:10 +00:00
JerryZhou54 c69ab42cc4 Refactor sf distill code 2025-10-07 09:41:55 +00:00
JerryZhou54 8c967268a3 checkpoint 2025-10-07 09:41:52 +00:00
15 changed files with 1128 additions and 96 deletions
+172
View File
@@ -0,0 +1,172 @@
#!/bin/bash
#SBATCH --job-name=prof_wan21_load
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1200G
#SBATCH --output=dmd_t2v_output/t2v_prof_%j.out
#SBATCH --error=dmd_t2v_output/t2v_prof_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
# 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='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
export FASTVIDEO_TORCH_PROFILER_DIR=/mnt/sharefs/users/hao.zhang/wl/torch_profiler/profiler_wan21_1_3_all_train_1n_8g/
export FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1
export FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1
export FASTVIDEO_TORCH_PROFILER_WITH_STACK=1
export FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1
export FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_training_train_one_step"
# Configs
NUM_GPUS=8
# Model paths for Self-Forcing DMD distillation with Wan2.2:
# GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # 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=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"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo3/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_1.json"
# export CUDA_VISIBLE_DEVICES=0
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "prof_wan21_train"
--max_train_steps 7
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--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
--enable_gradient_masking
--gradient_mask_last_n_frames 21
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
--init_weights_from_safetensors /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors
)
parallel_args=(
--num_gpus 8 # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 8
)
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"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
srun \
torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,180 @@
#!/bin/bash
#SBATCH --job-name=sf_distill_profile
#SBATCH --partition=main
#SBATCH --nodes=2
#SBATCH --ntasks=2
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=sf_distill_output/sf_distill_profile.out
#SBATCH --error=sf_distill_output/sf_distill_profile.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate wei-fv
export HOME="/mnt/weka/home/hao.zhang/wei"
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29503
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
export FASTVIDEO_TORCH_PROFILER_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/traces
export FASTVIDEO_TORCH_PROFILE_REGIONS=profiler_region_model_loading
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=16
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# DATA_DIR="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="data/crush-smol_processed_t2v/combined_parquet_dataset/"
# 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"
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_2.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan_t2v_finetune"
--wandb_run_name "sf_distill_2e-6_4e-7_4n"
# --use_sf_wan
# --sf_ode_init_path "checkpoints/ode_init.pt"
--max_train_steps 2
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# 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 $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 8
)
# Validation arguments
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--fake_score_learning_rate 4e-7
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/FastVideo2/vidprom_8b16k_1e-5_gn1/checkpoint-3500/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/sf2/checkpoint_v3_7k/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
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)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,165 @@
#!/bin/bash
#SBATCH --job-name=moe_4n_sf_distill
#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=moe_4n_sf_distill_output/moe_sf_distill_%j.out
#SBATCH --error=moe_4n_sf_distill_output/moe_sf_distill_%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
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation with Wan2.2:
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers"
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # 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=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 SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_dmd"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--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
--enable_gradient_masking
--gradient_mask_last_n_frames 21
--init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
--init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/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"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 40
--validation_sampling_steps "8"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 12
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
)
dmd_args=(
--dmd_denoising_steps '1000,850,700,550,350,275,200,125'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 4.0
--real_score_guidance_scale_2 3.0
--fake_score_learning_rate 4e-7
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,144 @@
#!/bin/bash
#SBATCH --job-name=moe_dmd_distill
#SBATCH --partition=main
#SBATCH --nodes=2
#SBATCH --ntasks=2
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=moe_dmd_distill_output/moe_dmd_distill_%j.out
#SBATCH --error=moe_dmd_distill_output/moe_dmd_distill_%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="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
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=16
# Model paths for Self-Forcing DMD distillation with Wan2.2:
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset"
# 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 SFwan2.2_t2v_distill_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_dmd"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--enable_gradient_checkpointing_type "full"
# --log_visualization
--simulate_generator_forward
--num_frames 81
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
)
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"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--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
--ema_start_step 100
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--real_score_guidance_scale 3.0
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -0,0 +1,131 @@
#!/bin/bash
#SBATCH --job-name=moe_finetune
#SBATCH --partition=main
#SBATCH --nodes=2
#SBATCH --ntasks=2
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=moe_output/moe_%j.out
#SBATCH --error=moe_output/moe_%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="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
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=16
# Model paths for Self-Forcing DMD distillation with Wan2.2:
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset"
# 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/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name VSA_finetune # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/wei/Wan2.2-MoE-finetune"
--max_train_steps 200
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--enable_gradient_checkpointing_type "full"
# --log_visualization
--num_frames 81
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
)
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
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "40"
--validation_guidance_scale "5.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--dit_cpu_offload True
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
+5
View File
@@ -726,6 +726,7 @@ class TrainingArgs(FastVideoArgs):
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
real_score_guidance_scale_2: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
@@ -1103,6 +1104,10 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.real_score_guidance_scale,
help="Teacher guidance scale")
parser.add_argument("--real-score-guidance-scale-2",
type=float,
default=TrainingArgs.real_score_guidance_scale_2,
help="Teacher guidance scale")
parser.add_argument("--fake-score-learning-rate",
type=float,
default=TrainingArgs.fake_score_learning_rate,
@@ -81,6 +81,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)
+50 -33
View File
@@ -1,4 +1,5 @@
import torch # type: ignore
import gc
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -39,8 +40,8 @@ class CausalDMDDenosingStage(DenoisingStage):
# KV and cross-attention cache state (initialized on first forward)
self.transformer = transformer
self.transformer_2 = transformer_2
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
# self.kv_cache1: list | None = None
# self.crossattn_cache: list | None = None
# 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
@@ -103,34 +104,43 @@ class CausalDMDDenosingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
if self.kv_cache1 is None:
self._initialize_kv_cache(batch_size=latents.shape[0],
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
self._initialize_crossattn_cache(
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device)
else:
assert self.crossattn_cache is not None
# reset cross-attention cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index][
"is_init"] = False # type: ignore
# reset kv cache pointers
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index][
"global_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
self.kv_cache1[block_index][
"local_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
# if self.kv_cache1 is None:
# self._initialize_kv_cache(batch_size=latents.shape[0],
# dtype=target_dtype,
# device=latents.device)
# self._initialize_crossattn_cache(
# batch_size=latents.shape[0],
# max_text_len=fastvideo_args.pipeline_config.
# text_encoder_configs[0].arch_config.text_len,
# dtype=target_dtype,
# device=latents.device)
# else:
# assert self.crossattn_cache is not None
# # reset cross-attention cache
# for block_index in range(self.num_transformer_blocks):
# self.crossattn_cache[block_index][
# "is_init"] = False # type: ignore
# # reset kv cache pointers
# for block_index in range(len(self.kv_cache1)):
# self.kv_cache1[block_index][
# "global_end_index"] = torch.tensor( # type: ignore
# [0],
# dtype=torch.long,
# device=latents.device)
# self.kv_cache1[block_index][
# "local_end_index"] = torch.tensor( # type: ignore
# [0],
# dtype=torch.long,
# device=latents.device)
# Optional: cache context features from provided image latents prior to generation
current_start_frame = 0
@@ -153,8 +163,8 @@ class CausalDMDDenosingStage(DenoisingStage):
image_first_btchw,
prompt_embeds,
t_zero,
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,
**image_kwargs,
@@ -179,8 +189,8 @@ class CausalDMDDenosingStage(DenoisingStage):
ref_btchw,
prompt_embeds,
t_zero,
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,
**image_kwargs,
@@ -280,8 +290,8 @@ class CausalDMDDenosingStage(DenoisingStage):
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -345,8 +355,8 @@ class CausalDMDDenosingStage(DenoisingStage):
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -355,6 +365,11 @@ class CausalDMDDenosingStage(DenoisingStage):
)
start_index += current_num_frames
del kv_cache1
del crossattn_cache
gc.collect()
torch.cuda.empty_cache()
batch.latents = latents
return batch
@@ -392,7 +407,8 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0], dtype=torch.long, device=device),
})
self.kv_cache1 = kv_cache1
# self.kv_cache1 = kv_cache1
return kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> None:
@@ -421,7 +437,8 @@ class CausalDMDDenosingStage(DenoisingStage):
"is_init":
False,
})
self.crossattn_cache = crossattn_cache
# self.crossattn_cache = crossattn_cache
return crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
+8
View File
@@ -202,12 +202,20 @@ def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
logger.info("FASTVIDEO_TORCH_PROFILE_REGIONS=%s",
envs.FASTVIDEO_TORCH_PROFILE_REGIONS)
def trace_handler(prof):
# print(prof.key_averages().table(
# sort_by="self_cuda_time_total", row_limit=-1))
logger.info("Profiling trace saved to: %s", "test_trace_" + str(prof.step_num) + ".json")
prof.export_chrome_trace("test2_trace_" + str(prof.step_num) + ".json")
profiler = torch.profiler.profile(
activities=_DEFAULT_ACTIVITIES,
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
schedule=torch.profiler.schedule(wait=2, warmup=1, active=2),
# on_trace_ready=trace_handler,
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir,
use_gzip=True),
)
+95 -25
View File
@@ -7,7 +7,7 @@ import time
from abc import abstractmethod
from collections import deque
from collections.abc import Iterator
from typing import Any
from typing import Any, Tuple
import imageio
import numpy as np
@@ -260,6 +260,7 @@ class DistillationPipeline(TrainingPipeline):
self.num_train_timestep)
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
self.real_score_guidance_scale_2 = self.training_args.real_score_guidance_scale_2
self.generator_ema: EMA_FSDP | None = None
self.generator_ema_2: EMA_FSDP | None = None
@@ -541,6 +542,15 @@ class DistillationPipeline(TrainingPipeline):
self.generator_ema_2.decay = original_decay_2
logger.info("EMA_2 reset completed")
def _get_real_score_guidance_scale(self, timestep: torch.Tensor):
"""
Get the appropriate real score guidance scale based on timestep and boundary logic.
"""
if self.boundary_timestep is not None:
if timestep.item() < self.boundary_timestep:
return self.real_score_guidance_scale_2
return self.real_score_guidance_scale
def _get_real_score_transformer(self, timestep: torch.Tensor):
"""
Get the appropriate real score transformer based on timestep and boundary logic.
@@ -615,7 +625,11 @@ class DistillationPipeline(TrainingPipeline):
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
if self.transformer_2 is not None and self.train_transformer_2:
current_transformer = self.transformer_2
else:
current_transformer = self.transformer
pred_noise = current_transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
@@ -626,7 +640,9 @@ class DistillationPipeline(TrainingPipeline):
return pred_video
def _generator_multi_step_simulation_forward(
self, training_batch: TrainingBatch) -> torch.Tensor:
self,
training_batch: TrainingBatch,
exit_flags: list[int] | None = None) -> Tuple[torch.Tensor, float]:
"""Forward pass through student transformer matching inference procedure."""
latents = training_batch.latents
dtype = latents.dtype
@@ -662,13 +678,17 @@ class DistillationPipeline(TrainingPipeline):
with torch.no_grad():
for step_idx in range(max_target_idx):
current_timestep = self.denoising_step_list[step_idx]
if self.transformer_2 is not None and current_timestep < self.boundary_timestep:
current_transformer = self.transformer_2
else:
current_transformer = self.transformer
current_timestep_tensor = current_timestep * torch.ones(
1, device=self.device, dtype=torch.long)
# Run student model to get flow prediction
training_batch_temp = self._build_distill_input_kwargs(
current_noise_latents, current_timestep_tensor,
training_batch.conditional_dict, training_batch)
pred_flow = self.transformer(
pred_flow = current_transformer(
**training_batch_temp.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_clean = pred_noise_to_pred_video(
@@ -711,8 +731,12 @@ class DistillationPipeline(TrainingPipeline):
training_batch = self._build_distill_input_kwargs(
noisy_input, target_timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
if self.transformer_2 is not None:
pred_noise = self.transformer_2(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
else:
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
@@ -720,15 +744,28 @@ class DistillationPipeline(TrainingPipeline):
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
training_batch.dmd_latent_vis_dict[
"generator_timestep"] = target_timestep.float().detach()
return pred_video
return pred_video, target_timestep.item()
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
training_batch: TrainingBatch,
exit_timestep: float) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
original_latent = generator_pred_video
with torch.no_grad():
timestep = torch.randint(0,
self.num_train_timestep, [1],
# Detect whether we're doing MoE DMD
if self.boundary_timestep is not None:
if exit_timestep < self.boundary_timestep:
start_timestep = 0
end_timestep = int(self.boundary_timestep)
else:
start_timestep = int(self.boundary_timestep)
end_timestep = self.num_train_timestep
else:
start_timestep = 0
end_timestep = self.num_train_timestep
timestep = torch.randint(start_timestep,
end_timestep, [1],
device=self.device,
dtype=torch.long)
world_group = get_world_group()
@@ -738,7 +775,8 @@ class DistillationPipeline(TrainingPipeline):
timestep = shift_timestep(
timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
min_timestep=start_timestep,
max_timestep=end_timestep)
timestep = timestep.clamp(self.min_timestep, self.max_timestep)
@@ -805,7 +843,7 @@ class DistillationPipeline(TrainingPipeline):
real_score_pred_video = pred_real_video_cond + (
pred_real_video_cond -
pred_real_video_uncond) * self.real_score_guidance_scale
pred_real_video_uncond) * self._get_real_score_guidance_scale(timestep)
grad = (faker_score_pred_video - real_score_pred_video) / torch.abs(
original_latent - real_score_pred_video).mean()
@@ -837,13 +875,26 @@ class DistillationPipeline(TrainingPipeline):
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video = self._generator_multi_step_simulation_forward(
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
training_batch)
else:
generator_pred_video = self._generator_forward(training_batch)
fake_score_timestep = torch.randint(0,
self.num_train_timestep, [1],
# Detect whether we're doing MoE DMD
# If yes, then we need to use the appropriate timestep range
if self.boundary_timestep is not None:
if exit_timestep < self.boundary_timestep:
start_timestep = 0
end_timestep = int(self.boundary_timestep)
else:
start_timestep = int(self.boundary_timestep)
end_timestep = self.num_train_timestep
else:
start_timestep = 0
end_timestep = self.num_train_timestep
fake_score_timestep = torch.randint(start_timestep,
end_timestep, [1],
device=self.device,
dtype=torch.long)
world_group = get_world_group()
@@ -853,7 +904,8 @@ class DistillationPipeline(TrainingPipeline):
fake_score_timestep = shift_timestep(
fake_score_timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
min_timestep=start_timestep,
max_timestep=end_timestep)
fake_score_timestep = fake_score_timestep.clamp(self.min_timestep,
self.max_timestep)
@@ -918,7 +970,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, prepare_timesteps=prepare_timesteps)
conditional_dict = {
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
@@ -929,10 +981,17 @@ class DistillationPipeline(TrainingPipeline):
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
training_batch.unconditional_dict = unconditional_dict
else:
unconditional_dict = {
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
}
training_batch.unconditional_dict = unconditional_dict
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
training_batch.timesteps = 0
training_batch.conditional_dict = conditional_dict
training_batch.raw_latent_shape = training_batch.latents.shape
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
@@ -967,6 +1026,8 @@ class DistillationPipeline(TrainingPipeline):
batches.append(batch)
self.optimizer.zero_grad()
if self.transformer_2 is not None:
self.optimizer_2.zero_grad()
total_dmd_loss = 0.0
dmd_latent_vis_dict = {}
fake_score_latent_vis_dict = {}
@@ -978,7 +1039,7 @@ class DistillationPipeline(TrainingPipeline):
current_timestep=batch_gen.timesteps,
attn_metadata=batch_gen.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video = self._generator_multi_step_simulation_forward(
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
batch_gen)
else:
generator_pred_video = self._generator_forward(
@@ -988,7 +1049,8 @@ class DistillationPipeline(TrainingPipeline):
attn_metadata=batch_gen.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=batch_gen)
training_batch=batch_gen,
exit_timestep=exit_timestep)
with set_forward_context(
current_timestep=batch_gen.timesteps,
@@ -997,12 +1059,20 @@ class DistillationPipeline(TrainingPipeline):
total_dmd_loss += dmd_loss.detach().item()
# Only clip gradients for the model that is currently training
self._clip_model_grad_norm_(batch_gen, self.transformer)
for param in self.transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.transformer_2 is not None and self.train_transformer_2:
self._clip_model_grad_norm_(batch_gen, self.transformer_2)
for param in self.transformer_2.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer_2.step()
self.optimizer_2.zero_grad(set_to_none=True)
else:
self._clip_model_grad_norm_(batch_gen, self.transformer)
for param in self.transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
@@ -2,7 +2,8 @@
import copy
import time
from collections import deque
from typing import Any
from typing import Any, Tuple
import gc
import numpy as np
import torch
@@ -77,8 +78,8 @@ 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
# 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)
@@ -128,17 +129,42 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
Compute generator loss using DMD-style approach.
The generator tries to fool the critic (fake_score_transformer).
"""
exit_flags = None
# If the generator is MoE, randomly sample an exit timestep, and turn on training for one of the dits
if self.boundary_timestep is not None and self.transformer_2 is not None:
assert self.same_step_across_blocks, "same_step_across_blocks must be True for MoE generator. Otherwise we might need to train both transformers which will cause OOM"
exit_flags = self.generate_and_sync_list(1, len(self.denoising_step_list), training_batch.latents.device)
exit_timestep = self.denoising_step_list[exit_flags[0]].item()
# logger.info("Exit timestep in generator_loss(): %s", exit_timestep)
if exit_timestep < self.boundary_timestep:
self.train_transformer_2 = True
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self.train_transformer_2 = False
self._enable_training(self.transformer, self.optimizer)
self._disable_training(self.transformer_2, self.optimizer_2)
else: # Default to normal single-dit generator
self.train_transformer_2 = False
self._enable_training(self.transformer, self.optimizer)
# Turns off training for fake score transformers
self._disable_training(self.fake_score_transformer, self.fake_score_optimizer)
if self.fake_score_transformer_2 is not None:
self._disable_training(self.fake_score_transformer_2, self.fake_score_optimizer_2)
with set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
generator_pred_video = self._generator_multi_step_simulation_forward(
training_batch)
generator_pred_video, exit_timestep = self._generator_multi_step_simulation_forward(
training_batch, exit_flags=exit_flags)
with set_forward_context(current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=training_batch)
training_batch=training_batch,
exit_timestep=exit_timestep)
log_dict = {
"dmdtrain_gradient_norm": torch.tensor(0.0, device=self.device)
@@ -153,6 +179,16 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
Compute critic loss using flow matching between noise and generator output.
The critic learns to predict the flow from noise to the generator's output.
"""
# Turns off training for all generators
self._disable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
# Turns on training for fake score transformers
self._enable_training(self.fake_score_transformer, self.fake_score_optimizer)
if self.fake_score_transformer_2 is not None:
self._enable_training(self.fake_score_transformer_2, self.fake_score_optimizer_2)
updated_batch, flow_matching_loss = self.faker_score_forward(
training_batch)
training_batch.fake_score_latent_vis_dict = updated_batch.fake_score_latent_vis_dict
@@ -163,7 +199,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
def _generator_multi_step_simulation_forward(
self,
training_batch: TrainingBatch,
return_sim_steps: bool = False) -> torch.Tensor:
exit_flags: list[int] | None = None) -> Tuple[torch.Tensor, float]:
"""Forward pass through student transformer matching inference procedure with KV cache management.
This function is adapted from the reference self-forcing implementation's inference_with_trajectory
@@ -245,7 +281,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
@@ -269,8 +305,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)
current_start_frame += 1
@@ -280,9 +316,13 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
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)
exit_flags = self.generate_and_sync_list(len(all_num_frames),
if exit_flags is None:
exit_flags = self.generate_and_sync_list(len(all_num_frames),
num_denoising_steps,
device=noise.device)
exit_timestep = self.denoising_step_list[exit_flags[0]].item()
start_gradient_frame_index = max(0, num_output_frames - 21)
for block_index, current_num_frames in enumerate(all_num_frames):
@@ -303,14 +343,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
if self.boundary_timestep is not None and current_timestep < self.boundary_timestep and self.transformer_2 is not None:
current_model = self.transformer_2
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
current_model = self.transformer
self._enable_training(self.transformer, self.optimizer)
if self.boundary_timestep is not None and self.transformer_2 is not None:
self._disable_training(self.transformer_2,
self.optimizer_2)
if not exit_flag:
with torch.no_grad():
@@ -328,8 +362,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(
@@ -352,6 +386,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
dtype=torch.long)).unflatten(
0, denoised_pred.shape[:2])
else:
logger.info("Exit timestep in _generator_multi_step_simulation_forward(): %s", current_timestep)
# Final prediction with gradient control
if current_start_frame < start_gradient_frame_index:
with torch.no_grad():
@@ -368,8 +403,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(
@@ -388,8 +423,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(
@@ -429,8 +464,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)
@@ -521,12 +556,18 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
training_batch.dmd_latent_vis_dict["min_num_frames"] = torch.tensor(
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)
# Clean up caches - properly free GPU memory
self._cleanup_simulation_caches(kv_cache1, crossattn_cache)
gc.collect()
torch.cuda.empty_cache()
# assert kv_cache1 is not None
# assert crossattn_cache is not None
# self._reset_simulation_caches(self.kv_cache1, self.crossattn_cache)
return final_output if gradient_mask is not None else pred_image_or_video
if gradient_mask is not None:
return final_output, exit_timestep
else:
return pred_image_or_video, exit_timestep
def _initialize_simulation_caches(
self,
@@ -623,6 +664,30 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
cache_dict["k"].zero_()
cache_dict["v"].zero_()
def _cleanup_simulation_caches(self, kv_cache: list[dict[str, Any]],
crossattn_cache: list[dict[str, Any]]) -> None:
"""Properly clean up KV cache and cross-attention cache GPU memory."""
if kv_cache is not None:
for cache_dict in kv_cache:
# Clear tensor references to free GPU memory
if "k" in cache_dict and cache_dict["k"] is not None:
cache_dict["k"] = None
if "v" in cache_dict and cache_dict["v"] is not None:
cache_dict["v"] = None
if "global_end_index" in cache_dict and cache_dict["global_end_index"] is not None:
cache_dict["global_end_index"] = None
if "local_end_index" in cache_dict and cache_dict["local_end_index"] is not None:
cache_dict["local_end_index"] = None
if crossattn_cache is not None:
for cache_dict in crossattn_cache:
# Clear tensor references to free GPU memory
if "k" in cache_dict and cache_dict["k"] is not None:
cache_dict["k"] = None
if "v" in cache_dict and cache_dict["v"] is not None:
cache_dict["v"] = None
cache_dict["is_init"] = False
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
@@ -710,9 +775,33 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
setattr(batch_gen, key, copy.deepcopy(value))
generator_loss, gen_log_dict = self.generator_loss(batch_gen)
logger.info("train_transformer_2: %s", self.train_transformer_2)
with set_forward_context(current_timestep=batch_gen.timesteps,
attn_metadata=batch_gen.attn_metadata):
(generator_loss / gradient_accumulation_steps).backward()
# Ensure that only one of the two transformers have received gradients
if self.train_transformer_2:
# Assert that all gradients are None for transformer
assert all(p.grad is None for p in self.transformer.parameters())
grad_sum = 0
for n, p in self.transformer_2.named_parameters():
if p.grad is not None:
grad_sum += p.grad.sum().item()
else:
raise ValueError("Transformer 2 param %s has no gradient", n)
assert grad_sum != 0, "Transformer 2 did not receive gradients"
else:
# assert all(p.grad is None for p in self.transformer_2.parameters())
grad_sum = 0
for n, p in self.transformer.named_parameters():
if p.grad is not None:
grad_sum += p.grad.sum().item()
else:
raise ValueError("Transformer param %s has no gradient", n)
assert grad_sum != 0, "Transformer did not receive gradients"
total_generator_loss += generator_loss.detach().item()
generator_log_dict.update(gen_log_dict)
# Store visualization data from generator training
@@ -749,11 +838,16 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
world_group.all_reduce(avg_generator_loss,
op=torch.distributed.ReduceOp.AVG)
training_batch.generator_loss = avg_generator_loss.item()
# training_batch.fake_score_loss = 0
# training_batch.total_loss = training_batch.generator_loss
# return training_batch
else:
training_batch.generator_loss = 0.0
logger.debug("Training critic at step %s", self.current_trainstep)
self.fake_score_optimizer.zero_grad()
if self.fake_score_transformer_2 is not None:
self.fake_score_optimizer_2.zero_grad()
total_critic_loss = 0.0
critic_log_dict = {}
@@ -778,6 +872,26 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
with set_forward_context(current_timestep=batch_critic.timesteps,
attn_metadata=batch_critic.attn_metadata):
(critic_loss / gradient_accumulation_steps).backward()
if self.train_fake_score_transformer_2:
assert all(p.grad is None for p in self.fake_score_transformer.parameters())
grad_sum = 0
for n, p in self.fake_score_transformer_2.named_parameters():
if p.grad is not None:
grad_sum += p.grad.sum().item()
else:
raise ValueError("Fake score transformer 2 param %s has no gradient", n)
assert grad_sum != 0, "Fake score transformer 2 did not receive gradients"
else:
# assert all(p.grad is None for p in self.fake_score_transformer_2.parameters())
grad_sum = 0
for n, p in self.fake_score_transformer.named_parameters():
if p.grad is not None:
grad_sum += p.grad.sum().item()
else:
raise ValueError("Fake score transformer param %s has no gradient", n)
assert grad_sum != 0, "Fake score transformer did not receive gradients"
total_critic_loss += critic_loss.detach().item()
critic_log_dict.update(crit_log_dict)
# Store visualization data from critic training
@@ -978,6 +1092,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self._log_training_info()
self._log_validation(self.transformer, self.training_args,
self.init_steps)
gc.collect()
torch.cuda.empty_cache()
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
@@ -1023,7 +1139,10 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
step, self.training_args.ema_decay)
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
with self.profiler_controller.region("profiler_region_training_train_one_step"):
training_batch = self.train_one_step(training_batch)
self.profiler.step()
total_loss = training_batch.total_loss
generator_loss = training_batch.generator_loss
@@ -1186,6 +1305,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
gc.collect()
torch.cuda.empty_cache()
wandb.finish()
@@ -1230,3 +1351,4 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
if get_sp_group():
cleanup_dist_env_and_memory()
+18 -2
View File
@@ -107,6 +107,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
self.transformer.train()
self.transformer.requires_grad_(True)
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
self.transformer,
@@ -235,14 +236,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
for param in model.parameters():
param.requires_grad = True
model.train()
optimizer.zero_grad()
# optimizer.zero_grad()
def _disable_training(self, model: torch.nn.Module,
optimizer: torch.optim.Optimizer) -> None:
"""Disable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = False
optimizer.zero_grad(set_to_none=True)
model.eval()
# optimizer.zero_grad(set_to_none=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
@@ -358,6 +360,20 @@ class TrainingPipeline(LoRAPipeline, ABC):
# pass
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
# timestep = self.noise_scheduler.timesteps[indices].to(device=device)
# if timestep < self.training_args.boundary_ratio * self.noise_scheduler.config.num_train_timesteps:
# self.train_transformer_2 = True
# else:
# self.train_transformer_2 = False
# # Broadcast the decision to all processes
# decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0,
# device=self.device)
# dist.broadcast(decision, src=0)
# dist.broadcast(timestep, src=0)
# self.train_transformer_2 = decision.item() == 1.0
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(
+4 -4
View File
@@ -1152,14 +1152,14 @@ def custom_to_hf_state_dict(
return new_state_dict
# More generalized version of shift_timestep
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
min_timestep: int = 0, max_timestep: int = 1000) -> torch.Tensor:
if shift == 1:
return timestep
t = timestep / num_train_timestep
t = (timestep - min_timestep) / (max_timestep - min_timestep)
denominator = 1 + (shift - 1) * t
return num_train_timestep * (shift * t / denominator)
return min_timestep + (max_timestep - min_timestep) * (shift * t / denominator)
# coding=utf-8
@@ -42,7 +42,7 @@ class WanDistillationPipeline(DistillationPipeline):
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
loaded_modules={"transformer": self.get_module("transformer"), "transformer_2": self.get_module("transformer_2")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
@@ -42,6 +42,7 @@ class WanTrainingPipeline(TrainingPipeline):
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
"transformer_2": self.get_module("transformer_2"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,