Compare commits
27
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
083c2258ea | ||
|
|
60260a6be9 | ||
|
|
06cc80ee93 | ||
|
|
c6b3475f77 | ||
|
|
1d34380c58 | ||
|
|
ddf8198e65 | ||
|
|
c38ee88294 | ||
|
|
ff1b77bdd6 | ||
|
|
22ba3759ee | ||
|
|
2c3a3893c0 | ||
|
|
6e9ba7ba9d | ||
|
|
64f1295878 | ||
|
|
d9cc3b7e42 | ||
|
|
cb3230f5e9 | ||
|
|
1e91cec93d | ||
|
|
5f0688084c | ||
|
|
f3023db253 | ||
|
|
c405941f21 | ||
|
|
f4e818305b | ||
|
|
4c4ab65020 | ||
|
|
6b8bf495ee | ||
|
|
ebd23ba8e7 | ||
|
|
0c1cb0c19d | ||
|
|
d87ebe051b | ||
|
|
0ae669cd4c | ||
|
|
63c124fec8 | ||
|
|
ececbe1839 |
@@ -0,0 +1,16 @@
|
||||
# Wan2.1-I2V-1.3B-InP Crush-Smol Example
|
||||
These are e2e example scripts for finetuning Wan2.1 T2V 1.3B InP on the crush-smol dataset.
|
||||
|
||||
## Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
|
||||
@@ -0,0 +1,111 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp_single/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_distill"
|
||||
--output_dir "checkpoints/wan_i2v_distill"
|
||||
--wandb_run_name "wan_i2v_distill"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
# --log_visualization
|
||||
--log_visualization
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "constant"
|
||||
--fake_score_learning_rate 8e-7
|
||||
--fake_score_lr_scheduler "constant"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 2000
|
||||
--weight_only_checkpointing_steps 2000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--simulate_generator_forward
|
||||
--real_score_guidance_scale 5
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_i2v_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,111 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
export HF_HUB_ENABLE_HF_TRANSFER=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_distill"
|
||||
--output_dir "checkpoints/wan_i2v_distill"
|
||||
--wandb_run_name "wan_i2v_distill"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
# --log_visualization
|
||||
--log_visualization
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "constant"
|
||||
--fake_score_learning_rate 8e-7
|
||||
--fake_score_lr_scheduler "constant"
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--simulate_generator_forward
|
||||
--real_score_guidance_scale 3.5
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,129 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=i2v
|
||||
#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=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_i2v_finetune
|
||||
--output_dir "checkpoints/wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 2
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
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_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v_1_3b_inp/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "i2v"
|
||||
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
# Wan2.1-I2V-1.3B-InP Crush-Smol Example
|
||||
These are e2e example scripts for finetuning Wan2.1 T2V 1.3B InP on the crush-smol dataset.
|
||||
|
||||
## Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
|
||||
@@ -0,0 +1,122 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
# MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_wan21_i2v_14b/combined_parquet_dataset/"
|
||||
# DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-I2V/crush_smol/validation_orig.json"
|
||||
NUM_GPUS=8
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_distill"
|
||||
--output_dir "checkpoints/wan_i2v_distill"
|
||||
--wandb_run_name "14b_i2v_cm"
|
||||
--max_train_steps 1500
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
# --log_visualization
|
||||
--log_visualization
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "constant"
|
||||
# --min_lr_ratio 0.5
|
||||
# --lr_warmup_steps 50
|
||||
--fake_score_learning_rate 7e-7
|
||||
--fake_score_lr_scheduler "constant"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 2000
|
||||
--weight_only_checkpointing_steps 2000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--warp_denoising_step True
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.96
|
||||
--generator_update_interval 5
|
||||
--simulate_generator_forward
|
||||
--simulate_forward_interval 1
|
||||
--real_score_guidance_scale 5
|
||||
--VSA_sparsity 0.8
|
||||
--regression_loss_weight 0.01
|
||||
--use_regression_loss False
|
||||
--cm_loss_weight 1
|
||||
--ema_decay 0.999
|
||||
--cm_use_ema_teacher
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_i2v_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,111 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_wan21_i2v_14b/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-I2V/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
export HF_HUB_ENABLE_HF_TRANSFER=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_distill"
|
||||
--output_dir "checkpoints/wan_i2v_distill"
|
||||
--wandb_run_name "wan21_i2v_14b_ode_50steps"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 8
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
# --log_visualization
|
||||
--log_visualization
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "constant"
|
||||
--fake_score_learning_rate 8e-7
|
||||
--fake_score_lr_scheduler "constant"
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--simulate_generator_forward
|
||||
--real_score_guidance_scale 3.5
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,129 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=i2v
|
||||
#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=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_i2v_finetune
|
||||
--output_dir "checkpoints/wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 2
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
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_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
# DATA_MERGE_PATH="data/single/merge.txt"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_wan21_t2v_14b/"
|
||||
# OUTPUT_DIR="data/single_processed_i2v_14b/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8\
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
# MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_wan21_t2v_14b/combined_parquet_dataset/"
|
||||
# DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-I2V/crush_smol/validation_orig.json"
|
||||
NUM_GPUS=8
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_distill"
|
||||
--output_dir "checkpoints/wan_i2v_distill"
|
||||
--wandb_run_name "cm_t2v"
|
||||
--max_train_steps 1500
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
# --log_visualization
|
||||
--log_visualization
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 4e-6
|
||||
--lr_scheduler "constant"
|
||||
# --min_lr_ratio 0.5
|
||||
# --lr_warmup_steps 50
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 2000
|
||||
--weight_only_checkpointing_steps 2000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--warp_denoising_step True
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.96
|
||||
--real_score_guidance_scale 5
|
||||
--VSA_sparsity 0.8
|
||||
--regression_loss_weight 0.01
|
||||
--use_regression_loss False
|
||||
--cm_loss_weight 1
|
||||
--ema_decay 0.999
|
||||
--cm_weighing_function "constant"
|
||||
# --cm_use_ema_teacher
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_cm_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -9,17 +9,19 @@ MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
|
||||
--output_dir "outputs/wan_i2v_finetune"
|
||||
--wandb_run_name "wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 4
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
@@ -30,10 +32,10 @@ training_args=(
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 4
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
@@ -59,7 +61,7 @@ validation_args=(
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
|
||||
@@ -34,6 +34,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ from fastvideo.configs.sample.wan import (
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -38,6 +39,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -107,13 +107,28 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
fps: int = 16
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.2 TI2V Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_2_Base_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
negative_prompt: str = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
+56
-10
@@ -161,6 +161,9 @@ class FastVideoArgs:
|
||||
# MoE parameters used by Wan2.2
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# scheduler parameters
|
||||
warp_denoising_step: bool = False
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
@@ -382,6 +385,15 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
|
||||
# scheduler parameters
|
||||
parser.add_argument(
|
||||
"--warp-denoising-step",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.warp_denoising_step,
|
||||
help="Warp denoising step for scheduler",
|
||||
)
|
||||
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -592,8 +604,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
dit_model_name_or_path: str = ""
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
training_cfg_rate: float = 0.0
|
||||
precondition_outputs: bool = False
|
||||
|
||||
@@ -674,6 +684,16 @@ class TrainingArgs(FastVideoArgs):
|
||||
log_visualization: bool = False
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
simulate_forward_interval: int = 1
|
||||
regression_loss_weight: float = 0.0
|
||||
use_regression_loss: bool = False
|
||||
|
||||
# consistency model parameters
|
||||
cm_loss_weight: float = 0.0
|
||||
cm_use_ema_teacher: bool = False
|
||||
ema_decay: float = 0.999
|
||||
ema_start_step: int = 0
|
||||
cm_weighing_function: str = "constant"
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -776,14 +796,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Directory to cache models")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
default=0.999,
|
||||
help="EMA decay rate")
|
||||
parser.add_argument("--ema-start-step",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Step to start EMA")
|
||||
parser.add_argument("--training-cfg-rate",
|
||||
type=float,
|
||||
help="Classifier-free guidance scale")
|
||||
@@ -1018,6 +1030,40 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--simulate-generator-forward",
|
||||
action=StoreBoolean,
|
||||
help="Whether to simulate generator forward to match inference")
|
||||
parser.add_argument("--simulate-forward-interval",
|
||||
type=int,
|
||||
default=TrainingArgs.simulate_forward_interval,
|
||||
help="Ratio of steps to simulate generator forward")
|
||||
parser.add_argument("--regression-loss-weight",
|
||||
type=float,
|
||||
default=TrainingArgs.regression_loss_weight,
|
||||
help="Weight for regression loss")
|
||||
parser.add_argument("--use-regression-loss",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use regression loss")
|
||||
|
||||
# Consistency model arguments
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
default=0.999,
|
||||
help="EMA decay rate")
|
||||
parser.add_argument("--ema-start-step",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Step to start EMA")
|
||||
parser.add_argument("--cm-loss-weight",
|
||||
type=float,
|
||||
default=TrainingArgs.cm_loss_weight,
|
||||
help="Weight for consistency loss")
|
||||
parser.add_argument(
|
||||
"--cm-use-ema-teacher",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA teacher for consistency loss")
|
||||
parser.add_argument("--cm-weighing-function",
|
||||
type=str,
|
||||
choices=["constant", "sigma_sqrt"],
|
||||
default=TrainingArgs.cm_weighing_function,
|
||||
help="Weighting function for consistency loss")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@@ -12,12 +12,10 @@ from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from fastvideo.pipelines.stages import (ImageEncodingStage, ConditioningStage,
|
||||
DecodingStage, DmdDenoisingStage,
|
||||
EncodingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.pipelines.stages import (
|
||||
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
|
||||
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
|
||||
TextEncodingStage, TimestepPreparationStage)
|
||||
# isort: on
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
@@ -67,7 +65,7 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DmdDenoisingStage(
|
||||
|
||||
@@ -118,6 +118,8 @@ class ForwardBatch:
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
return_trajectory_latents: bool = False
|
||||
trajectory_latents: list[torch.Tensor] = field(default_factory=list)
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -193,6 +195,8 @@ class TrainingBatch:
|
||||
# Distillation losses
|
||||
generator_loss: float = 0.0
|
||||
fake_score_loss: float = 0.0
|
||||
regression_loss: float = 0.0
|
||||
consistency_loss: float = 0.0
|
||||
|
||||
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@@ -197,6 +197,7 @@ class DenoisingStage(PipelineStage):
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
logger.info('timesteps: %s', timesteps)
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
@@ -380,7 +381,8 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
def progress_bar(self,
|
||||
iterable: Iterable | None = None,
|
||||
total: int | None = None) -> tqdm:
|
||||
total: int | None = None,
|
||||
disable: bool = False) -> tqdm:
|
||||
"""
|
||||
Create a progress bar for the denoising process.
|
||||
|
||||
@@ -393,7 +395,7 @@ class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
local_rank = get_world_group().local_rank
|
||||
if local_rank == 0:
|
||||
return tqdm(iterable=iterable, total=total)
|
||||
return tqdm(iterable=iterable, total=total, disable=disable)
|
||||
else:
|
||||
return tqdm(iterable=iterable, total=total, disable=True)
|
||||
|
||||
@@ -676,12 +678,22 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
# TODO(yongqi) hard code prepare latents
|
||||
latents = torch.randn(
|
||||
latents.permute(0, 2, 1, 3, 4).shape,
|
||||
# latents.shape,
|
||||
dtype=torch.bfloat16,
|
||||
device="cuda",
|
||||
generator=torch.Generator(device="cuda").manual_seed(42))
|
||||
video_raw_latent_shape = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
timesteps_50 = [
|
||||
999, 993, 986, 979, 971, 964, 956, 948, 940, 931, 922, 913, 904,
|
||||
895, 885, 874, 864, 853, 841, 830, 818, 805, 792, 778, 764, 749,
|
||||
734, 718, 702, 684, 666, 647, 627, 607, 585, 562, 538, 513, 486,
|
||||
458, 428, 396, 363, 328, 290, 249, 206, 160, 111, 57
|
||||
]
|
||||
# timesteps = torch.tensor(timesteps_50,
|
||||
# dtype=torch.long,
|
||||
# device=get_local_torch_device())
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
@@ -705,7 +717,9 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
batch.image_latent = image_latent
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=len(timesteps)) as progress_bar:
|
||||
with self.progress_bar(
|
||||
total=len(timesteps),
|
||||
disable=batch.return_trajectory_latents) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
@@ -805,6 +819,10 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
latents = self.scheduler.add_noise(
|
||||
pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
next_timestep).unflatten(0, pred_video.shape[:2])
|
||||
|
||||
if batch.return_trajectory_latents:
|
||||
batch.trajectory_latents.append(latents.clone())
|
||||
|
||||
else:
|
||||
latents = pred_video
|
||||
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
import types
|
||||
|
||||
import pytest # type: ignore
|
||||
import torch # type: ignore
|
||||
import torch.nn as nn # type: ignore
|
||||
|
||||
from fastvideo.training.training_utils import EMA_FSDP
|
||||
|
||||
|
||||
class TinyNet(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.lin = nn.Linear(4, 3, bias=True)
|
||||
# Non-trainable param should be ignored by EMA
|
||||
self.register_parameter("frozen", nn.Parameter(torch.randn(2), requires_grad=False))
|
||||
# Buffer should be ignored by EMA
|
||||
self.register_buffer("buf", torch.ones(1))
|
||||
|
||||
|
||||
def named_params_to_cpu_float_dict(module: nn.Module):
|
||||
return {n: p.detach().clone().float().cpu() for n, p in module.named_parameters() if p.requires_grad}
|
||||
|
||||
|
||||
def test_ema_init_local_shard_copies_params_cpu_float():
|
||||
torch.manual_seed(0)
|
||||
model = TinyNet()
|
||||
ema = EMA_FSDP(model, decay=0.9, mode="local_shard")
|
||||
|
||||
expected = named_params_to_cpu_float_dict(model)
|
||||
assert set(ema.shadow.keys()) == set(expected.keys())
|
||||
for n, v in expected.items():
|
||||
assert torch.equal(ema.shadow[n], v)
|
||||
|
||||
|
||||
def test_ema_update_local_shard_matches_formula():
|
||||
torch.manual_seed(0)
|
||||
model = TinyNet()
|
||||
decay = 0.9
|
||||
ema = EMA_FSDP(model, decay=decay, mode="local_shard")
|
||||
|
||||
# Save initial snapshot
|
||||
start = {n: p.detach().clone().float().cpu() for n, p in model.named_parameters() if p.requires_grad}
|
||||
|
||||
# Mutate model parameters to new values
|
||||
with torch.no_grad():
|
||||
for _, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.add_(torch.randn_like(p))
|
||||
|
||||
after = {n: p.detach().clone().float().cpu() for n, p in model.named_parameters() if p.requires_grad}
|
||||
|
||||
ema.update(model)
|
||||
|
||||
for n in start.keys():
|
||||
expected = start[n] * decay + after[n] * (1.0 - decay)
|
||||
assert torch.allclose(ema.shadow[n], expected, rtol=1e-6, atol=1e-8)
|
||||
|
||||
|
||||
def test_apply_to_model_swaps_and_restores():
|
||||
torch.manual_seed(0)
|
||||
model = TinyNet()
|
||||
ema = EMA_FSDP(model, decay=0.0, mode="local_shard") # decay 0 -> shadow becomes last param on update
|
||||
|
||||
# Change model and update EMA so shadow != current
|
||||
with torch.no_grad():
|
||||
for _, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.add_(1.2345)
|
||||
ema.update(model)
|
||||
|
||||
# Change model again so current != shadow
|
||||
with torch.no_grad():
|
||||
for _, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.add_(2.0)
|
||||
|
||||
# Snapshot of values at context entry
|
||||
entry_values = {n: p.detach().clone() for n, p in model.named_parameters() if p.requires_grad}
|
||||
|
||||
# Inside context, params should equal EMA shadow; outside, restored
|
||||
with ema.apply_to_model(model):
|
||||
for n, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
assert torch.allclose(p.detach().cpu().float(), ema.shadow[n])
|
||||
|
||||
for n, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
assert torch.equal(p.detach(), entry_values[n])
|
||||
|
||||
|
||||
def test_state_dict_roundtrip():
|
||||
torch.manual_seed(0)
|
||||
model = TinyNet()
|
||||
ema = EMA_FSDP(model, decay=0.8, mode="local_shard")
|
||||
|
||||
with torch.no_grad():
|
||||
for _, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.add_(torch.randn_like(p))
|
||||
ema.update(model)
|
||||
|
||||
sd = ema.state_dict()
|
||||
ema2 = EMA_FSDP(model, decay=0.8, mode="local_shard")
|
||||
ema2.load_state_dict(sd)
|
||||
|
||||
for k in sd.keys():
|
||||
assert torch.equal(ema2.shadow[k], sd[k])
|
||||
|
||||
|
||||
def test_copy_to_unwrapped_and_rank0_full_guard():
|
||||
torch.manual_seed(0)
|
||||
src = TinyNet()
|
||||
ema = EMA_FSDP(src, decay=0.7, mode="local_shard")
|
||||
|
||||
# Make EMA shadow distinct from a fresh target
|
||||
with torch.no_grad():
|
||||
for _, p in src.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.mul_(3.14)
|
||||
ema.update(src)
|
||||
|
||||
tgt = TinyNet()
|
||||
# copy in local_shard mode always applies
|
||||
ema.copy_to_unwrapped(tgt)
|
||||
for n, p in tgt.named_parameters():
|
||||
if p.requires_grad:
|
||||
assert torch.allclose(p.detach().cpu().float(), ema.shadow[n])
|
||||
|
||||
# If mode is rank0_full but rank != 0, copy is a no-op
|
||||
ema.mode = "rank0_full"
|
||||
ema.rank = 1
|
||||
before = {n: p.detach().clone() for n, p in tgt.named_parameters()}
|
||||
ema.copy_to_unwrapped(tgt)
|
||||
for n, p in tgt.named_parameters():
|
||||
assert torch.equal(p.detach(), before[n])
|
||||
|
||||
|
||||
def test_rank0_full_init_and_update_with_stubbed_gather(monkeypatch):
|
||||
torch.manual_seed(0)
|
||||
model = TinyNet()
|
||||
|
||||
# Stub gather_state_dict_on_cpu_rank0 to avoid requiring initialized dist
|
||||
def fake_gather(mod, device=None):
|
||||
return {n: p.detach().clone() for n, p in mod.named_parameters() if p.requires_grad}
|
||||
|
||||
# Patch in the module where EMA_FSDP is defined
|
||||
import fastvideo.training.training_utils as tu
|
||||
monkeypatch.setattr(tu, "gather_state_dict_on_cpu_rank0", fake_gather, raising=True)
|
||||
|
||||
ema = EMA_FSDP(model, decay=0.5, mode="rank0_full")
|
||||
assert len(ema.shadow) > 0
|
||||
|
||||
start = {n: t.clone().float().cpu() for n, t in ema.shadow.items()}
|
||||
|
||||
with torch.no_grad():
|
||||
for _, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.add_(1.0)
|
||||
|
||||
ema.update(model)
|
||||
|
||||
cpu_state = {n: p.detach().clone().float().cpu() for n, p in model.named_parameters() if p.requires_grad}
|
||||
for n in start.keys():
|
||||
expected = start[n] * 0.5 + cpu_state[n] * 0.5
|
||||
assert torch.allclose(ema.shadow[n], expected)
|
||||
|
||||
# state_dict should be empty if rank != 0 in rank0_full mode
|
||||
ema.rank = 1
|
||||
assert ema.state_dict() == {}
|
||||
|
||||
|
||||
def test_apply_to_model_raises_when_rank0_full():
|
||||
model = TinyNet()
|
||||
ema = EMA_FSDP(model, decay=0.9, mode="rank0_full")
|
||||
with pytest.raises(RuntimeError):
|
||||
with ema.apply_to_model(model):
|
||||
pass
|
||||
|
||||
|
||||
def test_copy_to_unwrapped_rank0_full_rank0(monkeypatch):
|
||||
torch.manual_seed(0)
|
||||
model = TinyNet()
|
||||
|
||||
# Stub gather to produce keys matching named_parameters
|
||||
def fake_gather(mod, device=None):
|
||||
return {n: p.detach().clone() for n, p in mod.named_parameters() if p.requires_grad}
|
||||
|
||||
import fastvideo.training.training_utils as tu
|
||||
monkeypatch.setattr(tu, "gather_state_dict_on_cpu_rank0", fake_gather, raising=True)
|
||||
|
||||
ema = EMA_FSDP(model, decay=0.0, mode="rank0_full")
|
||||
ema.rank = 0
|
||||
|
||||
# Modify model and update so EMA has distinct weights
|
||||
with torch.no_grad():
|
||||
for _, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
p.add_(1.0)
|
||||
ema.update(model)
|
||||
|
||||
tgt = TinyNet()
|
||||
ema.copy_to_unwrapped(tgt)
|
||||
for n, p in tgt.named_parameters():
|
||||
if p.requires_grad:
|
||||
assert torch.allclose(p.detach().cpu().float(), ema.shadow[n])
|
||||
|
||||
|
||||
@@ -0,0 +1,846 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import copy
|
||||
import gc
|
||||
import os
|
||||
import time
|
||||
from abc import abstractmethod
|
||||
from collections import deque
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from torchmetrics.image.lpip import (
|
||||
LearnedPerceptualImagePatchSimilarity as LPIPSimilarity)
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.distributed import (cleanup_dist_env_and_memory,
|
||||
get_local_torch_device, get_sp_group,
|
||||
get_world_group)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
get_scheduler, load_checkpoint, pred_noise_to_pred_video, save_checkpoint,
|
||||
shift_timestep)
|
||||
from fastvideo.utils import is_vsa_available, set_random_seed
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CMDistillationPipeline(TrainingPipeline):
|
||||
"""
|
||||
A distillation pipeline for training a 3 step model.
|
||||
Inherits from TrainingPipeline to reuse training infrastructure.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae"
|
||||
]
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[dict[str, Any]]
|
||||
current_epoch: int = 0
|
||||
init_steps: int
|
||||
current_trainstep: int
|
||||
num_generator_updates: int = 0
|
||||
video_latent_shape: tuple[int, ...]
|
||||
video_latent_shape_sp: tuple[int, ...]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize the distillation training pipeline with multiple models."""
|
||||
logger.info("Initializing distillation pipeline...")
|
||||
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=self.timestep_shift)
|
||||
|
||||
if self.training_args.warp_denoising_step:
|
||||
# timesteps = self.noise_scheduler.timesteps.cpu()
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0], dtype=torch.float32)))
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=torch.device("cpu"))
|
||||
self.denoising_step_list = timesteps[1000 -
|
||||
self.denoising_step_list].to(
|
||||
get_local_torch_device())
|
||||
logger.info(
|
||||
"Warp denoising step is enabled, using %s denoising steps",
|
||||
self.denoising_step_list)
|
||||
else:
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
logger.info(
|
||||
"Warp denoising step is disabled, using %s denoising steps",
|
||||
self.denoising_step_list)
|
||||
logger.info("Distillation generator model to %s denoising steps",
|
||||
len(self.denoising_step_list))
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
|
||||
self.min_timestep = int(self.training_args.min_timestep_ratio *
|
||||
self.num_train_timestep)
|
||||
self.max_timestep = int(self.training_args.max_timestep_ratio *
|
||||
self.num_train_timestep)
|
||||
|
||||
# self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
|
||||
# Initialize EMA teacher for CM if enabled
|
||||
self.use_cm_with_ema = getattr(self.training_args, "cm_use_ema_teacher",
|
||||
False)
|
||||
if self.use_cm_with_ema:
|
||||
# Use local_shard mode for teacher forward compatibility with FSDP2
|
||||
self.ema_teacher = EMA_FSDP(self.transformer,
|
||||
decay=self.training_args.ema_decay,
|
||||
mode="local_shard")
|
||||
else:
|
||||
self.ema_teacher = None
|
||||
self.lpips = LPIPSimilarity().to(self.device)
|
||||
|
||||
def _compute_lpips_loss(self, pred_video: torch.Tensor,
|
||||
latents: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Helper method to compute LPIPS loss between predicted video and ground truth latents.
|
||||
This method handles VAE scaling, shifting, decoding, and frame processing consistently.
|
||||
"""
|
||||
print("Computing LPIPS loss...")
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16), torch.no_grad():
|
||||
# Apply VAE scaling factor and shift factor before decoding (same as visualize_intermediate_latents)
|
||||
pred_video = pred_video.permute(0, 2, 1, 3, 4)
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
pred_video_scaled = pred_video / self.vae.scaling_factor.to(
|
||||
pred_video.device, pred_video.dtype)
|
||||
latents_scaled = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
pred_video_scaled = pred_video / self.vae.scaling_factor
|
||||
latents_scaled = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
pred_video_scaled += self.vae.shift_factor.to(
|
||||
pred_video.device, pred_video.dtype)
|
||||
latents_scaled += self.vae.shift_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
pred_video_scaled += self.vae.shift_factor
|
||||
latents_scaled += self.vae.shift_factor
|
||||
|
||||
# Permute from [batch, channels, frames, height, width] to [batch, frames, channels, height, width] for VAE decode
|
||||
# pred_video_scaled = pred_video_scaled.permute(0, 2, 1, 3, 4)
|
||||
# latents_scaled = latents_scaled.permute(0, 2, 1, 3, 4)
|
||||
|
||||
# print(f"pred_video_scaled shape after permute: {pred_video_scaled.shape}")
|
||||
# print(f"latents_scaled shape after permute: {latents_scaled.shape}")
|
||||
pred_video_frames = self.vae.decode(pred_video_scaled)
|
||||
latents_frames = self.vae.decode(latents_scaled)
|
||||
|
||||
# VAE output is already in [-1, 1] range, keep it for LPIPS
|
||||
# Just permute back to [B, T, C, H, W] format for frame processing
|
||||
pred_video_frames = pred_video_frames.permute(0, 2, 1, 3, 4)
|
||||
latents_frames = latents_frames.permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_video_frames = rearrange(pred_video_frames,
|
||||
"b n c h w -> (b n) c h w")
|
||||
latents_frames = rearrange(latents_frames, "b n c h w -> (b n) c h w")
|
||||
|
||||
return self.lpips(pred_video_frames, latents_frames)
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize validation pipeline - must be implemented by subclasses."""
|
||||
raise NotImplementedError(
|
||||
"Distillation pipelines must implement this method")
|
||||
|
||||
def _prepare_distillation(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Prepare training environment for distillation."""
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_distill_input_kwargs(
|
||||
self, noise_input: torch.Tensor, timestep: torch.Tensor,
|
||||
text_dict: dict[str, torch.Tensor] | None,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
if text_dict is None:
|
||||
raise ValueError(
|
||||
"text_dict cannot be None for distillation pipeline")
|
||||
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"timestep": timestep,
|
||||
"return_dict": False,
|
||||
}
|
||||
|
||||
return training_batch
|
||||
|
||||
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass through student transformer for a single randomly sampled
|
||||
denoising step; returns predicted clean video latents and records
|
||||
auxiliary info/losses on training_batch.
|
||||
"""
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
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
|
||||
|
||||
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).unflatten(
|
||||
0, (1, latents.shape[1]))
|
||||
|
||||
training_batch = self._build_distill_input_kwargs(noisy_latent, timestep,
|
||||
training_batch.conditional_dict,
|
||||
training_batch)
|
||||
pred_noise = self.transformer(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler,
|
||||
).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
# Compute regression loss (LPIPS) for optional auxiliary loss/logging
|
||||
# regression_loss = self._compute_lpips_loss(pred_video, latents)
|
||||
# training_batch.regression_loss = regression_loss
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"generator_pred_video": pred_video.detach().clone(),
|
||||
})
|
||||
return pred_video
|
||||
|
||||
def _calculate_consistency_loss(
|
||||
self, pred_video: torch.Tensor,
|
||||
training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""
|
||||
Consistency loss (CM-like): two variants
|
||||
- Without EMA: teacher = model at higher noise with stop-grad
|
||||
- With EMA: teacher = EMA(model) at lower noise or higher noise depending on schedule (use higher noise as teacher)
|
||||
"""
|
||||
# Need at least two steps to form a pair
|
||||
if len(self.denoising_step_list) < 2:
|
||||
return torch.tensor(0.0, device=self.device, dtype=pred_video.dtype)
|
||||
|
||||
# Choose a random index k in [1, N-1] so that (k-1, k) is valid
|
||||
idx_high = torch.randint(1,
|
||||
len(self.denoising_step_list), [1],
|
||||
device=self.device).item()
|
||||
t_high = self.denoising_step_list[idx_high]
|
||||
t_low = self.denoising_step_list[idx_high - 1]
|
||||
t_low = t_low * torch.ones(1, device=self.device, dtype=torch.long)
|
||||
t_high = t_high * torch.ones(1, device=self.device, dtype=torch.long)
|
||||
logger.info("t_high: %s, t_low: %s", t_high, t_low)
|
||||
|
||||
# Use student output as base; construct noisy sample at higher timestep
|
||||
base_student = pred_video
|
||||
base_noise = torch.randn(self.video_latent_shape, device=self.device, dtype=base_student.dtype)
|
||||
if self.sp_world_size > 1:
|
||||
base_noise = rearrange(base_noise, "b (n t) c h w -> b n t c h w", n=self.sp_world_size).contiguous()
|
||||
base_noise = base_noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
noisy_high = self.noise_scheduler.add_noise(base_student.flatten(0, 1),
|
||||
base_noise.flatten(0, 1),
|
||||
t_high).unflatten(0, (1, base_student.shape[1]))
|
||||
|
||||
# Compute ODE-adjacent lower-t sample via a single Euler step using teacher output at t_high
|
||||
# Temporarily switch to eval and no-grad for teacher forward(s)
|
||||
was_training = self.transformer.training
|
||||
self.transformer.eval()
|
||||
with torch.no_grad():
|
||||
tb_high_teacher = self._build_distill_input_kwargs(noisy_high, t_high,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
pred_flow_high_teacher = self.transformer(**tb_high_teacher.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Prepare flat tensors for scheduler stepping
|
||||
sample_flat = noisy_high.flatten(0, 1).to(dtype=torch.float32)
|
||||
flow_flat = pred_flow_high_teacher.flatten(0, 1).to(dtype=torch.float32)
|
||||
|
||||
# Reset step index so stepping starts at the provided timestep
|
||||
prev_step_index = getattr(self.noise_scheduler, "_step_index", None)
|
||||
self.noise_scheduler._step_index = None
|
||||
stepped = self.noise_scheduler.step(model_output=flow_flat,
|
||||
timestep=t_high,
|
||||
sample=sample_flat,
|
||||
return_dict=True)
|
||||
# Restore step index state
|
||||
self.noise_scheduler._step_index = prev_step_index
|
||||
|
||||
noisy_low_ode = stepped.prev_sample.unflatten(0, (1, noisy_high.shape[1])).to(noisy_high.dtype)
|
||||
|
||||
# Teacher target at lower t (clean latent)
|
||||
tb_low_teacher = self._build_distill_input_kwargs(noisy_low_ode, t_low,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
pred_noise_low_teacher = self.transformer(**tb_low_teacher.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
y_low_teacher = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_low_teacher.flatten(0, 1),
|
||||
noise_input_latent=noisy_low_ode.flatten(0, 1),
|
||||
timestep=t_low,
|
||||
scheduler=self.noise_scheduler,
|
||||
).unflatten(0, pred_noise_low_teacher.shape[:2])
|
||||
if was_training:
|
||||
self.transformer.train()
|
||||
|
||||
# Student prediction at lower t on the ODE-stepped sample
|
||||
tb_low_student = self._build_distill_input_kwargs(noisy_low_ode, t_low,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
pred_noise_low_student = self.transformer(**tb_low_student.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
y_low_student = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_low_student.flatten(0, 1),
|
||||
noise_input_latent=noisy_low_ode.flatten(0, 1),
|
||||
timestep=t_low,
|
||||
scheduler=self.noise_scheduler,
|
||||
).unflatten(0, pred_noise_low_student.shape[:2])
|
||||
|
||||
cm_loss_weight = self.training_args.cm_loss_weight
|
||||
weighing_fn = getattr(self.training_args, "cm_weighing_function", "constant")
|
||||
|
||||
if weighing_fn == "constant":
|
||||
cm_weighing_function = lambda x: cm_loss_weight
|
||||
elif weighing_fn == "sigma_sqrt":
|
||||
cm_weighing_function = lambda x: cm_loss_weight * (x**0.5)
|
||||
else:
|
||||
raise ValueError(f"Invalid cm_weighing_function: {weighing_fn}")
|
||||
|
||||
# Match student at lower t to teacher at lower t (adjacent PF ODE point)
|
||||
cm_loss = F.mse_loss(y_low_student, y_low_teacher.detach())
|
||||
cm_loss = cm_loss * cm_weighing_function(t_low)
|
||||
|
||||
# Optionally record for logging
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"cm_timestep_high": t_high,
|
||||
"cm_timestep_low": t_low,
|
||||
})
|
||||
|
||||
return cm_loss
|
||||
|
||||
def _clip_model_grad_norm_(self, training_batch: TrainingBatch,
|
||||
transformer) -> TrainingBatch:
|
||||
|
||||
max_grad_norm = self.training_args.max_grad_norm
|
||||
|
||||
if max_grad_norm is not None:
|
||||
model_parts = [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,
|
||||
foreach=None,
|
||||
)
|
||||
assert grad_norm is not float('nan') or grad_norm is not float(
|
||||
'inf')
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
training_batch.grad_norm = grad_norm
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
super()._prepare_dit_inputs(training_batch)
|
||||
conditional_dict = {
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
training_batch.conditional_dict = conditional_dict
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
self.video_latent_shape = training_batch.latents.shape
|
||||
|
||||
if self.sp_world_size > 1:
|
||||
training_batch.latents = rearrange(
|
||||
training_batch.latents,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
training_batch.latents = training_batch.latents[:, self.
|
||||
rank_in_sp_group, :, :, :, :]
|
||||
|
||||
self.video_latent_shape_sp = training_batch.latents.shape
|
||||
|
||||
return training_batch
|
||||
|
||||
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
gradient_accumulation_steps = getattr(self.training_args, 'gradient_accumulation_steps', 1)
|
||||
batches: list[TrainingBatch] = []
|
||||
# Collect N batches for gradient accumulation
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
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._build_attention_metadata(batch)
|
||||
batch.attn_metadata_vsa = copy.deepcopy(batch.attn_metadata)
|
||||
if batch.attn_metadata is not None:
|
||||
batch.attn_metadata.VSA_sparsity = 0.0 # type: ignore
|
||||
batches.append(batch)
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
total_loss_value = 0.0
|
||||
dmd_latent_vis_dict: dict[str, torch.Tensor] = {}
|
||||
batch_gen = None
|
||||
|
||||
for batch in batches:
|
||||
batch_gen = copy.deepcopy(batch)
|
||||
# Forward student once to obtain a clean prediction to anchor CM pair
|
||||
with set_forward_context(current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata_vsa):
|
||||
generator_pred_video = self._generator_forward(batch_gen)
|
||||
|
||||
# Consistency loss
|
||||
with set_forward_context(current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata):
|
||||
cm_loss = self._calculate_consistency_loss(
|
||||
pred_video=generator_pred_video, training_batch=batch_gen)
|
||||
|
||||
# Optional auxiliary regression loss
|
||||
if getattr(self.training_args, "use_regression_loss", False):
|
||||
cm_loss = cm_loss + batch_gen.regression_loss * self.training_args.regression_loss_weight
|
||||
|
||||
with set_forward_context(current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata_vsa):
|
||||
(cm_loss / gradient_accumulation_steps).backward()
|
||||
total_loss_value += cm_loss.detach().item()
|
||||
dmd_latent_vis_dict.update(batch_gen.dmd_latent_vis_dict)
|
||||
|
||||
# Clip and step
|
||||
assert batch_gen is not None
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
# Update EMA teacher after parameter update
|
||||
if self.use_cm_with_ema:
|
||||
assert self.ema_teacher is not None
|
||||
self.ema_teacher.update(self.transformer)
|
||||
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
avg_loss = torch.tensor(total_loss_value / max(1, gradient_accumulation_steps), device=self.device)
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
|
||||
training_batch.total_loss = avg_loss.item()
|
||||
training_batch.dmd_latent_vis_dict = dmd_latent_vis_dict
|
||||
training_batch.grad_norm = getattr(training_batch, "grad_norm", 0.0)
|
||||
return training_batch
|
||||
|
||||
def _resume_from_checkpoint(self) -> None:
|
||||
"""Resume training from checkpoint for the generator model only."""
|
||||
|
||||
logger.info("Loading checkpoint from %s",
|
||||
self.training_args.resume_from_checkpoint)
|
||||
|
||||
resumed_step = load_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.resume_from_checkpoint,
|
||||
self.optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.noise_random_generator)
|
||||
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
logger.info("Successfully resumed from step %s", resumed_step)
|
||||
else:
|
||||
logger.warning("Failed to load checkpoint, starting from step 0")
|
||||
self.init_steps = -1
|
||||
|
||||
def _log_training_info(self) -> None:
|
||||
"""Log distillation-specific training information."""
|
||||
# First call parent class method to get basic training info
|
||||
super()._log_training_info()
|
||||
|
||||
# Then add distillation-specific information
|
||||
logger.info("Distillation-specific settings:")
|
||||
assert isinstance(self.training_args, TrainingArgs)
|
||||
logger.info(" Max gradient norm: %s", self.training_args.max_grad_norm)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
training_args.inference_mode = True
|
||||
training_args.dit_cpu_offload = True
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
raise ValueError("Validation pipeline is not set")
|
||||
|
||||
logger.info("Starting validation")
|
||||
|
||||
# Create sampling parameters if not provided
|
||||
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
|
||||
|
||||
# Set deterministic seed for validation
|
||||
|
||||
logger.info("Using validation seed: %s", self.seed)
|
||||
|
||||
# Prepare validation prompts
|
||||
logger.info('rank: %s: fastvideo_args.validation_dataset_file: %s',
|
||||
self.global_rank,
|
||||
training_args.validation_dataset_file,
|
||||
local_main_process_only=False)
|
||||
validation_dataset = ValidationDataset(
|
||||
training_args.validation_dataset_file)
|
||||
validation_dataloader = DataLoader(validation_dataset,
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
transformer.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
validation_steps = [step for step in validation_steps if step > 0]
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
# Process each validation prompt for each validation step
|
||||
for num_inference_steps in validation_steps:
|
||||
logger.info("rank: %s: num_inference_steps: %s",
|
||||
self.global_rank,
|
||||
num_inference_steps,
|
||||
local_main_process_only=False)
|
||||
step_videos: list[np.ndarray] = []
|
||||
step_captions: list[str] = []
|
||||
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(sampling_param,
|
||||
training_args,
|
||||
validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
|
||||
batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
|
||||
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
self.global_rank,
|
||||
self.rank_in_sp_group,
|
||||
batch.prompt,
|
||||
local_main_process_only=False)
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
with torch.no_grad():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in video:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
step_videos.append(frames)
|
||||
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
|
||||
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
|
||||
# results to global rank 0
|
||||
if self.rank_in_sp_group == 0:
|
||||
if self.global_rank == 0:
|
||||
# Global rank 0 collects results from all sp_group leaders
|
||||
all_videos = step_videos # Start with own results
|
||||
all_captions = step_captions
|
||||
|
||||
# Receive from other sp_group leaders
|
||||
for sp_group_idx in range(1, num_sp_groups):
|
||||
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
|
||||
recv_videos = world_group.recv_object(src=src_rank)
|
||||
recv_captions = world_group.recv_object(src=src_rank)
|
||||
all_videos.extend(recv_videos)
|
||||
all_captions.extend(recv_captions)
|
||||
|
||||
video_filenames = []
|
||||
for i, (video, caption) in enumerate(
|
||||
zip(all_videos, all_captions, strict=True)):
|
||||
os.makedirs(training_args.output_dir, exist_ok=True)
|
||||
filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
|
||||
)
|
||||
imageio.mimsave(filename, video, fps=sampling_param.fps)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
f"validation_videos_{num_inference_steps}_steps": [
|
||||
wandb.Video(filename, caption=caption)
|
||||
for filename, caption in zip(
|
||||
video_filenames, all_captions, strict=True)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
else:
|
||||
# Other sp_group leaders send their results to global rank 0
|
||||
world_group.send_object(step_videos, dst=0)
|
||||
world_group.send_object(step_captions, dst=0)
|
||||
|
||||
# Re-enable gradients for training
|
||||
transformer.train()
|
||||
gc.collect()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
wandb_loss_dict = {}
|
||||
dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
|
||||
|
||||
# Only log generator output for CM pipeline
|
||||
if 'generator_pred_video' in dmd_latents_vis_dict:
|
||||
latents = dmd_latents_vis_dict['generator_pred_video']
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
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['generator_pred_video'] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
# Clean up references
|
||||
del video, latents
|
||||
|
||||
# Log to wandb
|
||||
if self.global_rank == 0:
|
||||
wandb.log(wandb_loss_dict, step=step)
|
||||
|
||||
def train(self) -> None:
|
||||
"""Main training loop with distillation-specific logging."""
|
||||
assert self.training_args.seed is not None, "seed must be set"
|
||||
seed = self.training_args.seed
|
||||
|
||||
# Set the same seed within each SP group to ensure reproducibility
|
||||
if self.sp_world_size > 1:
|
||||
# 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)
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
|
||||
self.seed)
|
||||
self.validation_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(self.seed)
|
||||
logger.info("Initialized random seeds with seed: %s", seed)
|
||||
|
||||
# Resume from checkpoint if specified (this will restore random states)
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
logger.info("Resumed from checkpoint, random states restored")
|
||||
else:
|
||||
logger.info("Starting training from scratch")
|
||||
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
self._log_training_info()
|
||||
self._log_validation(self.transformer, self.training_args,
|
||||
self.init_steps)
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, self.training_args.max_train_steps),
|
||||
initial=self.init_steps,
|
||||
desc="Steps",
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
|
||||
use_vsa = vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"
|
||||
for step in range(self.init_steps + 1,
|
||||
self.training_args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
if use_vsa:
|
||||
vsa_sparsity = self.training_args.VSA_sparsity
|
||||
vsa_decay_rate = self.training_args.VSA_decay_rate
|
||||
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
|
||||
if vsa_decay_interval_steps > 1:
|
||||
current_decay_times = min(step // vsa_decay_interval_steps,
|
||||
vsa_sparsity // vsa_decay_rate)
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
else:
|
||||
current_vsa_sparsity = vsa_sparsity
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
training_batch = TrainingBatch()
|
||||
self.current_trainstep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
total_loss = training_batch.total_loss
|
||||
grad_norm = training_batch.grad_norm
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"total_loss":
|
||||
f"{total_loss:.4f}",
|
||||
"lpips_loss":
|
||||
f"{training_batch.regression_loss:.4f}"
|
||||
if hasattr(training_batch, 'regression_loss')
|
||||
and training_batch.regression_loss is not None else "N/A",
|
||||
"step_time":
|
||||
f"{step_time:.2f}s",
|
||||
"grad_norm":
|
||||
grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
if self.global_rank == 0:
|
||||
# Prepare logging data
|
||||
log_data = {
|
||||
"train_total_loss": total_loss,
|
||||
"learning_rate": self.lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
}
|
||||
if use_vsa:
|
||||
log_data["VSA_train_sparsity"] = current_vsa_sparsity
|
||||
|
||||
if training_batch.dmd_latent_vis_dict:
|
||||
dmd_additional_logs = {
|
||||
"generator_timestep": training_batch.dmd_latent_vis_dict["generator_timestep"].item(),
|
||||
"regression_loss": training_batch.regression_loss,
|
||||
}
|
||||
log_data.update(dmd_additional_logs)
|
||||
|
||||
wandb.log(log_data, step=step)
|
||||
|
||||
# Save training state checkpoint (for resuming training)
|
||||
if (self.training_args.training_state_checkpointing_steps > 0
|
||||
and step % self.training_args.training_state_checkpointing_steps == 0):
|
||||
print("rank", self.global_rank,
|
||||
"save training state checkpoint at step", step)
|
||||
save_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.output_dir, step,
|
||||
self.optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.noise_random_generator)
|
||||
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
|
||||
# Save weight-only checkpoint (export consolidated generator weights)
|
||||
if (self.training_args.weight_only_checkpointing_steps > 0
|
||||
and step % self.training_args.weight_only_checkpointing_steps == 0):
|
||||
print("rank", self.global_rank,
|
||||
"save weight-only checkpoint at step", step)
|
||||
save_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}", self.optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.noise_random_generator)
|
||||
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
if self.training_args.log_visualization:
|
||||
self.visualize_intermediate_latents(training_batch,
|
||||
self.training_args,
|
||||
step)
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
|
||||
wandb.finish()
|
||||
|
||||
# Save final training state checkpoint
|
||||
print("rank", self.global_rank,
|
||||
"save final training state checkpoint at step",
|
||||
self.training_args.max_train_steps)
|
||||
save_checkpoint(self.transformer, self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -16,6 +16,8 @@ import torchvision
|
||||
from einops import rearrange
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from torchmetrics.image.lpip import (
|
||||
LearnedPerceptualImagePatchSimilarity as LPIPSimilarity)
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import fastvideo.envs as envs
|
||||
@@ -35,8 +37,8 @@ from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases, get_scheduler,
|
||||
load_distillation_checkpoint, pred_noise_to_pred_video,
|
||||
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
get_scheduler, load_distillation_checkpoint, pred_noise_to_pred_video,
|
||||
save_distillation_checkpoint, shift_timestep)
|
||||
from fastvideo.utils import is_vsa_available, set_random_seed
|
||||
|
||||
@@ -66,6 +68,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
current_epoch: int = 0
|
||||
init_steps: int
|
||||
current_trainstep: int
|
||||
num_generator_updates: int = 0
|
||||
video_latent_shape: tuple[int, ...]
|
||||
video_latent_shape_sp: tuple[int, ...]
|
||||
|
||||
@@ -143,10 +146,28 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"Distillation pipeline initialized with generator_update_interval=%s",
|
||||
self.generator_update_interval)
|
||||
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
if self.training_args.warp_denoising_step:
|
||||
# timesteps = self.noise_scheduler.timesteps.cpu()
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0], dtype=torch.float32)))
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=torch.device("cpu"))
|
||||
self.denoising_step_list = timesteps[1000 -
|
||||
self.denoising_step_list].to(
|
||||
get_local_torch_device())
|
||||
logger.info(
|
||||
"Warp denoising step is enabled, using %s denoising steps",
|
||||
self.denoising_step_list)
|
||||
else:
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
logger.info(
|
||||
"Warp denoising step is disabled, using %s denoising steps",
|
||||
self.denoising_step_list)
|
||||
logger.info("Distillation generator model to %s denoising steps",
|
||||
len(self.denoising_step_list))
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
@@ -158,6 +179,70 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
|
||||
# Initialize EMA teacher for CM if enabled
|
||||
self.use_cm_with_ema = getattr(self.training_args, "cm_use_ema_teacher",
|
||||
False)
|
||||
if self.use_cm_with_ema:
|
||||
# Use local_shard mode for teacher forward compatibility with FSDP2
|
||||
self.ema_teacher = EMA_FSDP(self.transformer,
|
||||
decay=self.training_args.ema_decay,
|
||||
mode="local_shard")
|
||||
else:
|
||||
self.ema_teacher = None
|
||||
self.lpips = LPIPSimilarity().to(self.device)
|
||||
|
||||
def _compute_lpips_loss(self, pred_video: torch.Tensor,
|
||||
latents: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Helper method to compute LPIPS loss between predicted video and ground truth latents.
|
||||
This method handles VAE scaling, shifting, decoding, and frame processing consistently.
|
||||
"""
|
||||
print("Computing LPIPS loss...")
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16), torch.no_grad():
|
||||
# Apply VAE scaling factor and shift factor before decoding (same as visualize_intermediate_latents)
|
||||
pred_video = pred_video.permute(0, 2, 1, 3, 4)
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
pred_video_scaled = pred_video / self.vae.scaling_factor.to(
|
||||
pred_video.device, pred_video.dtype)
|
||||
latents_scaled = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
pred_video_scaled = pred_video / self.vae.scaling_factor
|
||||
latents_scaled = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
pred_video_scaled += self.vae.shift_factor.to(
|
||||
pred_video.device, pred_video.dtype)
|
||||
latents_scaled += self.vae.shift_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
pred_video_scaled += self.vae.shift_factor
|
||||
latents_scaled += self.vae.shift_factor
|
||||
|
||||
# Permute from [batch, channels, frames, height, width] to [batch, frames, channels, height, width] for VAE decode
|
||||
# pred_video_scaled = pred_video_scaled.permute(0, 2, 1, 3, 4)
|
||||
# latents_scaled = latents_scaled.permute(0, 2, 1, 3, 4)
|
||||
|
||||
# print(f"pred_video_scaled shape after permute: {pred_video_scaled.shape}")
|
||||
# print(f"latents_scaled shape after permute: {latents_scaled.shape}")
|
||||
pred_video_frames = self.vae.decode(pred_video_scaled)
|
||||
latents_frames = self.vae.decode(latents_scaled)
|
||||
|
||||
# VAE output is already in [-1, 1] range, keep it for LPIPS
|
||||
# Just permute back to [B, T, C, H, W] format for frame processing
|
||||
pred_video_frames = pred_video_frames.permute(0, 2, 1, 3, 4)
|
||||
latents_frames = latents_frames.permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_video_frames = rearrange(pred_video_frames,
|
||||
"b n c h w -> (b n) c h w")
|
||||
latents_frames = rearrange(latents_frames, "b n c h w -> (b n) c h w")
|
||||
|
||||
return self.lpips(pred_video_frames, latents_frames)
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize validation pipeline - must be implemented by subclasses."""
|
||||
@@ -228,7 +313,18 @@ class DistillationPipeline(TrainingPipeline):
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
# regression_loss = F.mse_loss(pred_video, latents)
|
||||
|
||||
# LPIPS loss
|
||||
regression_loss = self._compute_lpips_loss(pred_video, latents)
|
||||
pred_video = pred_video.type_as(noisy_latent)
|
||||
training_batch.regression_loss = regression_loss
|
||||
# regression_log_dict = {
|
||||
# "regression_loss": regression_Loss.detach(),
|
||||
# "regression_clean_latents": latents,
|
||||
# }
|
||||
return pred_video
|
||||
# return pred_video
|
||||
|
||||
def _generator_multi_step_simulation_forward(
|
||||
self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
@@ -241,23 +337,28 @@ class DistillationPipeline(TrainingPipeline):
|
||||
len(self.denoising_step_list), [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
# target_timestep_idx = torch.tensor([2],
|
||||
# device=self.device,
|
||||
# dtype=torch.long)
|
||||
target_timestep_idx_int = target_timestep_idx.item()
|
||||
target_timestep = self.denoising_step_list[target_timestep_idx]
|
||||
|
||||
# Step 2: Simulate the multi-step inference process up to the target timestep
|
||||
# Start from pure noise like in inference
|
||||
current_noise_latents = torch.randn(self.video_latent_shape,
|
||||
device=self.device,
|
||||
dtype=dtype)
|
||||
noise_latent = torch.randn(self.video_latent_shape,
|
||||
device=self.device,
|
||||
dtype=dtype)
|
||||
if self.sp_world_size > 1:
|
||||
current_noise_latents = rearrange(
|
||||
current_noise_latents,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
current_noise_latents = current_noise_latents[:, self.
|
||||
rank_in_sp_group, :, :, :, :]
|
||||
noise_latents = rearrange(noise_latent,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise_latent = noise_latent[:, self.rank_in_sp_group, :, :, :, :]
|
||||
current_noise_latents = noise_latent.clone()
|
||||
|
||||
# Only run intermediate steps if target_timestep_idx > 0
|
||||
max_target_idx = len(self.denoising_step_list) - 1
|
||||
noise_latents = []
|
||||
noise_latent_index = target_timestep_idx_int - 1
|
||||
if max_target_idx > 0:
|
||||
# Run student model for all steps before the target timestep
|
||||
with torch.no_grad():
|
||||
@@ -294,11 +395,19 @@ class DistillationPipeline(TrainingPipeline):
|
||||
current_noise_latents = self.noise_scheduler.add_noise(
|
||||
pred_clean.flatten(0, 1), noise.flatten(0, 1),
|
||||
next_timestep_tensor).unflatten(0, pred_clean.shape[:2])
|
||||
latent_copy = current_noise_latents.clone()
|
||||
noise_latents.append(latent_copy)
|
||||
|
||||
# Step 3: Use the simulated noisy input for the final training step
|
||||
# For timestep index 0, this is pure noise
|
||||
# For timestep index k > 0, this is the result after k denoising steps + noise at target level
|
||||
noisy_input = current_noise_latents
|
||||
if noise_latent_index >= 0:
|
||||
assert noise_latent_index < len(
|
||||
self.denoising_step_list
|
||||
) - 1, "noise_latent_index is out of bounds"
|
||||
noisy_input = noise_latents[noise_latent_index]
|
||||
else:
|
||||
noisy_input = noise_latent
|
||||
|
||||
# Step 4: Final student prediction (this is what we train on)
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -311,10 +420,122 @@ class DistillationPipeline(TrainingPipeline):
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=target_timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
regression_loss = F.mse_loss(pred_video, latents)
|
||||
# LPIPS loss - decode latents to actual frames first
|
||||
regression_loss = self._compute_lpips_loss(pred_video, latents)
|
||||
training_batch.regression_loss = regression_loss
|
||||
training_batch.dmd_latent_vis_dict[
|
||||
"generator_timestep"] = target_timestep.float().detach()
|
||||
|
||||
return pred_video
|
||||
|
||||
def _calculate_consistency_loss(
|
||||
self, pred_video: torch.Tensor,
|
||||
training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""
|
||||
Consistency loss (CM-like): two variants
|
||||
- Without EMA: teacher = model at higher noise with stop-grad
|
||||
- With EMA: teacher = EMA(model) at lower noise or higher noise depending on schedule (use higher noise as teacher)
|
||||
"""
|
||||
# Need at least two steps to form a pair
|
||||
if len(self.denoising_step_list) < 2:
|
||||
return torch.tensor(0.0, device=self.device, dtype=pred_video.dtype)
|
||||
|
||||
# Choose a random index k in [1, N-1] so that (k-1, k) is valid
|
||||
idx_high = torch.randint(1,
|
||||
len(self.denoising_step_list), [1],
|
||||
device=self.device).item()
|
||||
t_high = self.denoising_step_list[idx_high]
|
||||
t_low = self.denoising_step_list[idx_high - 1]
|
||||
t_low = t_low * torch.ones(1, device=self.device, dtype=torch.long)
|
||||
t_high = t_high * torch.ones(1, device=self.device, dtype=torch.long)
|
||||
logger.info("t_high: %s, t_low: %s", t_high, t_low)
|
||||
|
||||
# Use student output as base for CM pairing (noisy inputs derived from student's prediction)
|
||||
base_student = pred_video
|
||||
# base_student = training_batch.latents
|
||||
# Shared base noise for both timesteps
|
||||
base_noise = torch.randn(self.video_latent_shape,
|
||||
device=self.device,
|
||||
dtype=base_student.dtype)
|
||||
if self.sp_world_size > 1:
|
||||
base_noise = rearrange(base_noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
base_noise = base_noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
# Build noisy inputs for the selected timesteps using the same student base and same noise
|
||||
noisy_high = self.noise_scheduler.add_noise(
|
||||
base_student.flatten(0, 1), base_noise.flatten(0, 1),
|
||||
t_high).unflatten(0, (1, base_student.shape[1]))
|
||||
|
||||
# Teacher ODE step: t_high → t_low and teacher target at t_low
|
||||
was_training = self.transformer.training
|
||||
self.transformer.eval()
|
||||
with torch.no_grad():
|
||||
tb_high_teacher = self._build_distill_input_kwargs(
|
||||
noisy_high, t_high, training_batch.conditional_dict, training_batch)
|
||||
pred_flow_high_teacher = self.transformer(
|
||||
**tb_high_teacher.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
sample_flat = noisy_high.flatten(0, 1).float()
|
||||
flow_flat = pred_flow_high_teacher.flatten(0, 1).float()
|
||||
|
||||
prev_step_index = getattr(self.noise_scheduler, "_step_index", None)
|
||||
self.noise_scheduler._step_index = None
|
||||
stepped = self.noise_scheduler.step(
|
||||
model_output=flow_flat, timestep=t_high, sample=sample_flat,
|
||||
return_dict=True)
|
||||
self.noise_scheduler._step_index = prev_step_index
|
||||
|
||||
noisy_low_ode = stepped.prev_sample.unflatten(0, (1, noisy_high.shape[1])).to(noisy_high.dtype)
|
||||
|
||||
tb_low_teacher = self._build_distill_input_kwargs(
|
||||
noisy_low_ode, t_low, training_batch.conditional_dict, training_batch)
|
||||
pred_noise_low_teacher = self.transformer(
|
||||
**tb_low_teacher.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
y_low_teacher = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_low_teacher.flatten(0, 1),
|
||||
noise_input_latent=noisy_low_ode.flatten(0, 1),
|
||||
timestep=t_low,
|
||||
scheduler=self.noise_scheduler,
|
||||
).unflatten(0, pred_noise_low_teacher.shape[:2])
|
||||
if was_training:
|
||||
self.transformer.train()
|
||||
|
||||
# Student prediction at lower t on the ODE-stepped sample
|
||||
tb_low_student = self._build_distill_input_kwargs(
|
||||
noisy_low_ode, t_low, training_batch.conditional_dict, training_batch)
|
||||
pred_noise_low_student = self.transformer(
|
||||
**tb_low_student.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
y_low_student = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_low_student.flatten(0, 1),
|
||||
noise_input_latent=noisy_low_ode.flatten(0, 1),
|
||||
timestep=t_low,
|
||||
scheduler=self.noise_scheduler,
|
||||
).unflatten(0, pred_noise_low_student.shape[:2])
|
||||
|
||||
cm_loss_weight = self.training_args.cm_loss_weight
|
||||
weighing_fn = getattr(self.training_args, "cm_weighing_function", "constant")
|
||||
if weighing_fn == "constant":
|
||||
cm_weighing_function = lambda x: cm_loss_weight
|
||||
elif weighing_fn == "sigma_sqrt":
|
||||
cm_weighing_function = lambda x: cm_loss_weight * (x**0.5)
|
||||
else:
|
||||
raise ValueError(f"Invalid cm_weighing_function: {weighing_fn}")
|
||||
|
||||
cm_loss = F.mse_loss(y_low_student, y_low_teacher.detach())
|
||||
cm_loss = cm_loss * cm_weighing_function(t_low)
|
||||
|
||||
# Optionally record for logging
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"cm_timestep_high": t_high,
|
||||
"cm_timestep_low": t_low,
|
||||
})
|
||||
|
||||
return cm_loss
|
||||
|
||||
def _dmd_forward(self, generator_pred_video: torch.Tensor,
|
||||
training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Compute DMD (Diffusion Model Distillation) loss."""
|
||||
@@ -398,6 +619,9 @@ class DistillationPipeline(TrainingPipeline):
|
||||
generator_pred_video.float(),
|
||||
(generator_pred_video.float() - grad.float()).detach())
|
||||
|
||||
if self.training_args.use_regression_loss:
|
||||
dmd_loss = dmd_loss + training_batch.regression_loss * self.training_args.regression_loss_weight
|
||||
|
||||
training_batch.dmd_latent_vis_dict.update({
|
||||
"training_batch_dmd_fwd_clean_latent":
|
||||
training_batch.latents,
|
||||
@@ -409,6 +633,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
faker_score_pred_video,
|
||||
"dmd_timestep":
|
||||
timestep,
|
||||
"regression_loss":
|
||||
training_batch.regression_loss,
|
||||
})
|
||||
|
||||
return dmd_loss
|
||||
@@ -420,8 +646,12 @@ class DistillationPipeline(TrainingPipeline):
|
||||
current_timestep=training_batch.timesteps,
|
||||
attn_metadata=training_batch.attn_metadata_vsa):
|
||||
if self.training_args.simulate_generator_forward:
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
training_batch)
|
||||
if False:
|
||||
generator_pred_video = self._generator_forward_with_validation_pipeline(
|
||||
training_batch)
|
||||
else:
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
training_batch)
|
||||
else:
|
||||
generator_pred_video = self._generator_forward(training_batch)
|
||||
|
||||
@@ -551,29 +781,57 @@ class DistillationPipeline(TrainingPipeline):
|
||||
for batch in batches:
|
||||
batch_gen = copy.deepcopy(batch)
|
||||
|
||||
# Generator forward and DMD loss after CM
|
||||
with set_forward_context(
|
||||
current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata_vsa):
|
||||
if self.training_args.simulate_generator_forward:
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
batch_gen)
|
||||
if self.training_args.simulate_generator_forward and self.num_generator_updates % self.training_args.simulate_forward_interval == 0:
|
||||
logger.info(
|
||||
"using simulate_generator_forward, num_generator_updates: %s",
|
||||
self.num_generator_updates)
|
||||
if False:
|
||||
generator_pred_video = self._generator_forward_with_validation_pipeline(
|
||||
batch_gen)
|
||||
else:
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
batch_gen)
|
||||
else:
|
||||
generator_pred_video = self._generator_forward(
|
||||
batch_gen)
|
||||
|
||||
# Consistency loss first (to avoid in-place param swaps after any grad-tracked forward)
|
||||
cm_weight = self.training_args.cm_loss_weight
|
||||
dmd_loss_extra = 0.0
|
||||
if cm_weight > 0.0:
|
||||
with set_forward_context(
|
||||
current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata):
|
||||
cm_loss = self._calculate_consistency_loss(
|
||||
pred_video=generator_pred_video,
|
||||
training_batch=batch_gen,
|
||||
)
|
||||
dmd_loss_extra = cm_loss
|
||||
|
||||
with set_forward_context(current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata):
|
||||
dmd_loss = self._dmd_forward(
|
||||
generator_pred_video=generator_pred_video,
|
||||
training_batch=batch_gen)
|
||||
logger.info("dmd_loss: %s", dmd_loss.item())
|
||||
dmd_loss = dmd_loss + dmd_loss_extra
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=batch_gen.timesteps,
|
||||
attn_metadata=batch_gen.attn_metadata_vsa):
|
||||
(dmd_loss / gradient_accumulation_steps).backward()
|
||||
total_dmd_loss += dmd_loss.detach().item()
|
||||
self.num_generator_updates += 1
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
self.optimizer.step()
|
||||
# Update EMA teacher after generator parameter update
|
||||
if self.use_cm_with_ema:
|
||||
assert self.ema_teacher is not None
|
||||
self.ema_teacher.update(self.transformer)
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
avg_dmd_loss = torch.tensor(total_dmd_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -802,7 +1060,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
|
||||
fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
|
||||
fake_score_log_keys = ['generator_pred_video']
|
||||
dmd_log_keys = ['faker_score_pred_video']
|
||||
dmd_log_keys = ['faker_score_pred_video', 'real_score_pred_video']
|
||||
|
||||
for latent_key in fake_score_log_keys:
|
||||
latents = fake_score_latents_vis_dict[latent_key]
|
||||
@@ -948,11 +1206,20 @@ class DistillationPipeline(TrainingPipeline):
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"total_loss": f"{total_loss:.4f}",
|
||||
"generator_loss": f"{generator_loss:.4f}",
|
||||
"fake_score_loss": f"{fake_score_loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
"total_loss":
|
||||
f"{total_loss:.4f}",
|
||||
"generator_loss":
|
||||
f"{generator_loss:.4f}",
|
||||
"fake_score_loss":
|
||||
f"{fake_score_loss:.4f}",
|
||||
"lpips_loss":
|
||||
f"{training_batch.regression_loss:.4f}"
|
||||
if hasattr(training_batch, 'regression_loss')
|
||||
and training_batch.regression_loss is not None else "N/A",
|
||||
"step_time":
|
||||
f"{step_time:.2f}s",
|
||||
"grad_norm":
|
||||
grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
@@ -988,6 +1255,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"dmd_timestep":
|
||||
training_batch.dmd_latent_vis_dict["dmd_timestep"].item(
|
||||
),
|
||||
"regression_loss":
|
||||
training_batch.regression_loss,
|
||||
}
|
||||
log_data.update(dmd_additional_logs)
|
||||
|
||||
|
||||
@@ -1299,3 +1299,156 @@ def get_scheduler(
|
||||
num_warmup_steps=num_warmup_steps,
|
||||
num_training_steps=num_training_steps,
|
||||
last_epoch=last_epoch)
|
||||
|
||||
|
||||
class EMA_FSDP:
|
||||
"""
|
||||
FSDP2-friendly EMA with two modes:
|
||||
- mode="local_shard" (default): maintain float32 CPU EMA of local parameter shards on every rank.
|
||||
Provides a context manager to temporarily swap EMA weights into the live model for teacher forward.
|
||||
- mode="rank0_full": maintain a consolidated float32 CPU EMA of full parameters on rank 0 only
|
||||
using gather_state_dict_on_cpu_rank0(). Useful for checkpoint export; not for teacher forward.
|
||||
|
||||
Usage (local_shard for CM teacher):
|
||||
ema = EMA_FSDP(model, decay=0.999, mode="local_shard")
|
||||
for step in ...:
|
||||
ema.update(model)
|
||||
with ema.apply_to_model(model):
|
||||
with torch.no_grad():
|
||||
y_teacher = model(...)
|
||||
|
||||
Usage (rank0_full for export):
|
||||
ema = EMA_FSDP(model, decay=0.999, mode="rank0_full")
|
||||
ema.update(model)
|
||||
ema.state_dict() # on rank 0
|
||||
"""
|
||||
def __init__(self, module, decay: float = 0.999, mode: str = "local_shard"):
|
||||
self.decay = float(decay)
|
||||
self.mode = mode
|
||||
self.shadow: dict[str, torch.Tensor] = {}
|
||||
self.rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
if self.mode not in {"local_shard", "rank0_full"}:
|
||||
raise ValueError(f"Unsupported EMA_FSDP mode: {self.mode}")
|
||||
self._init_shadow(module)
|
||||
|
||||
@staticmethod
|
||||
def _to_local_tensor(t: torch.Tensor) -> torch.Tensor:
|
||||
# DTensor-aware to_local fetch; fall back to raw tensor
|
||||
try:
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
if isinstance(t, DTensor):
|
||||
return t.to_local()
|
||||
except Exception:
|
||||
pass
|
||||
return t
|
||||
|
||||
@torch.no_grad()
|
||||
def _init_shadow(self, module):
|
||||
if self.mode == "rank0_full":
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
|
||||
if self.rank == 0:
|
||||
self.shadow = {k: v.detach().clone().float().cpu() for k, v in cpu_state.items()}
|
||||
else:
|
||||
self.shadow = {}
|
||||
return
|
||||
|
||||
# local_shard: maintain EMA of local shards for requires_grad params
|
||||
self.shadow = {}
|
||||
for name, p in module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
local = self._to_local_tensor(p.detach())
|
||||
self.shadow[name] = local.clone().float().cpu()
|
||||
|
||||
@torch.no_grad()
|
||||
def update(self, module):
|
||||
d = self.decay
|
||||
if self.mode == "rank0_full":
|
||||
if self.rank != 0:
|
||||
return
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
|
||||
for n, v in cpu_state.items():
|
||||
v_cpu = v.detach().float().cpu()
|
||||
if n not in self.shadow:
|
||||
self.shadow[n] = v_cpu.clone()
|
||||
else:
|
||||
self.shadow[n].mul_(d).add_(v_cpu, alpha=1.0 - d)
|
||||
return
|
||||
|
||||
# local_shard: update local shard EMA on every rank
|
||||
for name, p in module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
local = self._to_local_tensor(p.detach())
|
||||
v_cpu = local.float().cpu()
|
||||
if name not in self.shadow:
|
||||
self.shadow[name] = v_cpu.clone()
|
||||
else:
|
||||
self.shadow[name].mul_(d).add_(v_cpu, alpha=1.0 - d)
|
||||
|
||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||
if self.mode == "rank0_full":
|
||||
return {k: v.clone() for k, v in self.shadow.items()} if self.rank == 0 else {}
|
||||
return {k: v.clone() for k, v in self.shadow.items()}
|
||||
|
||||
def load_state_dict(self, sd: dict[str, torch.Tensor]):
|
||||
self.shadow = {k: v.clone() for k, v in sd.items()}
|
||||
|
||||
@torch.no_grad()
|
||||
def copy_to_unwrapped(self, module) -> None:
|
||||
"""
|
||||
Copy EMA weights into a non-sharded (unwrapped) module. Intended for export/eval.
|
||||
For mode="rank0_full", only rank 0 has the full EMA state.
|
||||
"""
|
||||
if self.mode == "rank0_full" and self.rank != 0:
|
||||
return
|
||||
name_to_param = dict(module.named_parameters())
|
||||
for n, w in self.shadow.items():
|
||||
if n in name_to_param:
|
||||
p = name_to_param[n]
|
||||
p.data.copy_(w.to(dtype=p.dtype, device=p.device))
|
||||
|
||||
class _ApplyEMACtx:
|
||||
def __init__(self, ema: "EMA_FSDP", module):
|
||||
self.ema = ema
|
||||
self.module = module
|
||||
self.saved: dict[str, torch.Tensor] = {}
|
||||
|
||||
def __enter__(self):
|
||||
if self.ema.mode != "local_shard":
|
||||
raise RuntimeError("EMA apply_to_model is only supported for mode='local_shard'")
|
||||
with torch.no_grad():
|
||||
for name, p in self.module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
# Save local shard
|
||||
p_local = EMA_FSDP._to_local_tensor(p.detach())
|
||||
if p_local.numel() == 0:
|
||||
# Nothing to swap on this rank for this param
|
||||
continue
|
||||
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
|
||||
if name in self.ema.shadow:
|
||||
ema_cpu = self.ema.shadow[name]
|
||||
if ema_cpu.numel() != p_local.numel():
|
||||
# Shard shape mismatch (e.g., empty shard here), skip
|
||||
continue
|
||||
# Copy EMA shard into local param shard
|
||||
p_local.copy_(ema_cpu.to(dtype=p_local.dtype, device=p_local.device))
|
||||
return self.module
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
with torch.no_grad():
|
||||
for name, p in self.module.named_parameters():
|
||||
if name in self.saved:
|
||||
p_local = EMA_FSDP._to_local_tensor(p.detach())
|
||||
if p_local.numel() == 0:
|
||||
continue
|
||||
saved_local = self.saved[name]
|
||||
if saved_local.numel() != p_local.numel():
|
||||
continue
|
||||
p_local.copy_(saved_local)
|
||||
self.saved.clear()
|
||||
return False
|
||||
|
||||
def apply_to_model(self, module):
|
||||
return EMA_FSDP._ApplyEMACtx(self, module)
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_dmd_pipeline import WanDMDPipeline
|
||||
from fastvideo.training.cm_distillation_pipeline import CMDistillationPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanCMDistillationPipeline(CMDistillationPipeline):
|
||||
"""
|
||||
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"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize Wan-specific scheduler."""
|
||||
shift_value = fastvideo_args.pipeline_config.flow_shift
|
||||
shift = 1.0 if shift_value is None else float(shift_value)
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
validation_pipeline = WanDMDPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={"transformer": self.get_module("transformer")},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting Wan distillation pipeline...")
|
||||
|
||||
# Create pipeline with original args
|
||||
pipeline = WanCMDistillationPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
|
||||
args = pipeline.training_args
|
||||
# Start training
|
||||
pipeline.train()
|
||||
logger.info("Wan distillation pipeline completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -185,6 +185,20 @@ class WanI2VDistillationPipeline(DistillationPipeline):
|
||||
assert torch.isnan(image_embeds).sum() == 0
|
||||
image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
# from fastvideo.models.vision_utils import load_video
|
||||
# image = load_video("examples/distill/Wan2.1-I2V/crush_smol/validation_dataset/1gGQy4nxyUo-Scene-016.mp4")[0]
|
||||
# image_processor = self.validation_pipeline.get_module("image_processor")
|
||||
# image_encoder = self.validation_pipeline.get_module("image_encoder")
|
||||
# image_inputs = image_processor(
|
||||
# images=image, return_tensors="pt").to(get_local_torch_device())
|
||||
# from fastvideo.forward_context import set_forward_context
|
||||
# with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
# outputs = image_encoder(**image_inputs)
|
||||
# image_embeds = outputs.last_hidden_state
|
||||
# image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
# dtype=torch.bfloat16)
|
||||
# assert torch.isnan(image_embeds).sum() == 0
|
||||
# self.validation_pipeline.
|
||||
|
||||
noisy_model_input = torch.cat(
|
||||
[noise_input,
|
||||
@@ -196,6 +210,7 @@ class WanI2VDistillationPipeline(DistillationPipeline):
|
||||
4), # bs, c, t, h, w
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"image_latent": training_batch.image_latents,
|
||||
"timestep": timestep,
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"return_dict": False,
|
||||
|
||||
Reference in New Issue
Block a user