Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eddc7efc64 |
@@ -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"
|
||||
|
||||
@@ -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?
|
||||
|
||||
@@ -36,8 +36,9 @@ VALIDATION_DATASET_FILE=your_validation_data_dir
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir your_output_dir
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
@@ -46,15 +47,16 @@ training_args=(
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frames 81
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 1
|
||||
@@ -63,18 +65,22 @@ parallel_args=(
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
@@ -83,6 +89,7 @@ validation_args=(
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
@@ -93,6 +100,7 @@ optimizer_args=(
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
@@ -106,6 +114,7 @@ miscellaneous_args=(
|
||||
--init_weights_from_safetensors your_ode_init_weights_path
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
@@ -117,11 +126,13 @@ dmd_args=(
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
|
||||
@@ -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"
|
||||
@@ -9,9 +9,9 @@ def main():
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
num_gpus=4,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
@@ -25,9 +25,7 @@ def main():
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
@@ -35,11 +33,7 @@ def main():
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
"The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
|
||||
+13
-41
@@ -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,8 +603,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
generator_model_path: str = "" # path for generator (student) model
|
||||
real_score_model_path: str = "" # path for real score (teacher) model
|
||||
fake_score_model_path: str = "" # path for fake score (critic) model
|
||||
|
||||
@@ -668,7 +630,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
# output
|
||||
output_dir: str = ""
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
@@ -740,6 +704,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
independent_first_frame: bool = False
|
||||
enable_gradient_masking: bool = True
|
||||
gradient_mask_last_n_frames: int = 21
|
||||
validate_cache_structure: bool = False # Debug flag for cache validation
|
||||
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
|
||||
@@ -913,6 +878,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,
|
||||
@@ -925,6 +893,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--resume-from-checkpoint",
|
||||
type=str,
|
||||
help="Path to checkpoint to resume from")
|
||||
parser.add_argument(
|
||||
"--init-weights-from-safetensors",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
parser.add_argument("--logging-dir",
|
||||
type=str,
|
||||
help="Directory for logging")
|
||||
|
||||
@@ -212,9 +212,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
modulated = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1.0 + scale) + shift).flatten(1, 2)
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
modulated = normalized * (1 + scale) + shift
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
@@ -267,13 +267,13 @@ class LayerNormScaleShift(nn.Module):
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
output = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1.0 + scale) + shift).flatten(1, 2)
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
# scale.shape: [batch_size, 1, inner_dim]
|
||||
# shift.shape: [batch_size, 1, inner_dim]
|
||||
output = normalized * (1.0 + scale) + shift
|
||||
output = normalized * (1 + scale) + shift
|
||||
|
||||
if self.compute_dtype == torch.float32:
|
||||
output = output.to(x.dtype)
|
||||
|
||||
return output
|
||||
return output
|
||||
@@ -262,9 +262,14 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
# assert shift_msa.dtype == torch.float32
|
||||
|
||||
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
|
||||
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
|
||||
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
|
||||
|
||||
# 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)
|
||||
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
@@ -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 = 3
|
||||
self.independent_first_frame = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
@@ -37,16 +39,14 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
dtype = encoder_hidden_states_image.dtype
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states).to(dtype)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
@@ -156,12 +156,12 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
|
||||
if crossattn_cache is not None:
|
||||
if not crossattn_cache["is_init"]:
|
||||
crossattn_cache["is_init"] = True
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
crossattn_cache["k"] = k
|
||||
crossattn_cache["v"] = v
|
||||
@@ -169,7 +169,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
k = crossattn_cache["k"]
|
||||
v = crossattn_cache["v"]
|
||||
else:
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
@@ -213,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
|
||||
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
|
||||
b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
@@ -247,7 +247,7 @@ class WanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -278,29 +278,29 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -319,12 +319,11 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
|
||||
if temb.dim() == 4:
|
||||
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb.float()
|
||||
self.scale_shift_table.unsqueeze(0) + temb
|
||||
).chunk(6, dim=2)
|
||||
# batch_size, seq_len, 1, inner_dim
|
||||
shift_msa = shift_msa.squeeze(2)
|
||||
@@ -335,22 +334,20 @@ class WanTransformerBlock(nn.Module):
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
|
||||
e = self.scale_shift_table + temb.float()
|
||||
e = self.scale_shift_table + temb
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -370,26 +367,20 @@ class WanTransformerBlock(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,
|
||||
context=encoder_hidden_states,
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
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
|
||||
|
||||
|
||||
class WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -406,7 +397,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -438,8 +429,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -459,8 +449,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -480,23 +469,22 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
e = self.scale_shift_table + temb
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
norm_hidden_states = (self.norm1(hidden_states) *
|
||||
(1 + scale_msa) + shift_msa)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -521,8 +509,6 @@ class WanTransformerBlock_VSA(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,
|
||||
@@ -530,17 +516,15 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
context_lens=None)
|
||||
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
|
||||
|
||||
|
||||
|
||||
class WanTransformer3DModel(CachableDiT):
|
||||
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = WanVideoConfig()._compile_conditions
|
||||
@@ -598,8 +582,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -659,10 +642,12 @@ class WanTransformer3DModel(CachableDiT):
|
||||
rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
|
||||
@@ -672,6 +657,8 @@ class WanTransformer3DModel(CachableDiT):
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
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, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
@@ -728,14 +715,35 @@ class WanTransformer3DModel(CachableDiT):
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output
|
||||
return torch.stack(output)
|
||||
|
||||
def 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
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
@@ -827,5 +835,4 @@ class WanTransformer3DModel(CachableDiT):
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
|
||||
return hidden_states + self.previous_residual_odd
|
||||
@@ -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
|
||||
@@ -438,21 +433,15 @@ class TransformerLoader(ComponentLoader):
|
||||
# Check if we should use custom initialization weights
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
|
||||
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
|
||||
fastvideo_args.training_mode and
|
||||
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
|
||||
|
||||
if use_custom_weights:
|
||||
if 'transformer_2' in model_path:
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors_2', None)
|
||||
assert custom_weights_path is not None, "Custom initialization weights must be provided"
|
||||
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("Using custom initialization weights from: %s", custom_weights_path)
|
||||
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]
|
||||
@@ -479,9 +468,7 @@ class TransformerLoader(ComponentLoader):
|
||||
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())
|
||||
|
||||
@@ -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],
|
||||
@@ -70,8 +70,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.
|
||||
@@ -141,14 +139,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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -49,7 +49,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")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -116,8 +116,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.
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_ori_causal_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = CausalWanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
|
||||
new_state_dict = {}
|
||||
for k, v in causal_state_dict.items():
|
||||
if k.startswith("model."):
|
||||
new_state_dict[k.replace("model.", "")] = v
|
||||
causal_state_dict = new_state_dict
|
||||
model1.load_state_dict(causal_state_dict)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
12,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
block_sizes = [3 for _ in range(4)]
|
||||
timesteps = [1000, 750, 500, 250]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
output1 = _causal_inference(model1, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
|
||||
logger.info("Finish inference for model1")
|
||||
output2 = _causal_inference(model2, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
logger.info("Output 1 Sum: %s", output1.float().sum().item())
|
||||
logger.info("Output 2 Sum: %s", output2.float().sum().item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
def _causal_inference(transformer, latents, prompt_embeds, block_sizes, timesteps, target_dtype):
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
start_index = 0
|
||||
pos_start_base = 0
|
||||
frame_seq_length = latents.shape[-1] * latents.shape[-2] // (WanVideoConfig().arch_config.patch_size[-1] * WanVideoConfig().arch_config.patch_size[-2])
|
||||
seq_len = frame_seq_length * latents.shape[2]
|
||||
kv_cache1 = _initialize_kv_cache(transformer, batch_size=latents.shape[0],
|
||||
kv_cache_size=frame_seq_length * latents.shape[2],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
crossattn_cache = _initialize_crossattn_cache(
|
||||
transformer,
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=WanVideoConfig().arch_config.text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
for current_num_frames, t_cur in zip(block_sizes, timesteps):
|
||||
# logger.info(f"Current frame idx: {start_index}, Current timestep: {t_cur}")
|
||||
# logger.info(f"k cache sum: {sum(kv_cache['k'].float().sum().item() for kv_cache in kv_cache1)}, v cache sum: {sum(kv_cache['v'].float().sum().item() for kv_cache in kv_cache1)}")
|
||||
# logger.info(f"latents sum: {latents.float().sum().item()}, encoder_hidden_states sum: {prompt_embeds.float().sum().item()}")
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
|
||||
attn_metadata = None
|
||||
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(current_latents.shape[0], 1),
|
||||
device=current_latents.device,
|
||||
dtype=torch.long)
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
pred_noise_btchw = transformer(
|
||||
x=current_latents,
|
||||
context=prompt_embeds,
|
||||
t=t_expanded_noise,
|
||||
seq_len=seq_len,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length
|
||||
)
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
pred_noise_btchw = transformer(
|
||||
current_latents,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length,
|
||||
start_frame=start_index
|
||||
)
|
||||
|
||||
# Write back and advance
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = pred_noise_btchw.clone()
|
||||
|
||||
# Re-run with context timestep to update KV cache using clean context
|
||||
context_noise = 0
|
||||
t_context = torch.ones([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = pred_noise_btchw.to(target_dtype)
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
_ = transformer(
|
||||
x=context_bcthw,
|
||||
context=prompt_embeds,
|
||||
t=t_expanded_context,
|
||||
seq_len=seq_len,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length
|
||||
)
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
_ = transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=kv_cache1,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
frame_seq_length,
|
||||
start_frame=start_index
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
return latents
|
||||
|
||||
def _initialize_kv_cache(transformer, batch_size, kv_cache_size, dtype, device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
num_attention_heads = transformer.num_heads
|
||||
attention_head_dim = transformer.dim // transformer.num_heads
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
num_attention_heads = transformer.num_attention_heads
|
||||
attention_head_dim = transformer.attention_head_dim
|
||||
|
||||
for _ in range(len(transformer.blocks)):
|
||||
kv_cache1.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
return kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(transformer, batch_size, max_text_len, dtype,
|
||||
device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
crossattn_cache = []
|
||||
if isinstance(transformer, CausalWanModel):
|
||||
num_attention_heads = transformer.num_heads
|
||||
attention_head_dim = transformer.dim // transformer.num_heads
|
||||
elif isinstance(transformer, CausalWanTransformer3DModel):
|
||||
num_attention_heads = transformer.num_attention_heads
|
||||
attention_head_dim = transformer.attention_head_dim
|
||||
|
||||
for _ in range(len(transformer.blocks)):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
return crossattn_cache
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.model import WanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_ori_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = WanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
seq_len = math.ceil((160 * 90) /
|
||||
(2 * 2) *
|
||||
21)
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
# with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
x=hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
t=timestep,
|
||||
seq_len=seq_len,
|
||||
)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -0,0 +1,144 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from fastvideo.wan.modules.causal_model import CausalWanModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_train_ori_causal_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = CausalWanModel.from_pretrained(
|
||||
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
|
||||
new_state_dict = {}
|
||||
for k, v in causal_state_dict.items():
|
||||
if k.startswith("model."):
|
||||
new_state_dict[k.replace("model.", "")] = v
|
||||
causal_state_dict = new_state_dict
|
||||
model1.load_state_dict(causal_state_dict)
|
||||
|
||||
model1.num_frame_per_block = 3
|
||||
model2.num_frame_per_block = 3
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
text_seq_len = 30
|
||||
seq_len = math.ceil((160 * 90) /
|
||||
(2 * 2) *
|
||||
21)
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
text_seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.randint(0, 1000, (batch_size, 21), device=device, dtype=torch.long)
|
||||
logger.info("timestep: %s", timestep)
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
data_type="dummy",
|
||||
)
|
||||
|
||||
# with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
x=hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
t=timestep,
|
||||
seq_len=seq_len,
|
||||
)
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -12,6 +12,7 @@ from typing import Any
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
@@ -56,10 +57,13 @@ class DistillationPipeline(TrainingPipeline):
|
||||
Inherits from TrainingPipeline to reuse training infrastructure.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"transformer",
|
||||
"vae",
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
]
|
||||
_extra_config_module_map = {
|
||||
"real_score_transformer": "transformer",
|
||||
"fake_score_transformer": "transformer"
|
||||
}
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[dict[str, Any]]
|
||||
@@ -68,7 +72,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
current_trainstep: int
|
||||
video_latent_shape: tuple[int, ...]
|
||||
video_latent_shape_sp: tuple[int, ...]
|
||||
train_fake_score_transformer_2: bool = False
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
@@ -88,90 +91,40 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=self.timestep_shift)
|
||||
|
||||
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
|
||||
|
||||
if training_args.real_score_model_path:
|
||||
logger.info("Loading real score transformer from: %s",
|
||||
training_args.real_score_model_path)
|
||||
training_args.override_transformer_cls_name = "WanTransformer3DModel"
|
||||
self.real_score_transformer = self.load_module_from_path(
|
||||
training_args.real_score_model_path, "transformer",
|
||||
training_args)
|
||||
try:
|
||||
self.real_score_transformer_2 = self.load_module_from_path(
|
||||
training_args.real_score_model_path, "transformer_2",
|
||||
training_args)
|
||||
logger.info("Loaded real score transformer_2 for MoE support")
|
||||
except Exception:
|
||||
logger.info(
|
||||
"real score transformer_2 not found, using single transformer"
|
||||
)
|
||||
self.real_score_transformer_2 = None
|
||||
else:
|
||||
self.real_score_transformer = self.get_module(
|
||||
"real_score_transformer")
|
||||
self.real_score_transformer_2 = self.get_module(
|
||||
"real_score_transformer_2")
|
||||
|
||||
if training_args.fake_score_model_path:
|
||||
logger.info("Loading fake score transformer from: %s",
|
||||
training_args.fake_score_model_path)
|
||||
training_args.override_transformer_cls_name = "WanTransformer3DModel"
|
||||
self.fake_score_transformer = self.load_module_from_path(
|
||||
training_args.fake_score_model_path, "transformer",
|
||||
training_args)
|
||||
try:
|
||||
self.fake_score_transformer_2 = self.load_module_from_path(
|
||||
training_args.fake_score_model_path, "transformer_2",
|
||||
training_args)
|
||||
logger.info("Loaded fake score transformer_2 for MoE support")
|
||||
except Exception:
|
||||
logger.info(
|
||||
"fake score transformer_2 not found, using single transformer"
|
||||
)
|
||||
self.fake_score_transformer_2 = None
|
||||
else:
|
||||
self.fake_score_transformer = self.get_module(
|
||||
"fake_score_transformer")
|
||||
self.fake_score_transformer_2 = self.get_module(
|
||||
"fake_score_transformer_2")
|
||||
|
||||
self.real_score_transformer.requires_grad_(False)
|
||||
self.real_score_transformer.eval()
|
||||
if self.real_score_transformer_2 is not None:
|
||||
self.real_score_transformer_2.requires_grad_(False)
|
||||
self.real_score_transformer_2.eval()
|
||||
|
||||
# Set training modes for fake score transformers (trainable)
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
self.fake_score_transformer_2.requires_grad_(True)
|
||||
self.fake_score_transformer_2.train()
|
||||
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
self.fake_score_transformer = apply_activation_checkpointing(
|
||||
self.fake_score_transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
self.fake_score_transformer_2 = apply_activation_checkpointing(
|
||||
self.fake_score_transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
self.real_score_transformer = apply_activation_checkpointing(
|
||||
self.real_score_transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.real_score_transformer_2 is not None:
|
||||
self.real_score_transformer_2 = apply_activation_checkpointing(
|
||||
self.real_score_transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
# Initialize optimizers
|
||||
fake_score_params = list(
|
||||
@@ -205,28 +158,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
fake_score_params_2 = list(
|
||||
filter(lambda p: p.requires_grad,
|
||||
self.fake_score_transformer_2.parameters()))
|
||||
self.fake_score_optimizer_2 = torch.optim.AdamW(
|
||||
fake_score_params_2,
|
||||
lr=fake_score_lr,
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
self.fake_score_lr_scheduler_2 = get_scheduler(
|
||||
training_args.fake_score_lr_scheduler,
|
||||
optimizer=self.fake_score_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,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Distillation optimizers initialized: generator and fake_score")
|
||||
|
||||
@@ -262,20 +193,12 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
self.generator_ema_2: EMA_FSDP | None = None
|
||||
if (self.training_args.ema_decay
|
||||
is not None) and (self.training_args.ema_decay > 0.0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer,
|
||||
decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA with decay=%s",
|
||||
self.training_args.ema_decay)
|
||||
|
||||
# Initialize EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None:
|
||||
self.generator_ema_2 = EMA_FSDP(
|
||||
self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info("Initialized generator EMA_2 with decay=%s",
|
||||
self.training_args.ema_decay)
|
||||
else:
|
||||
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
|
||||
|
||||
@@ -356,25 +279,16 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"""Prepare training environment for distillation."""
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2.requires_grad_(True)
|
||||
self.transformer_2.train()
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
self.fake_score_transformer_2.requires_grad_(True)
|
||||
self.fake_score_transformer_2.train()
|
||||
|
||||
return training_batch
|
||||
|
||||
def apply_ema_to_model(self, model):
|
||||
"""Apply EMA weights to the model for validation or inference."""
|
||||
if model is self.transformer and self.generator_ema is not None:
|
||||
if self.generator_ema is not None:
|
||||
with self.generator_ema.apply_to_model(model):
|
||||
return model
|
||||
elif model is self.transformer_2 and self.generator_ema_2 is not None:
|
||||
with self.generator_ema_2.apply_to_model(model):
|
||||
return model
|
||||
return model
|
||||
|
||||
def get_ema_model_copy(self) -> torch.nn.Module | None:
|
||||
@@ -385,14 +299,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
return ema_model
|
||||
return None
|
||||
|
||||
def get_ema_2_model_copy(self) -> torch.nn.Module | None:
|
||||
"""Get a copy of the transformer_2 model with EMA weights applied."""
|
||||
if self.generator_ema_2 is not None and self.transformer_2 is not None:
|
||||
ema_2_model = copy.deepcopy(self.transformer_2)
|
||||
self.generator_ema_2.copy_to_unwrapped(ema_2_model)
|
||||
return ema_2_model
|
||||
return None
|
||||
|
||||
def is_ema_ready(self, current_step: int | None = None):
|
||||
"""Check if EMA is ready for use (after ema_start_step)."""
|
||||
if current_step is None:
|
||||
@@ -402,8 +308,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
def save_ema_weights(self, output_dir: str, step: int):
|
||||
"""Save EMA weights separately for inference purposes."""
|
||||
if self.generator_ema is None and self.generator_ema_2 is None:
|
||||
logger.warning("Cannot save EMA weights: No EMA initialized")
|
||||
if self.generator_ema is None:
|
||||
logger.warning("Cannot save EMA weights: EMA not initialized")
|
||||
return
|
||||
|
||||
if not self.is_ema_ready():
|
||||
@@ -413,107 +319,58 @@ class DistillationPipeline(TrainingPipeline):
|
||||
return
|
||||
|
||||
try:
|
||||
# Save main transformer EMA
|
||||
if self.generator_ema is not None:
|
||||
ema_model = self.get_ema_model_copy()
|
||||
if ema_model is None:
|
||||
logger.warning("Failed to create EMA model copy")
|
||||
else:
|
||||
ema_save_dir = os.path.join(output_dir,
|
||||
f"ema_checkpoint-{step}")
|
||||
os.makedirs(ema_save_dir, exist_ok=True)
|
||||
ema_model = self.get_ema_model_copy()
|
||||
if ema_model is None:
|
||||
logger.warning("Failed to create EMA model copy")
|
||||
return
|
||||
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
|
||||
os.makedirs(ema_save_dir, exist_ok=True)
|
||||
|
||||
from fastvideo.training.training_utils import (
|
||||
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(ema_model,
|
||||
device=None)
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
|
||||
if self.global_rank == 0:
|
||||
weight_path = os.path.join(
|
||||
ema_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict = custom_to_hf_state_dict(
|
||||
cpu_state, ema_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, weight_path)
|
||||
from fastvideo.training.training_utils import (
|
||||
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(ema_model, device=None)
|
||||
|
||||
config_dict = ema_model.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"]
|
||||
config_path = os.path.join(ema_save_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
if self.global_rank == 0:
|
||||
weight_path = os.path.join(
|
||||
ema_save_dir, "diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict = custom_to_hf_state_dict(
|
||||
cpu_state, ema_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict, weight_path)
|
||||
|
||||
logger.info("EMA weights saved to %s", weight_path)
|
||||
config_dict = ema_model.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"]
|
||||
config_path = os.path.join(ema_save_dir, "config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
|
||||
del ema_model
|
||||
logger.info("EMA weights saved to %s", weight_path)
|
||||
|
||||
# Save transformer_2 EMA
|
||||
if self.generator_ema_2 is not None:
|
||||
ema_2_model = self.get_ema_2_model_copy()
|
||||
if ema_2_model is None:
|
||||
logger.warning("Failed to create EMA_2 model copy")
|
||||
else:
|
||||
ema_2_save_dir = os.path.join(output_dir,
|
||||
f"ema_2_checkpoint-{step}")
|
||||
os.makedirs(ema_2_save_dir, exist_ok=True)
|
||||
|
||||
# save as diffusers format
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from fastvideo.training.training_utils import (
|
||||
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
|
||||
cpu_state_2 = gather_state_dict_on_cpu_rank0(ema_2_model,
|
||||
device=None)
|
||||
|
||||
if self.global_rank == 0:
|
||||
weight_path_2 = os.path.join(
|
||||
ema_2_save_dir,
|
||||
"diffusion_pytorch_model.safetensors")
|
||||
diffusers_state_dict_2 = custom_to_hf_state_dict(
|
||||
cpu_state_2,
|
||||
ema_2_model.reverse_param_names_mapping)
|
||||
save_file(diffusers_state_dict_2, weight_path_2)
|
||||
|
||||
config_dict_2 = ema_2_model.hf_config
|
||||
if "dtype" in config_dict_2:
|
||||
del config_dict_2["dtype"]
|
||||
config_path_2 = os.path.join(ema_2_save_dir,
|
||||
"config.json")
|
||||
with open(config_path_2, "w") as f:
|
||||
json.dump(config_dict_2, f, indent=4)
|
||||
|
||||
logger.info("EMA_2 weights saved to %s", weight_path_2)
|
||||
|
||||
del ema_2_model
|
||||
del ema_model
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to save EMA weights: %s", str(e))
|
||||
|
||||
def get_ema_stats(self) -> dict[str, Any]:
|
||||
"""Get EMA statistics for monitoring."""
|
||||
ema_enabled = self.generator_ema is not None
|
||||
ema_2_enabled = self.generator_ema_2 is not None
|
||||
|
||||
if not ema_enabled and not ema_2_enabled:
|
||||
if self.generator_ema is None:
|
||||
return {
|
||||
"ema_enabled": False,
|
||||
"ema_2_enabled": False,
|
||||
"ema_decay": None,
|
||||
"ema_start_step": self.training_args.ema_start_step,
|
||||
"ema_ready": False,
|
||||
"ema_2_ready": False,
|
||||
"ema_step": self.current_trainstep,
|
||||
}
|
||||
|
||||
return {
|
||||
"ema_enabled": ema_enabled,
|
||||
"ema_2_enabled": ema_2_enabled,
|
||||
"ema_enabled": True,
|
||||
"ema_decay": self.training_args.ema_decay,
|
||||
"ema_start_step": self.training_args.ema_start_step,
|
||||
"ema_ready": self.is_ema_ready() if ema_enabled else False,
|
||||
"ema_2_ready": self.is_ema_ready() if ema_2_enabled else False,
|
||||
"ema_ready": self.is_ema_ready(),
|
||||
"ema_step": self.current_trainstep,
|
||||
}
|
||||
|
||||
@@ -531,43 +388,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
else:
|
||||
logger.warning("Cannot reset EMA: EMA not initialized")
|
||||
|
||||
if self.generator_ema_2 is not None:
|
||||
logger.info("Resetting EMA_2 to current model weights")
|
||||
self.generator_ema_2.update(self.transformer_2)
|
||||
# Force update to current weights by setting decay to 0 temporarily
|
||||
original_decay_2 = self.generator_ema_2.decay
|
||||
self.generator_ema_2.decay = 0.0
|
||||
self.generator_ema_2.update(self.transformer_2)
|
||||
self.generator_ema_2.decay = original_decay_2
|
||||
logger.info("EMA_2 reset completed")
|
||||
|
||||
def _get_real_score_transformer(self, timestep: torch.Tensor):
|
||||
"""
|
||||
Get the appropriate real score transformer based on timestep and boundary logic.
|
||||
"""
|
||||
if self.real_score_transformer_2 is not None and self.boundary_timestep is not None:
|
||||
if timestep.item() < self.boundary_timestep:
|
||||
return self.real_score_transformer_2 # Low noise expert
|
||||
else:
|
||||
return self.real_score_transformer # High noise expert
|
||||
else:
|
||||
return self.real_score_transformer
|
||||
|
||||
def _get_fake_score_transformer(self, timestep: torch.Tensor):
|
||||
"""
|
||||
Get the appropriate fake score transformer based on timestep and boundary logic.
|
||||
"""
|
||||
if self.fake_score_transformer_2 is not None and self.boundary_timestep is not None:
|
||||
if timestep.item() < self.boundary_timestep:
|
||||
self.train_fake_score_transformer_2 = True
|
||||
return self.fake_score_transformer_2 # Low noise expert
|
||||
else:
|
||||
self.train_fake_score_transformer_2 = False
|
||||
return self.fake_score_transformer # High noise expert
|
||||
else:
|
||||
self.train_fake_score_transformer_2 = False
|
||||
return self.fake_score_transformer
|
||||
|
||||
def _build_distill_input_kwargs(
|
||||
self, noise_input: torch.Tensor, timestep: torch.Tensor,
|
||||
text_dict: dict[str, torch.Tensor] | None,
|
||||
@@ -614,7 +434,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
|
||||
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
|
||||
0, 2, 1, 3, 4)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
@@ -731,9 +550,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(timestep, src=0)
|
||||
|
||||
timestep = shift_timestep(
|
||||
timestep,
|
||||
@@ -760,9 +576,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
current_fake_score_transformer = self._get_fake_score_transformer(
|
||||
timestep)
|
||||
fake_score_pred_noise = current_fake_score_transformer(
|
||||
fake_score_pred_noise = self.fake_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
faker_score_pred_video = pred_noise_to_pred_video(
|
||||
@@ -776,9 +590,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
current_real_score_transformer = self._get_real_score_transformer(
|
||||
timestep)
|
||||
real_score_pred_noise_cond = current_real_score_transformer(
|
||||
real_score_pred_noise_cond = self.real_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_real_video_cond = pred_noise_to_pred_video(
|
||||
@@ -792,8 +604,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.unconditional_dict,
|
||||
training_batch)
|
||||
# Use same transformer as conditional forward for consistency
|
||||
real_score_pred_noise_uncond = current_real_score_transformer(
|
||||
real_score_pred_noise_uncond = self.real_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
pred_real_video_uncond = pred_noise_to_pred_video(
|
||||
@@ -846,9 +657,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.num_train_timestep, [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
world_group = get_world_group()
|
||||
if world_group.world_size > 1:
|
||||
world_group.broadcast(fake_score_timestep, src=0)
|
||||
|
||||
fake_score_timestep = shift_timestep(
|
||||
fake_score_timestep,
|
||||
@@ -879,9 +687,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
noisy_generator_pred_video, fake_score_timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
current_fake_score_transformer = self._get_fake_score_transformer(
|
||||
fake_score_timestep)
|
||||
fake_score_pred_noise = current_fake_score_transformer(
|
||||
fake_score_pred_noise = self.fake_score_transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
target = fake_score_noise - generator_pred_video
|
||||
@@ -995,8 +801,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
attn_metadata=batch_gen.attn_metadata_vsa):
|
||||
(dmd_loss / gradient_accumulation_steps).backward()
|
||||
total_dmd_loss += dmd_loss.detach().item()
|
||||
|
||||
# Only clip gradients for the model that is currently training
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
for param in self.transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
@@ -1006,8 +810,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
if self.generator_ema is not None:
|
||||
self.generator_ema.update(self.transformer)
|
||||
if self.generator_ema_2 is not None:
|
||||
self.generator_ema_2.update(self.transformer_2)
|
||||
|
||||
avg_dmd_loss = torch.tensor(total_dmd_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -1021,8 +823,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch.generator_loss = 0.0
|
||||
|
||||
self.fake_score_optimizer.zero_grad()
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
self.fake_score_optimizer_2.zero_grad()
|
||||
total_fake_score_loss = 0.0
|
||||
for batch in batches:
|
||||
batch_fake = copy.deepcopy(batch)
|
||||
@@ -1033,36 +833,14 @@ class DistillationPipeline(TrainingPipeline):
|
||||
total_fake_score_loss += fake_score_loss.detach().item()
|
||||
fake_score_latent_vis_dict.update(
|
||||
batch_fake.fake_score_latent_vis_dict)
|
||||
if self.train_fake_score_transformer_2 and self.fake_score_transformer_2 is not None:
|
||||
self._clip_model_grad_norm_(batch_fake,
|
||||
self.fake_score_transformer_2)
|
||||
else:
|
||||
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
|
||||
|
||||
# Check gradients for fake score transformer
|
||||
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
|
||||
for param in self.fake_score_transformer.parameters():
|
||||
if param.requires_grad:
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
|
||||
# Check gradients for fake score transformer_2 if available
|
||||
if self.train_fake_score_transformer_2 and self.fake_score_transformer_2 is not None:
|
||||
for param in self.fake_score_transformer_2.parameters():
|
||||
if param.requires_grad:
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
|
||||
if self.train_fake_score_transformer_2 and self.fake_score_transformer_2 is not None:
|
||||
self.fake_score_optimizer_2.step()
|
||||
self.fake_score_lr_scheduler_2.step()
|
||||
else:
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
|
||||
# Step the appropriate scheduler
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
self.fake_score_optimizer.zero_grad(set_to_none=True)
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
self.fake_score_optimizer_2.zero_grad(set_to_none=True)
|
||||
avg_fake_score_loss = torch.tensor(total_fake_score_loss /
|
||||
gradient_accumulation_steps,
|
||||
device=self.device)
|
||||
@@ -1083,30 +861,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.training_args.resume_from_checkpoint)
|
||||
|
||||
resumed_step = load_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.resume_from_checkpoint,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2', None),
|
||||
real_score_transformer_2=getattr(self, 'real_score_transformer_2',
|
||||
None),
|
||||
fake_score_transformer_2=getattr(self, 'fake_score_transformer_2',
|
||||
None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self, 'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self, 'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
self.transformer, self.fake_score_transformer, self.global_rank,
|
||||
self.training_args.resume_from_checkpoint, self.optimizer,
|
||||
self.fake_score_optimizer, self.train_dataloader, self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator,
|
||||
self.generator_ema)
|
||||
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
@@ -1128,39 +887,15 @@ class DistillationPipeline(TrainingPipeline):
|
||||
logger.info(" Max gradient norm: %s", self.training_args.max_grad_norm)
|
||||
|
||||
logger.info(
|
||||
" Real score transformer (high noise expert) parameters: %s B",
|
||||
" Real score transformer parameters: %s B",
|
||||
sum(p.numel()
|
||||
for p in self.real_score_transformer.parameters()) / 1e9)
|
||||
|
||||
if self.real_score_transformer_2 is not None:
|
||||
logger.info(
|
||||
" Real score transformer_2 (low noise expert) parameters: %s B",
|
||||
sum(p.numel()
|
||||
for p in self.real_score_transformer_2.parameters()) / 1e9)
|
||||
logger.info(" Real score MoE enabled with boundary_timestep: %s",
|
||||
self.boundary_timestep)
|
||||
|
||||
logger.info(
|
||||
" Fake score transformer (high noise expert) parameters: %s B",
|
||||
" Fake score transformer parameters: %s B",
|
||||
sum(p.numel()
|
||||
for p in self.fake_score_transformer.parameters()) / 1e9)
|
||||
|
||||
if self.fake_score_transformer_2 is not None:
|
||||
logger.info(
|
||||
" Fake score transformer_2 (low noise expert) parameters: %s B",
|
||||
sum(p.numel()
|
||||
for p in self.fake_score_transformer_2.parameters()) / 1e9)
|
||||
logger.info(" Fake score MoE enabled with boundary_timestep: %s",
|
||||
self.boundary_timestep)
|
||||
|
||||
if self.generator_ema is not None:
|
||||
logger.info(" Generator EMA enabled with decay: %s",
|
||||
self.training_args.ema_decay)
|
||||
logger.info(" Generator EMA start step: %s",
|
||||
self.training_args.ema_start_step)
|
||||
else:
|
||||
logger.info(" Generator EMA disabled")
|
||||
|
||||
if self.generator_ema is not None:
|
||||
logger.info(" Generator EMA enabled with decay: %s",
|
||||
self.training_args.ema_decay)
|
||||
@@ -1198,35 +933,19 @@ class DistillationPipeline(TrainingPipeline):
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
# Set both transformers to eval mode
|
||||
transformer.eval()
|
||||
if hasattr(self, 'transformer_2') and self.transformer_2 is not None:
|
||||
self.transformer_2.eval()
|
||||
|
||||
# Optionally use EMA model for validation if available and ready
|
||||
use_ema_for_validation = (self.training_args.use_ema
|
||||
and self.is_ema_ready(global_step))
|
||||
ema_context = None
|
||||
ema_2_context = None
|
||||
|
||||
if use_ema_for_validation:
|
||||
logger.info("Using EMA model for validation")
|
||||
# Use self.transformer for consistency (the passed transformer should be self.transformer anyway)
|
||||
validation_transformer = self.transformer
|
||||
if self.generator_ema is not None:
|
||||
ema_context = self.generator_ema.apply_to_model(
|
||||
validation_transformer)
|
||||
|
||||
# Handle transformer_2 EMA if available
|
||||
if hasattr(
|
||||
self, 'transformer_2'
|
||||
) and self.transformer_2 is not None and self.generator_ema_2 is not None:
|
||||
ema_2_context = self.generator_ema_2.apply_to_model(
|
||||
self.transformer_2)
|
||||
logger.info("Using EMA_2 model for transformer_2 validation")
|
||||
ema_context = self.generator_ema.apply_to_model(
|
||||
validation_transformer)
|
||||
else:
|
||||
# Use self.transformer for consistency, but the passed transformer should be the same
|
||||
validation_transformer = self.transformer
|
||||
validation_transformer = transformer
|
||||
ema_context = None
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
@@ -1243,14 +962,58 @@ class DistillationPipeline(TrainingPipeline):
|
||||
step_videos: list[np.ndarray] = []
|
||||
step_captions: list[str] = []
|
||||
|
||||
# Helper function to run validation with optional EMA contexts
|
||||
def run_validation_with_ema(
|
||||
steps: int) -> tuple[list[np.ndarray], list[str]]:
|
||||
videos: list[np.ndarray] = []
|
||||
captions: list[str] = []
|
||||
if ema_context is not None:
|
||||
with ema_context:
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(
|
||||
sampling_param, training_args, validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
|
||||
batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
|
||||
logger.info(
|
||||
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
self.global_rank,
|
||||
self.rank_in_sp_group,
|
||||
batch.prompt,
|
||||
local_main_process_only=False)
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
with torch.no_grad():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in video:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
step_videos.append(frames)
|
||||
else:
|
||||
# Use original transformer without EMA
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(
|
||||
sampling_param, training_args, validation_batch, steps)
|
||||
sampling_param, training_args, validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
@@ -1273,7 +1036,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
captions.append(batch.prompt)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
with torch.no_grad():
|
||||
@@ -1290,26 +1053,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
videos.append(frames)
|
||||
|
||||
return videos, captions
|
||||
|
||||
# Apply EMA contexts if available (nested context managers)
|
||||
if ema_context is not None and ema_2_context is not None:
|
||||
with ema_context, ema_2_context:
|
||||
step_videos, step_captions = run_validation_with_ema(
|
||||
num_inference_steps)
|
||||
elif ema_context is not None:
|
||||
with ema_context:
|
||||
step_videos, step_captions = run_validation_with_ema(
|
||||
num_inference_steps)
|
||||
elif ema_2_context is not None:
|
||||
with ema_2_context:
|
||||
step_videos, step_captions = run_validation_with_ema(
|
||||
num_inference_steps)
|
||||
else:
|
||||
step_videos, step_captions = run_validation_with_ema(
|
||||
num_inference_steps)
|
||||
step_videos.append(frames)
|
||||
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
@@ -1355,10 +1099,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
world_group.send_object(step_videos, dst=0)
|
||||
world_group.send_object(step_captions, dst=0)
|
||||
|
||||
# Re-enable gradients for training - set both transformers back to train mode
|
||||
# Re-enable gradients for training
|
||||
transformer.train()
|
||||
if hasattr(self, 'transformer_2') and self.transformer_2 is not None:
|
||||
self.transformer_2.train()
|
||||
gc.collect()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
@@ -1512,14 +1254,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
logger.info("Created generator EMA at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
# Create EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None and self.generator_ema_2 is None:
|
||||
self.generator_ema_2 = EMA_FSDP(
|
||||
self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info(
|
||||
"Created generator EMA_2 at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -1546,9 +1280,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"ema":
|
||||
"✓" if (self.generator_ema is not None and self.is_ema_ready())
|
||||
else "✗",
|
||||
"ema2":
|
||||
"✓" if (self.generator_ema_2 is not None
|
||||
and self.is_ema_ready()) else "✗",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
@@ -1576,13 +1307,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
if use_vsa:
|
||||
log_data["VSA_train_sparsity"] = current_vsa_sparsity
|
||||
|
||||
if self.generator_ema is not None or self.generator_ema_2 is not None:
|
||||
log_data["ema_enabled"] = self.generator_ema is not None
|
||||
log_data["ema_2_enabled"] = self.generator_ema_2 is not None
|
||||
if self.generator_ema is not None:
|
||||
log_data["ema_enabled"] = True
|
||||
log_data["ema_decay"] = self.training_args.ema_decay
|
||||
else:
|
||||
log_data["ema_enabled"] = False
|
||||
log_data["ema_2_enabled"] = False
|
||||
|
||||
ema_stats = self.get_ema_stats()
|
||||
log_data.update(ema_stats)
|
||||
@@ -1614,34 +1343,12 @@ class DistillationPipeline(TrainingPipeline):
|
||||
print("rank", self.global_rank,
|
||||
"save training state checkpoint at step", step)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
step,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2',
|
||||
None),
|
||||
real_score_transformer_2=getattr(
|
||||
self, 'real_score_transformer_2', None),
|
||||
fake_score_transformer_2=getattr(
|
||||
self, 'fake_score_transformer_2', None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self,
|
||||
'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self,
|
||||
'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
self.transformer, self.fake_score_transformer,
|
||||
self.global_rank, self.training_args.output_dir, step,
|
||||
self.optimizer, self.fake_score_optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator,
|
||||
self.generator_ema)
|
||||
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
@@ -1653,30 +1360,13 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.training_args.weight_only_checkpointing_steps == 0):
|
||||
print("rank", self.global_rank,
|
||||
"save weight-only checkpoint at step", step)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2',
|
||||
None),
|
||||
real_score_transformer_2=getattr(
|
||||
self, 'real_score_transformer_2', None),
|
||||
fake_score_transformer_2=getattr(
|
||||
self, 'fake_score_transformer_2', None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self,
|
||||
'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self,
|
||||
'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
save_distillation_checkpoint(self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema)
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir, step)
|
||||
@@ -1695,35 +1385,15 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"save final training state checkpoint at step",
|
||||
self.training_args.max_train_steps)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2', None),
|
||||
real_score_transformer_2=getattr(self, 'real_score_transformer_2',
|
||||
None),
|
||||
fake_score_transformer_2=getattr(self, 'fake_score_transformer_2',
|
||||
None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self, 'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self, 'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
self.transformer, self.fake_score_transformer, self.global_rank,
|
||||
self.training_args.output_dir, self.training_args.max_train_steps,
|
||||
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator, self.generator_ema)
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir,
|
||||
self.training_args.max_train_steps)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -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)
|
||||
@@ -48,16 +48,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
logger.info("Initializing self-forcing distillation pipeline...")
|
||||
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
self.generator_ema_2: EMA_FSDP | None = None
|
||||
|
||||
super().initialize_training_pipeline(training_args)
|
||||
try:
|
||||
logger.info("RANK: %s, entered initialize_training_pipeline",
|
||||
self.global_rank,
|
||||
local_main_process_only=False)
|
||||
except Exception:
|
||||
logger.info("Entered initialize_training_pipeline (rank unknown)")
|
||||
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
num_inference_steps=1000,
|
||||
shift=5.0,
|
||||
@@ -67,6 +59,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.dfake_gen_update_ratio = getattr(training_args,
|
||||
'dfake_gen_update_ratio', 5)
|
||||
|
||||
# Self-forcing specific properties
|
||||
self.num_frame_per_block = getattr(training_args, 'num_frame_per_block',
|
||||
3)
|
||||
self.independent_first_frame = getattr(training_args,
|
||||
@@ -76,25 +69,19 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.last_step_only = getattr(training_args, 'last_step_only', False)
|
||||
self.context_noise = getattr(training_args, 'context_noise', 0)
|
||||
|
||||
# Calculate frame sequence length - this will be set properly in _prepare_dit_inputs
|
||||
self.frame_seq_length = 1560 # TODO: Calculate this dynamically based on patch size
|
||||
|
||||
# Cache references (will be initialized per forward pass)
|
||||
self.kv_cache1: list[dict[str, Any]] | None = None
|
||||
self.crossattn_cache: list[dict[str, Any]] | None = None
|
||||
|
||||
logger.info("Self-forcing generator update ratio: %s",
|
||||
self.dfake_gen_update_ratio)
|
||||
logger.info("RANK: %s, exiting initialize_training_pipeline",
|
||||
self.global_rank,
|
||||
local_main_process_only=False)
|
||||
|
||||
def generate_and_sync_list(self, num_blocks: int, num_denoising_steps: int,
|
||||
device: torch.device) -> list[int]:
|
||||
"""Generate and synchronize random exit flags across distributed processes."""
|
||||
logger.info(
|
||||
"RANK: %s, enter generate_and_sync_list blocks=%s steps=%s device=%s",
|
||||
self.global_rank,
|
||||
num_blocks,
|
||||
num_denoising_steps,
|
||||
str(device),
|
||||
local_main_process_only=False)
|
||||
rank = dist.get_rank() if dist.is_initialized() else 0
|
||||
|
||||
if rank == 0:
|
||||
@@ -111,14 +98,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
if dist.is_initialized():
|
||||
dist.broadcast(indices,
|
||||
src=0) # Broadcast the random indices to all ranks
|
||||
flags = indices.tolist()
|
||||
logger.info(
|
||||
"RANK: %s, exit generate_and_sync_list flags_len=%s first=%s",
|
||||
self.global_rank,
|
||||
len(flags),
|
||||
flags[0] if len(flags) > 0 else None,
|
||||
local_main_process_only=False)
|
||||
return flags
|
||||
return indices.tolist()
|
||||
|
||||
def generator_loss(
|
||||
self, training_batch: TrainingBatch
|
||||
@@ -130,8 +110,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps,
|
||||
attn_metadata=training_batch.attn_metadata_vsa):
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
training_batch)
|
||||
if self.training_args.simulate_generator_forward:
|
||||
generator_pred_video = self._generator_multi_step_simulation_forward(
|
||||
training_batch)
|
||||
else:
|
||||
generator_pred_video = self._generator_forward(training_batch)
|
||||
|
||||
with set_forward_context(current_timestep=training_batch.timesteps,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
@@ -159,6 +142,78 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
return flow_matching_loss, log_dict
|
||||
|
||||
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Forward pass through generator with KV cache support for causal generation."""
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
batch_size = latents.shape[0]
|
||||
|
||||
# Step 1: Sample a timestep from denoising_step_list
|
||||
index = torch.randint(0,
|
||||
len(self.denoising_step_list), [1],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
timestep = self.denoising_step_list[index]
|
||||
training_batch.dmd_latent_vis_dict["generator_timestep"] = timestep
|
||||
|
||||
# Step 2: Initialize KV cache and cross-attention cache for causal generation
|
||||
kv_cache, crossattn_cache = self._initialize_simulation_caches(
|
||||
batch_size, dtype, self.device)
|
||||
|
||||
if getattr(self.training_args, 'validate_cache_structure', False):
|
||||
self._validate_cache_structure(kv_cache, crossattn_cache,
|
||||
batch_size)
|
||||
|
||||
# Step 3: Add noise to latents
|
||||
noise = torch.randn(self.video_latent_shape,
|
||||
device=self.device,
|
||||
dtype=dtype)
|
||||
if self.sp_world_size > 1:
|
||||
noise = rearrange(noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
latents.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep * torch.ones([latents.shape[0] * latents.shape[1]],
|
||||
device=noise.device,
|
||||
dtype=torch.long))
|
||||
|
||||
# Step 4: Build input kwargs with KV cache support
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
|
||||
# Step 5: Forward pass with KV cache if available
|
||||
if hasattr(self.transformer, '_forward_inference'):
|
||||
# Use causal inference forward with KV cache
|
||||
pred_noise = self.transformer(
|
||||
hidden_states=training_batch.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch.
|
||||
input_kwargs['encoder_hidden_states'],
|
||||
timestep=training_batch.input_kwargs['timestep'],
|
||||
encoder_hidden_states_image=training_batch.input_kwargs.get(
|
||||
'encoder_hidden_states_image'),
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=0, # Start from beginning for single-step
|
||||
cache_start=0).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
# Fallback to regular forward
|
||||
pred_noise = self.transformer(
|
||||
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Step 6: Convert noise prediction to video prediction
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=torch.tensor([timestep], device=noisy_latent.device),
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
self._reset_simulation_caches(kv_cache, crossattn_cache)
|
||||
|
||||
return pred_video
|
||||
|
||||
def _generator_multi_step_simulation_forward(
|
||||
self,
|
||||
training_batch: TrainingBatch,
|
||||
@@ -234,18 +289,14 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
device=noise.device,
|
||||
dtype=noise.dtype)
|
||||
|
||||
def get_model_device(model):
|
||||
if model is None:
|
||||
return "None"
|
||||
try:
|
||||
return next(model.parameters()).device
|
||||
except (StopIteration, AttributeError):
|
||||
return "Unknown"
|
||||
|
||||
# Step 1: Initialize KV cache to all zeros
|
||||
cache_frames = num_generated_frames + num_input_frames
|
||||
self.kv_cache1, self.crossattn_cache = self._initialize_simulation_caches(
|
||||
batch_size, dtype, self.device, max_num_frames=cache_frames)
|
||||
batch_size, dtype, self.device)
|
||||
|
||||
# Validate cache structure (can be disabled in production)
|
||||
if getattr(self.training_args, 'validate_cache_structure', False):
|
||||
self._validate_cache_structure(self.kv_cache1, self.crossattn_cache,
|
||||
batch_size)
|
||||
|
||||
# Step 2: Cache context feature
|
||||
current_start_frame = 0
|
||||
@@ -258,9 +309,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
initial_latent, timestep * 0,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
# we process the image latent with self.transformer_2 (low-noise expert)
|
||||
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
|
||||
current_model(
|
||||
|
||||
self.transformer(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -300,17 +350,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
device=noise.device,
|
||||
dtype=torch.int64) * current_timestep
|
||||
|
||||
if self.boundary_timestep is not None and current_timestep < self.boundary_timestep and self.transformer_2 is not None:
|
||||
current_model = self.transformer_2
|
||||
self._enable_training(self.transformer_2, self.optimizer_2)
|
||||
self._disable_training(self.transformer, self.optimizer)
|
||||
else:
|
||||
current_model = self.transformer
|
||||
self._enable_training(self.transformer, self.optimizer)
|
||||
if self.boundary_timestep is not None and self.transformer_2 is not None:
|
||||
self._disable_training(self.transformer_2,
|
||||
self.optimizer_2)
|
||||
|
||||
if not exit_flag:
|
||||
with torch.no_grad():
|
||||
# Build input kwargs
|
||||
@@ -318,7 +357,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noisy_input, timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = current_model(
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -358,7 +397,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noisy_input, timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = current_model(
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -378,7 +417,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noisy_input, timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_flow = current_model(
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -418,9 +457,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
denoised_pred, context_timestep,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
|
||||
# context_timestep is 0 so we use transformer_2
|
||||
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
|
||||
current_model(
|
||||
self.transformer(
|
||||
hidden_states=training_batch_temp.
|
||||
input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.
|
||||
@@ -520,7 +557,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.dmd_latent_vis_dict["min_num_frames"] = torch.tensor(
|
||||
min_num_frames, dtype=torch.float32, device=self.device)
|
||||
|
||||
# Clean up caches
|
||||
assert self.kv_cache1 is not None
|
||||
assert self.crossattn_cache is not None
|
||||
self._reset_simulation_caches(self.kv_cache1, self.crossattn_cache)
|
||||
@@ -528,35 +564,42 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
return final_output if gradient_mask is not None else pred_image_or_video
|
||||
|
||||
def _initialize_simulation_caches(
|
||||
self,
|
||||
batch_size: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
*,
|
||||
max_num_frames: int | None = None,
|
||||
self, batch_size: int, dtype: torch.dtype, device: torch.device
|
||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
|
||||
num_transformer_blocks = len(self.transformer.blocks)
|
||||
latent_shape = self.video_latent_shape_sp
|
||||
_, num_frames, _, height, width = latent_shape
|
||||
|
||||
_, p_h, p_w = self.transformer.patch_size
|
||||
# Calculate frame sequence length based on input dimensions and patch size
|
||||
# From the training batch, we can get the actual latent dimensions
|
||||
latent_shape = self.video_latent_shape_sp # This is set in _prepare_dit_inputs
|
||||
batch_size_actual, num_frames, num_channels, height, width = latent_shape
|
||||
|
||||
# Get patch size from transformer config
|
||||
p_t, p_h, p_w = self.transformer.patch_size
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# Frame sequence length is the spatial sequence length per frame
|
||||
frame_seq_length = post_patch_height * post_patch_width
|
||||
self.frame_seq_length = frame_seq_length
|
||||
|
||||
# Get local attention size from transformer config
|
||||
# local_attn_size = getattr(self.transformer, 'local_attn_size', -1)
|
||||
|
||||
# Get model configuration parameters - handle FSDP wrapping
|
||||
num_attention_heads = getattr(self.transformer, 'num_attention_heads',
|
||||
None)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
None)
|
||||
text_len = getattr(self.transformer, 'text_len', None)
|
||||
if hasattr(self.transformer, 'config'):
|
||||
config = self.transformer.config
|
||||
num_attention_heads = config.num_attention_heads
|
||||
attention_head_dim = config.attention_head_dim
|
||||
text_len = config.text_len
|
||||
else:
|
||||
# Fallback to direct attribute access for non-FSDP models
|
||||
num_attention_heads = getattr(self.transformer,
|
||||
'num_attention_heads', 40)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
128)
|
||||
text_len = getattr(self.transformer, 'text_len', 512)
|
||||
|
||||
if max_num_frames is None:
|
||||
max_num_frames = num_frames
|
||||
num_max_frames = max(max_num_frames, num_frames)
|
||||
num_max_frames = getattr(self.training_args, "num_frames", num_frames)
|
||||
kv_cache_size = num_max_frames * frame_seq_length
|
||||
|
||||
kv_cache = []
|
||||
@@ -622,10 +665,64 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
cache_dict["k"].zero_()
|
||||
cache_dict["v"].zero_()
|
||||
|
||||
def _validate_cache_structure(self, kv_cache, crossattn_cache,
|
||||
batch_size: int):
|
||||
"""Validate that cache structures are correctly initialized."""
|
||||
num_transformer_blocks = len(self.transformer.blocks)
|
||||
|
||||
# Get model configuration parameters - handle FSDP wrapping
|
||||
if hasattr(self.transformer, 'config'):
|
||||
config = self.transformer.config
|
||||
num_attention_heads = config.num_attention_heads
|
||||
attention_head_dim = config.attention_head_dim
|
||||
text_len = config.text_len
|
||||
else:
|
||||
# Fallback to direct attribute access for non-FSDP models
|
||||
num_attention_heads = getattr(self.transformer,
|
||||
'num_attention_heads', 40)
|
||||
attention_head_dim = getattr(self.transformer, 'attention_head_dim',
|
||||
128)
|
||||
text_len = getattr(self.transformer, 'text_len', 512)
|
||||
|
||||
if kv_cache is not None:
|
||||
assert len(
|
||||
kv_cache
|
||||
) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(kv_cache)}"
|
||||
for i, cache_dict in enumerate(kv_cache):
|
||||
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in kv_cache block {i}"
|
||||
assert "global_end_index" in cache_dict and "local_end_index" in cache_dict, f"Missing indices in kv_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
|
||||
assert cache_dict["v"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
2] == num_attention_heads, f"Attention heads mismatch in kv_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
3] == attention_head_dim, f"Attention head dim mismatch in kv_cache block {i}"
|
||||
|
||||
if crossattn_cache is not None:
|
||||
assert len(
|
||||
crossattn_cache
|
||||
) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(crossattn_cache)}"
|
||||
for i, cache_dict in enumerate(crossattn_cache):
|
||||
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in crossattn_cache block {i}"
|
||||
assert "is_init" in cache_dict, f"Missing is_init in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["v"].shape[
|
||||
0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
1] == text_len, f"Text length mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
2] == num_attention_heads, f"Attention heads mismatch in crossattn_cache block {i}"
|
||||
assert cache_dict["k"].shape[
|
||||
3] == attention_head_dim, f"Attention head dim mismatch in crossattn_cache block {i}"
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
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)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
@@ -656,6 +753,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
|
||||
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
@@ -683,11 +781,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
if train_generator:
|
||||
logger.debug("Training generator at step %s",
|
||||
self.current_trainstep)
|
||||
self.optimizer.zero_grad()
|
||||
if self.transformer_2 is not None:
|
||||
self.optimizer_2.zero_grad()
|
||||
total_generator_loss = 0.0
|
||||
generator_log_dict = {}
|
||||
|
||||
@@ -719,27 +813,12 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.dmd_latent_vis_dict.update(
|
||||
batch_gen.dmd_latent_vis_dict)
|
||||
|
||||
# Only clip gradients and step optimizer for the model that is currently training
|
||||
if hasattr(
|
||||
self, 'train_transformer_2'
|
||||
) and self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer_2)
|
||||
self.optimizer_2.step()
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if self.generator_ema is not None:
|
||||
if hasattr(
|
||||
self, 'train_transformer_2'
|
||||
) and self.train_transformer_2 and self.transformer_2 is not None:
|
||||
# Update EMA for transformer_2 when training it
|
||||
if self.generator_ema_2 is not None:
|
||||
self.generator_ema_2.update(self.transformer_2)
|
||||
else:
|
||||
self.generator_ema.update(self.transformer)
|
||||
self.generator_ema.update(self.transformer)
|
||||
|
||||
avg_generator_loss = torch.tensor(total_generator_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -751,7 +830,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
else:
|
||||
training_batch.generator_loss = 0.0
|
||||
|
||||
logger.debug("Training critic at step %s", self.current_trainstep)
|
||||
self.fake_score_optimizer.zero_grad()
|
||||
total_critic_loss = 0.0
|
||||
critic_log_dict = {}
|
||||
@@ -784,16 +862,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.fake_score_latent_vis_dict.update(
|
||||
batch_critic.fake_score_latent_vis_dict)
|
||||
|
||||
if self.train_fake_score_transformer_2 and self.fake_score_transformer_2 is not None:
|
||||
self._clip_model_grad_norm_(batch_critic,
|
||||
self.fake_score_transformer_2)
|
||||
self.fake_score_optimizer_2.step()
|
||||
self.fake_score_lr_scheduler_2.step()
|
||||
else:
|
||||
self._clip_model_grad_norm_(batch_critic,
|
||||
self.fake_score_transformer)
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
self._clip_model_grad_norm_(batch_critic, self.fake_score_transformer)
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
|
||||
avg_critic_loss = torch.tensor(total_critic_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -804,6 +875,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch.fake_score_loss = avg_critic_loss.item()
|
||||
|
||||
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_loss
|
||||
|
||||
return training_batch
|
||||
|
||||
def _log_training_info(self) -> None:
|
||||
@@ -818,6 +890,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
wandb_loss_dict = {}
|
||||
|
||||
# Debug logging
|
||||
logger.info("Step %s: Starting visualization", step)
|
||||
if hasattr(training_batch, 'dmd_latent_vis_dict'):
|
||||
logger.info("DMD latent keys: %s",
|
||||
list(training_batch.dmd_latent_vis_dict.keys()))
|
||||
@@ -861,14 +934,20 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
try:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
logger.info("Successfully processed DMD latent: %s",
|
||||
latent_key)
|
||||
except Exception as e:
|
||||
logger.error("Error processing DMD latent %s: %s",
|
||||
latent_key, str(e))
|
||||
del video, latents
|
||||
|
||||
# Process critic predictions
|
||||
@@ -903,14 +982,20 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
try:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(
|
||||
video, fps=24, format="mp4")
|
||||
logger.info("Successfully processed critic latent: %s",
|
||||
latent_key)
|
||||
except Exception as e:
|
||||
logger.error("Error processing critic latent %s: %s",
|
||||
latent_key, str(e))
|
||||
del video, latents
|
||||
|
||||
# Log metadata
|
||||
@@ -950,6 +1035,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
# Use the same seed for all processes within the same SP group
|
||||
sp_group_seed = seed + (self.global_rank // self.sp_world_size)
|
||||
set_random_seed(sp_group_seed)
|
||||
logger.info("Rank %s: Using SP group seed %s", self.global_rank,
|
||||
sp_group_seed)
|
||||
else:
|
||||
set_random_seed(seed + self.global_rank)
|
||||
|
||||
@@ -1012,14 +1099,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
logger.info("Created generator EMA at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
# Create EMA for transformer_2 if it exists
|
||||
if self.transformer_2 is not None and self.generator_ema_2 is None:
|
||||
self.generator_ema_2 = EMA_FSDP(
|
||||
self.transformer_2, decay=self.training_args.ema_decay)
|
||||
logger.info(
|
||||
"Created generator EMA_2 at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -1046,9 +1125,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
"ema":
|
||||
"✓" if (self.generator_ema is not None and self.is_ema_ready())
|
||||
else "✗",
|
||||
"ema2":
|
||||
"✓" if (self.generator_ema_2 is not None
|
||||
and self.is_ema_ready()) else "✗",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
@@ -1074,13 +1150,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
if use_vsa:
|
||||
log_data["VSA_train_sparsity"] = current_vsa_sparsity
|
||||
|
||||
if self.generator_ema is not None or self.generator_ema_2 is not None:
|
||||
log_data["ema_enabled"] = self.generator_ema is not None
|
||||
log_data["ema_2_enabled"] = self.generator_ema_2 is not None
|
||||
if self.generator_ema is not None:
|
||||
log_data["ema_enabled"] = True
|
||||
log_data["ema_decay"] = self.training_args.ema_decay
|
||||
else:
|
||||
log_data["ema_enabled"] = False
|
||||
log_data["ema_2_enabled"] = False
|
||||
|
||||
ema_stats = self.get_ema_stats()
|
||||
log_data.update(ema_stats)
|
||||
@@ -1116,34 +1190,12 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
print("rank", self.global_rank,
|
||||
"save training state checkpoint at step", step)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
step,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2',
|
||||
None),
|
||||
real_score_transformer_2=getattr(
|
||||
self, 'real_score_transformer_2', None),
|
||||
fake_score_transformer_2=getattr(
|
||||
self, 'fake_score_transformer_2', None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self,
|
||||
'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self,
|
||||
'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
self.transformer, self.fake_score_transformer,
|
||||
self.global_rank, self.training_args.output_dir, step,
|
||||
self.optimizer, self.fake_score_optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler, self.noise_random_generator,
|
||||
self.generator_ema)
|
||||
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
@@ -1154,30 +1206,13 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
self.training_args.weight_only_checkpointing_steps == 0):
|
||||
print("rank", self.global_rank,
|
||||
"save weight-only checkpoint at step", step)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2',
|
||||
None),
|
||||
real_score_transformer_2=getattr(
|
||||
self, 'real_score_transformer_2', None),
|
||||
fake_score_transformer_2=getattr(
|
||||
self, 'fake_score_transformer_2', None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self,
|
||||
'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self,
|
||||
'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
save_distillation_checkpoint(self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
f"{step}_weight_only",
|
||||
only_save_generator_weight=True,
|
||||
generator_ema=self.generator_ema)
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir, step)
|
||||
@@ -1191,31 +1226,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
"save final training state checkpoint at step",
|
||||
self.training_args.max_train_steps)
|
||||
save_distillation_checkpoint(
|
||||
self.transformer,
|
||||
self.fake_score_transformer,
|
||||
self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps,
|
||||
self.optimizer,
|
||||
self.fake_score_optimizer,
|
||||
self.train_dataloader,
|
||||
self.lr_scheduler,
|
||||
self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator,
|
||||
self.generator_ema,
|
||||
# MoE support
|
||||
generator_transformer_2=getattr(self, 'transformer_2', None),
|
||||
real_score_transformer_2=getattr(self, 'real_score_transformer_2',
|
||||
None),
|
||||
fake_score_transformer_2=getattr(self, 'fake_score_transformer_2',
|
||||
None),
|
||||
generator_optimizer_2=getattr(self, 'optimizer_2', None),
|
||||
fake_score_optimizer_2=getattr(self, 'fake_score_optimizer_2',
|
||||
None),
|
||||
generator_scheduler_2=getattr(self, 'lr_scheduler_2', None),
|
||||
fake_score_scheduler_2=getattr(self, 'fake_score_lr_scheduler_2',
|
||||
None),
|
||||
generator_ema_2=getattr(self, 'generator_ema_2', None))
|
||||
self.transformer, self.fake_score_transformer, self.global_rank,
|
||||
self.training_args.output_dir, self.training_args.max_train_steps,
|
||||
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.fake_score_lr_scheduler,
|
||||
self.noise_random_generator, self.generator_ema)
|
||||
|
||||
if self.training_args.use_ema and self.is_ema_ready():
|
||||
self.save_ema_weights(self.training_args.output_dir,
|
||||
|
||||
@@ -22,7 +22,7 @@ from tqdm.auto import tqdm
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
# from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
@@ -39,20 +39,26 @@ from fastvideo.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, count_trainable, get_scheduler,
|
||||
get_sigmas, load_checkpoint, normalize_dit_input, save_checkpoint,
|
||||
compute_density_for_timestep_sampling, get_scheduler, get_sigmas,
|
||||
load_checkpoint, normalize_dit_input, save_checkpoint,
|
||||
shard_latents_across_sp)
|
||||
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
# from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
# set_random_seed, shallow_asdict)
|
||||
from fastvideo.utils import (is_vsa_available,
|
||||
set_random_seed, shallow_asdict)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
vmoba_available = is_vmoba_available()
|
||||
# vmoba_available = is_vmoba_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _get_trainable_params(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class TrainingPipeline(LoRAPipeline, ABC):
|
||||
"""
|
||||
A pipeline for training a model. All training pipelines should inherit from this class.
|
||||
@@ -63,7 +69,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 +104,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,11 +116,6 @@ 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"]
|
||||
self.set_trainable()
|
||||
@@ -126,7 +125,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# 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,
|
||||
@@ -148,30 +147,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 +161,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 +187,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 +233,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 +266,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
|
||||
@@ -374,20 +281,20 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=current_vsa_sparsity,
|
||||
device=get_local_torch_device())
|
||||
elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
moba_params = self.training_args.moba_config.copy()
|
||||
moba_params.update({
|
||||
"current_timestep":
|
||||
training_batch.timesteps,
|
||||
"raw_latent_shape":
|
||||
training_batch.raw_latent_shape[2:5],
|
||||
"patch_size":
|
||||
self.training_args.pipeline_config.dit_config.patch_size,
|
||||
"device":
|
||||
get_local_torch_device(),
|
||||
})
|
||||
training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
|
||||
).build(**moba_params)
|
||||
# elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
# moba_params = self.training_args.moba_config.copy()
|
||||
# moba_params.update({
|
||||
# "current_timestep":
|
||||
# training_batch.timesteps,
|
||||
# "raw_latent_shape":
|
||||
# training_batch.raw_latent_shape[2:5],
|
||||
# "patch_size":
|
||||
# self.training_args.pipeline_config.dit_config.patch_size,
|
||||
# "device":
|
||||
# get_local_torch_device(),
|
||||
# })
|
||||
# training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
|
||||
# ).build(**moba_params)
|
||||
else:
|
||||
training_batch.attn_metadata = None
|
||||
|
||||
@@ -412,7 +319,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
def _transformer_forward_and_compute_loss(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
# if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
assert training_batch.attn_metadata is not None
|
||||
else:
|
||||
assert training_batch.attn_metadata is None
|
||||
@@ -423,12 +331,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 +366,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 +411,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
|
||||
@@ -544,16 +441,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
local_main_process_only=False)
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
num_trainable_params = count_trainable(self.transformer)
|
||||
num_trainable_params = _get_trainable_params(self.transformer)
|
||||
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)
|
||||
@@ -595,9 +486,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
current_decay_times = min(step // vsa_decay_interval_steps,
|
||||
vsa_sparsity // vsa_decay_rate)
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
elif vmoba_available:
|
||||
#TODO: add vmoba sparsity scheduling here
|
||||
current_vsa_sparsity = 0.0
|
||||
# elif vmoba_available:
|
||||
# # TODO: add vmoba sparsity scheduling here
|
||||
# pass
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
@@ -631,7 +522,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,14 +530,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
if self.training_args.log_visualization:
|
||||
self.visualize_intermediate_latents(training_batch,
|
||||
self.training_args,
|
||||
step)
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
trainable_params = round(
|
||||
count_trainable(self.transformer) / 1e9, 3)
|
||||
_get_trainable_params(self.transformer) / 1e9, 3)
|
||||
logger.info(
|
||||
"GPU memory usage after validation: %s MB, trainable params: %sB",
|
||||
gpu_memory_usage, trainable_params)
|
||||
@@ -682,7 +569,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
logger.info(" Total optimization steps = %s",
|
||||
self.training_args.max_train_steps)
|
||||
logger.info(" Total training parameters per FSDP shard = %s B",
|
||||
round(count_trainable(self.transformer) / 1e9, 3))
|
||||
round(_get_trainable_params(self.transformer) / 1e9, 3))
|
||||
# print dtype
|
||||
logger.info(" Master weight dtype: %s",
|
||||
self.transformer.parameters().__next__().dtype)
|
||||
@@ -729,7 +616,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:
|
||||
@@ -751,9 +638,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
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 +730,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,27 @@ 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,
|
||||
generator_ema=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)
|
||||
@@ -275,41 +254,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 +283,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 +335,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 +396,20 @@ 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,
|
||||
generator_ema=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",
|
||||
@@ -655,63 +474,6 @@ def load_distillation_checkpoint(
|
||||
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 +512,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,14 +1298,8 @@ 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)
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class EMA_FSDP:
|
||||
@@ -1638,7 +1323,6 @@ class EMA_FSDP:
|
||||
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
|
||||
@@ -1732,7 +1416,6 @@ class EMA_FSDP:
|
||||
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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -21,9 +21,8 @@ class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
|
||||
with DMD for video generation.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"transformer",
|
||||
"vae",
|
||||
"scheduler", "transformer", "vae", "real_score_transformer",
|
||||
"fake_score_transformer"
|
||||
]
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
@@ -41,10 +40,7 @@ class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
|
||||
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")
|
||||
},
|
||||
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,
|
||||
|
||||
@@ -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