Compare commits

..
Author SHA1 Message Date
Yongqi Chen 5125256d4b update 2025-09-20 21:10:28 -04:00
Yongqi Chen 6bf030dcf5 update 2025-09-20 21:10:09 -04:00
76 changed files with 345 additions and 4601 deletions
-12
View File
@@ -104,18 +104,6 @@ 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,10 +110,6 @@ 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"
+7 -4
View File
@@ -12,6 +12,9 @@ exclude: |
scripts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/distill/.*|
fastvideo/distill\.py|
fastvideo/distill_adv\.py|
fastvideo/models/.*|
fastvideo/sample/.*|
fastvideo/train\.py|
@@ -41,10 +44,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:
+1 -1
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/q46BbX6" 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/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
+2 -42
View File
@@ -1,40 +1,32 @@
(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
@@ -42,7 +34,6 @@ os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
```
#### 2. In CLI
You can also set the environment variable on the command line:
```bash
@@ -50,7 +41,6 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
```
(optimizations-flash)=
### Flash Attention
**`FLASH_ATTN`**
@@ -67,7 +57,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
```
@@ -76,9 +66,7 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
:::
(optimizations-sta)=
### Sliding Tile Attention
**`SLIDING_TILE_ATTN`**
```bash
@@ -88,9 +76,7 @@ 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
@@ -101,45 +87,19 @@ 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?
@@ -1,140 +0,0 @@
#!/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[@]}"
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,24 +0,0 @@
#!/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"
@@ -1,157 +0,0 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=8
# Model paths for Self-Forcing DMD distillation with Wan2.2:
# GENERATOR_MODEL_PATH="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/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 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/
)
parallel_args=(
--num_gpus 32 # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 32
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -39,18 +39,15 @@ 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 $OUTPUT_DIR
--output_dir"checkpoints/wan_t2v_finetune"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -75,8 +72,6 @@ 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
@@ -96,7 +91,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -139,4 +134,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -39,18 +39,15 @@ 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 "$OUTPUT_DIR"
--output_dir "checkpoints/wan_t2v_finetune"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -75,8 +72,6 @@ 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
@@ -96,7 +91,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -139,4 +134,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -39,18 +39,15 @@ 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 "$OUTPUT_DIR"
--output_dir "checkpoints/wan_t2v_finetune"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -75,8 +72,6 @@ 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
@@ -96,7 +91,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -138,4 +133,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"
@@ -40,18 +40,15 @@ 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 "$OUTPUT_DIR"
--output_dir "your_output_dir"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -76,8 +73,6 @@ 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
@@ -97,11 +92,11 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--learning_rate 2e-5
--lr_scheduler "cosine_with_min_lr"
--min_lr_ratio 0.5
--lr_warmup_steps 100
--fake_score_learning_rate 2e-6
--fake_score_learning_rate 1e-5
--fake_score_lr_scheduler "cosine_with_min_lr"
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
@@ -146,4 +141,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -40,8 +40,6 @@ 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
@@ -75,8 +73,6 @@ 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
@@ -96,11 +92,11 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--learning_rate 2e-5
--lr_scheduler "cosine_with_min_lr"
--min_lr_ratio 0.5
--lr_warmup_steps 100
--fake_score_learning_rate 2e-6
--fake_score_learning_rate 1e-5
--fake_score_lr_scheduler "cosine_with_min_lr"
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
@@ -146,4 +142,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -14,29 +14,26 @@ 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 "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--output_dir="checkpoints/wan_t2v_finetune"
--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
@@ -52,8 +49,6 @@ 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
@@ -73,8 +68,8 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--mixed_precision "bf16"
--learning_rate=1e-5
--mixed_precision="bf16"
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -112,4 +107,4 @@ torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -14,8 +14,6 @@ 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
@@ -53,8 +51,6 @@ 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
@@ -113,4 +109,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,8 +62,7 @@ validation_args=(
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -1,136 +0,0 @@
#!/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,5 +1,14 @@
{
"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,
@@ -19,52 +28,7 @@
"num_frames": 77
},
{
"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. ",
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
@@ -61,8 +61,7 @@ validation_args=(
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 2000
--training_state_checkpointing_steps 2000
--checkpointing_steps 2000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -95,8 +95,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -93,8 +93,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -91,10 +91,9 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -91,10 +91,9 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -61,8 +61,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -95,8 +95,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,8 +61,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,8 +61,7 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -92,8 +92,7 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 400
--training_state_checkpointing_steps 400
--checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,8 +61,7 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 400
--training_state_checkpointing_steps 400
--checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -1,72 +0,0 @@
# 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
+7 -4
View File
@@ -15,10 +15,13 @@ 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.SAGE_ATTN_THREE)
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
)
hidden_size: int = 0
num_attention_heads: int = 0
-8
View File
@@ -54,9 +54,6 @@ 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):
@@ -136,11 +133,6 @@ 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
-1
View File
@@ -175,7 +175,6 @@ 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),
+6 -130
View File
@@ -133,7 +133,6 @@ class FastVideoArgs:
# Compilation
enable_torch_compile: bool = False
torch_compile_kwargs: dict[str, Any] = field(default_factory=dict)
disable_autocast: bool = False
@@ -159,15 +158,12 @@ class FastVideoArgs:
"transformer": True,
"vae": True,
})
override_transformer_cls_name: str | None = None
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
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
# MoE parameters used by Wan2.2
boundary_ratio: float | None = 0.875
boundary_ratio: float | None = None
@property
def training_mode(self) -> bool:
@@ -333,13 +329,6 @@ 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",
@@ -407,20 +396,6 @@ 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",
)
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")
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -449,21 +424,6 @@ 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',
@@ -643,10 +603,7 @@ class TrainingArgs(FastVideoArgs):
# text encoder & vae & diffusion model
pretrained_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
dit_model_name_or_path: str = ""
# diffusion setting
ema_decay: float = 0.0
@@ -668,6 +625,7 @@ 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
# optimizer & scheduler
@@ -700,7 +658,6 @@ 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
@@ -721,28 +678,16 @@ 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
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":
@@ -844,20 +789,6 @@ 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,
@@ -913,6 +844,9 @@ 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,
@@ -1029,10 +963,6 @@ 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")
@@ -1083,13 +1013,6 @@ 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,
@@ -1106,11 +1029,6 @@ class TrainingArgs(FastVideoArgs):
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,
@@ -1123,48 +1041,6 @@ 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
+38 -59
View File
@@ -147,9 +147,6 @@ 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(
@@ -179,7 +176,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -212,7 +209,8 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -225,7 +223,8 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -250,29 +249,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
e = self.scale_shift_table + temb.float()
# 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).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
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)
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.forward_native(query)
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -286,6 +285,8 @@ 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,
@@ -294,10 +295,13 @@ 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
@@ -360,7 +364,8 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -370,8 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.num_frame_per_block = 1
self.independent_first_frame = False
self.__post_init__()
@@ -483,16 +487,12 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) 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,9 +539,14 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
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)
return torch.stack(output)
return output
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -582,8 +587,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
@@ -596,12 +601,8 @@ 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)
@@ -636,9 +637,14 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
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)
return torch.stack(output)
return output
def forward(
self,
@@ -649,30 +655,3 @@ 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
+6 -29
View File
@@ -416,11 +416,6 @@ 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
@@ -435,24 +430,8 @@ class TransformerLoader(ComponentLoader):
if not safetensors_list:
raise ValueError(f"No safetensors files found in {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
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"
if os.path.isdir(custom_weights_path):
safetensors_list = glob.glob(
os.path.join(str(custom_weights_path), "*.safetensors"))
else:
assert custom_weights_path.endswith(".safetensors"), "Custom initialization weights must be a safetensors file"
safetensors_list = [custom_weights_path]
logger.info("Loading model from %s safetensors files: %s",
len(safetensors_list), safetensors_list)
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
default_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
@@ -475,20 +454,18 @@ 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,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs)
training_mode=fastvideo_args.training_mode)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
dtypes = set(param.dtype for param in model.parameters())
if len(dtypes) > 1:
model = model.to(default_dtype)
model = model.eval()
return model
+3 -15
View File
@@ -54,7 +54,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
torch.set_default_dtype(old_dtype)
# Supports optional torch.compile for FSDP-wrapped models during training
# TODO(PY): add compile option
def maybe_load_fsdp_model(
model_cls: type[nn.Module],
init_params: dict[str, Any],
@@ -62,7 +62,6 @@ 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,
@@ -70,8 +69,6 @@ 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.
@@ -90,8 +87,7 @@ def maybe_load_fsdp_model(
mp_policy=mp_policy,
)
logger.info("Loading model with default_dtype: %s", default_dtype)
with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(param_dtype), torch.device("meta"):
model = model_cls(**init_params)
# Check if we should use FSDP
@@ -129,7 +125,7 @@ def maybe_load_fsdp_model(
model,
weight_iterator,
device,
default_dtype,
param_dtype,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
@@ -141,14 +137,6 @@ 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,31 +635,8 @@ 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,10 +22,8 @@ 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
+4 -4
View File
@@ -171,10 +171,10 @@ 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.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(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)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
@@ -7,6 +7,8 @@ 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
@@ -26,6 +28,10 @@ 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."""
@@ -49,7 +55,6 @@ 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",
+3 -24
View File
@@ -99,30 +99,9 @@ class ComposedPipelineBase(ABC):
self.initialize_pipeline(self.fastvideo_args)
if self.fastvideo_args.enable_torch_compile:
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")
self.modules["transformer"] = torch.compile(
self.modules["transformer"])
logger.info("Torch Compile enabled for DiT")
if not self.fastvideo_args.training_mode:
logger.info("Creating pipeline stages...")
@@ -246,7 +246,6 @@ 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)
+10 -13
View File
@@ -34,15 +34,13 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
super().__init__(transformer, scheduler, transformer_2)
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
# KV and cross-attention cache state (initialized on first forward)
self.transformer = transformer
self.transformer_2 = transformer_2
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
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
@@ -67,18 +65,21 @@ 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 if hasattr(
self.transformer, 'independent_first_frame') else False
independent_first_frame = self.transformer.independent_first_frame
# 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 = {}
@@ -222,10 +223,6 @@ 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)
@@ -276,7 +273,7 @@ class CausalDMDDenosingStage(DenoisingStage):
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw = current_model(
pred_noise_btchw = self.transformer(
latent_model_input,
prompt_embeds,
t_expanded_noise,
@@ -341,7 +338,7 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = current_model(
_ = self.transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
+2 -3
View File
@@ -85,9 +85,8 @@ class DenoisingStage(PipelineStage):
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
)
def forward(
-15
View File
@@ -115,7 +115,6 @@ 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
@@ -145,20 +144,6 @@ 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,7 +18,6 @@ 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()
+1 -5
View File
@@ -62,7 +62,7 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="H100:1", image=image, timeout=900)
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_encoder_tests():
run_test("pytest ./fastvideo/tests/encoders -vs")
@@ -118,10 +118,6 @@ 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")
@@ -1,184 +0,0 @@
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,8 +111,7 @@ def run_training():
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -62,11 +62,6 @@ 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",
@@ -116,8 +111,7 @@ def run_training():
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -1 +1 @@
{"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}
{"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}
@@ -46,8 +46,7 @@ def run_worker():
"--max_train_steps", "5",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "30",
"--training_state_checkpointing_steps", "30",
"--checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -111,7 +110,7 @@ def test_distributed_training():
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 1.0,
'train_loss': 0.005
'train_loss': 0.001
}
failures = []
@@ -1 +1 @@
{"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}}
{"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}}
@@ -50,8 +50,7 @@ def run_worker():
"--max_train_steps", "5",
"--learning_rate", "1e-6",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "30",
"--training_state_checkpointing_steps", "30",
"--checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "8",
"--log_validation",
@@ -33,8 +33,6 @@ 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,8 +55,7 @@ def test_lora_training():
"--max_train_steps", "5",
"--learning_rate", "5e-5",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
"--checkpointing_steps", "6000",
"--validation_steps", "50",
"--validation_sampling_steps", "50",
"--log_validation",
@@ -94,10 +93,10 @@ def test_lora_training():
# Define thresholds for LoRA training based on the provided console outputs
fields_and_thresholds = {
'avg_step_time': 20.0, # something up with modal
'avg_step_time': 2.0,
# 'grad_norm': 0.05, # too volatile for now. TODO: fix nondeterminism in training
'step_time': 20.0, # something up with modal
'train_loss': 0.05
'step_time': 2.0,
'train_loss': 0.03
}
failures = []
@@ -1,149 +0,0 @@
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
@@ -1,406 +0,0 @@
# 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
+22 -157
View File
@@ -63,7 +63,6 @@ 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,
@@ -99,7 +98,6 @@ 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()
@@ -112,25 +110,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
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=betas,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -148,30 +138,6 @@ 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,
@@ -186,17 +152,6 @@ 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) /
@@ -223,27 +178,9 @@ 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
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:
@@ -287,17 +224,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
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)
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=latents.device)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
@@ -320,45 +257,6 @@ 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()
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
@@ -423,12 +321,11 @@ 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 = current_model(**input_kwargs)
model_pred = self.transformer(**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
@@ -459,12 +356,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
# 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]
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,
@@ -509,13 +401,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch = self._clip_grad_norm(training_batch)
# 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()
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
@@ -548,12 +435,6 @@ 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)
@@ -596,7 +477,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
@@ -631,7 +512,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
},
step=step,
)
if step % self.training_args.training_state_checkpointing_steps == 0:
if step % self.training_args.checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
@@ -639,10 +520,6 @@ 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(
@@ -729,7 +606,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 = False
training_args.dit_cpu_offload = True
if not training_args.log_validation:
return
if self.validation_pipeline is None:
@@ -750,10 +627,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
self.transformer.eval()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.eval()
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
@@ -845,13 +719,4 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
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"
)
transformer.train()
+23 -522
View File
@@ -191,48 +191,26 @@ def save_checkpoint(transformer,
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
def save_distillation_checkpoint(
generator_transformer,
fake_score_transformer,
rank,
output_dir,
step,
generator_optimizer=None,
fake_score_optimizer=None,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None,
only_save_generator_weight=False,
# 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:
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:
"""
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)
@@ -255,8 +233,6 @@ def save_distillation_checkpoint(
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")
@@ -275,41 +251,6 @@ def save_distillation_checkpoint(
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),
@@ -339,67 +280,6 @@ def save_distillation_checkpoint(
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),
@@ -452,47 +332,6 @@ def save_distillation_checkpoint(
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,
@@ -554,43 +393,19 @@ def load_checkpoint(transformer,
return step
def load_distillation_checkpoint(
generator_transformer,
fake_score_transformer,
rank,
checkpoint_path,
generator_optimizer=None,
fake_score_optimizer=None,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None,
# 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:
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:
"""
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",
@@ -641,77 +456,6 @@ def load_distillation_checkpoint(
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")
@@ -750,77 +494,6 @@ def load_distillation_checkpoint(
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")
@@ -1607,177 +1280,5 @@ 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(_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)
return sum(p.numel() for p in model.parameters() if p.requires_grad)
@@ -20,7 +20,10 @@ 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"]
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
@@ -29,7 +29,10 @@ 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"]
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
@@ -1,76 +0,0 @@
# 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)
+1 -3
View File
@@ -12,8 +12,6 @@ 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" \
@@ -30,7 +28,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 2e-6 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 400 \
--validation_steps 100 \
+1 -3
View File
@@ -13,8 +13,6 @@ 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" \
@@ -31,7 +29,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 2e-6 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 400 \
--validation_steps 100 \
+1 -1
View File
@@ -28,7 +28,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--dataloader_num_workers 10\
--gradient_accumulation_steps=1 \
--max_train_steps=5000 \
--learning_rate=1e-6\
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=6000 \
--validation_steps 200\
+1 -1
View File
@@ -34,7 +34,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--dataloader_num_workers 4 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 1e-6 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 6000 \
--validation_steps 100 \