Compare commits

...
Author SHA1 Message Date
RandNMR73 731923b619 fix hardcode 2025-10-15 11:26:24 +00:00
RandNMR73 e80611c5c4 checkpoint 2025-10-11 10:54:06 +00:00
JerryZhou54 fba8b61c4f Fix MoE recipe and Add 1.3B MoE script 2025-10-08 03:58:07 +00:00
JerryZhou54 55d8c1e5fb Fix small runtime issues 2025-10-07 05:03:36 +00:00
JerryZhou54 aeb8f2e5ac Add matthew's timestep change & add real score guidance scale 2 2025-10-06 22:24:44 +00:00
JerryZhou54 f755dd1ad5 Refactor sf distill code 2025-10-05 23:46:35 +00:00
JerryZhou54 256f788d34 checkpoint 2025-10-04 23:56:36 +00:00
dc7596b973 [self-forcing][8/n] Self-Forcing For Wan2.2-A14B + torch.compile training and distillation support (#818)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-10-02 15:01:45 -07:00
William Lin 335afa4457 [bugfix] Use training_state_checkpointing_steps instead of checkpointing_steps (#821) 2025-09-28 15:22:43 -07:00
Yongqi Chen 3f77a6805a [Feature]Update count trainable param for FSDP2 (#820) 2025-09-28 15:22:04 -07:00
RandNMR73 13d0aae706 Add Sage Attention 3 Backend (#815) 2025-09-24 15:11:38 -07:00
William Lin 404cbf4f3c [self-forcing] [6/n] Add Ode Init training (#811) 2025-09-22 17:58:19 -07:00
William Lin 958ffec844 [bugfix] Update learning rates for sparse distillation recipe (#812) 2025-09-22 12:07:03 -07:00
31f000d1cc [self-forcing] [5/n] Add Self-Forcing distillation pipeline (#808)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-09-20 19:32:10 -07:00
Yongqi Chen cd32b3e02f Update example files and readme (#809) 2025-09-20 18:15:59 -07:00
Zhang Peiyuan bf27908095 Update WeChat Link 2025-09-20 14:16:20 -07:00
82 changed files with 6105 additions and 412 deletions
+12
View File
@@ -104,6 +104,18 @@ steps:
- TEST_TYPE=distillation_dmd
agents:
queue: "default"
- path:
- "fastvideo/training/*self_forcing_distillation_pipeline.py"
- "fastvideo/tests/training/self-forcing/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Self-Forcing Tests"
env:
- TEST_TYPE=self_forcing
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
+4
View File
@@ -110,6 +110,10 @@ case "$TEST_TYPE" in
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
;;
# run_inference_tests_vmoba
"self_forcing")
log "Running self-forcing tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_self_forcing_tests"
;;
"inference_vmoba")
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
+4 -7
View File
@@ -12,9 +12,6 @@ exclude: |
scripts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/distill/.*|
fastvideo/distill\.py|
fastvideo/distill_adv\.py|
fastvideo/models/.*|
fastvideo/sample/.*|
fastvideo/train\.py|
@@ -44,10 +41,10 @@ repos:
- id: codespell
additional_dependencies: ['tomli']
args: ['--toml', 'pyproject.toml']
- repo: https://github.com/PyCQA/isort
rev: 6.0.1
hooks:
- id: isort
# - repo: https://github.com/PyCQA/isort
# rev: 6.0.1
# hooks:
# - id: isort
- repo: https://github.com/jackdewinter/pymarkdown
rev: v0.9.30
hooks:
+3 -3
View File
@@ -7,7 +7,7 @@
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/q46BbX6" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -155,8 +155,8 @@ If you find FastVideo useful, please considering citing our work:
}
@article{zhang2025vsa,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
title={Vsa: Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
+42 -2
View File
@@ -1,32 +1,40 @@
(inference-optimizations)=
# Optimizations
This page describes the various options for speeding up generation times in FastVideo.
## Table of Contents
- Optimized Attention Backends
- [Flash Attention](#optimizations-flash)
- [Sliding Tile Attention](#optimizations-sta)
- [Sage Attention](#optimizations-sage)
- [Sage Attention 3](#optimizations-sage3)
- Caching Techniques
- [TeaCache](#optimizations-teacache)
(optimizations-backends)=
## Attention Backends
### Available Backends
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
### Configuring Backends
There are two ways to configure the attention backend in FastVideo.
#### 1. In Python
In python, set the `FASTVIDEO_ATTENTION_BACKEND` environment variable before instantiating `VideoGenerator` like this:
```python
@@ -34,6 +42,7 @@ os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
```
#### 2. In CLI
You can also set the environment variable on the command line:
```bash
@@ -41,6 +50,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
```
(optimizations-flash)=
### Flash Attention
**`FLASH_ATTN`**
@@ -57,7 +67,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
cd hopper
pip install ninja
pip install ninja
python setup.py install
```
@@ -66,7 +76,9 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
:::
(optimizations-sta)=
### Sliding Tile Attention
**`SLIDING_TILE_ATTN`**
```bash
@@ -76,7 +88,9 @@ pip install st_attn==0.0.4
Please see [this page](#sta-installation) for more installation instructions.
(optimizations-vsa)=
### Video Sparse Attention
**`VIDEO_SPARSE_ATTN`**
```bash
@@ -87,19 +101,45 @@ python setup_vsa.py install
Please see [this page](#vsa-installation) for more installation instructions.
(optimizations-sage)=
### Sage Attention
**`SAGE_ATTN`**
To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please compile from source:
```bash
git clone https://github.com/thu-ml/SageAttention.git
cd sageattention
cd sageattention
python setup.py install # or pip install -e .
```
(optimizations-sage3)=
### Sage Attention 3
**`SAGE_ATTN_THREE`**
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
#### Hardware Requirements
- RTX5090
#### Installation
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
```bash
python setup.py install
```
(optimizations-teacache)=
## Teacache
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
### What is TeaCache?
@@ -0,0 +1,182 @@
#!/bin/bash
#SBATCH --job-name=1.3B_moe_4n_sf_distill_no_2nd_update
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=1.3B_moe_4n_sf_distill_output/moe_sf_distill_no_2nd_update.out
#SBATCH --error=1.3B_moe_4n_sf_distill_output/moe_sf_distill_no_2nd_update.err
#SBATCH --exclusive
set -e -x
# Environment Setup
# source ~/conda/miniconda/bin/activate
# conda activate wei-fv
# export HOME="/mnt/weka/home/hao.zhang/wei"
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation:
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_2.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/wei/1.3B_MoE_SFwan_t2v_finetune"
--wandb_run_name "1.3B_moe_sf_distill"
# --use_sf_wan
# --sf_ode_init_path "checkpoints/ode_init.pt"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 8
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--fake_score_learning_rate 4e-7
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
--init_weights_from_safetensors_2 "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/FastVideo2/vidprom_8b16k_1e-5_gn1/checkpoint-3500/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/sf2/checkpoint_v3_7k/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 4.0
--real_score_guidance_scale_2 3.0
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,140 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29503
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY=your_wandb_api_key
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=1
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_data_dir
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
--output_dir your_output_dir
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
parallel_args=(
--num_gpus $NUM_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
--init_weights_from_safetensors your_ode_init_weights_path
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,180 @@
#!/bin/bash
#SBATCH --job-name=sf_distill_profile
#SBATCH --partition=main
#SBATCH --nodes=2
#SBATCH --ntasks=2
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=sf_distill_output/sf_distill_profile.out
#SBATCH --error=sf_distill_output/sf_distill_profile.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate wei-fv
export HOME="/mnt/weka/home/hao.zhang/wei"
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29503
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
export FASTVIDEO_TORCH_PROFILER_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/traces
export FASTVIDEO_TORCH_PROFILE_REGIONS=profiler_region_model_loading
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=16
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_2.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan_t2v_finetune"
--wandb_run_name "sf_distill_2e-6_4e-7_4n"
# --use_sf_wan
# --sf_ode_init_path "checkpoints/ode_init.pt"
--max_train_steps 2
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 8
)
# Validation arguments
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--fake_score_learning_rate 4e-7
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/FastVideo2/vidprom_8b16k_1e-5_gn1/checkpoint-3500/transformer/diffusion_pytorch_model.safetensors"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/sf2/checkpoint_v3_7k/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,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,24 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
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 81 \
--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,180 @@
#!/bin/bash
#SBATCH --job-name=14B_moe_4n_sf_distill_no_2nd_update
#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=14B_moe_4n_sf_distill_output/moe_sf_distill_no_2nd_update.out
#SBATCH --error=14B_moe_4n_sf_distill_output/moe_sf_distill_no_2nd_update.err
#SBATCH --exclusive
set -e -x
# Environment Setup
# source ~/conda/miniconda/bin/activate
# conda activate wei-fv
# export HOME="/mnt/weka/home/hao.zhang/wei"
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation:
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_2.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/wei/14B_MoE_SFwan_t2v_finetune"
--wandb_run_name "14B_moe_sf_distill"
# --use_sf_wan
# --sf_ode_init_path "checkpoints/ode_init.pt"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 8
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--fake_score_learning_rate 4e-7
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 12
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
--init_weights_from_safetensors "/mnt/sharefs/users/hao.zhang/wl/release/self_forcing_ode_init_wan22_high_bz128_1e-5/diffusers_2000/"
--init_weights_from_safetensors_2 "/mnt/sharefs/users/hao.zhang/wl/release/self_forcing_ode_init_wan22_low_bz128_1e-5/diffusers_2000/"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 4.0
--real_score_guidance_scale_2 3.0
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,165 @@
#!/bin/bash
#SBATCH --job-name=moe_4n_sf_distill
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=moe_4n_sf_distill_output/moe_sf_distill_%j.out
#SBATCH --error=moe_4n_sf_distill_output/moe_sf_distill_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation with Wan2.2:
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers"
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_dmd"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--enable_gradient_checkpointing_type "full"
--simulate_generator_forward
--log_visualization
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
--init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
--init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
# --init_weights_from_safetensors /mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors
)
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 40
--validation_sampling_steps "8"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 12
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
)
dmd_args=(
--dmd_denoising_steps '1000,850,700,550,350,275,200,125'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 4.0
--real_score_guidance_scale_2 3.0
--fake_score_learning_rate 4e-7
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,181 @@
#!/bin/bash
#SBATCH --job-name=14B_moe_4n_sf_distill_no_2nd_update
#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=14B_moe_4n_sf_distill_output/moe_sf_distill_no_2nd_update.out
#SBATCH --error=14B_moe_4n_sf_distill_output/moe_sf_distill_no_2nd_update.err
#SBATCH --exclusive
set -e -x
# Environment Setup
# source ~/conda/miniconda/bin/activate
# conda activate wei-fv
# export HOME="/mnt/weka/home/hao.zhang/wei"
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=32
# Model paths for Self-Forcing DMD distillation:
# GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
GENERATOR_MODEL_PATH="rand0nmr/SFWan2.2-T2V-A14B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_2.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/wei/14B_MoE_SFwan_t2v_finetune"
--wandb_run_name "14B_moe_sf_distill_low_dmd"
# --use_sf_wan
# --sf_ode_init_path "checkpoints/ode_init.pt"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 8
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--fake_score_learning_rate 4e-7
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 10.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 12
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 200
--init_weights_from_safetensors "/mnt/sharefs/users/hao.zhang/wl/release/self_forcing_ode_init_wan22_low_bz128_1e-5/diffusers_2000/"
--init_weights_from_safetensors_2 "/mnt/sharefs/users/hao.zhang/wl/release/self_forcing_ode_init_wan22_low_bz128_1e-5/diffusers_2000/"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '350,275,200,125'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 4.0
--real_score_guidance_scale_2 3.0
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir"checkpoints/wan_t2v_finetune"
--output_dir $OUTPUT_DIR
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -72,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -91,7 +96,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -134,4 +139,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir "checkpoints/wan_t2v_finetune"
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -72,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -91,7 +96,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -134,4 +139,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd
--output_dir "checkpoints/wan_t2v_finetune"
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -72,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -91,7 +96,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -133,4 +138,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -1,3 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
@@ -0,0 +1,144 @@
#!/bin/bash
#SBATCH --job-name=moe_dmd_distill
#SBATCH --partition=main
#SBATCH --nodes=2
#SBATCH --ntasks=2
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=moe_dmd_distill_output/moe_dmd_distill_%j.out
#SBATCH --error=moe_dmd_distill_output/moe_dmd_distill_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=16
# Model paths for Self-Forcing DMD distillation with Wan2.2:
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan2.2_t2v_distill_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_dmd"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--enable_gradient_checkpointing_type "full"
# --log_visualization
--simulate_generator_forward
--num_frames 81
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
)
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--ema_start_step 100
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--real_score_guidance_scale 3.0
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -40,15 +40,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name Wan_distillation
--output_dir "your_output_dir"
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -73,6 +76,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -92,11 +97,11 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--learning_rate 4e-6
--lr_scheduler "cosine_with_min_lr"
--min_lr_ratio 0.5
--lr_warmup_steps 100
--fake_score_learning_rate 1e-5
--fake_score_learning_rate 2e-6
--fake_score_lr_scheduler "cosine_with_min_lr"
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
@@ -141,4 +146,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -40,6 +40,8 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
# export CUDA_VISIBLE_DEVICES=4,5
@@ -73,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -92,11 +96,11 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--learning_rate 4e-6
--lr_scheduler "cosine_with_min_lr"
--min_lr_ratio 0.5
--lr_warmup_steps 100
--fake_score_learning_rate 1e-5
--fake_score_learning_rate 2e-6
--fake_score_lr_scheduler "cosine_with_min_lr"
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
@@ -142,4 +146,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -14,26 +14,29 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# Configs
NUM_GPUS=1
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir="checkpoints/wan_t2v_finetune"
--max_train_steps=4000
--train_batch_size=1
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--gradient_accumulation_steps 1
--num_latent_t 31
--num_height 704
--num_width 1280
--num_frames 121
--enable_gradient_checkpointing_type "full"
--training_state_checkpointing_steps=500
--weight_only_checkpointing_steps=500
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
)
# Parallel arguments
@@ -49,6 +52,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -68,8 +73,8 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-5
--mixed_precision="bf16"
--learning_rate 2e-6
--mixed_precision "bf16"
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -107,4 +112,4 @@ torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -14,6 +14,8 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# Configs
NUM_GPUS=1
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
@@ -51,6 +53,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -109,4 +113,4 @@ torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -1,3 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
--preprocess_task "t2v"
@@ -62,7 +62,8 @@ validation_args=(
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -0,0 +1,136 @@
#!/bin/bash
#SBATCH --job-name=2e6B8_16kFV_ode_vidprom
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=ode_vidprom16k/ode_vidprom8b16k_2e-6.out
#SBATCH --error=ode_vidprom16k/ode_vidprom8b16k_2e-6.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate your-conda-env
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_API_KEY=your-wandb-api-key
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
DATA_DIR="your-data-dir"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/causal_ode_init/validation.json"
OUTPUT_DIR="your-output-dir"
INIT_WEIGHTS_FROM_SAFETENSORS="your-init-weights-from-safetensors" # bidirectional weights from Wan2.1-T2V-1.3B-Diffusers
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir $OUTPUT_DIR
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "vidprom_8b16k_ode_init_2e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
--warp_denoising_step
--log_visualization
--max_train_steps 6001
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
--dmd_denoising_steps "1000,750,500,250"
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim $NUM_GPUS
--hsdp_shard_dim 1
)
# 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
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
--init_weights_from_safetensors $INIT_WEIGHTS_FROM_SAFETENSORS
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 500
--training_state_checkpointing_steps 500
--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"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
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/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,14 +1,5 @@
{
"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": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
"image_path": null,
@@ -28,7 +19,52 @@
"num_frames": 77
},
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"caption": "Elon Musk, dressed in a sleek white spacesuit with a reflective visor, walks confidently across the lunar surface. His posture is upright, and he moves steadily with purpose. The moon's rocky terrain and scattered boulders surround him, casting shadows under the dim sunlight. The background shows vast stretches of the moon's barren landscape with craters and dust clouds kicked up by his boots. The scene captures a wide shot, emphasizing the vastness and desolation of the lunar environment. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "In a dynamic action-packed sequence set in the Marvel multiverse, Spider-Man and Venom engage in an intense battle. Spider-Man, in his classic red and blue suit, swings and dodges venomous attacks from the black symbiote-covered Venom. Both characters display a range of acrobatic moves and powerful strikes. The environment is a chaotic urban landscape with crumbling buildings and neon lights, reflecting the multiversal theme. The camera captures the epic fight from various angles, including wide shots to show the scale of destruction and close-ups to highlight their fierce expressions and physical combat. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A warm, family-oriented scene depicting a father getting ready to leave the house to buy milk. The father, a middle-aged man with a kind face and a casual outfit, picks up a jacket from the coat rack. His posture is upright as he bends down slightly to put on his shoes. In the background, there are glimpses of a cozy living room with a family photograph on the wall. The camera focuses closely on the father, capturing his gentle smile and reassuring nod towards the camera before he opens the front door and steps outside. Static medium close-up shot. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Close-up shot of a man with a prosthetic hand that functions as a rocket launcher. He looks at his new hand with a mix of amazement and concern, his facial expression showing a blend of curiosity and apprehension. The prosthetic hand is sleek and metallic, with intricate details that resemble a high-tech weapon. The background is a dimly lit laboratory with various scientific equipment and monitors displaying data. The man stands in a relaxed posture, his other hand resting on his hip, as he inspects his new limb. The scene is rendered in a realistic sci-fi style, emphasizing the futuristic technology and the man's emotional response to his new appendage. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Realistic CCTV footage style, Kim Taehyung from the band BTS is involved in a drug deal, caught on camera. Kim Taehyung appears nervous and cautious, wearing casual clothing typical of a public space. He exchanges items discreetly with another person, who is partially obscured. Both individuals maintain a watchful demeanor, occasionally glancing around to ensure no one is watching them. The lighting is dim, with flickering fluorescent lights casting shadows on their faces. The background shows a typical urban setting with blurred figures moving in the distance. Static camera angle, medium close-up shot focusing on the interaction between Taehyung and the other individual. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Photorealistic studio setup with professional lighting, showcasing detailed cubic dissections of experimental plastic and felt-like materials on a pristine white background. Each cube reveals intricate layers and textures of the materials, emphasizing their unique properties. The scene has a shallow depth of field initially, then slowly pulls out to reveal the full arrangement of cubes, maintaining a wide depth of field throughout the transition. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--checkpointing_steps 2000
--weight_only_checkpointing_steps 2000
--training_state_checkpointing_steps 2000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -95,7 +95,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -93,7 +93,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -0,0 +1,134 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=8
#SBATCH --ntasks=8
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=VSA_t2v_output/t2v_%j.out
#SBATCH --error=VSA_t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate your_env
# 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=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
# --enable_gradient_checkpointing_type "full" # if OOM enable this
)
# Parallel arguments
parallel_args=(
--num_gpus 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 64
--hsdp_shard_dim 1
)
# 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 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "5.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 1
--seed 1000
)
# VSA arguments
vsa_args=(
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 50 \
--VSA_sparsity 0.9 \
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
@@ -93,7 +93,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -0,0 +1,131 @@
#!/bin/bash
#SBATCH --job-name=moe_finetune
#SBATCH --partition=main
#SBATCH --nodes=2
#SBATCH --ntasks=2
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=moe_output/moe_%j.out
#SBATCH --error=moe_output/moe_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv
export HOME="/mnt/weka/home/hao.zhang/wei"
# Configs
NUM_GPUS=8
NUM_TOTAL_GPUS=16
# Model paths for Self-Forcing DMD distillation with Wan2.2:
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name VSA_finetune # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/wei/Wan2.2-MoE-finetune"
--max_train_steps 200
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--enable_gradient_checkpointing_type "full"
# --log_visualization
--num_frames 81
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
)
parallel_args=(
--num_gpus $NUM_TOTAL_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_TOTAL_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
# --log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "40"
--validation_guidance_scale "5.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--dit_cpu_offload True
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -95,7 +95,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -92,7 +92,8 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 400
--weight_only_checkpointing_steps 400
--training_state_checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 400
--weight_only_checkpointing_steps 400
--training_state_checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -0,0 +1,72 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class SageAttention3Backend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128, 256]
@staticmethod
def get_name() -> str:
return "SAGE_ATTN_THREE"
@staticmethod
def get_impl_cls() -> type["SageAttention3Impl"]:
return SageAttention3Impl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
# @staticmethod
# def get_metadata_cls() -> Type["AttentionMetadata"]:
# return FlashAttentionMetadata
class SageAttention3Impl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
output = output.transpose(1, 2)
return output
+4 -7
View File
@@ -15,13 +15,10 @@ class DiTArchConfig(ArchConfig):
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
)
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.SAGE_ATTN_THREE)
hidden_size: int = 0
num_attention_heads: int = 0
+4 -2
View File
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig)
SelfForcingWanT2V480PConfig, SelfForcingMoEWanT2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -36,6 +36,8 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers": SelfForcingMoEWanT2V480PConfig,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingMoEWanT2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
@@ -137,4 +139,4 @@ def get_pipeline_config_cls_from_name(
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
)
return pipeline_config_cls
return pipeline_config_cls
+15
View File
@@ -54,6 +54,9 @@ class WanT2V480PConfig(PipelineConfig):
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# self-forcing params
warp_denoising_step: bool = True
# WanConfig-specific added parameters
def __post_init__(self):
@@ -133,6 +136,11 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
# self-forcing params
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio
@@ -157,3 +165,10 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
@dataclass
class SelfForcingMoEWanT2V480PConfig(SelfForcingWanT2V480PConfig):
boundary_ratio: float | None = 0.875
def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio
+2
View File
@@ -55,6 +55,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.1-T2V-A1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWanT2V480PConfig,
# Add other specific weight variants
}
+1
View File
@@ -175,6 +175,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# - "SLIDING_TILE_ATTN" : use Sliding Tile Attention
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
# - "SAGE_ATTN": use Sage Attention
# - "SAGE_ATTN_THREE": use Sage Attention 3
"FASTVIDEO_ATTENTION_BACKEND":
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
+135 -6
View File
@@ -133,6 +133,7 @@ class FastVideoArgs:
# Compilation
enable_torch_compile: bool = False
torch_compile_kwargs: dict[str, Any] = field(default_factory=dict)
disable_autocast: bool = False
@@ -158,12 +159,13 @@ class FastVideoArgs:
"transformer": True,
"vae": True,
})
override_transformer_cls_name: str | None = None
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
# MoE parameters used by Wan2.2
boundary_ratio: float | None = None
boundary_ratio: float | None = 0.875
@property
def training_mode(self) -> bool:
@@ -329,6 +331,13 @@ class FastVideoArgs:
help="Use torch.compile to speed up DiT inference." +
"However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)",
)
parser.add_argument(
"--torch-compile-kwargs",
type=str,
default=None,
help=
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
)
parser.add_argument(
"--dit-cpu-offload",
@@ -396,6 +405,12 @@ class FastVideoArgs:
default=FastVideoArgs.enable_stage_verification,
help="Enable input/output verification for pipeline stages",
)
parser.add_argument(
"--override-transformer-cls-name",
type=str,
default=FastVideoArgs.override_transformer_cls_name,
help="Override transformer cls name",
)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -424,6 +439,21 @@ class FastVideoArgs:
mode_value = getattr(args, attr, FastVideoArgs.mode.value)
kwargs['mode'] = ExecutionMode.from_string(
mode_value) if isinstance(mode_value, str) else mode_value
elif attr == 'torch_compile_kwargs':
# Parse JSON string for torch.compile kwargs
torch_compile_kwargs_str = getattr(args, 'torch_compile_kwargs',
None)
if torch_compile_kwargs_str:
try:
import json
kwargs['torch_compile_kwargs'] = json.loads(
torch_compile_kwargs_str)
except json.JSONDecodeError as e:
raise ValueError(
f"Invalid JSON for torch_compile_kwargs: {e}"
) from e
else:
kwargs['torch_compile_kwargs'] = {}
elif attr == 'workload_type':
# Convert string to WorkloadType enum
workload_type_value = getattr(args, 'workload_type',
@@ -603,7 +633,10 @@ class TrainingArgs(FastVideoArgs):
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
# DMD model paths - separate paths for each network
real_score_model_path: str = "" # path for real score (teacher) model
fake_score_model_path: str = "" # path for fake score (critic) model
# diffusion setting
ema_decay: float = 0.0
@@ -625,8 +658,9 @@ class TrainingArgs(FastVideoArgs):
# output
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
# optimizer & scheduler
num_train_epochs: int = 0
@@ -658,6 +692,7 @@ class TrainingArgs(FastVideoArgs):
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
@@ -678,16 +713,29 @@ class TrainingArgs(FastVideoArgs):
# distillation args
generator_update_interval: int = 5
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
real_score_guidance_scale_2: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
independent_first_frame: bool = False
enable_gradient_masking: bool = True
gradient_mask_last_n_frames: int = 21
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -789,6 +837,20 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Directory to cache models")
# DMD model paths - separate paths for each network
parser.add_argument(
"--generator-model-path",
type=str,
help="Path to generator (student) model for DMD distillation")
parser.add_argument(
"--real-score-model-path",
type=str,
help="Path to real score (teacher) model for DMD distillation")
parser.add_argument(
"--fake-score-model-path",
type=str,
help="Path to fake score (critic) model for DMD distillation")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
@@ -844,9 +906,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--checkpoints-total-limit",
type=int,
help="Maximum number of checkpoints to keep")
parser.add_argument("--checkpointing-steps",
type=int,
help="Steps between checkpoints")
parser.add_argument(
"--training-state-checkpointing-steps",
type=int,
@@ -859,6 +918,14 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument(
"--init-weights-from-safetensors",
type=str,
help="Path to safetensors file for initial weight loading")
parser.add_argument(
"--init-weights-from-safetensors-2",
type=str,
help="Path to safetensors file for initial weight loading")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
@@ -963,6 +1030,10 @@ class TrainingArgs(FastVideoArgs):
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--betas",
type=str,
default=TrainingArgs.betas,
help="Betas for optimizer (format: 'beta1,beta2')")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
@@ -1013,6 +1084,13 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument(
"--dfake-gen-update-ratio",
type=int,
default=TrainingArgs.dfake_gen_update_ratio,
help=
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
)
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
@@ -1025,10 +1103,19 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.real_score_guidance_scale,
help="Teacher guidance scale")
parser.add_argument("--real-score-guidance-scale-2",
type=float,
default=TrainingArgs.real_score_guidance_scale_2,
help="Teacher guidance scale")
parser.add_argument("--fake-score-learning-rate",
type=float,
default=TrainingArgs.fake_score_learning_rate,
help="Learning rate for fake score transformer")
parser.add_argument(
"--fake-score-betas",
type=str,
default=TrainingArgs.fake_score_betas,
help="Betas for fake score optimizer (format: 'beta1,beta2')")
parser.add_argument(
"--fake-score-lr-scheduler",
type=str,
@@ -1041,6 +1128,48 @@ class TrainingArgs(FastVideoArgs):
"--simulate-generator-forward",
action=StoreBoolean,
help="Whether to simulate generator forward to match inference")
parser.add_argument(
"--warp-denoising-step",
action=StoreBoolean,
help=
"Whether to warp denoising step according to the scheduler time shift"
)
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=TrainingArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
parser.add_argument(
"--independent-first-frame",
action=StoreBoolean,
help="Whether the first frame is independent in causal generation")
parser.add_argument(
"--enable-gradient-masking",
action=StoreBoolean,
help="Whether to enable frame-level gradient masking")
parser.add_argument(
"--gradient-mask-last-n-frames",
type=int,
default=TrainingArgs.gradient_mask_last_n_frames,
help="Number of last frames to enable gradients for")
parser.add_argument(
"--validate-cache-structure",
action=StoreBoolean,
help="Whether to validate KV cache structure (debug flag)")
parser.add_argument(
"--same-step-across-blocks",
action=StoreBoolean,
help="Whether to use the same exit timestep for all blocks")
parser.add_argument(
"--last-step-only",
action=StoreBoolean,
help="Whether to only use the last timestep for training")
parser.add_argument("--context-noise",
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
return parser
+59 -38
View File
@@ -147,6 +147,9 @@ class CausalWanSelfAttention(nn.Module):
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
@@ -176,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -209,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -223,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -249,29 +250,29 @@ class CausalWanTransformerBlock(nn.Module):
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
# e.shape: [batch_size, num_frames, 6, inner_dim]
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
assert shift_msa.dtype == torch.float32
# assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -285,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -295,13 +294,10 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -364,8 +360,7 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -375,7 +370,8 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.independent_first_frame = False
self.__post_init__()
@@ -487,12 +483,16 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -539,14 +539,9 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -587,8 +582,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
@@ -601,8 +596,12 @@ class CausalWanTransformer3DModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -637,14 +636,9 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def forward(
self,
@@ -655,3 +649,30 @@ class CausalWanTransformer3DModel(BaseDiT):
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
+38 -6
View File
@@ -416,6 +416,11 @@ class TransformerLoader(ComponentLoader):
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
logger.info("transformer cls_name: %s", cls_name)
if fastvideo_args.override_transformer_cls_name is not None:
cls_name = fastvideo_args.override_transformer_cls_name
logger.info("Overriding transformer cls_name to %s", cls_name)
fastvideo_args.model_paths["transformer"] = model_path
# Config from Diffusers supersedes fastvideo's model config
@@ -430,8 +435,33 @@ class TransformerLoader(ComponentLoader):
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
# Check if we should use custom initialization weights
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
fastvideo_args.training_mode and
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
if use_custom_weights:
if 'transformer_2' in model_path:
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors_2', None)
assert custom_weights_path is not None, "Custom initialization weights must be provided"
# Handle both directory and single file cases
if os.path.isdir(custom_weights_path):
# Directory: collect all safetensors files
safetensors_list = glob.glob(
os.path.join(str(custom_weights_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in directory {custom_weights_path}")
elif os.path.isfile(custom_weights_path):
# Single file: verify it's a safetensors file
assert custom_weights_path.endswith(".safetensors"), "Custom initialization weights must be a safetensors file"
safetensors_list = [custom_weights_path]
else:
raise ValueError(f"Custom weights path does not exist or is not accessible: {custom_weights_path}")
logger.info("Loading model from %s safetensors files: %s",
len(safetensors_list), safetensors_list)
default_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
@@ -454,18 +484,20 @@ class TransformerLoader(ComponentLoader):
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
fsdp_inference=fastvideo_args.use_fsdp_inference,
# TODO(will): make these configurable
default_dtype=default_dtype,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
training_mode=fastvideo_args.training_mode)
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
dtypes = set(param.dtype for param in model.parameters())
if len(dtypes) > 1:
model = model.to(default_dtype)
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
model = model.eval()
return model
+15 -3
View File
@@ -54,7 +54,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
torch.set_default_dtype(old_dtype)
# TODO(PY): add compile option
# Supports optional torch.compile for FSDP-wrapped models during training
def maybe_load_fsdp_model(
model_cls: type[nn.Module],
init_params: dict[str, Any],
@@ -62,6 +62,7 @@ def maybe_load_fsdp_model(
device: torch.device,
hsdp_replicate_dim: int,
hsdp_shard_dim: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
@@ -69,6 +70,8 @@ def maybe_load_fsdp_model(
output_dtype: torch.dtype | None = None,
training_mode: bool = True,
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -87,7 +90,8 @@ def maybe_load_fsdp_model(
mp_policy=mp_policy,
)
with set_default_dtype(param_dtype), torch.device("meta"):
logger.info("Loading model with default_dtype: %s", default_dtype)
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
# Check if we should use FSDP
@@ -125,7 +129,7 @@ def maybe_load_fsdp_model(
model,
weight_iterator,
device,
param_dtype,
default_dtype,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
@@ -137,6 +141,14 @@ def maybe_load_fsdp_model(
# Avoid unintended computation graph accumulation during inference
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
compile_in_loader = enable_torch_compile and training_mode
if compile_in_loader:
compile_kwargs = torch_compile_kwargs or {}
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s",
compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
return model
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
noise: torch.Tensor,
timestep: torch.IntTensor,
) -> torch.Tensor:
"""
Args:
clean_latent: the clean latent with shape [B, C, H, W],
where B is batch_size or batch_size * num_frames
noise: the noise with shape [B, C, H, W]
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
Returns:
the corrupted latent with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == clean_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(clean_latent.shape[0])
else:
assert timestep.numel() == clean_latent.shape[0]
else:
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
self.sigmas = self.sigmas.to(noise.device)
timestep = timestep.expand(clean_latent.shape[0])
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
@@ -22,8 +22,10 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
config_name = "scheduler_config.json"
order = 1
@register_to_config
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
+51 -4
View File
@@ -171,12 +171,59 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
noise_input_latent = noise_input_latent.float().to(device)
sigmas = scheduler.sigmas.float().to(device)
timesteps = scheduler.timesteps.float().to(device)
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def pred_noise_to_x_bound(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
boundary_timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
boundary_timestep: the boundary timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_t_boundary = sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - (sigma_t - sigma_t_boundary) * pred_noise
return pred_video.to(dtype)
@@ -7,8 +7,6 @@ This module wires the causal DMD denoising stage into the modular pipeline.
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
@@ -28,10 +26,6 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
@@ -55,6 +49,7 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="denoising_stage",
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
+26 -4
View File
@@ -71,6 +71,7 @@ class ComposedPipelineBase(ABC):
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
fastvideo_args.dit_cpu_offload = False
self.modules = self.load_modules(fastvideo_args, loaded_modules)
def set_trainable(self) -> None:
@@ -99,9 +100,30 @@ class ComposedPipelineBase(ABC):
self.initialize_pipeline(self.fastvideo_args)
if self.fastvideo_args.enable_torch_compile:
self.modules["transformer"] = torch.compile(
self.modules["transformer"])
logger.info("Torch Compile enabled for DiT")
transformer_module = self.modules["transformer"]
if self.fastvideo_args.training_mode:
logger.info(
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
)
else:
fsdp_module_cls = None
try:
from torch.distributed.fsdp import FSDPModule # type: ignore
fsdp_module_cls = FSDPModule
except Exception: # pragma: no cover - FSDP not always available
fsdp_module_cls = None
if fsdp_module_cls is not None and isinstance(
transformer_module, fsdp_module_cls):
logger.info(
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
)
else:
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
logger.info("Enabling torch.compile for DiT with kwargs=%s",
compile_kwargs)
self.modules["transformer"] = torch.compile(
transformer_module, **compile_kwargs)
logger.info("Torch Compile enabled for DiT")
if not self.fastvideo_args.training_mode:
logger.info("Creating pipeline stages...")
@@ -144,7 +166,7 @@ class ComposedPipelineBase(ABC):
for key, value in kwargs.items():
setattr(fastvideo_args, key, value)
fastvideo_args.dit_cpu_offload = False
fastvideo_args.dit_cpu_offload = True
# we hijack the precision to be the master weight type so that the
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
@@ -246,6 +246,7 @@ class TrainingBatch:
fake_score_loss: float = 0.0
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
+71 -51
View File
@@ -1,4 +1,5 @@
import torch # type: ignore
import gc
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -34,13 +35,15 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
super().__init__(transformer, scheduler, transformer_2)
# KV and cross-attention cache state (initialized on first forward)
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
self.transformer = transformer
self.transformer_2 = transformer_2
# self.kv_cache1: list | None = None
# self.crossattn_cache: list | None = None
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
@@ -65,21 +68,18 @@ class CausalDMDDenosingStage(DenoisingStage):
-1] * self.transformer.config.arch_config.patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
# TODO(will): make this a parameter once we add i2v support
independent_first_frame = self.transformer.independent_first_frame
independent_first_frame = self.transformer.independent_first_frame if hasattr(
self.transformer, 'independent_first_frame') else False
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
logger.info("Warping timesteps...")
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
logger.info("Using timesteps: %s", timesteps)
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -104,34 +104,43 @@ class CausalDMDDenosingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
if self.kv_cache1 is None:
self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device)
else:
assert self.crossattn_cache is not None
# reset cross-attention cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index][
"is_init"] = False # type: ignore
# reset kv cache pointers
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index][
"global_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
self.kv_cache1[block_index][
"local_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.text_encoder_configs[0].
arch_config.text_len,
dtype=target_dtype,
device=latents.device)
# if self.kv_cache1 is None:
# self._initialize_kv_cache(batch_size=latents.shape[0],
# dtype=target_dtype,
# device=latents.device)
# self._initialize_crossattn_cache(
# batch_size=latents.shape[0],
# max_text_len=fastvideo_args.pipeline_config.
# text_encoder_configs[0].arch_config.text_len,
# dtype=target_dtype,
# device=latents.device)
# else:
# assert self.crossattn_cache is not None
# # reset cross-attention cache
# for block_index in range(self.num_transformer_blocks):
# self.crossattn_cache[block_index][
# "is_init"] = False # type: ignore
# # reset kv cache pointers
# for block_index in range(len(self.kv_cache1)):
# self.kv_cache1[block_index][
# "global_end_index"] = torch.tensor( # type: ignore
# [0],
# dtype=torch.long,
# device=latents.device)
# self.kv_cache1[block_index][
# "local_end_index"] = torch.tensor( # type: ignore
# [0],
# dtype=torch.long,
# device=latents.device)
# Optional: cache context features from provided image latents prior to generation
current_start_frame = 0
@@ -154,8 +163,8 @@ class CausalDMDDenosingStage(DenoisingStage):
image_first_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
@@ -180,8 +189,8 @@ class CausalDMDDenosingStage(DenoisingStage):
ref_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
@@ -223,6 +232,10 @@ class CausalDMDDenosingStage(DenoisingStage):
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
if self.transformer_2 is not None and fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None and t_cur < fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps:
current_model = self.transformer_2
else:
current_model = self.transformer
# Copy for pred conversion
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
@@ -273,12 +286,12 @@ class CausalDMDDenosingStage(DenoisingStage):
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw = self.transformer(
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -338,12 +351,12 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = self.transformer(
_ = current_model(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -352,6 +365,11 @@ class CausalDMDDenosingStage(DenoisingStage):
)
start_index += current_num_frames
del kv_cache1
del crossattn_cache
gc.collect()
torch.cuda.empty_cache()
batch.latents = latents
return batch
@@ -389,7 +407,8 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0], dtype=torch.long, device=device),
})
self.kv_cache1 = kv_cache1
# self.kv_cache1 = kv_cache1
return kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> None:
@@ -418,7 +437,8 @@ class CausalDMDDenosingStage(DenoisingStage):
"is_init":
False,
})
self.crossattn_cache = crossattn_cache
# self.crossattn_cache = crossattn_cache
return crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
@@ -442,4 +462,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result
+3 -2
View File
@@ -85,8 +85,9 @@ class DenoisingStage(PipelineStage):
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
)
def forward(
+15
View File
@@ -115,6 +115,7 @@ class CudaPlatformBase(Platform):
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
envs.FASTVIDEO_ATTENTION_BACKEND)
logger.info("Selected backend: %s", selected_backend)
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
try:
from st_attn import sliding_tile_attention # noqa: F401
@@ -144,6 +145,20 @@ class CudaPlatformBase(Platform):
logger.info(
"Sage Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
try:
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
SageAttention3Backend)
logger.info("Using Sage Attention 3 backend.")
return "fastvideo.attention.backends.sage_attn3.SageAttention3Backend"
except ImportError as e:
logger.info(e)
logger.info(
"Sage Attention 3 backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
try:
from vsa import block_sparse_attn # noqa: F401
+1
View File
@@ -18,6 +18,7 @@ class AttentionBackendEnum(enum.Enum):
SLIDING_TILE_ATTN = enum.auto()
TORCH_SDPA = enum.auto()
SAGE_ATTN = enum.auto()
SAGE_ATTN_THREE = enum.auto()
VIDEO_SPARSE_ATTN = enum.auto()
VMOBA_ATTN = enum.auto()
NO_ATTENTION = enum.auto()
+5 -1
View File
@@ -62,7 +62,7 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=900)
@app.function(gpu="H100:1", image=image, timeout=900)
def run_encoder_tests():
run_test("pytest ./fastvideo/tests/encoders -vs")
@@ -118,6 +118,10 @@ def run_inference_lora_tests():
def run_distill_dmd_tests():
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_self_forcing_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
@@ -0,0 +1,184 @@
import os
from pathlib import Path
from huggingface_hub import snapshot_download
import subprocess
import sys
from fastvideo.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
import shutil
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
NUM_NODES = "1"
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
# preprocessing
DATA_DIR = "data"
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol"))
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/pipelines/preprocess/v1_preprocess.py"
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol_processed_t2v"))
# training
NUM_GPUS_PER_NODE_TRAINING = "4"
TRAINING_ENTRY_FILE_PATH = "fastvideo/training/wan_distillation_pipeline.py"
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
LOCAL_VALIDATION_DATASET_FILE = "examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir, exist_ok=True)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
result = snapshot_download(
repo_id="wlsaidhi/crush-smol-merged",
local_dir=str(LOCAL_RAW_DATA_DIR),
repo_type="dataset",
resume_download=True,
token=os.environ.get("HF_TOKEN"), # In case authentication is needed
)
print(f"Download completed successfully. Files downloaded to: {result}")
# Verify the download
if not LOCAL_RAW_DATA_DIR.exists():
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
# List downloaded files
print("Downloaded files:")
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
if file.is_file():
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
except Exception as e:
print(f"Error during download: {str(e)}")
raise
def run_preprocessing():
# remove the local_preprocessed_data_dir if it exists
if LOCAL_PREPROCESSED_DATA_DIR.exists():
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
PREPROCESSING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--seed", "42",
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge.txt"),
"--preprocess_video_batch_size", "1",
"--max_height", "480",
"--max_width", "832",
"--num_frames", "81",
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
"--preprocess_task", "t2v",
]
process = subprocess.run(cmd, check=True)
def run_training():
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
TRAINING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--sp_size", "1",
"--tp_size", "1",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "501",
"--learning_rate", "2e-6",
"--fake_score_learning_rate", "2e-6",
"--mixed_precision", "bf16",
"--training_state_checkpointing_steps", "1000",
"--weight_only_checkpointing_steps", "1000",
"--validation_steps", "50",
"--validation_sampling_steps", "3",
"--log_validation",
"--checkpoints_total_limit", "3",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "ci_wan_t2v_dmd_overfit",
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--flow_shift", "8",
"--validation_guidance_scale", "6.0",
"--weight_decay", "0.01",
"--generator_update_interval", "5",
"--dmd_denoising_steps", "1000,757,522",
"--min_timestep_ratio", "0.02",
"--max_timestep_ratio", "0.98",
"--seed", "1000",
"--real_score_guidance_scale", "3.5",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
"--enable_gradient_checkpointing_type", "full",
]
print(f"Running training with command: {cmd}")
process = subprocess.run(cmd, check=True)
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
print(f"reference_video_file: {reference_video_file}")
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
print(f"final_validation_video_file: {final_validation_video_file}")
# Ensure both files exist
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
# Compute SSIM
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
reference_video_file,
final_validation_video_file,
use_ms_ssim=True # Using MS-SSIM for better quality assessment
)
print("\n===== SSIM Results for Step 900 Validation =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
if __name__ == "__main__":
test_e2e_overfit_single_sample()
@@ -111,7 +111,8 @@ def run_training():
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -62,6 +62,11 @@ def download_data():
def run_preprocessing():
# remove the local_preprocessed_data_dir if it exists
if LOCAL_PREPROCESSED_DATA_DIR.exists():
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
# Run torchrun command
cmd = [
"torchrun",
@@ -111,7 +116,8 @@ def run_training():
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -1 +1 @@
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":1.260593056678772,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.2620866410434246,"_runtime":107.325113071}
@@ -46,7 +46,8 @@ def run_worker():
"--max_train_steps", "5",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "30",
"--weight_only_checkpointing_steps", "30",
"--training_state_checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -110,7 +111,7 @@ def test_distributed_training():
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 1.0,
'train_loss': 0.001
'train_loss': 0.005
}
failures = []
@@ -1 +1 @@
{"train_loss":0.10545400530099869,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.53125,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"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.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"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.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","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."],"_type":"videos","count":3}}
{"train_loss":0.1021774671971798,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.11582941561937332,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"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.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"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.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","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."],"_type":"videos","count":3}}
@@ -50,7 +50,8 @@ def run_worker():
"--max_train_steps", "5",
"--learning_rate", "1e-6",
"--mixed_precision", "bf16",
"--checkpointing_steps", "30",
"--weight_only_checkpointing_steps", "30",
"--training_state_checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "8",
"--log_validation",
@@ -33,6 +33,8 @@ def run_worker():
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--real_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--fake_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--data_path", "data/crush-smol_processed_t2v/combined_parquet_dataset",
"--validation_dataset_file", "examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json",
"--train_batch_size", "1",
@@ -55,7 +55,8 @@ def test_lora_training():
"--max_train_steps", "5",
"--learning_rate", "5e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
"--validation_steps", "50",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -93,10 +94,10 @@ def test_lora_training():
# Define thresholds for LoRA training based on the provided console outputs
fields_and_thresholds = {
'avg_step_time': 2.0,
'avg_step_time': 20.0, # something up with modal
# 'grad_norm': 0.05, # too volatile for now. TODO: fix nondeterminism in training
'step_time': 2.0,
'train_loss': 0.03
'step_time': 20.0, # something up with modal
'train_loss': 0.05
}
failures = []
@@ -0,0 +1,149 @@
import os
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29513"
import sys
import subprocess
from pathlib import Path
import torch
import json
from huggingface_hub import snapshot_download
from fastvideo.utils import logger
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
from fastvideo.training.wan_self_forcing_distillation_pipeline import WanSelfForcingDistillationPipeline
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
wandb_name = "test_self_forcing_distill"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
def run_worker():
"""Worker function that will be run on each GPU"""
# Create and populate args
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
# Set the arguments based on the distill_dmd_t2v_1.3B.sh script
args = parser.parse_args([
"--model_path", "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
"--real_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--fake_score_model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--data_path", "data/crush-smol_processed_t2v/combined_parquet_dataset",
"--validation_dataset_file", "examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json",
"--train_batch_size", "1",
"--num_latent_t", "21",
"--num_gpus", "2",
"--sp_size", "1",
"--tp_size", "1",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "2",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "1",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "2",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--training_state_checkpointing_steps", "30",
"--weight_only_checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "3",
"--log_validation",
"--checkpoints_total_limit", "3",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--output_dir", "data/wan_self_forcing_test",
"--tracker_project_name", "wan_self_forcing_ci",
"--wandb_run_name", wandb_name,
"--num_height", "480",
"--num_width", "832",
"--num_frames", "21",
"--flow_shift", "5",
"--validation_guidance_scale", "1.0",
"--weight_decay", "0.01",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
# DMD args
"--dmd_denoising_steps", "1000,750,500", # Reduced steps for testing
"--min_timestep_ratio", "0.02",
"--max_timestep_ratio", "0.98",
"--dfake_gen_update_ratio", "5",
"--real_score_guidance_scale", "3.0",
"--fake_score_learning_rate", "8e-6",
"--fake_score_betas", "0.0,0.999",
"--warp_denoising_step",
"--enable_gradient_checkpointing_type", "full",
# Self-forcing specific args
"--log_visualization",
"--simulate_generator_forward",
"--num_frame_per_block", "3",
"--enable_gradient_masking",
"--gradient_mask_last_n_frames", "21",
"--independent_first_frame", "False",
"--same_step_across_blocks", "True",
"--last_step_only", "False",
"--context_noise", "0",
"--use_ema", "True",
"--ema_decay", "0.99",
"--ema_start_step", "100",
])
# Call the main training function
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Self-forcing distillation training pipeline done")
def test_distributed_training():
"""Test the distributed self-forcing training setup"""
os.environ["WANDB_MODE"] = "offline"
data_dir = Path("data/crush-smol_processed_t2v")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="wlsaidhi/crush-smol_processed_t2v",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
)
# Get the current file path
current_file = Path(__file__).resolve()
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
"--master_port", os.environ["MASTER_PORT"],
str(current_file)
]
process = subprocess.run(cmd, capture_output=True, text=True)
# Print stdout and stderr for debugging
if process.stdout:
print("STDOUT:", process.stdout)
if process.stderr:
print("STDERR:", process.stderr)
# Check if the process failed
if process.returncode != 0:
print(f"Process failed with return code: {process.returncode}")
raise subprocess.CalledProcessError(process.returncode, cmd, process.stdout, process.stderr)
if __name__ == "__main__":
if os.environ.get("LOCAL_RANK") is not None:
# We're being run by torchrun
run_worker()
else:
# We're being run directly
test_distributed_training()
File diff suppressed because it is too large Load Diff
+406
View File
@@ -0,0 +1,406 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import cast
import numpy as np
import torch
import torch.nn.functional as F
import wandb
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.distributed import get_local_torch_device
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_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
logger = init_logger(__name__)
class ODEInitTrainingPipeline(TrainingPipeline):
"""
Training pipeline for ODE-init using precomputed denoising trajectories.
Supervision: predict the next latent in the stored trajectory by
- feeding current latent at timestep t into the transformer to predict noise
- stepping the scheduler with the predicted noise
- minimizing MSE to the stored next latent at timestep t_next
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Match the preprocess/generation scheduler for consistent stepping
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
training=True)
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
def initialize_training_pipeline(self, training_args: TrainingArgs):
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
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32))).cuda()
logger.info("timesteps: %s", timesteps)
self.dmd_denoising_steps = timesteps[1000 -
self.dmd_denoising_steps]
logger.info("warped self.dmd_denoising_steps: %s",
self.dmd_denoising_steps)
else:
raise ValueError("warp_denoising_step must be true")
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
get_local_torch_device())
logger.info("denoising_step_list: %s", self.dmd_denoising_steps)
logger.info(
"Initialized ODE-init training pipeline with %s denoising steps",
len(self.dmd_denoising_steps))
# Cache for nearest trajectory index per DMD step (computed lazily on first batch)
self._cached_closest_idx_per_dmd = None
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.manual_idx = 0
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
# Warm start validation with current transformer
self.validation_pipeline = WanCausalDMDPipeline.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)
def _get_next_batch(
self,
training_batch) -> tuple[TrainingBatch, torch.Tensor, torch.Tensor]:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
# Trajectory tensors may include a leading singleton batch dim per row
trajectory_latents = batch['trajectory_latents']
if trajectory_latents.dim() == 7:
# [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
trajectory_latents = trajectory_latents[:, 0]
elif trajectory_latents.dim() == 6:
# already [B, S, C, T, H, W]
pass
else:
raise ValueError(
f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
)
trajectory_timesteps = batch['trajectory_timesteps']
if trajectory_timesteps.dim() == 3:
# [B, 1, S] -> [B, S]
trajectory_timesteps = trajectory_timesteps[:, 0]
elif trajectory_timesteps.dim() == 2:
# [B, S]
pass
else:
raise ValueError(
f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
)
# [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
# Move to device
device = get_local_torch_device()
training_batch.encoder_hidden_states = encoder_hidden_states.to(
device, dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
device, dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
## TEMP Used for loading the sf .pt files directly
"""
self.manual_idx = self.manual_idx % 155
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
logger.info("path: %s", path)
self.manual_idx += 1
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
b = torch.load(path)
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
0).to(device, dtype=torch.bfloat16)
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
logger.info("trajectory_latents: %s", trajectory_latents.shape)
logger.info("encoder_hidden_states: %s",
training_batch.encoder_hidden_states.shape)
assert trajectory_latents.shape[1] <= 10, "trajectory_latents.shape[1] must be <= 10"
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
"""
def _get_timestep(self,
min_timestep: int,
max_timestep: int,
batch_size: int,
num_frame: int,
num_frame_per_block: int,
uniform_timestep: bool = False) -> torch.Tensor:
if uniform_timestep:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, 1],
device=self.device,
dtype=torch.long).repeat(1, num_frame)
return timestep
else:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, num_frame],
device=self.device,
dtype=torch.long)
# logger.info(f"individual timestep: {timestep}")
# make the noise level the same within every block
timestep = timestep.reshape(timestep.shape[0], -1,
num_frame_per_block)
timestep[:, :, 1:] = timestep[:, :, 0:1]
timestep = timestep.reshape(timestep.shape[0], -1)
return timestep
def _step_predict_next_latent(
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
torch.Tensor]]:
latent_vis_dict: dict[str, torch.Tensor] = {}
device = get_local_torch_device()
target_latent = traj_latents[:, -1]
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
B, S, num_frames, num_channels, height, width = traj_latents.shape
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36], dtype=torch.long).cpu()
# [0, 1, 2, 3], dtype=torch.long).cpu()
logger.info("self._cached_closest_idx_per_dmd: %s",
self._cached_closest_idx_per_dmd)
logger.info(
"corresponding timesteps: %s", self.noise_scheduler.timesteps[
self._cached_closest_idx_per_dmd])
# Select the K indexes from traj_latents using self._cached_closest_idx_per_dmd
# traj_latents: [B, S, C, T, H, W], self._cached_closest_idx_per_dmd: [K]
# Output: [B, K, C, T, H, W]
assert self._cached_closest_idx_per_dmd is not None
relevant_traj_latents = torch.index_select(
traj_latents,
dim=1,
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
logger.info("relevant_traj_latents: %s", relevant_traj_latents.shape)
# assert relevant_traj_latents.shape[0] == 1
indexes = self._get_timestep( # [B, num_frames]
0,
len(self.dmd_denoising_steps),
B,
num_frames,
3,
uniform_timestep=False)
logger.info("indexes: %s", indexes.shape)
logger.info("indexes: %s", indexes)
# noisy_input = relevant_traj_latents[indexes]
noisy_input = torch.gather(
relevant_traj_latents,
dim=1,
index=indexes.reshape(B, 1, num_frames, 1, 1,
1).expand(-1, -1, -1, num_channels, height,
width).to(self.device)).squeeze(1)
timestep = self.dmd_denoising_steps[indexes]
logger.info("selected timestep for rank %s: %s",
self.global_rank,
timestep,
local_main_process_only=False)
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
4).detach().clone().cpu()
input_kwargs = {
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep.to(device, dtype=torch.bfloat16),
"return_dict": False,
}
# Predict noise and step the scheduler to obtain next latent
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
from fastvideo.models.utils import pred_noise_to_pred_video
pred_video = pred_noise_to_pred_video(
pred_noise=noise_pred.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=timestep.to(dtype=torch.bfloat16).flatten(0, 1),
scheduler=self.modules["scheduler"]).unflatten(
0, noise_pred.shape[:2])
latent_vis_dict["pred_video"] = pred_video.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
return pred_video, target_latent, timestep, latent_vis_dict
def train_one_step(self, training_batch): # type: ignore[override]
self.transformer.train()
self.optimizer.zero_grad()
training_batch.total_loss = 0.0
args = cast(TrainingArgs, self.training_args)
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
for _ in range(args.gradient_accumulation_steps):
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
training_batch)
text_embeds = training_batch.encoder_hidden_states
text_attention_mask = training_batch.encoder_attention_mask
assert traj_latents.shape[0] == 1
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
_, S = traj_latents.shape[0], traj_latents.shape[1]
if S < 2:
raise ValueError("Trajectory must contain at least 2 steps")
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
training_batch.latent_vis_dict.update(latent_vis_dict)
mask = t != 0
# Compute loss
loss = F.mse_loss(noise_pred[mask],
target_latent[mask],
reduction="mean")
loss = loss / args.gradient_accumulation_steps
with set_forward_context(current_timestep=t,
attn_metadata=None,
forward_batch=None):
loss.backward()
avg_loss = loss.detach().clone()
training_batch.total_loss += avg_loss.item()
# Clip grad and step optimizers
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in self.transformer.parameters() if p.requires_grad],
args.max_grad_norm if args.max_grad_norm is not None else 0.0)
self.optimizer.step()
self.lr_scheduler.step()
if grad_norm is None:
grad_value = 0.0
else:
try:
if isinstance(grad_norm, torch.Tensor):
grad_value = float(grad_norm.detach().float().item())
else:
grad_value = float(grad_norm)
except Exception:
grad_value = 0.0
training_batch.grad_norm = grad_value
return training_batch
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 = {}
latents_vis_dict = training_batch.latent_vis_dict
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
for latent_key in latent_log_keys:
assert latent_key in latents_vis_dict and latents_vis_dict[
latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video, fps=16, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, pixel_latent, latent
# Log to wandb
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
def main(args) -> None:
logger.info("Starting ODE-init training pipeline...")
pipeline = ODEInitTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("ODE-init training pipeline done")
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()
args.dit_cpu_offload = False
main(args)
File diff suppressed because it is too large Load Diff
+173 -22
View File
@@ -63,6 +63,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
train_transformer_2: bool = False
def __init__(
self,
@@ -98,6 +99,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
self.transformer_2 = self.get_module("transformer_2", None)
self.seed = training_args.seed
self.set_schemas()
@@ -105,22 +107,31 @@ class TrainingPipeline(LoRAPipeline, ABC):
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
self.transformer.train()
self.transformer.requires_grad_(True)
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
# Set grads for proper modules based on the training mode (Distill, LoRA, etc.)
self.set_trainable()
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
# Parse betas from string format "beta1,beta2"
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -138,6 +149,30 @@ class TrainingPipeline(LoRAPipeline, ABC):
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
if self.transformer_2 is not None:
# Ensure transformer_2 has trainable parameters before creating optimizer
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
@@ -152,6 +187,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
seed=self.seed)
self.noise_scheduler = noise_scheduler
if self.training_args.boundary_ratio is not None:
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
else:
self.boundary_timestep = None
logger.info("train_dataloader length: %s", len(self.train_dataloader))
logger.info("train_sp_batch_size: %s",
training_args.train_sp_batch_size)
logger.info("gradient_accumulation_steps: %s",
training_args.gradient_accumulation_steps)
logger.info("sp_size: %s", training_args.sp_size)
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
@@ -178,9 +224,28 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
self.transformer.train()
self.optimizer.zero_grad()
if self.transformer_2 is not None:
self.transformer_2.train()
self.optimizer_2.zero_grad()
training_batch.total_loss = 0.0
return training_batch
def _enable_training(self, model: torch.nn.Module,
optimizer: torch.optim.Optimizer) -> None:
"""Enable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = True
model.train()
# optimizer.zero_grad()
def _disable_training(self, model: torch.nn.Module,
optimizer: torch.optim.Optimizer) -> None:
"""Disable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = False
model.eval()
# optimizer.zero_grad(set_to_none=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
@@ -224,17 +289,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=latents.device)
timesteps = self._sample_timesteps(batch_size, latents.device)
# Enable training for the model that will be trained next and disable the other
if self.train_transformer_2:
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self._enable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
@@ -257,6 +322,59 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _sample_timesteps(self, batch_size: int,
device: torch.device) -> torch.Tensor:
# Determine which model to train based on the boundary timestep
if (self.transformer_2 is not None
and self.boundary_timestep is not None
and torch.rand(1, generator=self.noise_random_generator).item()
<= self.training_args.boundary_ratio):
self.train_transformer_2 = True
else:
self.train_transformer_2 = False
# Broadcast the decision to all processes
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0,
device=self.device)
dist.broadcast(decision, src=0)
self.train_transformer_2 = decision.item() == 1.0
# Sample u from the appropriate range
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
boundary_ratio = self.training_args.boundary_ratio
if self.train_transformer_2:
u = (1 - boundary_ratio
) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
# elif self.transformer_2 is not None:
# u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
# else: # patch for now to align with non-MoE timestep logic
# pass
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
# timestep = self.noise_scheduler.timesteps[indices].to(device=device)
# if timestep < self.training_args.boundary_ratio * self.noise_scheduler.config.num_train_timesteps:
# self.train_transformer_2 = True
# else:
# self.train_transformer_2 = False
# # Broadcast the decision to all processes
# decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0,
# device=self.device)
# dist.broadcast(decision, src=0)
# dist.broadcast(timestep, src=0)
# self.train_transformer_2 = decision.item() == 1.0
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
latents_shape = training_batch.raw_latent_shape
@@ -321,11 +439,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
# [1000.0],
# device=training_batch.noisy_model_input.device,
# dtype=torch.bfloat16)
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
with set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
model_pred = current_model(**input_kwargs)
if self.training_args.precondition_outputs:
assert training_batch.sigmas is not None
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
@@ -356,7 +475,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
model_parts = [self.transformer]
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
@@ -401,8 +525,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch = self._clip_grad_norm(training_batch)
self.optimizer.step()
self.lr_scheduler.step()
# Only step the optimizer and scheduler for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
@@ -435,6 +564,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
if getattr(self, "transformer_2", None) is not None:
num_trainable_params = count_trainable(self.transformer_2)
logger.info(
"Transformer 2: Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
@@ -477,7 +612,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
elif vmoba_available:
# TODO: add vmoba sparsity scheduling here
#TODO: add vmoba sparsity scheduling here
current_vsa_sparsity = 0.0
else:
current_vsa_sparsity = 0.0
@@ -512,7 +647,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
},
step=step,
)
if step % self.training_args.checkpointing_steps == 0:
if step % self.training_args.training_state_checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
@@ -520,6 +655,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer.train()
self.sp_group.barrier()
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)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
@@ -606,7 +745,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
Generate a validation video and log it to wandb to check the quality during training.
"""
training_args.inference_mode = True
training_args.dit_cpu_offload = True
training_args.dit_cpu_offload = False
if not training_args.log_validation:
return
if self.validation_pipeline is None:
@@ -627,7 +766,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
transformer.eval()
self.transformer.eval()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
@@ -719,4 +861,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
self.transformer.train()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.train()
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
raise NotImplementedError(
"Visualize intermediate latents is not implemented for training pipeline"
)
+526 -27
View File
@@ -191,26 +191,48 @@ def save_checkpoint(transformer,
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
def save_distillation_checkpoint(generator_transformer,
fake_score_transformer,
rank,
output_dir,
step,
generator_optimizer=None,
fake_score_optimizer=None,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
only_save_generator_weight=False) -> None:
def save_distillation_checkpoint(
generator_transformer,
fake_score_transformer,
rank,
output_dir,
step,
generator_optimizer=None,
fake_score_optimizer=None,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None,
only_save_generator_weight=False,
# MoE support
generator_transformer_2=None,
real_score_transformer_2=None,
fake_score_transformer_2=None,
generator_optimizer_2=None,
fake_score_optimizer_2=None,
generator_scheduler_2=None,
fake_score_scheduler_2=None,
generator_ema_2=None) -> None:
"""
Save distillation checkpoint with both generator and fake_score models.
Supports MoE (Mixture of Experts) models with transformer_2 variants.
Saves both distributed checkpoint and consolidated model weights.
Only saves the generator model for inference (consolidated weights).
Args:
generator_transformer: Main generator transformer model
fake_score_transformer: Main fake score transformer model
only_save_generator_weight: If True, only save the generator model weights for inference
without saving distributed checkpoint for training resume.
generator_transformer_2: Secondary generator transformer for MoE (optional)
real_score_transformer_2: Secondary real score transformer for MoE (optional)
fake_score_transformer_2: Secondary fake score transformer for MoE (optional)
generator_optimizer_2: Optimizer for generator_transformer_2 (optional)
fake_score_optimizer_2: Optimizer for fake_score_transformer_2 (optional)
generator_scheduler_2: Scheduler for generator_transformer_2 (optional)
fake_score_scheduler_2: Scheduler for fake_score_transformer_2 (optional)
generator_ema_2: EMA for generator_transformer_2 (optional)
"""
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
@@ -233,6 +255,8 @@ def save_distillation_checkpoint(generator_transformer,
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -251,6 +275,41 @@ def save_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Save generator_2 distributed checkpoint (MoE support)
if generator_transformer_2 is not None:
generator_2_states = {
"model": ModelWrapper(generator_transformer_2),
}
if generator_optimizer_2 is not None:
generator_2_states["optimizer"] = OptimizerWrapper(
generator_transformer_2, generator_optimizer_2)
if dataloader is not None:
generator_2_states["dataloader"] = dataloader
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
if generator_ema_2 is not None:
generator_2_states["ema"] = generator_ema_2.state_dict()
generator_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
"generator_2")
logger.info(
"rank: %s, saving generator_2 distributed checkpoint to %s",
rank,
generator_2_dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.save(generator_2_states, checkpoint_id=generator_2_dcp_dir)
end_time = time.perf_counter()
logger.info(
"rank: %s, generator_2 distributed checkpoint saved in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
# Save critic distributed checkpoint
critic_states = {
"model": ModelWrapper(fake_score_transformer),
@@ -280,6 +339,67 @@ def save_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Save critic_2 distributed checkpoint (MoE support)
if fake_score_transformer_2 is not None:
critic_2_states = {
"model": ModelWrapper(fake_score_transformer_2),
}
if fake_score_optimizer_2 is not None:
critic_2_states["optimizer"] = OptimizerWrapper(
fake_score_transformer_2, fake_score_optimizer_2)
if dataloader is not None:
critic_2_states["dataloader"] = dataloader
if fake_score_scheduler_2 is not None:
critic_2_states["scheduler"] = SchedulerWrapper(
fake_score_scheduler_2)
critic_2_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"critic_2")
logger.info(
"rank: %s, saving critic_2 distributed checkpoint to %s",
rank,
critic_2_dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.save(critic_2_states, checkpoint_id=critic_2_dcp_dir)
end_time = time.perf_counter()
logger.info(
"rank: %s, critic_2 distributed checkpoint saved in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
# Save real_score_transformer_2 distributed checkpoint (MoE support)
if real_score_transformer_2 is not None:
real_score_2_states = {
"model": ModelWrapper(real_score_transformer_2),
}
# Note: real_score_transformer_2 typically doesn't have optimizer/scheduler
# since it's used for inference only, but we include dataloader for consistency
if dataloader is not None:
real_score_2_states["dataloader"] = dataloader
real_score_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
"real_score_2")
logger.info(
"rank: %s, saving real_score_2 distributed checkpoint to %s",
rank,
real_score_2_dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.save(real_score_2_states, checkpoint_id=real_score_2_dcp_dir)
end_time = time.perf_counter()
logger.info(
"rank: %s, real_score_2 distributed checkpoint saved in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
# Save shared random state separately
shared_states = {
"random_state": RandomStateWrapper(noise_generator),
@@ -332,6 +452,47 @@ def save_distillation_checkpoint(generator_transformer,
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(
generator_transformer_2, device=None)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2,
"config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
def load_checkpoint(transformer,
rank,
@@ -393,19 +554,43 @@ def load_checkpoint(transformer,
return step
def load_distillation_checkpoint(generator_transformer,
fake_score_transformer,
rank,
checkpoint_path,
generator_optimizer=None,
fake_score_optimizer=None,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None) -> int:
def load_distillation_checkpoint(
generator_transformer,
fake_score_transformer,
rank,
checkpoint_path,
generator_optimizer=None,
fake_score_optimizer=None,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None,
# MoE support
generator_transformer_2=None,
real_score_transformer_2=None,
fake_score_transformer_2=None,
generator_optimizer_2=None,
fake_score_optimizer_2=None,
generator_scheduler_2=None,
fake_score_scheduler_2=None,
generator_ema_2=None) -> int:
"""
Load distillation checkpoint with both generator and fake_score models.
Supports MoE (Mixture of Experts) models with transformer_2 variants.
Returns the step number from which training should resume.
Args:
generator_transformer: Main generator transformer model
fake_score_transformer: Main fake score transformer model
generator_transformer_2: Secondary generator transformer for MoE (optional)
real_score_transformer_2: Secondary real score transformer for MoE (optional)
fake_score_transformer_2: Secondary fake score transformer for MoE (optional)
generator_optimizer_2: Optimizer for generator_transformer_2 (optional)
fake_score_optimizer_2: Optimizer for fake_score_transformer_2 (optional)
generator_scheduler_2: Scheduler for generator_transformer_2 (optional)
fake_score_scheduler_2: Scheduler for fake_score_transformer_2 (optional)
generator_ema_2: EMA for generator_transformer_2 (optional)
"""
if not os.path.exists(checkpoint_path):
logger.warning("Distillation checkpoint path %s does not exist",
@@ -456,6 +641,77 @@ def load_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
if generator_ema is not None:
try:
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully",
rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank,
str(e))
# Load generator_2 distributed checkpoint (MoE support)
if generator_transformer_2 is not None:
generator_2_dcp_dir = os.path.join(checkpoint_path,
"distributed_checkpoint",
"generator_2")
if os.path.exists(generator_2_dcp_dir):
generator_2_states = {
"model": ModelWrapper(generator_transformer_2),
}
if generator_optimizer_2 is not None:
generator_2_states["optimizer"] = OptimizerWrapper(
generator_transformer_2, generator_optimizer_2)
if dataloader is not None:
generator_2_states["dataloader"] = dataloader
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
logger.info(
"rank: %s, loading generator_2 distributed checkpoint from %s",
rank,
generator_2_dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.load(generator_2_states, checkpoint_id=generator_2_dcp_dir)
end_time = time.perf_counter()
logger.info(
"rank: %s, generator_2 distributed checkpoint loaded in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
# Load EMA_2 state if available and generator_ema_2 is provided
if generator_ema_2 is not None:
try:
ema_2_state = generator_2_states.get("ema")
if ema_2_state is not None:
generator_ema_2.load_state_dict(ema_2_state)
logger.info(
"rank: %s, generator_2 EMA state loaded successfully",
rank)
else:
logger.info(
"rank: %s, no EMA_2 state found in checkpoint",
rank)
except Exception as e:
logger.warning("rank: %s, failed to load EMA_2 state: %s",
rank, str(e))
else:
logger.info("rank: %s, generator_2 checkpoint not found, skipping",
rank)
# Load critic distributed checkpoint
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
"critic")
@@ -494,6 +750,77 @@ def load_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Load critic_2 distributed checkpoint (MoE support)
if fake_score_transformer_2 is not None:
critic_2_dcp_dir = os.path.join(checkpoint_path,
"distributed_checkpoint", "critic_2")
if os.path.exists(critic_2_dcp_dir):
critic_2_states = {
"model": ModelWrapper(fake_score_transformer_2),
}
if fake_score_optimizer_2 is not None:
critic_2_states["optimizer"] = OptimizerWrapper(
fake_score_transformer_2, fake_score_optimizer_2)
if dataloader is not None:
critic_2_states["dataloader"] = dataloader
if fake_score_scheduler_2 is not None:
critic_2_states["scheduler"] = SchedulerWrapper(
fake_score_scheduler_2)
logger.info(
"rank: %s, loading critic_2 distributed checkpoint from %s",
rank,
critic_2_dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.load(critic_2_states, checkpoint_id=critic_2_dcp_dir)
end_time = time.perf_counter()
logger.info(
"rank: %s, critic_2 distributed checkpoint loaded in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
else:
logger.info("rank: %s, critic_2 checkpoint not found, skipping",
rank)
# Load real_score_2 distributed checkpoint (MoE support)
if real_score_transformer_2 is not None:
real_score_2_dcp_dir = os.path.join(checkpoint_path,
"distributed_checkpoint",
"real_score_2")
if os.path.exists(real_score_2_dcp_dir):
real_score_2_states = {
"model": ModelWrapper(real_score_transformer_2),
}
if dataloader is not None:
real_score_2_states["dataloader"] = dataloader
logger.info(
"rank: %s, loading real_score_2 distributed checkpoint from %s",
rank,
real_score_2_dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.load(real_score_2_states, checkpoint_id=real_score_2_dcp_dir)
end_time = time.perf_counter()
logger.info(
"rank: %s, real_score_2 distributed checkpoint loaded in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
else:
logger.info("rank: %s, real_score_2 checkpoint not found, skipping",
rank)
# Load shared random state
shared_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
"shared")
@@ -825,14 +1152,14 @@ def custom_to_hf_state_dict(
return new_state_dict
# More generalized version of shift_timestep
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
min_timestep: int = 0, max_timestep: int = 1000) -> torch.Tensor:
if shift == 1:
return timestep
t = timestep / num_train_timestep
t = (timestep - min_timestep) / (max_timestep - min_timestep)
denominator = 1 + (shift - 1) * t
return num_train_timestep * (shift * t / denominator)
return min_timestep + (max_timestep - min_timestep) * (shift * t / denominator)
# coding=utf-8
@@ -1280,5 +1607,177 @@ def get_scheduler(
last_epoch=last_epoch)
def _local_numel(p: torch.Tensor) -> int:
if hasattr(p, "to_local"):
return p.to_local().numel()
return p.numel()
def count_trainable(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
return sum(_local_numel(p) for p in model.parameters() if p.requires_grad)
class EMA_FSDP:
"""
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)
@@ -20,10 +20,7 @@ class WanDistillationPipeline(DistillationPipeline):
A distillation pipeline for Wan that uses a single transformer model.
The main transformer serves as the student model, and copies are made for teacher and critic.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
@@ -45,7 +42,7 @@ class WanDistillationPipeline(DistillationPipeline):
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
loaded_modules={"transformer": self.get_module("transformer"), "transformer_2": self.get_module("transformer_2")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
@@ -29,10 +29,7 @@ class WanI2VDistillationPipeline(DistillationPipeline):
A distillation pipeline for Wan that uses a single transformer model.
The main transformer serves as the student model, and copies are made for teacher and critic.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
@@ -0,0 +1,76 @@
# 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.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.training.self_forcing_distillation_pipeline import (
SelfForcingDistillationPipeline)
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
"""
A self-forcing distillation pipeline for Wan that uses the self-forcing methodology
with DMD for video generation.
"""
_required_config_modules = [
"scheduler",
"transformer",
"vae",
]
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 = WanCausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
"transformer_2": self.get_module("transformer_2")
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
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 self-forcing distillation pipeline...")
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Wan self-forcing 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)
@@ -42,6 +42,7 @@ class WanTrainingPipeline(TrainingPipeline):
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
"transformer_2": self.get_module("transformer_2"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
+2
View File
@@ -12,6 +12,8 @@ export TOKENIZERS_PARALLELISM=false
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_distillation_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--real_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--fake_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
@@ -13,6 +13,8 @@ export TOKENIZERS_PARALLELISM=false
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_distillation_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--real_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--fake_score_model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \