Compare commits
56
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
80ce083ecb | ||
|
|
1abb7bb38c | ||
|
|
2211938f8c | ||
|
|
2d12867ce2 | ||
|
|
632a06ee6d | ||
|
|
9947c26b68 | ||
|
|
112495b8a9 | ||
|
|
5449d034e3 | ||
|
|
93703cb521 | ||
|
|
cb28697d29 | ||
|
|
814c99310b | ||
|
|
4d392445dc | ||
|
|
9ee229bb68 | ||
|
|
a274a90c88 | ||
|
|
be01fbe536 | ||
|
|
22d6f0cb8d | ||
|
|
b835c0cf5d | ||
|
|
c8fce0c21c | ||
|
|
1037f81960 | ||
|
|
9717ffe017 | ||
|
|
fb64178611 | ||
|
|
ee736ecb48 | ||
|
|
c0c77c7e03 | ||
|
|
793ce61c30 | ||
|
|
cbaaabf360 | ||
|
|
add1463b0b | ||
|
|
6fb1bdb0a0 | ||
|
|
d2e091150a | ||
|
|
920184403f | ||
|
|
d016119f2e | ||
|
|
066118dada | ||
|
|
a547fc1005 | ||
|
|
f27e28ad87 | ||
|
|
2b8e695721 | ||
|
|
6d28a05a1c | ||
|
|
b55a78be3c | ||
|
|
f468651f4e | ||
|
|
6b0f456d6f | ||
|
|
63eeef039d | ||
|
|
133533adb4 | ||
|
|
abdd8c851d | ||
|
|
d03eacc1c6 | ||
|
|
03747310c6 | ||
|
|
a444862561 | ||
|
|
e9802cc5f2 | ||
|
|
6ff7686a9d | ||
|
|
92d3ce1612 | ||
|
|
2b9d8e91b5 | ||
|
|
b666de819a | ||
|
|
2df82f0ee4 | ||
|
|
4f45b6e515 | ||
|
|
4c11512cde | ||
|
|
32a820c735 | ||
|
|
d6165981b4 | ||
|
|
db2e2d79ca | ||
|
|
aabf1fe135 |
@@ -104,6 +104,18 @@ steps:
|
||||
- TEST_TYPE=distillation_dmd
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/training/*self_forcing_distillation_pipeline.py"
|
||||
- "fastvideo/tests/training/self-forcing/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Self-Forcing Tests"
|
||||
env:
|
||||
- TEST_TYPE=self_forcing
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -110,6 +110,10 @@ case "$TEST_TYPE" in
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
|
||||
;;
|
||||
# run_inference_tests_vmoba
|
||||
"self_forcing")
|
||||
log "Running self-forcing tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_self_forcing_tests"
|
||||
;;
|
||||
"inference_vmoba")
|
||||
log "Running V-MoBA inference tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
|
||||
|
||||
@@ -36,9 +36,8 @@ VALIDATION_DATASET_FILE=your_validation_data_dir
|
||||
# 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
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
|
||||
--output_dir your_output_dir
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
@@ -47,16 +46,15 @@ training_args=(
|
||||
--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_frames 81
|
||||
--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_GPUS # 64
|
||||
--sp_size 1
|
||||
@@ -65,22 +63,18 @@ parallel_args=(
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--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 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
@@ -89,7 +83,6 @@ validation_args=(
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
@@ -100,7 +93,6 @@ optimizer_args=(
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
@@ -111,10 +103,9 @@ miscellaneous_args=(
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors your_ode_init_weights_path
|
||||
# --init_weights_from_safetensors your_ode_init_weights_path
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
@@ -126,13 +117,11 @@ dmd_args=(
|
||||
--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
|
||||
)
|
||||
|
||||
torchrun \
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%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="2f25ad37933894dbf0966c838c0b8494987f9f2f"
|
||||
# export WANDB_API_KEY='your_wandb_api_key_here'
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
|
||||
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
|
||||
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
|
||||
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
|
||||
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
|
||||
|
||||
# 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/SFwan2.2_t2v_finetune"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--simulate_generator_forward
|
||||
--log_visualization
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
|
||||
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
|
||||
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus 32 # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim 32
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -39,6 +39,8 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
@@ -73,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -39,6 +39,8 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
@@ -73,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -39,6 +39,8 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
@@ -73,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -40,6 +40,8 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
@@ -74,6 +76,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -40,6 +40,8 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
@@ -73,6 +75,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -14,6 +14,8 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
|
||||
@@ -50,6 +52,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -14,6 +14,8 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
@@ -51,6 +53,8 @@ parallel_args=(
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
|
||||
@@ -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.2-T2V-A14B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"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,
|
||||
|
||||
@@ -54,6 +54,9 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp32", ))
|
||||
|
||||
# self-forcing params
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
@@ -133,6 +136,11 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
# self-forcing params
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
@@ -55,6 +55,9 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
|
||||
# Causal Self-Forcing Wan2.2
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -133,6 +133,7 @@ class FastVideoArgs:
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
torch_compile_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
disable_autocast: bool = False
|
||||
|
||||
@@ -164,7 +165,7 @@ class FastVideoArgs:
|
||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||
|
||||
# MoE parameters used by Wan2.2
|
||||
boundary_ratio: float | None = None
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
@@ -330,6 +331,13 @@ class FastVideoArgs:
|
||||
help="Use torch.compile to speed up DiT inference." +
|
||||
"However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--torch-compile-kwargs",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--dit-cpu-offload",
|
||||
@@ -431,6 +439,21 @@ class FastVideoArgs:
|
||||
mode_value = getattr(args, attr, FastVideoArgs.mode.value)
|
||||
kwargs['mode'] = ExecutionMode.from_string(
|
||||
mode_value) if isinstance(mode_value, str) else mode_value
|
||||
elif attr == 'torch_compile_kwargs':
|
||||
# Parse JSON string for torch.compile kwargs
|
||||
torch_compile_kwargs_str = getattr(args, 'torch_compile_kwargs',
|
||||
None)
|
||||
if torch_compile_kwargs_str:
|
||||
try:
|
||||
import json
|
||||
kwargs['torch_compile_kwargs'] = json.loads(
|
||||
torch_compile_kwargs_str)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(
|
||||
f"Invalid JSON for torch_compile_kwargs: {e}"
|
||||
) from e
|
||||
else:
|
||||
kwargs['torch_compile_kwargs'] = {}
|
||||
elif attr == 'workload_type':
|
||||
# Convert string to WorkloadType enum
|
||||
workload_type_value = getattr(args, 'workload_type',
|
||||
@@ -610,10 +633,8 @@ class TrainingArgs(FastVideoArgs):
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
generator_model_path: str = "" # path for generator (student) model
|
||||
real_score_model_path: str = "" # path for real score (teacher) model
|
||||
fake_score_model_path: str = "" # path for fake score (critic) model
|
||||
|
||||
@@ -640,6 +661,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
@@ -711,7 +733,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
independent_first_frame: bool = False
|
||||
enable_gradient_masking: bool = True
|
||||
gradient_mask_last_n_frames: int = 21
|
||||
validate_cache_structure: bool = False # Debug flag for cache validation
|
||||
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
@@ -904,6 +925,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--init-weights-from-safetensors",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
parser.add_argument(
|
||||
"--init-weights-from-safetensors-2",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
parser.add_argument("--logging-dir",
|
||||
type=str,
|
||||
help="Directory for logging")
|
||||
|
||||
@@ -179,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -212,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
# Only T2V for now
|
||||
@@ -226,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -252,29 +250,29 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
num_frames = temb.shape[1]
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
e = self.scale_shift_table + temb
|
||||
# e.shape: [batch_size, num_frames, 6, inner_dim]
|
||||
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
assert shift_msa.dtype == torch.float32
|
||||
# assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
|
||||
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -288,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -298,13 +294,10 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
crossattn_cache=crossattn_cache)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -367,8 +360,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -491,12 +483,16 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -543,14 +539,9 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output
|
||||
return torch.stack(output)
|
||||
|
||||
def _forward_train(self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -591,8 +582,8 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.block_mask is None:
|
||||
@@ -605,8 +596,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
@@ -641,14 +636,9 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output
|
||||
return torch.stack(output)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -659,3 +649,30 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
return self._forward_inference(*args, **kwargs)
|
||||
else:
|
||||
return self._forward_train(*args, **kwargs)
|
||||
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
@@ -442,7 +442,8 @@ class TransformerLoader(ComponentLoader):
|
||||
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
|
||||
|
||||
if use_custom_weights:
|
||||
logger.info("Using custom initialization weights from: %s", custom_weights_path)
|
||||
if 'transformer_2' in model_path:
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors_2', None)
|
||||
assert custom_weights_path is not None, "Custom initialization weights must be provided"
|
||||
if os.path.isdir(custom_weights_path):
|
||||
safetensors_list = glob.glob(
|
||||
@@ -461,6 +462,7 @@ class TransformerLoader(ComponentLoader):
|
||||
logger.info("Loading model from %s, default_dtype: %s", cls_name,
|
||||
default_dtype)
|
||||
assert fastvideo_args.hsdp_shard_dim is not None
|
||||
logger.info("Loading model with dit_cpu_offload: %s", fastvideo_args.dit_cpu_offload)
|
||||
model = maybe_load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
init_params={
|
||||
@@ -479,7 +481,9 @@ class TransformerLoader(ComponentLoader):
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
training_mode=fastvideo_args.training_mode)
|
||||
training_mode=fastvideo_args.training_mode,
|
||||
enable_torch_compile=fastvideo_args.enable_torch_compile,
|
||||
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs)
|
||||
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
|
||||
@@ -54,7 +54,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
|
||||
torch.set_default_dtype(old_dtype)
|
||||
|
||||
|
||||
# TODO(PY): add compile option
|
||||
# Supports optional torch.compile for FSDP-wrapped models during training
|
||||
def maybe_load_fsdp_model(
|
||||
model_cls: type[nn.Module],
|
||||
init_params: dict[str, Any],
|
||||
@@ -70,6 +70,8 @@ def maybe_load_fsdp_model(
|
||||
output_dtype: torch.dtype | None = None,
|
||||
training_mode: bool = True,
|
||||
pin_cpu_memory: bool = True,
|
||||
enable_torch_compile: bool = False,
|
||||
torch_compile_kwargs: dict[str, Any] | None = None,
|
||||
) -> torch.nn.Module:
|
||||
"""
|
||||
Load the model with FSDP if is training, else load the model without FSDP.
|
||||
@@ -139,6 +141,14 @@ def maybe_load_fsdp_model(
|
||||
# Avoid unintended computation graph accumulation during inference
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
p.requires_grad = False
|
||||
|
||||
compile_in_loader = enable_torch_compile and training_mode
|
||||
if compile_in_loader:
|
||||
compile_kwargs = torch_compile_kwargs or {}
|
||||
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s",
|
||||
compile_kwargs)
|
||||
model = torch.compile(model, **compile_kwargs)
|
||||
logger.info("torch.compile enabled for %s", type(model).__name__)
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -49,6 +49,7 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=CausalDMDDenosingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
|
||||
@@ -99,9 +99,30 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
self.initialize_pipeline(self.fastvideo_args)
|
||||
if self.fastvideo_args.enable_torch_compile:
|
||||
self.modules["transformer"] = torch.compile(
|
||||
self.modules["transformer"])
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
transformer_module = self.modules["transformer"]
|
||||
if self.fastvideo_args.training_mode:
|
||||
logger.info(
|
||||
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
|
||||
)
|
||||
else:
|
||||
fsdp_module_cls = None
|
||||
try:
|
||||
from torch.distributed.fsdp import FSDPModule # type: ignore
|
||||
fsdp_module_cls = FSDPModule
|
||||
except Exception: # pragma: no cover - FSDP not always available
|
||||
fsdp_module_cls = None
|
||||
if fsdp_module_cls is not None and isinstance(
|
||||
transformer_module, fsdp_module_cls):
|
||||
logger.info(
|
||||
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
|
||||
)
|
||||
else:
|
||||
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
|
||||
logger.info("Enabling torch.compile for DiT with kwargs=%s",
|
||||
compile_kwargs)
|
||||
self.modules["transformer"] = torch.compile(
|
||||
transformer_module, **compile_kwargs)
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
|
||||
if not self.fastvideo_args.training_mode:
|
||||
logger.info("Creating pipeline stages...")
|
||||
@@ -144,7 +165,7 @@ class ComposedPipelineBase(ABC):
|
||||
for key, value in kwargs.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
|
||||
fastvideo_args.dit_cpu_offload = False
|
||||
fastvideo_args.dit_cpu_offload = True # TODO: fix this so it isn't a hardcode
|
||||
# we hijack the precision to be the master weight type so that the
|
||||
# model is loaded with the correct precision. Subsequently we will
|
||||
# use FSDP2's MixedPrecisionPolicy to set the precision for the
|
||||
|
||||
@@ -34,13 +34,15 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
Denoising stage for causal diffusion.
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__(transformer, scheduler)
|
||||
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
|
||||
super().__init__(transformer, scheduler, transformer_2)
|
||||
# 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
|
||||
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
||||
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
|
||||
self.num_transformer_blocks = len(self.transformer.blocks)
|
||||
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
|
||||
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
|
||||
|
||||
@@ -65,21 +67,18 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
-1] * self.transformer.config.arch_config.patch_size[-2]
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
# TODO(will): make this a parameter once we add i2v support
|
||||
independent_first_frame = self.transformer.independent_first_frame
|
||||
|
||||
independent_first_frame = self.transformer.independent_first_frame if hasattr(
|
||||
self.transformer, 'independent_first_frame') else False
|
||||
# Timesteps for DMD
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
logger.info("Warping timesteps...")
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
logger.info("Using timesteps: %s", timesteps)
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_kwargs: dict = {}
|
||||
@@ -223,6 +222,10 @@ 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:
|
||||
current_model = self.transformer_2
|
||||
else:
|
||||
current_model = self.transformer
|
||||
# Copy for pred conversion
|
||||
noise_latents = noise_latents_btchw.clone()
|
||||
latent_model_input = current_latents.to(target_dtype)
|
||||
@@ -273,7 +276,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
pred_noise_btchw = self.transformer(
|
||||
pred_noise_btchw = current_model(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
@@ -338,7 +341,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
_ = self.transformer(
|
||||
_ = current_model(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
@@ -360,8 +363,17 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
# Handle FSDP-wrapped models by using getattr
|
||||
num_attention_heads = getattr(self.transformer, 'num_attention_heads',
|
||||
None)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
None)
|
||||
|
||||
if num_attention_heads is None or attention_head_dim is None:
|
||||
raise AttributeError(
|
||||
f"Could not access num_attention_heads or attention_head_dim from transformer. "
|
||||
f"Model type: {type(self.transformer).__name__}")
|
||||
|
||||
if self.local_attn_size != -1:
|
||||
kv_cache_size = self.local_attn_size * self.frame_seq_length
|
||||
else:
|
||||
@@ -397,8 +409,17 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
crossattn_cache = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
# Handle FSDP-wrapped models by using getattr
|
||||
num_attention_heads = getattr(self.transformer, 'num_attention_heads',
|
||||
None)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
None)
|
||||
|
||||
if num_attention_heads is None or attention_head_dim is None:
|
||||
raise AttributeError(
|
||||
f"Could not access num_attention_heads or attention_head_dim from transformer. "
|
||||
f"Model type: {type(self.transformer).__name__}")
|
||||
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
|
||||
@@ -274,7 +274,7 @@ class DenoisingStage(PipelineStage):
|
||||
current_guidance_scale = batch.guidance_scale
|
||||
else:
|
||||
# low-noise stage in wan2.2
|
||||
if fastvideo_args.dit_cpu_offload and next(
|
||||
if fastvideo_args.dit_cpu_offload and self.transformer_2 is not None and next(
|
||||
self.transformer.parameters(
|
||||
)).device.type == 'cuda':
|
||||
self.transformer.to('cpu')
|
||||
|
||||
@@ -62,7 +62,7 @@ def run_test(pytest_command: str):
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
def run_encoder_tests():
|
||||
run_test("pytest ./fastvideo/tests/encoders -vs")
|
||||
|
||||
@@ -118,6 +118,10 @@ def run_inference_lora_tests():
|
||||
def run_distill_dmd_tests():
|
||||
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_self_forcing_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
|
||||
|
||||
@@ -33,6 +33,8 @@ def run_worker():
|
||||
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--real_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--fake_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--data_path", "data/crush-smol_processed_t2v/combined_parquet_dataset",
|
||||
"--validation_dataset_file", "examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json",
|
||||
"--train_batch_size", "1",
|
||||
|
||||
@@ -93,10 +93,10 @@ def test_lora_training():
|
||||
|
||||
# Define thresholds for LoRA training based on the provided console outputs
|
||||
fields_and_thresholds = {
|
||||
'avg_step_time': 2.0,
|
||||
'avg_step_time': 20.0, # something up with modal
|
||||
# 'grad_norm': 0.05, # too volatile for now. TODO: fix nondeterminism in training
|
||||
'step_time': 2.0,
|
||||
'train_loss': 0.03
|
||||
'step_time': 20.0, # something up with modal
|
||||
'train_loss': 0.05
|
||||
}
|
||||
|
||||
failures = []
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
import os
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29513"
|
||||
import sys
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
import torch
|
||||
import json
|
||||
from huggingface_hub import snapshot_download
|
||||
from fastvideo.utils import logger
|
||||
# Import the training pipeline
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
from fastvideo.training.wan_self_forcing_distillation_pipeline import WanSelfForcingDistillationPipeline
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
wandb_name = "test_self_forcing_distill"
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "2"
|
||||
|
||||
|
||||
def run_worker():
|
||||
"""Worker function that will be run on each GPU"""
|
||||
# Create and populate args
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
|
||||
# Set the arguments based on the distill_dmd_t2v_1.3B.sh script
|
||||
args = parser.parse_args([
|
||||
"--model_path", "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
"--real_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--fake_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"--data_path", "data/crush-smol_processed_t2v/combined_parquet_dataset",
|
||||
"--validation_dataset_file", "examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json",
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "21",
|
||||
"--num_gpus", "2",
|
||||
"--sp_size", "1",
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", "2",
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "1",
|
||||
"--gradient_accumulation_steps", "1",
|
||||
"--max_train_steps", "2",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--training_state_checkpointing_steps", "30",
|
||||
"--weight_only_checkpointing_steps", "30",
|
||||
"--validation_steps", "10",
|
||||
"--validation_sampling_steps", "3",
|
||||
"--log_validation",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--ema_start_step", "0",
|
||||
"--training_cfg_rate", "0.0",
|
||||
"--output_dir", "data/wan_self_forcing_test",
|
||||
"--tracker_project_name", "wan_self_forcing_ci",
|
||||
"--wandb_run_name", wandb_name,
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "21",
|
||||
"--flow_shift", "5",
|
||||
"--validation_guidance_scale", "1.0",
|
||||
"--weight_decay", "0.01",
|
||||
"--dit_precision", "fp32",
|
||||
"--max_grad_norm", "1.0",
|
||||
# DMD args
|
||||
"--dmd_denoising_steps", "1000,750,500", # Reduced steps for testing
|
||||
"--min_timestep_ratio", "0.02",
|
||||
"--max_timestep_ratio", "0.98",
|
||||
"--dfake_gen_update_ratio", "5",
|
||||
"--real_score_guidance_scale", "3.0",
|
||||
"--fake_score_learning_rate", "8e-6",
|
||||
"--fake_score_betas", "0.0,0.999",
|
||||
"--warp_denoising_step",
|
||||
"--enable_gradient_checkpointing_type", "full",
|
||||
# Self-forcing specific args
|
||||
"--log_visualization",
|
||||
"--simulate_generator_forward",
|
||||
"--num_frame_per_block", "3",
|
||||
"--enable_gradient_masking",
|
||||
"--gradient_mask_last_n_frames", "21",
|
||||
"--independent_first_frame", "False",
|
||||
"--same_step_across_blocks", "True",
|
||||
"--last_step_only", "False",
|
||||
"--context_noise", "0",
|
||||
"--use_ema", "True",
|
||||
"--ema_decay", "0.99",
|
||||
"--ema_start_step", "100",
|
||||
])
|
||||
|
||||
# Call the main training function
|
||||
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Self-forcing distillation training pipeline done")
|
||||
|
||||
def test_distributed_training():
|
||||
"""Test the distributed self-forcing training setup"""
|
||||
os.environ["WANDB_MODE"] = "offline"
|
||||
|
||||
data_dir = Path("data/crush-smol_processed_t2v")
|
||||
|
||||
if not data_dir.exists():
|
||||
print(f"Downloading test dataset to {data_dir}...")
|
||||
snapshot_download(
|
||||
repo_id="wlsaidhi/crush-smol_processed_t2v",
|
||||
local_dir=str(data_dir),
|
||||
repo_type="dataset",
|
||||
local_dir_use_symlinks=False
|
||||
)
|
||||
|
||||
# Get the current file path
|
||||
current_file = Path(__file__).resolve()
|
||||
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE,
|
||||
"--master_port", os.environ["MASTER_PORT"],
|
||||
str(current_file)
|
||||
]
|
||||
process = subprocess.run(cmd, capture_output=True, text=True)
|
||||
|
||||
# Print stdout and stderr for debugging
|
||||
if process.stdout:
|
||||
print("STDOUT:", process.stdout)
|
||||
if process.stderr:
|
||||
print("STDERR:", process.stderr)
|
||||
|
||||
# Check if the process failed
|
||||
if process.returncode != 0:
|
||||
print(f"Process failed with return code: {process.returncode}")
|
||||
raise subprocess.CalledProcessError(process.returncode, cmd, process.stdout, process.stderr)
|
||||
|
||||
if __name__ == "__main__":
|
||||
if os.environ.get("LOCAL_RANK") is not None:
|
||||
# We're being run by torchrun
|
||||
run_worker()
|
||||
else:
|
||||
# We're being run directly
|
||||
test_distributed_training()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -3,6 +3,7 @@ import copy
|
||||
import time
|
||||
from collections import deque
|
||||
from typing import Any
|
||||
import gc
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -48,8 +49,16 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
logger.info("Initializing self-forcing distillation pipeline...")
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
self.generator_ema_2: EMA_FSDP | None = None
|
||||
|
||||
super().initialize_training_pipeline(training_args)
|
||||
try:
|
||||
logger.info("RANK: %s, entered initialize_training_pipeline",
|
||||
self.global_rank,
|
||||
local_main_process_only=False)
|
||||
except Exception:
|
||||
logger.info("Entered initialize_training_pipeline (rank unknown)")
|
||||
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
num_inference_steps=1000,
|
||||
shift=5.0,
|
||||
@@ -59,7 +68,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.dfake_gen_update_ratio = getattr(training_args,
|
||||
'dfake_gen_update_ratio', 5)
|
||||
|
||||
# Self-forcing specific properties
|
||||
self.num_frame_per_block = getattr(training_args, 'num_frame_per_block',
|
||||
3)
|
||||
self.independent_first_frame = getattr(training_args,
|
||||
@@ -69,19 +77,25 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.last_step_only = getattr(training_args, 'last_step_only', False)
|
||||
self.context_noise = getattr(training_args, 'context_noise', 0)
|
||||
|
||||
# Calculate frame sequence length - this will be set properly in _prepare_dit_inputs
|
||||
self.frame_seq_length = 1560 # TODO: Calculate this dynamically based on patch size
|
||||
|
||||
# Cache references (will be initialized per forward pass)
|
||||
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)
|
||||
logger.info("RANK: %s, exiting initialize_training_pipeline",
|
||||
self.global_rank,
|
||||
local_main_process_only=False)
|
||||
|
||||
def generate_and_sync_list(self, num_blocks: int, num_denoising_steps: int,
|
||||
device: torch.device) -> list[int]:
|
||||
"""Generate and synchronize random exit flags across distributed processes."""
|
||||
logger.info(
|
||||
"RANK: %s, enter generate_and_sync_list blocks=%s steps=%s device=%s",
|
||||
self.global_rank,
|
||||
num_blocks,
|
||||
num_denoising_steps,
|
||||
str(device),
|
||||
local_main_process_only=False)
|
||||
rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
|
||||
if rank == 0:
|
||||
@@ -98,7 +112,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(indices,
|
||||
src=0) # Broadcast the random indices to all ranks
|
||||
return indices.tolist()
|
||||
flags = indices.tolist()
|
||||
logger.info(
|
||||
"RANK: %s, exit generate_and_sync_list flags_len=%s first=%s",
|
||||
self.global_rank,
|
||||
len(flags),
|
||||
flags[0] if len(flags) > 0 else None,
|
||||
local_main_process_only=False)
|
||||
return flags
|
||||
|
||||
def generator_loss(
|
||||
self, training_batch: TrainingBatch
|
||||
@@ -107,20 +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_timestep = self.denoising_step_list[exit_flags[0]]
|
||||
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)
|
||||
exit_timestep = None
|
||||
|
||||
# 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):
|
||||
if self.training_args.simulate_generator_forward:
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
training_batch)
|
||||
else:
|
||||
generator_pred_video = self._generator_forward(training_batch)
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
training_batch, exit_flags=exit_flags)
|
||||
|
||||
with set_forward_context(current_timestep=training_batch.timesteps,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
dmd_loss = self._dmd_forward(
|
||||
generator_pred_video=generator_pred_video,
|
||||
training_batch=training_batch)
|
||||
training_batch=training_batch,
|
||||
exit_timestep=exit_timestep)
|
||||
|
||||
log_dict = {
|
||||
"dmdtrain_gradient_norm": torch.tensor(0.0, device=self.device)
|
||||
@@ -135,6 +182,18 @@ 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
|
||||
@@ -142,82 +201,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
return flow_matching_loss, log_dict
|
||||
|
||||
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Forward pass through generator with KV cache support for causal generation."""
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
batch_size = latents.shape[0]
|
||||
|
||||
# Step 1: Sample a timestep from denoising_step_list
|
||||
index = torch.randint(0,
|
||||
len(self.denoising_step_list), [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
timestep = self.denoising_step_list[index]
|
||||
training_batch.dmd_latent_vis_dict["generator_timestep"] = timestep
|
||||
|
||||
# Step 2: Initialize KV cache and cross-attention cache for causal generation
|
||||
kv_cache, crossattn_cache = self._initialize_simulation_caches(
|
||||
batch_size, dtype, self.device)
|
||||
|
||||
if getattr(self.training_args, 'validate_cache_structure', False):
|
||||
self._validate_cache_structure(kv_cache, crossattn_cache,
|
||||
batch_size)
|
||||
|
||||
# Step 3: Add noise to latents
|
||||
noise = torch.randn(self.video_latent_shape,
|
||||
device=self.device,
|
||||
dtype=dtype)
|
||||
if self.sp_world_size > 1:
|
||||
noise = rearrange(noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
latents.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep * torch.ones([latents.shape[0] * latents.shape[1]],
|
||||
device=noise.device,
|
||||
dtype=torch.long))
|
||||
|
||||
# Step 4: Build input kwargs with KV cache support
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
|
||||
# Step 5: Forward pass with KV cache if available
|
||||
if hasattr(self.transformer, '_forward_inference'):
|
||||
# Use causal inference forward with KV cache
|
||||
pred_noise = self.transformer(
|
||||
hidden_states=training_batch.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch.
|
||||
input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch.input_kwargs.get(
|
||||
'encoder_hidden_states_image'),
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=0, # Start from beginning for single-step
|
||||
cache_start=0).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
# Fallback to regular forward
|
||||
pred_noise = self.transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Step 6: Convert noise prediction to video prediction
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=torch.tensor([timestep], device=noisy_latent.device),
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
self._reset_simulation_caches(kv_cache, crossattn_cache)
|
||||
|
||||
return pred_video
|
||||
|
||||
def _generator_multi_step_simulation_forward(
|
||||
self,
|
||||
training_batch: TrainingBatch,
|
||||
return_sim_steps: bool = False) -> torch.Tensor:
|
||||
return_sim_steps: bool = False,
|
||||
exit_flags: list[int] | None = None) -> torch.Tensor:
|
||||
"""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
|
||||
@@ -290,13 +278,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
dtype=noise.dtype)
|
||||
|
||||
# Step 1: Initialize KV cache to all zeros
|
||||
self.kv_cache1, self.crossattn_cache = self._initialize_simulation_caches(
|
||||
batch_size, dtype, self.device)
|
||||
|
||||
# Validate cache structure (can be disabled in production)
|
||||
if getattr(self.training_args, 'validate_cache_structure', False):
|
||||
self._validate_cache_structure(self.kv_cache1, self.crossattn_cache,
|
||||
batch_size)
|
||||
cache_frames = num_generated_frames + num_input_frames
|
||||
kv_cache1, crossattn_cache = self._initialize_simulation_caches(
|
||||
batch_size, dtype, self.device, max_num_frames=cache_frames)
|
||||
|
||||
# Step 2: Cache context feature
|
||||
current_start_frame = 0
|
||||
@@ -309,8 +293,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
initial_latent, timestep * 0,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
self.transformer(
|
||||
# we process the image latent with self.transformer_2 (low-noise expert)
|
||||
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'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -318,8 +303,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
|
||||
@@ -329,9 +314,11 @@ 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),
|
||||
num_denoising_steps,
|
||||
device=noise.device)
|
||||
|
||||
if exit_flags is None:
|
||||
exit_flags = self.generate_and_sync_list(len(all_num_frames),
|
||||
num_denoising_steps,
|
||||
device=noise.device)
|
||||
start_gradient_frame_index = max(0, num_output_frames - 21)
|
||||
|
||||
for block_index, current_num_frames in enumerate(all_num_frames):
|
||||
@@ -350,6 +337,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
device=noise.device,
|
||||
dtype=torch.int64) * current_timestep
|
||||
|
||||
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
|
||||
else:
|
||||
current_model = self.transformer
|
||||
|
||||
if not exit_flag:
|
||||
with torch.no_grad():
|
||||
# Build input kwargs
|
||||
@@ -357,7 +349,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noisy_input, timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = self.transformer(
|
||||
pred_flow = current_model(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -366,8 +358,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(
|
||||
@@ -390,6 +382,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
dtype=torch.long)).unflatten(
|
||||
0, denoised_pred.shape[:2])
|
||||
else:
|
||||
logger.info(
|
||||
"Exit timestep in _generator_multi_step_simulation_forward(): %s",
|
||||
current_timestep)
|
||||
# Final prediction with gradient control
|
||||
if current_start_frame < start_gradient_frame_index:
|
||||
with torch.no_grad():
|
||||
@@ -397,7 +392,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noisy_input, timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = self.transformer(
|
||||
pred_flow = current_model(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -406,8 +401,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(
|
||||
@@ -417,7 +412,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noisy_input, timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = self.transformer(
|
||||
pred_flow = current_model(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -426,8 +421,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(
|
||||
@@ -457,7 +452,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
denoised_pred, context_timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
self.transformer(
|
||||
# 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(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -465,8 +462,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)
|
||||
|
||||
@@ -557,49 +554,46 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.dmd_latent_vis_dict["min_num_frames"] = torch.tensor(
|
||||
min_num_frames, dtype=torch.float32, device=self.device)
|
||||
|
||||
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
|
||||
|
||||
def _initialize_simulation_caches(
|
||||
self, batch_size: int, dtype: torch.dtype, device: torch.device
|
||||
self,
|
||||
batch_size: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
*,
|
||||
max_num_frames: int | None = None,
|
||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
|
||||
num_transformer_blocks = len(self.transformer.blocks)
|
||||
latent_shape = self.video_latent_shape_sp
|
||||
_, num_frames, _, height, width = latent_shape
|
||||
|
||||
# Calculate frame sequence length based on input dimensions and patch size
|
||||
# From the training batch, we can get the actual latent dimensions
|
||||
latent_shape = self.video_latent_shape_sp # This is set in _prepare_dit_inputs
|
||||
batch_size_actual, num_frames, num_channels, height, width = latent_shape
|
||||
|
||||
# Get patch size from transformer config
|
||||
p_t, p_h, p_w = self.transformer.patch_size
|
||||
_, p_h, p_w = self.transformer.patch_size
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# Frame sequence length is the spatial sequence length per frame
|
||||
frame_seq_length = post_patch_height * post_patch_width
|
||||
|
||||
# Get local attention size from transformer config
|
||||
# local_attn_size = getattr(self.transformer, 'local_attn_size', -1)
|
||||
self.frame_seq_length = frame_seq_length
|
||||
|
||||
# Get model configuration parameters - handle FSDP wrapping
|
||||
if hasattr(self.transformer, 'config'):
|
||||
config = self.transformer.config
|
||||
num_attention_heads = config.num_attention_heads
|
||||
attention_head_dim = config.attention_head_dim
|
||||
text_len = config.text_len
|
||||
else:
|
||||
# Fallback to direct attribute access for non-FSDP models
|
||||
num_attention_heads = getattr(self.transformer,
|
||||
'num_attention_heads', 40)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
128)
|
||||
text_len = getattr(self.transformer, 'text_len', 512)
|
||||
num_attention_heads = getattr(self.transformer, 'num_attention_heads',
|
||||
None)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
None)
|
||||
text_len = getattr(self.transformer, 'text_len', None)
|
||||
|
||||
num_max_frames = getattr(self.training_args, "num_frames", num_frames)
|
||||
if max_num_frames is None:
|
||||
max_num_frames = num_frames
|
||||
num_max_frames = max(max_num_frames, num_frames)
|
||||
kv_cache_size = num_max_frames * frame_seq_length
|
||||
|
||||
kv_cache = []
|
||||
@@ -665,64 +659,37 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
cache_dict["k"].zero_()
|
||||
cache_dict["v"].zero_()
|
||||
|
||||
def _validate_cache_structure(self, kv_cache, crossattn_cache,
|
||||
batch_size: int):
|
||||
"""Validate that cache structures are correctly initialized."""
|
||||
num_transformer_blocks = len(self.transformer.blocks)
|
||||
|
||||
# Get model configuration parameters - handle FSDP wrapping
|
||||
if hasattr(self.transformer, 'config'):
|
||||
config = self.transformer.config
|
||||
num_attention_heads = config.num_attention_heads
|
||||
attention_head_dim = config.attention_head_dim
|
||||
text_len = config.text_len
|
||||
else:
|
||||
# Fallback to direct attribute access for non-FSDP models
|
||||
num_attention_heads = getattr(self.transformer,
|
||||
'num_attention_heads', 40)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
128)
|
||||
text_len = getattr(self.transformer, 'text_len', 512)
|
||||
|
||||
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:
|
||||
assert len(
|
||||
kv_cache
|
||||
) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(kv_cache)}"
|
||||
for i, cache_dict in enumerate(kv_cache):
|
||||
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in kv_cache block {i}"
|
||||
assert "global_end_index" in cache_dict and "local_end_index" in cache_dict, f"Missing indices in kv_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
|
||||
assert cache_dict["v"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
2] == num_attention_heads, f"Attention heads mismatch in kv_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
3] == attention_head_dim, f"Attention head dim mismatch in kv_cache block {i}"
|
||||
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:
|
||||
assert len(
|
||||
crossattn_cache
|
||||
) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(crossattn_cache)}"
|
||||
for i, cache_dict in enumerate(crossattn_cache):
|
||||
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in crossattn_cache block {i}"
|
||||
assert "is_init" in cache_dict, f"Missing is_init in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["v"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
1] == text_len, f"Text length mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
2] == num_attention_heads, f"Attention heads mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
3] == attention_head_dim, f"Attention head dim mismatch in crossattn_cache block {i}"
|
||||
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:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
@@ -753,7 +720,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
|
||||
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
@@ -770,7 +736,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
batch = self._prepare_distillation(training_batch)
|
||||
batch = self._get_next_batch(batch)
|
||||
batch = self._normalize_dit_input(batch)
|
||||
batch = self._prepare_dit_inputs(batch)
|
||||
batch = self._prepare_dit_inputs(batch, prepare_timesteps=False)
|
||||
batch = self._build_attention_metadata(batch)
|
||||
batch.attn_metadata_vsa = copy.deepcopy(batch.attn_metadata)
|
||||
if batch.attn_metadata is not None:
|
||||
@@ -781,7 +747,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
if train_generator:
|
||||
logger.debug("Training generator at step %s",
|
||||
self.current_trainstep)
|
||||
self.optimizer.zero_grad()
|
||||
if self.transformer_2 is not None:
|
||||
self.optimizer_2.zero_grad()
|
||||
total_generator_loss = 0.0
|
||||
generator_log_dict = {}
|
||||
|
||||
@@ -803,9 +773,39 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
setattr(batch_gen, key, copy.deepcopy(value))
|
||||
|
||||
generator_loss, gen_log_dict = self.generator_loss(batch_gen)
|
||||
logger.info("train_transformer_2: %s", self.train_transformer_2)
|
||||
|
||||
with set_forward_context(current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata):
|
||||
(generator_loss / gradient_accumulation_steps).backward()
|
||||
|
||||
# Ensure that only one of the two transformers have received gradients
|
||||
if self.train_transformer_2:
|
||||
# Assert that all gradients are None for transformer
|
||||
assert all(p.grad is None
|
||||
for p in self.transformer.parameters())
|
||||
grad_sum = 0
|
||||
for n, p in self.transformer_2.named_parameters():
|
||||
if p.grad is not None:
|
||||
grad_sum += p.grad.sum().item()
|
||||
else:
|
||||
raise ValueError(
|
||||
"Transformer 2 param %s has no gradient", n)
|
||||
logger.info("Transformer 2 gradient sum for all params %s",
|
||||
grad_sum)
|
||||
else:
|
||||
assert all(p.grad is None
|
||||
for p in self.transformer_2.parameters())
|
||||
grad_sum = 0
|
||||
for n, p in self.transformer.named_parameters():
|
||||
if p.grad is not None:
|
||||
grad_sum += p.grad.sum().item()
|
||||
else:
|
||||
raise ValueError(
|
||||
"Transformer param %s has no gradient", n)
|
||||
logger.info("Transformer gradient sum for all params %s",
|
||||
grad_sum)
|
||||
|
||||
total_generator_loss += generator_loss.detach().item()
|
||||
generator_log_dict.update(gen_log_dict)
|
||||
# Store visualization data from generator training
|
||||
@@ -813,12 +813,27 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.dmd_latent_vis_dict.update(
|
||||
batch_gen.dmd_latent_vis_dict)
|
||||
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
# Only clip gradients and step optimizer for the model that is currently training
|
||||
if hasattr(
|
||||
self, 'train_transformer_2'
|
||||
) and self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer_2)
|
||||
self.optimizer_2.step()
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if self.generator_ema is not None:
|
||||
self.generator_ema.update(self.transformer)
|
||||
if hasattr(
|
||||
self, 'train_transformer_2'
|
||||
) and self.train_transformer_2 and self.transformer_2 is not None:
|
||||
# Update EMA for transformer_2 when training it
|
||||
if self.generator_ema_2 is not None:
|
||||
self.generator_ema_2.update(self.transformer_2)
|
||||
else:
|
||||
self.generator_ema.update(self.transformer)
|
||||
|
||||
avg_generator_loss = torch.tensor(total_generator_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -830,6 +845,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
else:
|
||||
training_batch.generator_loss = 0.0
|
||||
|
||||
logger.debug("Training critic at step %s", self.current_trainstep)
|
||||
self.fake_score_optimizer.zero_grad()
|
||||
total_critic_loss = 0.0
|
||||
critic_log_dict = {}
|
||||
@@ -862,9 +878,16 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.fake_score_latent_vis_dict.update(
|
||||
batch_critic.fake_score_latent_vis_dict)
|
||||
|
||||
self._clip_model_grad_norm_(batch_critic, self.fake_score_transformer)
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
if self.train_fake_score_transformer_2 and self.fake_score_transformer_2 is not None:
|
||||
self._clip_model_grad_norm_(batch_critic,
|
||||
self.fake_score_transformer_2)
|
||||
self.fake_score_optimizer_2.step()
|
||||
self.fake_score_lr_scheduler_2.step()
|
||||
else:
|
||||
self._clip_model_grad_norm_(batch_critic,
|
||||
self.fake_score_transformer)
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
|
||||
avg_critic_loss = torch.tensor(total_critic_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -875,7 +898,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.fake_score_loss = avg_critic_loss.item()
|
||||
|
||||
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_loss
|
||||
|
||||
return training_batch
|
||||
|
||||
def _log_training_info(self) -> None:
|
||||
@@ -890,7 +912,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
wandb_loss_dict = {}
|
||||
|
||||
# Debug logging
|
||||
logger.info("Step %s: Starting visualization", step)
|
||||
if hasattr(training_batch, 'dmd_latent_vis_dict'):
|
||||
logger.info("DMD latent keys: %s",
|
||||
list(training_batch.dmd_latent_vis_dict.keys()))
|
||||
@@ -934,20 +955,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
try:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
logger.info("Successfully processed DMD latent: %s",
|
||||
latent_key)
|
||||
except Exception as e:
|
||||
logger.error("Error processing DMD latent %s: %s",
|
||||
latent_key, str(e))
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
del video, latents
|
||||
|
||||
# Process critic predictions
|
||||
@@ -982,20 +997,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
try:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
logger.info("Successfully processed critic latent: %s",
|
||||
latent_key)
|
||||
except Exception as e:
|
||||
logger.error("Error processing critic latent %s: %s",
|
||||
latent_key, str(e))
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
del video, latents
|
||||
|
||||
# Log metadata
|
||||
@@ -1035,8 +1044,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
# Use the same seed for all processes within the same SP group
|
||||
sp_group_seed = seed + (self.global_rank // self.sp_world_size)
|
||||
set_random_seed(sp_group_seed)
|
||||
logger.info("Rank %s: Using SP group seed %s", self.global_rank,
|
||||
sp_group_seed)
|
||||
else:
|
||||
set_random_seed(seed + self.global_rank)
|
||||
|
||||
@@ -1063,6 +1070,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),
|
||||
@@ -1099,6 +1108,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
logger.info("Created generator EMA at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
# Create EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None and self.generator_ema_2 is None:
|
||||
self.generator_ema_2 = EMA_FSDP(
|
||||
self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info(
|
||||
"Created generator EMA_2 at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -1125,6 +1142,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
"ema":
|
||||
"✓" if (self.generator_ema is not None and self.is_ema_ready())
|
||||
else "✗",
|
||||
"ema2":
|
||||
"✓" if (self.generator_ema_2 is not None
|
||||
and self.is_ema_ready()) else "✗",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
@@ -1150,11 +1170,13 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
if use_vsa:
|
||||
log_data["VSA_train_sparsity"] = current_vsa_sparsity
|
||||
|
||||
if self.generator_ema is not None:
|
||||
log_data["ema_enabled"] = True
|
||||
if self.generator_ema is not None or self.generator_ema_2 is not None:
|
||||
log_data["ema_enabled"] = self.generator_ema is not None
|
||||
log_data["ema_2_enabled"] = self.generator_ema_2 is not None
|
||||
log_data["ema_decay"] = self.training_args.ema_decay
|
||||
else:
|
||||
log_data["ema_enabled"] = False
|
||||
log_data["ema_2_enabled"] = False
|
||||
|
||||
ema_stats = self.get_ema_stats()
|
||||
log_data.update(ema_stats)
|
||||
@@ -1190,12 +1212,34 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
print("rank", self.global_rank,
|
||||
"save training state checkpoint at step", step)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer, self.fake_score_transformer,
|
||||
self.global_rank, self.training_args.output_dir, step,
|
||||
self.optimizer, self.fake_score_optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator,
|
||||
self.generator_ema)
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
step,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2',
|
||||
None),
|
||||
real_score_transformer_2=getattr(
|
||||
self, 'real_score_transformer_2', None),
|
||||
fake_score_transformer_2=getattr(
|
||||
self, 'fake_score_transformer_2', None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self,
|
||||
'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self,
|
||||
'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
@@ -1206,19 +1250,38 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.training_args.weight_only_checkpointing_steps == 0):
|
||||
print("rank", self.global_rank,
|
||||
"save weight-only checkpoint at step", step)
|
||||
save_distillation_checkpoint(self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2',
|
||||
None),
|
||||
real_score_transformer_2=getattr(
|
||||
self, 'real_score_transformer_2', None),
|
||||
fake_score_transformer_2=getattr(
|
||||
self, 'fake_score_transformer_2', None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self,
|
||||
'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self,
|
||||
'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir, step)
|
||||
|
||||
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()
|
||||
|
||||
@@ -1226,11 +1289,31 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
"save final training state checkpoint at step",
|
||||
self.training_args.max_train_steps)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer, self.fake_score_transformer, self.global_rank,
|
||||
self.training_args.output_dir, self.training_args.max_train_steps,
|
||||
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator, self.generator_ema)
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2', None),
|
||||
real_score_transformer_2=getattr(self, 'real_score_transformer_2',
|
||||
None),
|
||||
fake_score_transformer_2=getattr(self, 'fake_score_transformer_2',
|
||||
None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self, 'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self, 'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir,
|
||||
|
||||
@@ -63,6 +63,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[dict[str, Any]]
|
||||
current_epoch: int = 0
|
||||
train_transformer_2: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -98,6 +99,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.sp_world_size = self.sp_group.world_size
|
||||
self.local_rank = world_group.local_rank
|
||||
self.transformer = self.get_module("transformer")
|
||||
self.transformer_2 = self.get_module("transformer_2", None)
|
||||
self.seed = training_args.seed
|
||||
self.set_schemas()
|
||||
|
||||
@@ -105,22 +107,31 @@ 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,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2 = apply_activation_checkpointing(
|
||||
self.transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
# Set grads for proper modules based on the training mode (Distill, LoRA, etc.)
|
||||
self.set_trainable()
|
||||
params_to_optimize = self.transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
# Parse betas from string format "beta1,beta2"
|
||||
betas_str = training_args.betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
@@ -138,6 +149,30 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
if self.transformer_2 is not None:
|
||||
# Ensure transformer_2 has trainable parameters before creating optimizer
|
||||
self.transformer_2.train()
|
||||
self.transformer_2.requires_grad_(True)
|
||||
params_to_optimize_2 = self.transformer_2.parameters()
|
||||
params_to_optimize_2 = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize_2))
|
||||
self.optimizer_2 = torch.optim.AdamW(
|
||||
params_to_optimize_2,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
self.lr_scheduler_2 = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer_2,
|
||||
num_warmup_steps=training_args.lr_warmup_steps,
|
||||
num_training_steps=training_args.max_train_steps,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
@@ -152,6 +187,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
if self.training_args.boundary_ratio is not None:
|
||||
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
|
||||
else:
|
||||
self.boundary_timestep = None
|
||||
|
||||
logger.info("train_dataloader length: %s", len(self.train_dataloader))
|
||||
logger.info("train_sp_batch_size: %s",
|
||||
training_args.train_sp_batch_size)
|
||||
logger.info("gradient_accumulation_steps: %s",
|
||||
training_args.gradient_accumulation_steps)
|
||||
logger.info("sp_size: %s", training_args.sp_size)
|
||||
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
@@ -178,9 +224,34 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
self.transformer.train()
|
||||
self.optimizer.zero_grad()
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2.train()
|
||||
self.optimizer_2.zero_grad()
|
||||
training_batch.total_loss = 0.0
|
||||
return training_batch
|
||||
|
||||
def _enable_training(self, model: torch.nn.Module,
|
||||
optimizer: torch.optim.Optimizer) -> None:
|
||||
"""Enable training mode and gradients for the specified model."""
|
||||
for param in model.parameters():
|
||||
param.requires_grad = True
|
||||
model.train()
|
||||
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
|
||||
model.eval()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
def _get_moe_timestep_type(self, timestep: torch.Tensor) -> str:
|
||||
if timestep.item() < self.boundary_timestep:
|
||||
return "low" # low noise
|
||||
else:
|
||||
return "high" # high noise
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
@@ -217,24 +288,28 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
training_batch: TrainingBatch,
|
||||
prepare_timesteps: bool = True) -> TrainingBatch:
|
||||
latents = training_batch.latents
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn(latents.shape,
|
||||
generator=self.noise_gen_cuda,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
timesteps = self.noise_scheduler.timesteps[indices].to(
|
||||
device=latents.device)
|
||||
if prepare_timesteps:
|
||||
timesteps = self._sample_timesteps(batch_size, latents.device)
|
||||
|
||||
# Enable training for the model that will be trained next and disable the other
|
||||
if self.train_transformer_2:
|
||||
self._enable_training(self.transformer_2, self.optimizer_2)
|
||||
self._disable_training(self.transformer, self.optimizer)
|
||||
else:
|
||||
self._enable_training(self.transformer, self.optimizer)
|
||||
if self.transformer_2 is not None:
|
||||
self._disable_training(self.transformer_2, self.optimizer_2)
|
||||
else:
|
||||
timesteps = 0 # Fill in a dummy value
|
||||
|
||||
if self.training_args.sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
sp_group = get_sp_group()
|
||||
@@ -257,6 +332,59 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
return training_batch
|
||||
|
||||
def _sample_timesteps(self, batch_size: int,
|
||||
device: torch.device) -> torch.Tensor:
|
||||
# Determine which model to train based on the boundary timestep
|
||||
if (self.transformer_2 is not None
|
||||
and self.boundary_timestep is not None
|
||||
and torch.rand(1, generator=self.noise_random_generator).item()
|
||||
<= self.training_args.boundary_ratio):
|
||||
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)
|
||||
self.train_transformer_2 = decision.item() == 1.0
|
||||
|
||||
# Sample u from the appropriate range
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
|
||||
boundary_ratio = self.training_args.boundary_ratio
|
||||
if self.train_transformer_2:
|
||||
u = (1 - boundary_ratio
|
||||
) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
|
||||
# elif self.transformer_2 is not None:
|
||||
# u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
|
||||
# else: # patch for now to align with non-MoE timestep logic
|
||||
# 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(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
latents_shape = training_batch.raw_latent_shape
|
||||
@@ -321,11 +449,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# [1000.0],
|
||||
# device=training_batch.noisy_model_input.device,
|
||||
# dtype=torch.bfloat16)
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.current_timestep,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
model_pred = self.transformer(**input_kwargs)
|
||||
model_pred = current_model(**input_kwargs)
|
||||
if self.training_args.precondition_outputs:
|
||||
assert training_batch.sigmas is not None
|
||||
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
|
||||
@@ -356,7 +485,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
model_parts = [self.transformer]
|
||||
# Only clip gradients for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
model_parts = [self.transformer_2]
|
||||
else:
|
||||
model_parts = [self.transformer]
|
||||
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
@@ -401,8 +535,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
training_batch = self._clip_grad_norm(training_batch)
|
||||
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
# Only step the optimizer and scheduler for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self.optimizer_2.step()
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
training_batch.total_loss = training_batch.total_loss
|
||||
training_batch.grad_norm = training_batch.grad_norm
|
||||
@@ -435,6 +574,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
logger.info("Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
num_trainable_params = count_trainable(self.transformer_2)
|
||||
logger.info(
|
||||
"Transformer 2: Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
@@ -477,7 +622,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
vsa_sparsity // vsa_decay_rate)
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
elif vmoba_available:
|
||||
# TODO: add vmoba sparsity scheduling here
|
||||
#TODO: add vmoba sparsity scheduling here
|
||||
current_vsa_sparsity = 0.0
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
@@ -512,7 +657,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
},
|
||||
step=step,
|
||||
)
|
||||
if step % self.training_args.checkpointing_steps == 0:
|
||||
if step % self.training_args.training_state_checkpointing_steps == 0:
|
||||
save_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.output_dir, step,
|
||||
self.optimizer, self.train_dataloader,
|
||||
@@ -610,7 +755,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
Generate a validation video and log it to wandb to check the quality during training.
|
||||
"""
|
||||
training_args.inference_mode = True
|
||||
training_args.dit_cpu_offload = True
|
||||
training_args.dit_cpu_offload = False
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
@@ -631,7 +776,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
validation_dataloader = DataLoader(validation_dataset,
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
transformer.eval()
|
||||
|
||||
self.transformer.eval()
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
self.transformer_2.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
@@ -723,7 +871,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
self.transformer.train()
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
self.transformer_2.train()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
|
||||
@@ -191,27 +191,48 @@ def save_checkpoint(transformer,
|
||||
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
|
||||
|
||||
|
||||
def save_distillation_checkpoint(generator_transformer,
|
||||
fake_score_transformer,
|
||||
rank,
|
||||
output_dir,
|
||||
step,
|
||||
generator_optimizer=None,
|
||||
fake_score_optimizer=None,
|
||||
dataloader=None,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None,
|
||||
generator_ema=None,
|
||||
only_save_generator_weight=False) -> None:
|
||||
def save_distillation_checkpoint(
|
||||
generator_transformer,
|
||||
fake_score_transformer,
|
||||
rank,
|
||||
output_dir,
|
||||
step,
|
||||
generator_optimizer=None,
|
||||
fake_score_optimizer=None,
|
||||
dataloader=None,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None,
|
||||
generator_ema=None,
|
||||
only_save_generator_weight=False,
|
||||
# MoE support
|
||||
generator_transformer_2=None,
|
||||
real_score_transformer_2=None,
|
||||
fake_score_transformer_2=None,
|
||||
generator_optimizer_2=None,
|
||||
fake_score_optimizer_2=None,
|
||||
generator_scheduler_2=None,
|
||||
fake_score_scheduler_2=None,
|
||||
generator_ema_2=None) -> None:
|
||||
"""
|
||||
Save distillation checkpoint with both generator and fake_score models.
|
||||
Supports MoE (Mixture of Experts) models with transformer_2 variants.
|
||||
Saves both distributed checkpoint and consolidated model weights.
|
||||
Only saves the generator model for inference (consolidated weights).
|
||||
|
||||
Args:
|
||||
generator_transformer: Main generator transformer model
|
||||
fake_score_transformer: Main fake score transformer model
|
||||
only_save_generator_weight: If True, only save the generator model weights for inference
|
||||
without saving distributed checkpoint for training resume.
|
||||
generator_transformer_2: Secondary generator transformer for MoE (optional)
|
||||
real_score_transformer_2: Secondary real score transformer for MoE (optional)
|
||||
fake_score_transformer_2: Secondary fake score transformer for MoE (optional)
|
||||
generator_optimizer_2: Optimizer for generator_transformer_2 (optional)
|
||||
fake_score_optimizer_2: Optimizer for fake_score_transformer_2 (optional)
|
||||
generator_scheduler_2: Scheduler for generator_transformer_2 (optional)
|
||||
fake_score_scheduler_2: Scheduler for fake_score_transformer_2 (optional)
|
||||
generator_ema_2: EMA for generator_transformer_2 (optional)
|
||||
"""
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
@@ -254,6 +275,41 @@ def save_distillation_checkpoint(generator_transformer,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save generator_2 distributed checkpoint (MoE support)
|
||||
if generator_transformer_2 is not None:
|
||||
generator_2_states = {
|
||||
"model": ModelWrapper(generator_transformer_2),
|
||||
}
|
||||
if generator_optimizer_2 is not None:
|
||||
generator_2_states["optimizer"] = OptimizerWrapper(
|
||||
generator_transformer_2, generator_optimizer_2)
|
||||
if dataloader is not None:
|
||||
generator_2_states["dataloader"] = dataloader
|
||||
if generator_scheduler_2 is not None:
|
||||
generator_2_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler_2)
|
||||
if generator_ema_2 is not None:
|
||||
generator_2_states["ema"] = generator_ema_2.state_dict()
|
||||
|
||||
generator_2_dcp_dir = os.path.join(save_dir,
|
||||
"distributed_checkpoint",
|
||||
"generator_2")
|
||||
logger.info(
|
||||
"rank: %s, saving generator_2 distributed checkpoint to %s",
|
||||
rank,
|
||||
generator_2_dcp_dir,
|
||||
local_main_process_only=False)
|
||||
|
||||
begin_time = time.perf_counter()
|
||||
dcp.save(generator_2_states, checkpoint_id=generator_2_dcp_dir)
|
||||
end_time = time.perf_counter()
|
||||
|
||||
logger.info(
|
||||
"rank: %s, generator_2 distributed checkpoint saved in %.2f seconds",
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save critic distributed checkpoint
|
||||
critic_states = {
|
||||
"model": ModelWrapper(fake_score_transformer),
|
||||
@@ -283,6 +339,67 @@ def save_distillation_checkpoint(generator_transformer,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save critic_2 distributed checkpoint (MoE support)
|
||||
if fake_score_transformer_2 is not None:
|
||||
critic_2_states = {
|
||||
"model": ModelWrapper(fake_score_transformer_2),
|
||||
}
|
||||
if fake_score_optimizer_2 is not None:
|
||||
critic_2_states["optimizer"] = OptimizerWrapper(
|
||||
fake_score_transformer_2, fake_score_optimizer_2)
|
||||
if dataloader is not None:
|
||||
critic_2_states["dataloader"] = dataloader
|
||||
if fake_score_scheduler_2 is not None:
|
||||
critic_2_states["scheduler"] = SchedulerWrapper(
|
||||
fake_score_scheduler_2)
|
||||
|
||||
critic_2_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
|
||||
"critic_2")
|
||||
logger.info(
|
||||
"rank: %s, saving critic_2 distributed checkpoint to %s",
|
||||
rank,
|
||||
critic_2_dcp_dir,
|
||||
local_main_process_only=False)
|
||||
|
||||
begin_time = time.perf_counter()
|
||||
dcp.save(critic_2_states, checkpoint_id=critic_2_dcp_dir)
|
||||
end_time = time.perf_counter()
|
||||
|
||||
logger.info(
|
||||
"rank: %s, critic_2 distributed checkpoint saved in %.2f seconds",
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save real_score_transformer_2 distributed checkpoint (MoE support)
|
||||
if real_score_transformer_2 is not None:
|
||||
real_score_2_states = {
|
||||
"model": ModelWrapper(real_score_transformer_2),
|
||||
}
|
||||
# Note: real_score_transformer_2 typically doesn't have optimizer/scheduler
|
||||
# since it's used for inference only, but we include dataloader for consistency
|
||||
if dataloader is not None:
|
||||
real_score_2_states["dataloader"] = dataloader
|
||||
|
||||
real_score_2_dcp_dir = os.path.join(save_dir,
|
||||
"distributed_checkpoint",
|
||||
"real_score_2")
|
||||
logger.info(
|
||||
"rank: %s, saving real_score_2 distributed checkpoint to %s",
|
||||
rank,
|
||||
real_score_2_dcp_dir,
|
||||
local_main_process_only=False)
|
||||
|
||||
begin_time = time.perf_counter()
|
||||
dcp.save(real_score_2_states, checkpoint_id=real_score_2_dcp_dir)
|
||||
end_time = time.perf_counter()
|
||||
|
||||
logger.info(
|
||||
"rank: %s, real_score_2 distributed checkpoint saved in %.2f seconds",
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save shared random state separately
|
||||
shared_states = {
|
||||
"random_state": RandomStateWrapper(noise_generator),
|
||||
@@ -335,6 +452,47 @@ def save_distillation_checkpoint(generator_transformer,
|
||||
logger.info("--> distillation checkpoint saved at step %s to %s", step,
|
||||
weight_path)
|
||||
|
||||
# Save generator_2 model weights (consolidated) for inference (MoE support)
|
||||
if generator_transformer_2 is not None:
|
||||
inference_save_dir_2 = os.path.join(
|
||||
save_dir, "generator_2_inference_transformer")
|
||||
cpu_state_2 = gather_state_dict_on_cpu_rank0(
|
||||
generator_transformer_2, device=None)
|
||||
|
||||
if rank == 0:
|
||||
os.makedirs(inference_save_dir_2, exist_ok=True)
|
||||
weight_path_2 = os.path.join(
|
||||
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
|
||||
logger.info(
|
||||
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
|
||||
rank,
|
||||
weight_path_2,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Convert training format to diffusers format and save
|
||||
diffusers_state_dict_2 = custom_to_hf_state_dict(
|
||||
cpu_state_2,
|
||||
generator_transformer_2.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict_2, weight_path_2)
|
||||
|
||||
logger.info(
|
||||
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
|
||||
rank,
|
||||
weight_path_2,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save model config
|
||||
config_dict_2 = generator_transformer_2.hf_config
|
||||
if "dtype" in config_dict_2:
|
||||
del config_dict_2["dtype"] # TODO
|
||||
config_path_2 = os.path.join(inference_save_dir_2,
|
||||
"config.json")
|
||||
with open(config_path_2, "w") as f:
|
||||
json.dump(config_dict_2, f, indent=4)
|
||||
logger.info(
|
||||
"--> generator_2 distillation checkpoint saved at step %s to %s",
|
||||
step, weight_path_2)
|
||||
|
||||
|
||||
def load_checkpoint(transformer,
|
||||
rank,
|
||||
@@ -396,20 +554,43 @@ def load_checkpoint(transformer,
|
||||
return step
|
||||
|
||||
|
||||
def load_distillation_checkpoint(generator_transformer,
|
||||
fake_score_transformer,
|
||||
rank,
|
||||
checkpoint_path,
|
||||
generator_optimizer=None,
|
||||
fake_score_optimizer=None,
|
||||
dataloader=None,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None,
|
||||
generator_ema=None) -> int:
|
||||
def load_distillation_checkpoint(
|
||||
generator_transformer,
|
||||
fake_score_transformer,
|
||||
rank,
|
||||
checkpoint_path,
|
||||
generator_optimizer=None,
|
||||
fake_score_optimizer=None,
|
||||
dataloader=None,
|
||||
generator_scheduler=None,
|
||||
fake_score_scheduler=None,
|
||||
noise_generator=None,
|
||||
generator_ema=None,
|
||||
# MoE support
|
||||
generator_transformer_2=None,
|
||||
real_score_transformer_2=None,
|
||||
fake_score_transformer_2=None,
|
||||
generator_optimizer_2=None,
|
||||
fake_score_optimizer_2=None,
|
||||
generator_scheduler_2=None,
|
||||
fake_score_scheduler_2=None,
|
||||
generator_ema_2=None) -> int:
|
||||
"""
|
||||
Load distillation checkpoint with both generator and fake_score models.
|
||||
Supports MoE (Mixture of Experts) models with transformer_2 variants.
|
||||
Returns the step number from which training should resume.
|
||||
|
||||
Args:
|
||||
generator_transformer: Main generator transformer model
|
||||
fake_score_transformer: Main fake score transformer model
|
||||
generator_transformer_2: Secondary generator transformer for MoE (optional)
|
||||
real_score_transformer_2: Secondary real score transformer for MoE (optional)
|
||||
fake_score_transformer_2: Secondary fake score transformer for MoE (optional)
|
||||
generator_optimizer_2: Optimizer for generator_transformer_2 (optional)
|
||||
fake_score_optimizer_2: Optimizer for fake_score_transformer_2 (optional)
|
||||
generator_scheduler_2: Scheduler for generator_transformer_2 (optional)
|
||||
fake_score_scheduler_2: Scheduler for fake_score_transformer_2 (optional)
|
||||
generator_ema_2: EMA for generator_transformer_2 (optional)
|
||||
"""
|
||||
if not os.path.exists(checkpoint_path):
|
||||
logger.warning("Distillation checkpoint path %s does not exist",
|
||||
@@ -474,6 +655,63 @@ def load_distillation_checkpoint(generator_transformer,
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Load generator_2 distributed checkpoint (MoE support)
|
||||
if generator_transformer_2 is not None:
|
||||
generator_2_dcp_dir = os.path.join(checkpoint_path,
|
||||
"distributed_checkpoint",
|
||||
"generator_2")
|
||||
if os.path.exists(generator_2_dcp_dir):
|
||||
generator_2_states = {
|
||||
"model": ModelWrapper(generator_transformer_2),
|
||||
}
|
||||
|
||||
if generator_optimizer_2 is not None:
|
||||
generator_2_states["optimizer"] = OptimizerWrapper(
|
||||
generator_transformer_2, generator_optimizer_2)
|
||||
|
||||
if dataloader is not None:
|
||||
generator_2_states["dataloader"] = dataloader
|
||||
|
||||
if generator_scheduler_2 is not None:
|
||||
generator_2_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler_2)
|
||||
|
||||
logger.info(
|
||||
"rank: %s, loading generator_2 distributed checkpoint from %s",
|
||||
rank,
|
||||
generator_2_dcp_dir,
|
||||
local_main_process_only=False)
|
||||
|
||||
begin_time = time.perf_counter()
|
||||
dcp.load(generator_2_states, checkpoint_id=generator_2_dcp_dir)
|
||||
end_time = time.perf_counter()
|
||||
|
||||
logger.info(
|
||||
"rank: %s, generator_2 distributed checkpoint loaded in %.2f seconds",
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA_2 state if available and generator_ema_2 is provided
|
||||
if generator_ema_2 is not None:
|
||||
try:
|
||||
ema_2_state = generator_2_states.get("ema")
|
||||
if ema_2_state is not None:
|
||||
generator_ema_2.load_state_dict(ema_2_state)
|
||||
logger.info(
|
||||
"rank: %s, generator_2 EMA state loaded successfully",
|
||||
rank)
|
||||
else:
|
||||
logger.info(
|
||||
"rank: %s, no EMA_2 state found in checkpoint",
|
||||
rank)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA_2 state: %s",
|
||||
rank, str(e))
|
||||
else:
|
||||
logger.info("rank: %s, generator_2 checkpoint not found, skipping",
|
||||
rank)
|
||||
|
||||
# Load critic distributed checkpoint
|
||||
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
|
||||
"critic")
|
||||
@@ -512,6 +750,77 @@ def load_distillation_checkpoint(generator_transformer,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load critic_2 distributed checkpoint (MoE support)
|
||||
if fake_score_transformer_2 is not None:
|
||||
critic_2_dcp_dir = os.path.join(checkpoint_path,
|
||||
"distributed_checkpoint", "critic_2")
|
||||
if os.path.exists(critic_2_dcp_dir):
|
||||
critic_2_states = {
|
||||
"model": ModelWrapper(fake_score_transformer_2),
|
||||
}
|
||||
|
||||
if fake_score_optimizer_2 is not None:
|
||||
critic_2_states["optimizer"] = OptimizerWrapper(
|
||||
fake_score_transformer_2, fake_score_optimizer_2)
|
||||
|
||||
if dataloader is not None:
|
||||
critic_2_states["dataloader"] = dataloader
|
||||
|
||||
if fake_score_scheduler_2 is not None:
|
||||
critic_2_states["scheduler"] = SchedulerWrapper(
|
||||
fake_score_scheduler_2)
|
||||
|
||||
logger.info(
|
||||
"rank: %s, loading critic_2 distributed checkpoint from %s",
|
||||
rank,
|
||||
critic_2_dcp_dir,
|
||||
local_main_process_only=False)
|
||||
|
||||
begin_time = time.perf_counter()
|
||||
dcp.load(critic_2_states, checkpoint_id=critic_2_dcp_dir)
|
||||
end_time = time.perf_counter()
|
||||
|
||||
logger.info(
|
||||
"rank: %s, critic_2 distributed checkpoint loaded in %.2f seconds",
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
else:
|
||||
logger.info("rank: %s, critic_2 checkpoint not found, skipping",
|
||||
rank)
|
||||
|
||||
# Load real_score_2 distributed checkpoint (MoE support)
|
||||
if real_score_transformer_2 is not None:
|
||||
real_score_2_dcp_dir = os.path.join(checkpoint_path,
|
||||
"distributed_checkpoint",
|
||||
"real_score_2")
|
||||
if os.path.exists(real_score_2_dcp_dir):
|
||||
real_score_2_states = {
|
||||
"model": ModelWrapper(real_score_transformer_2),
|
||||
}
|
||||
|
||||
if dataloader is not None:
|
||||
real_score_2_states["dataloader"] = dataloader
|
||||
|
||||
logger.info(
|
||||
"rank: %s, loading real_score_2 distributed checkpoint from %s",
|
||||
rank,
|
||||
real_score_2_dcp_dir,
|
||||
local_main_process_only=False)
|
||||
|
||||
begin_time = time.perf_counter()
|
||||
dcp.load(real_score_2_states, checkpoint_id=real_score_2_dcp_dir)
|
||||
end_time = time.perf_counter()
|
||||
|
||||
logger.info(
|
||||
"rank: %s, real_score_2 distributed checkpoint loaded in %.2f seconds",
|
||||
rank,
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
else:
|
||||
logger.info("rank: %s, real_score_2 checkpoint not found, skipping",
|
||||
rank)
|
||||
|
||||
# Load shared random state
|
||||
shared_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
|
||||
"shared")
|
||||
@@ -1298,8 +1607,14 @@ def get_scheduler(
|
||||
last_epoch=last_epoch)
|
||||
|
||||
|
||||
def _local_numel(p: torch.Tensor) -> int:
|
||||
if hasattr(p, "to_local"):
|
||||
return p.to_local().numel()
|
||||
return p.numel()
|
||||
|
||||
|
||||
def count_trainable(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
return sum(_local_numel(p) for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class EMA_FSDP:
|
||||
|
||||
@@ -20,10 +20,7 @@ class WanDistillationPipeline(DistillationPipeline):
|
||||
A distillation pipeline for Wan that uses a single transformer model.
|
||||
The main transformer serves as the student model, and copies are made for teacher and critic.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
]
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize Wan-specific scheduler."""
|
||||
|
||||
@@ -29,10 +29,7 @@ class WanI2VDistillationPipeline(DistillationPipeline):
|
||||
A distillation pipeline for Wan that uses a single transformer model.
|
||||
The main transformer serves as the student model, and copies are made for teacher and critic.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
]
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize Wan-specific scheduler."""
|
||||
|
||||
@@ -21,8 +21,9 @@ class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
|
||||
with DMD for video generation.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
"scheduler",
|
||||
"transformer",
|
||||
"vae",
|
||||
]
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
@@ -40,7 +41,10 @@ class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
|
||||
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,
|
||||
|
||||
@@ -12,6 +12,8 @@ export TOKENIZERS_PARALLELISM=false
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--real_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--fake_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
|
||||
@@ -13,6 +13,8 @@ export TOKENIZERS_PARALLELISM=false
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--real_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--fake_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
|
||||
Reference in New Issue
Block a user