Compare commits

...
Author SHA1 Message Date
SolitaryThinker 083c2258ea fix cm 2025-08-21 05:18:30 +00:00
SolitaryThinker 60260a6be9 fix preprocessing before vae deocde 2025-08-20 01:10:53 +00:00
SolitaryThinker 06cc80ee93 init lpips 2025-08-19 23:48:06 +00:00
Matthew Noto c6b3475f77 LPIPs fix 2025-08-19 16:40:42 -07:00
Matthew Noto 1d34380c58 LPIPS 2025-08-19 15:40:29 -07:00
SolitaryThinker ddf8198e65 remove validation back sim; reg loss for back sim 2025-08-19 07:19:40 +00:00
SolitaryThinker c38ee88294 cm 2025-08-18 23:22:37 +00:00
SolitaryThinker ff1b77bdd6 update 2025-08-18 21:44:47 +00:00
SolitaryThinker 22ba3759ee cm 2025-08-18 03:14:18 +00:00
SolitaryThinker 2c3a3893c0 cm 2025-08-18 02:46:24 +00:00
SolitaryThinker 6e9ba7ba9d sim interval arg 2025-08-17 02:49:03 +00:00
SolitaryThinker 64f1295878 fix 2025-08-15 03:45:32 +00:00
SolitaryThinker d9cc3b7e42 reg loss 2025-08-15 02:08:15 +00:00
SolitaryThinker cb3230f5e9 reg loss 2025-08-15 01:30:46 +00:00
SolitaryThinker 1e91cec93d fix 2025-08-13 07:05:51 +00:00
SolitaryThinker 5f0688084c update 2025-08-13 07:05:24 +00:00
SolitaryThinker f3023db253 fix 2025-08-13 07:01:18 +00:00
SolitaryThinker c405941f21 fix 2025-08-13 04:05:01 +00:00
SolitaryThinker f4e818305b fix 2025-08-13 03:11:22 +00:00
SolitaryThinker 4c4ab65020 fix 2025-08-13 03:10:12 +00:00
SolitaryThinker 6b8bf495ee fix 2025-08-13 02:19:52 +00:00
SolitaryThinker ebd23ba8e7 add orig val 2025-08-12 20:00:33 +00:00
SolitaryThinker 0c1cb0c19d runnable 2025-08-12 19:58:36 +00:00
SolitaryThinker d87ebe051b fix 2025-08-11 03:42:08 +00:00
SolitaryThinker 0ae669cd4c fix 2025-08-11 03:31:02 +00:00
SolitaryThinker 63c124fec8 update 2025-08-10 23:37:01 +00:00
SolitaryThinker ececbe1839 wip 2025-08-10 08:33:03 +00:00
55 changed files with 2827 additions and 56 deletions
@@ -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
}
]
}
@@ -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
}
]
}
@@ -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[@]}"
@@ -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
+1
View File
@@ -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
}
+3
View File
@@ -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
}
+16 -1
View File
@@ -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
View File
@@ -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)
+21 -3
View File
@@ -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()
+296 -27
View File
@@ -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)
+153
View File
@@ -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,