Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5125256d4b | ||
|
|
6bf030dcf5 |
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
-136
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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")
|
||||
|
||||
Binary file not shown.
@@ -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",
|
||||
|
||||
BIN
Binary file not shown.
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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\
|
||||
|
||||
@@ -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 \
|
||||
|
||||
Reference in New Issue
Block a user