Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e8cbcbd757 | ||
|
|
df7381e2b0 | ||
|
|
de11ea3020 | ||
|
|
4a527bccf2 | ||
|
|
dabf87fefa | ||
|
|
8996a391a1 | ||
|
|
fba8b61c4f | ||
|
|
55d8c1e5fb | ||
|
|
aeb8f2e5ac | ||
|
|
f755dd1ad5 | ||
|
|
256f788d34 |
@@ -0,0 +1,183 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=1.3B_high_noise_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=1.3B_moe_4n_sf_distill_output/moe_sf_distill_2e-6_4e-7_4n.out
|
||||
#SBATCH --error=1.3B_moe_4n_sf_distill_output/moe_sf_distill_2e-6_4e-7_4n.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=29500
|
||||
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=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
NUM_TOTAL_GPUS=32
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers"
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
# FAKE_SCORE_MODEL_PATH="rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers" # Critic model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
|
||||
# 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/1.3B_MoE_SFwan_t2v_finetune"
|
||||
--wandb_run_name "1.3B_high_noise_sf_distill"
|
||||
# --use_sf_wan
|
||||
# --sf_ode_init_path "checkpoints/ode_init.pt"
|
||||
--max_train_steps 5000
|
||||
--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 50
|
||||
--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/wei/FastVideo/diffusers_1.3B_SF/model.safetensors"
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
# --init_weights_from_safetensors_2 "/mnt/weka/home/hao.zhang/wei/FastVideo/diffusers_1.3B_SF/model.safetensors"
|
||||
--init_weights_from_safetensors_2 "/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
|
||||
--real_score_guidance_scale_2 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,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[@]}"
|
||||
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig)
|
||||
SelfForcingWanT2V480PConfig, SelfForcingMoEWanT2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -36,6 +36,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers": SelfForcingMoEWanT2V480PConfig,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
|
||||
@@ -165,3 +165,10 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
@dataclass
|
||||
class SelfForcingMoEWanT2V480PConfig(SelfForcingWanT2V480PConfig):
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
@@ -55,6 +55,7 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -717,6 +717,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"
|
||||
@@ -1102,6 +1103,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,
|
||||
|
||||
@@ -86,6 +86,12 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
return (prev_sample, )
|
||||
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
@staticmethod
|
||||
def calculate_alpha_beta_high(sigma, sigma_bound):
|
||||
alpha = (1 - sigma) / (1 - sigma_bound)
|
||||
beta = torch.sqrt(sigma ** 2 - (alpha * sigma_bound) ** 2)
|
||||
return alpha, beta
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
@@ -105,6 +111,32 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def add_noise_high(self, original_samples, noise, timestep, boundary_timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B*T, C, H, W]
|
||||
- noise: the noise with shape [B*T, C, H, W]
|
||||
- timestep: the timestep with shape [B*T]
|
||||
Output: the corrupted latent with shape [B*T, C, H, W]
|
||||
"""
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
if boundary_timestep.ndim == 2:
|
||||
boundary_timestep = boundary_timestep.flatten(0, 1)
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
boundary_timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_boundary = self.sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
|
||||
alpha, beta = self.calculate_alpha_beta_high(sigma, sigma_boundary)
|
||||
sample = alpha * original_samples + beta * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def training_target(self, sample, noise, timestep):
|
||||
target = noise - sample
|
||||
return target
|
||||
|
||||
@@ -180,3 +180,51 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
pred_video = noise_input_latent - sigma_t * pred_noise
|
||||
return pred_video.to(dtype)
|
||||
|
||||
def pred_noise_to_x_bound(pred_noise: torch.Tensor,
|
||||
noise_input_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
boundary_timestep: torch.Tensor,
|
||||
scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert predicted noise to clean latent.
|
||||
|
||||
Args:
|
||||
pred_noise: the predicted noise with shape [B, C, H, W]
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise_input_latent: the noisy latent with shape [B, C, H, W],
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
boundary_timestep: the boundary timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
scheduler: the scheduler
|
||||
|
||||
Returns:
|
||||
the predicted video with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == noise_input_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(noise_input_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == noise_input_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
pred_noise = pred_noise.double().to(device)
|
||||
noise_input_latent = noise_input_latent.double().to(device)
|
||||
sigmas = scheduler.sigmas.double().to(device)
|
||||
timesteps = scheduler.timesteps.double().to(device)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
boundary_timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t_boundary = sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
|
||||
pred_video = noise_input_latent - (sigma_t - sigma_t_boundary) * pred_noise
|
||||
return pred_video.to(dtype)
|
||||
@@ -71,6 +71,7 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
fastvideo_args.dit_cpu_offload = False
|
||||
self.modules = self.load_modules(fastvideo_args, loaded_modules)
|
||||
|
||||
def set_trainable(self) -> None:
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
import torch # type: ignore
|
||||
import gc
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
from fastvideo.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.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
@@ -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,
|
||||
@@ -211,6 +221,15 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
block_sizes = [1] + [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
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
|
||||
else:
|
||||
boundary_timestep = None
|
||||
boundary_timestep = timesteps[2] + 1 # Hardcode for now
|
||||
|
||||
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
|
||||
assert len(high_noise_timesteps) == 2, "only support two high noise timesteps"
|
||||
|
||||
# DMD loop in causal blocks
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
@@ -222,7 +241,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
video_raw_latent_shape = noise_latents_btchw.shape
|
||||
|
||||
for i, t_cur in enumerate(timesteps):
|
||||
if self.transformer_2 is not None and fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None and t_cur < fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps:
|
||||
if boundary_timestep is not None and t_cur < boundary_timestep:
|
||||
current_model = self.transformer_2
|
||||
else:
|
||||
current_model = self.transformer
|
||||
@@ -280,8 +299,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,
|
||||
@@ -290,12 +309,20 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
).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 boundary_timestep is not None and t_cur >= boundary_timestep:
|
||||
pred_video_btchw = pred_noise_to_x_bound(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
boundary_timestep=torch.ones_like(t_expand) * boundary_timestep,
|
||||
scheduler=self.scheduler).unflatten(0, pred_noise_btchw.shape[:2])
|
||||
else:
|
||||
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(
|
||||
@@ -309,11 +336,20 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
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])
|
||||
if boundary_timestep is not None and i < len(high_noise_timesteps) - 1:
|
||||
noise_latents_btchw = self.scheduler.add_noise_high(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1),
|
||||
next_timestep,
|
||||
torch.ones_like(next_timestep) * boundary_timestep).unflatten(0, pred_video_btchw.shape[:2])
|
||||
elif boundary_timestep is not None and i == len(high_noise_timesteps) - 1:
|
||||
noise_latents_btchw = pred_video_btchw
|
||||
else:
|
||||
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:
|
||||
@@ -345,8 +381,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 +391,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 +433,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 +463,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:
|
||||
|
||||
@@ -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
|
||||
@@ -30,7 +30,7 @@ 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.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
|
||||
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.training.activation_checkpoint import (
|
||||
@@ -252,6 +252,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
get_local_torch_device())
|
||||
logger.info("Distillation generator model to %s denoising steps: %s",
|
||||
len(self.denoising_step_list), self.denoising_step_list)
|
||||
self.boundary_timestep = self.denoising_step_list[2] + 1 # Hardcode for now
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
|
||||
self.min_timestep = int(self.training_args.min_timestep_ratio *
|
||||
@@ -260,6 +261,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 +543,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.real_score_transformer_2 is not None and 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 +626,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 +641,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 +679,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 +732,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 +745,32 @@ 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:
|
||||
clean_generator_pred_video: torch.Tensor,
|
||||
training_batch: TrainingBatch,
|
||||
exit_timestep: float) -> torch.Tensor:
|
||||
"""Compute DMD (Diffusion Model Distillation) loss."""
|
||||
original_latent = generator_pred_video
|
||||
clean_original_latent = clean_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)
|
||||
# end_timestep = self.num_train_timestep
|
||||
# else:
|
||||
# # raise ValueError("Exit timestep is greater than boundary timestep")
|
||||
# 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 +780,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)
|
||||
|
||||
@@ -751,10 +794,23 @@ class DistillationPipeline(TrainingPipeline):
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).detach().unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and exit_timestep >= self.boundary_timestep:
|
||||
if timestep < self.boundary_timestep:
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
clean_generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).detach().unflatten(0,
|
||||
(1, clean_original_latent.shape[1]))
|
||||
else:
|
||||
noisy_latent = self.noise_scheduler.add_noise_high(
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep,
|
||||
self.boundary_timestep * torch.ones_like(timestep)).detach().unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
else:
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).detach().unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
|
||||
# fake_score_transformer forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -765,12 +821,27 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fake_score_pred_noise = current_fake_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
faker_score_pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, fake_score_pred_noise.shape[:2])
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and exit_timestep >= self.boundary_timestep:
|
||||
if timestep < self.boundary_timestep:
|
||||
faker_score_pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, fake_score_pred_noise.shape[:2])
|
||||
else:
|
||||
faker_score_pred_video = pred_noise_to_x_bound(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
boundary_timestep=torch.ones_like(timestep) * self.boundary_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, fake_score_pred_noise.shape[:2])
|
||||
else:
|
||||
faker_score_pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, fake_score_pred_noise.shape[:2])
|
||||
|
||||
# real_score_transformer cond forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -781,12 +852,27 @@ class DistillationPipeline(TrainingPipeline):
|
||||
real_score_pred_noise_cond = current_real_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_real_video_cond = pred_noise_to_pred_video(
|
||||
pred_noise=real_score_pred_noise_cond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, real_score_pred_noise_cond.shape[:2])
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and exit_timestep >= self.boundary_timestep:
|
||||
if timestep < self.boundary_timestep:
|
||||
pred_real_video_cond = pred_noise_to_pred_video(
|
||||
pred_noise=real_score_pred_noise_cond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, real_score_pred_noise_cond.shape[:2])
|
||||
else:
|
||||
pred_real_video_cond = pred_noise_to_x_bound(
|
||||
pred_noise=real_score_pred_noise_cond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
boundary_timestep=torch.ones_like(timestep) * self.boundary_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, real_score_pred_noise_cond.shape[:2])
|
||||
else:
|
||||
pred_real_video_cond = pred_noise_to_pred_video(
|
||||
pred_noise=real_score_pred_noise_cond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, real_score_pred_noise_cond.shape[:2])
|
||||
|
||||
# real_score_transformer uncond forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -796,16 +882,31 @@ class DistillationPipeline(TrainingPipeline):
|
||||
real_score_pred_noise_uncond = current_real_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_real_video_uncond = pred_noise_to_pred_video(
|
||||
pred_noise=real_score_pred_noise_uncond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, real_score_pred_noise_uncond.shape[:2])
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and exit_timestep >= self.boundary_timestep:
|
||||
if timestep < self.boundary_timestep:
|
||||
pred_real_video_uncond = pred_noise_to_pred_video(
|
||||
pred_noise=real_score_pred_noise_uncond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, real_score_pred_noise_uncond.shape[:2])
|
||||
else:
|
||||
pred_real_video_uncond = pred_noise_to_x_bound(
|
||||
pred_noise=real_score_pred_noise_uncond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
boundary_timestep=torch.ones_like(timestep) * self.boundary_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, real_score_pred_noise_uncond.shape[:2])
|
||||
else:
|
||||
pred_real_video_uncond = pred_noise_to_pred_video(
|
||||
pred_noise=real_score_pred_noise_uncond.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, real_score_pred_noise_uncond.shape[:2])
|
||||
|
||||
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 +938,28 @@ 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, clean_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 = self.num_train_timestep
|
||||
# # end_timestep = int(self.boundary_timestep)
|
||||
# else:
|
||||
# # raise ValueError("Exit timestep is greater than boundary timestep")
|
||||
# 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 +969,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)
|
||||
@@ -868,10 +985,23 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fake_score_noise = fake_score_noise[:, self.
|
||||
rank_in_sp_group, :, :, :, :]
|
||||
|
||||
noisy_generator_pred_video = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
|
||||
fake_score_timestep).unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and exit_timestep >= self.boundary_timestep:
|
||||
if fake_score_timestep < self.boundary_timestep:
|
||||
noisy_generator_pred_video = self.noise_scheduler.add_noise(
|
||||
clean_generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
|
||||
fake_score_timestep).unflatten(0,
|
||||
(1, clean_generator_pred_video.shape[1]))
|
||||
else:
|
||||
noisy_generator_pred_video = self.noise_scheduler.add_noise_high(
|
||||
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
|
||||
fake_score_timestep,
|
||||
self.boundary_timestep * torch.ones_like(fake_score_timestep)).unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
else:
|
||||
noisy_generator_pred_video = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
|
||||
fake_score_timestep).unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
|
||||
with set_forward_context(current_timestep=training_batch.timesteps,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
@@ -884,8 +1014,33 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fake_score_pred_noise = current_fake_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
target = fake_score_noise - generator_pred_video
|
||||
flow_matching_loss = torch.mean((fake_score_pred_noise - target)**2)
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and exit_timestep >= self.boundary_timestep:
|
||||
if fake_score_timestep < self.boundary_timestep:
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_generator_pred_video.flatten(0, 1),
|
||||
timestep=fake_score_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, fake_score_pred_noise.shape[:2])
|
||||
flow_matching_loss = torch.mean((pred_video - clean_generator_pred_video)**2)
|
||||
else:
|
||||
pred_video = pred_noise_to_x_bound(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_generator_pred_video.flatten(0, 1),
|
||||
timestep=fake_score_timestep,
|
||||
boundary_timestep=torch.ones_like(fake_score_timestep) * self.boundary_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, fake_score_pred_noise.shape[:2])
|
||||
flow_matching_loss = torch.mean((pred_video - generator_pred_video)**2)
|
||||
else:
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=fake_score_pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_generator_pred_video.flatten(0, 1),
|
||||
timestep=fake_score_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, fake_score_pred_noise.shape[:2])
|
||||
flow_matching_loss = torch.mean((pred_video - generator_pred_video)**2)
|
||||
|
||||
# This is mathematically equivalent to pred_video - generator_pred_video for the low noise case
|
||||
# target = fake_score_noise - generator_pred_video
|
||||
# flow_matching_loss = torch.mean((fake_score_pred_noise - target)**2)
|
||||
|
||||
training_batch.fake_score_latent_vis_dict = {
|
||||
"training_batch_fakerscore_fwd_clean_latent":
|
||||
@@ -918,7 +1073,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 +1084,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 +1129,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 +1142,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 +1152,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 +1162,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
|
||||
@@ -20,7 +21,7 @@ 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.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
|
||||
from fastvideo.pipelines import TrainingBatch
|
||||
from fastvideo.training.distillation_pipeline import DistillationPipeline
|
||||
from fastvideo.training.training_utils import (EMA_FSDP,
|
||||
@@ -60,7 +61,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
num_inference_steps=1000,
|
||||
shift=5.0,
|
||||
shift=self.timestep_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True,
|
||||
training=True)
|
||||
@@ -76,8 +77,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)
|
||||
@@ -127,17 +128,46 @@ 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_flags = self.generate_and_sync_list(1, 2, training_batch.latents.device)
|
||||
# exit_flags[0] += 2
|
||||
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)
|
||||
# If the exit timestep is in the high noise region, the generator_pred_video is x_boundary, otherwise it's x_0
|
||||
generator_pred_video, clean_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)
|
||||
clean_generator_pred_video=clean_generator_pred_video,
|
||||
training_batch=training_batch,
|
||||
exit_timestep=exit_timestep)
|
||||
|
||||
log_dict = {
|
||||
"dmdtrain_gradient_norm": torch.tensor(0.0, device=self.device)
|
||||
@@ -152,6 +182,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
|
||||
@@ -162,7 +202,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
|
||||
@@ -233,6 +273,10 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
[batch_size, num_output_frames, num_channels, height, width],
|
||||
device=noise.device,
|
||||
dtype=noise.dtype)
|
||||
clean_output = torch.zeros(
|
||||
[batch_size, num_output_frames, num_channels, height, width],
|
||||
device=noise.device,
|
||||
dtype=noise.dtype)
|
||||
|
||||
def get_model_device(model):
|
||||
if model is None:
|
||||
@@ -244,7 +288,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
|
||||
@@ -268,8 +312,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
|
||||
@@ -279,11 +323,20 @@ 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)
|
||||
# if exit_flags is None:
|
||||
# exit_flags = self.generate_and_sync_list(len(all_num_frames), 2, device=noise.device)
|
||||
# exit_flags[0] += 2
|
||||
exit_timestep = self.denoising_step_list[exit_flags[0]].item()
|
||||
|
||||
start_gradient_frame_index = max(0, num_output_frames - 21)
|
||||
|
||||
high_noise_timesteps = self.denoising_step_list[self.denoising_step_list >= self.boundary_timestep]
|
||||
|
||||
for block_index, current_num_frames in enumerate(all_num_frames):
|
||||
noisy_input = noise[:, current_start_frame -
|
||||
num_input_frames:current_start_frame +
|
||||
@@ -302,14 +355,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():
|
||||
@@ -327,30 +374,57 @@ 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(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, pred_flow.shape[:2])
|
||||
if self.boundary_timestep is not None and current_timestep >= self.boundary_timestep:
|
||||
denoised_pred = pred_noise_to_x_bound(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
boundary_timestep=torch.ones_like(timestep).flatten(0, 1) * self.boundary_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, pred_flow.shape[:2])
|
||||
else:
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, pred_flow.shape[:2])
|
||||
|
||||
next_timestep = self.denoising_step_list[index + 1]
|
||||
noisy_input = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
next_timestep *
|
||||
torch.ones([batch_size * current_num_frames],
|
||||
device=noise.device,
|
||||
dtype=torch.long)).unflatten(
|
||||
0, denoised_pred.shape[:2])
|
||||
|
||||
if self.boundary_timestep is not None and index < len(high_noise_timesteps) - 1:
|
||||
noisy_input = self.noise_scheduler.add_noise_high(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
next_timestep *
|
||||
torch.ones([batch_size * current_num_frames],
|
||||
device=noise.device,
|
||||
dtype=torch.long),
|
||||
self.boundary_timestep *
|
||||
torch.ones([batch_size * current_num_frames],
|
||||
device=noise.device,
|
||||
dtype=torch.long)).unflatten(
|
||||
0, denoised_pred.shape[:2])
|
||||
elif self.boundary_timestep is not None and index == len(high_noise_timesteps) - 1:
|
||||
noisy_input = denoised_pred
|
||||
else:
|
||||
noisy_input = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
next_timestep *
|
||||
torch.ones([batch_size * current_num_frames],
|
||||
device=noise.device,
|
||||
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():
|
||||
@@ -367,8 +441,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(
|
||||
@@ -387,31 +461,66 @@ 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(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, pred_flow.shape[:2])
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and current_timestep >= self.boundary_timestep:
|
||||
denoised_pred = pred_noise_to_x_bound(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
boundary_timestep=torch.ones_like(timestep).flatten(0, 1) * self.boundary_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
clean_denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
else:
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(
|
||||
0, pred_flow.shape[:2])
|
||||
clean_denoised_pred = denoised_pred
|
||||
break
|
||||
|
||||
# Step 3.2: record the model's output
|
||||
output[:, current_start_frame:current_start_frame +
|
||||
current_num_frames] = denoised_pred
|
||||
clean_output[:, current_start_frame:current_start_frame +
|
||||
current_num_frames] = clean_denoised_pred
|
||||
|
||||
# Step 3.3: rerun with timestep zero to update the cache
|
||||
context_timestep = torch.ones_like(timestep) * self.context_noise
|
||||
denoised_pred = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
context_timestep).unflatten(0, denoised_pred.shape[:2])
|
||||
if self.boundary_timestep is not None:
|
||||
if exit_timestep < self.boundary_timestep:
|
||||
context_timestep = torch.ones_like(timestep) * self.context_noise
|
||||
current_model = self.transformer_2
|
||||
else:
|
||||
# For high noise expert, the context timestep should be the boundary timestep
|
||||
context_timestep = torch.ones_like(timestep) * self.boundary_timestep
|
||||
current_model = self.transformer
|
||||
else:
|
||||
context_timestep = torch.ones_like(timestep) * self.context_noise
|
||||
current_model = self.transformer
|
||||
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None and current_timestep >= self.boundary_timestep:
|
||||
# raise ValueError("Exit timestep is greater than boundary timestep")
|
||||
denoised_pred = self.noise_scheduler.add_noise_high(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
context_timestep,
|
||||
self.boundary_timestep * torch.ones_like(context_timestep)).unflatten(0, denoised_pred.shape[:2])
|
||||
else:
|
||||
denoised_pred = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
context_timestep).unflatten(0, denoised_pred.shape[:2])
|
||||
|
||||
with torch.no_grad():
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
@@ -419,7 +528,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
# context_timestep is 0 so we use transformer_2
|
||||
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
|
||||
# current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
|
||||
current_model(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
@@ -428,8 +537,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)
|
||||
|
||||
@@ -520,12 +629,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, clean_output, exit_timestep
|
||||
|
||||
def _initialize_simulation_caches(
|
||||
self,
|
||||
@@ -622,6 +737,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:
|
||||
@@ -709,9 +848,35 @@ 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
|
||||
# raise ValueError("Transformer 2 is training")
|
||||
# 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:
|
||||
# # raise ValueError("Transformer 2 is not training")
|
||||
# assert all(p.grad is None for p in self.transformer_2.parameters())
|
||||
# grad_norm = 0.0
|
||||
# for n, p in self.transformer.named_parameters():
|
||||
# if p.grad is not None:
|
||||
# grad_norm += p.grad.data.norm(2).item() ** 2
|
||||
# else:
|
||||
# raise ValueError("Transformer param %s has no gradient", n)
|
||||
# assert grad_norm > 0.0, "Transformer has 0 gradient norm"
|
||||
|
||||
total_generator_loss += generator_loss.detach().item()
|
||||
generator_log_dict.update(gen_log_dict)
|
||||
# Store visualization data from generator training
|
||||
@@ -748,11 +913,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 = {}
|
||||
|
||||
@@ -777,6 +947,27 @@ 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:
|
||||
# raise ValueError("Fake score transformer 2 is not training")
|
||||
# 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
|
||||
@@ -976,6 +1167,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),
|
||||
@@ -1184,6 +1377,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()
|
||||
|
||||
@@ -1223,3 +1418,4 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -357,6 +359,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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user