Compare commits

..
Author SHA1 Message Date
SolitaryThinker 352e3c31fe tests 2025-09-10 01:57:42 +00:00
SolitaryThinker 4f5e79c41f update 2025-09-10 01:52:21 +00:00
SolitaryThinker a2d303b067 fix rebase 2025-09-09 23:30:19 +00:00
SolitaryThinker e07111b0de update parquet handling 2025-09-09 23:29:37 +00:00
SolitaryThinker aa49d2a5c8 fix num_inferenc_steps 2025-09-09 23:29:37 +00:00
SolitaryThinker b337d03e82 disable trajectory deocding 2025-09-09 23:29:37 +00:00
SolitaryThinker b92da9e912 hack to get it running 2025-09-09 23:29:36 +00:00
SolitaryThinkerandkevin314 d615271814 add kevin as coauthor
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-09-09 23:29:36 +00:00
SolitaryThinker 720cfe39ca update 2025-09-09 23:29:36 +00:00
SolitaryThinker 03da4b1cdc update 2025-09-09 23:29:36 +00:00
SolitaryThinker ce52c3e87e rename 2025-09-09 23:29:35 +00:00
SolitaryThinker 7b7a895e77 checkpoint 2025-09-09 23:29:35 +00:00
SolitaryThinker 4744ec2c0b checkpoint 2025-09-09 23:29:33 +00:00
129 changed files with 573 additions and 14054 deletions
+1 -12
View File
@@ -198,15 +198,4 @@ steps:
env:
- TEST_TYPE=inference_vmoba
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Unit Tests"
env:
- TEST_TYPE=unit_test
agents:
queue: "default"
queue: "default"
-4
View File
@@ -118,10 +118,6 @@ case "$TEST_TYPE" in
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
+9 -34
View File
@@ -62,8 +62,8 @@ on:
required: false
default: false
type: boolean
run_unit_test:
description: "Run unit-test"
run_nightly_test:
description: "Run nightly-test"
required: false
default: false
type: boolean
@@ -93,7 +93,6 @@ jobs:
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
unit-test: ${{ steps.filter.outputs.unit-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -103,8 +102,6 @@ jobs:
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.10'
- 'docker/Dockerfile.python3.11'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/sliding_tile_attn/**'
@@ -158,9 +155,6 @@ jobs:
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
unit-test:
- 'fastvideo/**'
- *common-paths
encoder-test:
needs: change-filter
@@ -339,42 +333,23 @@ jobs:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
unit-test:
needs: change-filter
nightly-test:
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "unit-test"
gpu_type: "NVIDIA L40S"
gpu_count: 1
job_id: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
# nightly-test:
# if: >-
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
# uses: ./.github/workflows/runpod-test.yml
# with:
# job_id: "nightly-test"
# gpu_type: "NVIDIA A40"
# gpu_count: 4
# volume_size: 100
# disk_size: 100
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
# timeout_minutes: 30
# secrets:
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
# Add other jobs to this list as you create them
-3
View File
@@ -64,6 +64,3 @@ docs/source/distillation/examples/
!docs/source/_static/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
preprocess_output_text/
+1 -3
View File
@@ -20,7 +20,5 @@ setup(
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[
"flash-attn >= 2.7.1",
]
install_requires=[]
)
+2 -10
View File
@@ -6,16 +6,8 @@ import time
import os
import torch
from typing import Tuple
try:
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
from functools import lru_cache
from einops import rearrange
-9
View File
@@ -1,9 +0,0 @@
# VidProm Dataset
From [Self-Forcing](https://github.com/gdhe17/Self-Forcing) repository.
## Download the dataset
```bash
./download_dataset.sh
```
@@ -1,3 +0,0 @@
#! /bin/bash
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,76 +0,0 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
"image_path": null,
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -1,13 +0,0 @@
{
"data": [
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -1,151 +0,0 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29501
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=4
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="data/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
VALIDATION_DATASET_FILE="data/mixkit-64_processed/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--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"
--validation_steps 100
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
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 arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,24 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
+4 -5
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
@@ -17,13 +17,12 @@ def main():
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
ti2v_task=True,
# image_encoder_cpu_offload=False,
)
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
sampling_param.image_path = "test.jpg"
# 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 "
@@ -45,4 +44,4 @@ def main():
if __name__ == "__main__":
main()
main()
@@ -19,15 +19,13 @@ def main():
)
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.num_frames = 81
prompts = [
"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.",
"A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
]
for prompt in prompts:
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
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."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
main()
@@ -15,7 +15,6 @@ torchrun --nproc_per_node=$GPU_NUM \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
@@ -1,99 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-extended-t2v-1-3b/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_70k"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "fixed_wan_ode_init_70k_6e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
--max_train_steps 6000
--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 77
--warp_denoising_step
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 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"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,135 +0,0 @@
#!/bin/bash
#SBATCH --job-name=1e5B2_16kFV_warp_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_warp/Dode_vidprom8b16k_1e-5.out
#SBATCH --error=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# 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"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b-81/"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
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 "Dwarp_vidprom_8b16k_test_warp_1e-5"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "Dwarp_vidprom_8b16k_wan_ode_init_1e-5"
# --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 77
--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 "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--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,131 +0,0 @@
#!/bin/bash
#SBATCH --job-name=ode_vidprom2k
#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_vidprom2k_output/ode_vidprom2k.out
#SBATCH --error=ode_vidprom2k_output/ode_vidprom2k.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# 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"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing/"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
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 "wan_ode_init_vidprom2k"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "vidprom2k_wan_ode_init_5e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
--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 77
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 8
--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 100
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-6
--mixed_precision "bf16"
--checkpointing_steps 2000
--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,132 +0,0 @@
#!/bin/bash
#SBATCH --job-name=ode_crush
#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_crush_output/ode_crush.out
#SBATCH --error=ode_crush_output/ode_crush.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# 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"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
NUM_GPUS=2
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_warp_2"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "2warp_fixed_wan_ode_init_5e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
# --warp_denoising_step
--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 77
# --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 20
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-6
--mixed_precision "bf16"
--checkpointing_steps 2000
--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,103 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export TOKENIZERS_PARALLELISM=false
export MASTER_PORT=29501
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
# DATA_DIR="/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_lmdb_sf/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "Dwarp_vidprom_8b16k_test_warp_1e-5"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "Dwarp_vidprom_8b16k_wan_ode_init_1e-5"
# --resume_from_checkpoint "ode_init_diffusers/"
--warp_denoising_step
--log_visualization
--max_train_steps 10
--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 77
--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 "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--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
--seed 1024
# --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
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,98 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/weka/home/hao.zhang/wl/FastVideo2/data/crush-smol_processed_t2v_1_3b_ode_init_single"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "overfitwan_ode_init_crush_smol"
# --resume_from_checkpoint "ode_init_diffusers/"
--max_train_steps 2001
# --warp_denoising_step
--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 77
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 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 20
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--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
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,100 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "debug_ode_init"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "debug_ode_init"
# --resume_from_checkpoint "ode_init_diffusers/"
--max_train_steps 1000
--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
--warp_denoising_step
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 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 10
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,25 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/vidprom_1.txt"
OUTPUT_DIR="data/ode_vidprom_1_fv/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 1 \
--flush_frequency 1 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory"
@@ -1,76 +0,0 @@
{
"data": [
{
"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,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"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. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -6,8 +6,8 @@ export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_old"
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
@@ -52,7 +52,7 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 50
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
@@ -4,7 +4,7 @@ GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_old/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
@@ -1,6 +1,6 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
GPU_NUM=2 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
@@ -14,7 +14,7 @@ torchrun --nproc_per_node=$GPU_NUM \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 8 \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
@@ -1,94 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_old"
VALIDATION_DATASET_FILE="examples/datasets/crush_smol/validation.json"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_i2v_finetune"
--output_dir "checkpoints/wan_t2v_i2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 20
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 4
--tp_size 1
--hsdp_replicate_dim 2
--hsdp_shard_dim 4
)
# 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"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--enable_gradient_checkpointing_type "full"
--t2v_as_i2v_task True
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_t2v_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,24 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_i2v_1_3b/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 2 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v_ode_trajectory"
+4 -4
View File
@@ -5,6 +5,7 @@ from dataclasses import dataclass
import torch
from einops import rearrange
from flash_attn.bert_padding import pad_input
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
@@ -133,8 +134,6 @@ class VMOBAAttentionImpl(AttentionImpl):
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
from flash_attn.bert_padding import pad_input
self.pad_input = pad_input
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
@@ -170,6 +169,7 @@ class VMOBAAttentionImpl(AttentionImpl):
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
# torch.distributed.breakpoint()
query, chunk_size = process_moba_input(query,
attn_metadata.patch_resolution,
moba_chunk_size)
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = process_moba_output(hidden_states,
attn_metadata.patch_resolution,
moba_chunk_size)
+2 -3
View File
@@ -3,7 +3,6 @@
import torch
import torch.nn as nn
import fastvideo.envs as envs
from fastvideo.attention.selector import backend_name_to_enum, get_attn_backend
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
@@ -37,7 +36,7 @@ class DistributedAttention(nn.Module):
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = torch.bfloat16 if envs.FASTVIDEO_FORCE_ATTN_BF16 else get_compute_dtype()
dtype = get_compute_dtype()
attn_backend = get_attn_backend(
head_size,
dtype,
@@ -222,7 +221,7 @@ class LocalAttention(nn.Module):
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = torch.bfloat16 if envs.FASTVIDEO_FORCE_ATTN_BF16 else get_compute_dtype()
dtype = get_compute_dtype()
attn_backend = get_attn_backend(
head_size,
dtype,
-1
View File
@@ -27,7 +27,6 @@ class DiTArchConfig(ArchConfig):
num_attention_heads: int = 0
num_channels_latents: int = 0
exclude_lora_layers: list[str] = field(default_factory=list)
boundary_ratio: float | None = None
def __post_init__(self) -> None:
if not self._compile_conditions:
@@ -92,9 +92,6 @@ class WanVideoArchConfig(DiTArchConfig):
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
-24
View File
@@ -45,13 +45,10 @@ class PipelineConfig:
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
disable_autocast: bool = False
ti2v_task: bool = False
t2v_as_i2v_task: bool = False
# Model configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
dit_precision: str = "bf16"
dit_forward_precision: str = "bf16"
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
@@ -90,7 +87,6 @@ class PipelineConfig:
# Wan2.2 TI2V parameters
ti2v_task: bool = False
boundary_ratio: float | None = None
# Compilation
# enable_torch_compile: bool = False
@@ -218,24 +214,6 @@ class PipelineConfig:
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
)
# TI2V task
parser.add_argument(
f"--{prefix_with_dot}ti2v-task",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}ti2v_task",
default=PipelineConfig.ti2v_task,
help="Enable TI2V",
)
# T2V to I2V task
parser.add_argument(
f"--{prefix_with_dot}t2v-as-i2v-task",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}t2v_as_i2v_task",
default=PipelineConfig.t2v_as_i2v_task,
help="Enable T2V to I2V task",
)
# Add VAE configuration arguments
from fastvideo.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
@@ -267,9 +245,7 @@ class PipelineConfig:
"""
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
logger.info("WTF model_path: %s", model_path)
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
logger.info("pipeline_config_cls: %s", pipeline_config_cls)
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
-1
View File
@@ -50,7 +50,6 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
# Add other pipeline architecture detectors
}
+8 -15
View File
@@ -82,7 +82,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
default_factory=CLIPVisionConfig)
image_encoder_precision: str = "fp32"
def __post_init__(self) -> None:
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -108,17 +108,19 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
flow_shift: float | None = 5.0
ti2v_task: bool = True
expand_timesteps: bool = True
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self.dit_config.expand_timesteps = self.expand_timesteps
@dataclass
@@ -130,21 +132,12 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
@dataclass
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio
pass
@dataclass
class Wan2_2_I2V_A14B_Config(WanI2V480PConfig):
flow_shift: float | None = 5.0
boundary_ratio: float | None = 0.900
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.boundary_ratio = self.boundary_ratio
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
pass
# =============================================
-7
View File
@@ -40,7 +40,6 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
# TeaCache parameters
enable_teacache: bool = False
@@ -170,12 +169,6 @@ class SamplingParam:
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--boundary-ratio",
type=float,
default=SamplingParam.boundary_ratio,
help="Boundary timestep ratio",
)
parser.add_argument(
"--save-video",
action="store_true",
+4 -8
View File
@@ -144,22 +144,18 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 4.0 # high_noise
guidance_scale_2: float = 3.0 # low_noise
guidance_scale: float = 4.0
guidance_scale_2: float = 3.0
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 3.5 # high_noise
guidance_scale_2: float = 3.5 # low_noise
guidance_scale: float = 3.5
guidance_scale_2: float = 3.5
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
# =============================================
-14
View File
@@ -102,17 +102,3 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # Always 'text' for text-only
])
pyarrow_schema_text_only = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("caption", pa.string()),
])
-43
View File
@@ -1,43 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.dataset.lmdb_utils import get_array_shape_from_lmdb, retrieve_row_from_lmdb
from torch.utils.data import Dataset
import numpy as np
import torch
import lmdb
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/dataset.py
class ODERegressionLMDBDataset(Dataset):
def __init__(self, data_path: str, max_pair: int = int(1e8)):
print(f"data_path: {data_path}")
self.env = lmdb.open(data_path, readonly=True,
lock=False, readahead=False, meminit=False)
self.latents_shape = get_array_shape_from_lmdb(self.env, 'latents')
self.max_pair = max_pair
def __len__(self):
return min(self.latents_shape[0], self.max_pair)
def __getitem__(self, idx):
"""
Outputs:
- prompts: List of Strings
- latents: Tensor of shape (num_denoising_steps, num_frames, num_channels, height, width). It is ordered from pure noise to clean image.
"""
latents = retrieve_row_from_lmdb(
self.env,
"latents", np.float16, idx, shape=self.latents_shape[1:]
)
if len(latents.shape) == 4:
latents = latents[None, ...]
prompts = retrieve_row_from_lmdb(
self.env,
"prompts", str, idx
)
return {
"prompts": prompts,
"ode_latent": torch.tensor(latents, dtype=torch.float32)
}
-75
View File
@@ -1,75 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/lmdb.py
import numpy as np
def get_array_shape_from_lmdb(env, array_name):
with env.begin() as txn:
image_shape = txn.get(f"{array_name}_shape".encode()).decode()
image_shape = tuple(map(int, image_shape.split()))
return image_shape
def store_arrays_to_lmdb(env, arrays_dict, start_index=0):
"""
Store rows of multiple numpy arrays in a single LMDB.
Each row is stored separately with a naming convention.
"""
with env.begin(write=True) as txn:
for array_name, array in arrays_dict.items():
for i, row in enumerate(array):
# Convert row to bytes
if isinstance(row, str):
row_bytes = row.encode()
else:
row_bytes = row.tobytes()
data_key = f'{array_name}_{start_index + i}_data'.encode()
txn.put(data_key, row_bytes)
def process_data_dict(data_dict, seen_prompts):
output_dict = {}
all_videos = []
all_prompts = []
for prompt, video in data_dict.items():
if prompt in seen_prompts:
continue
else:
seen_prompts.add(prompt)
video = video.half().numpy()
all_videos.append(video)
all_prompts.append(prompt)
if len(all_videos) == 0:
return {"latents": np.array([]), "prompts": np.array([])}
all_videos = np.concatenate(all_videos, axis=0)
output_dict['latents'] = all_videos
output_dict['prompts'] = np.array(all_prompts)
return output_dict
def retrieve_row_from_lmdb(lmdb_env, array_name, dtype, row_index, shape=None):
"""
Retrieve a specific row from a specific array in the LMDB.
"""
data_key = f'{array_name}_{row_index}_data'.encode()
with lmdb_env.begin() as txn:
row_bytes = txn.get(data_key)
if dtype == str:
array = row_bytes.decode()
else:
array = np.frombuffer(row_bytes, dtype=dtype)
if shape is not None and len(shape) > 0:
array = array.reshape(shape)
return array
-3
View File
@@ -3,9 +3,6 @@ from typing import Any, cast
import numpy as np
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
-5
View File
@@ -18,7 +18,6 @@ if TYPE_CHECKING:
FASTVIDEO_LOGGING_PREFIX: str = ""
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_FORCE_ATTN_BF16: bool = False
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
@@ -169,10 +168,6 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_TRACE_FUNCTION":
lambda: int(os.getenv("FASTVIDEO_TRACE_FUNCTION", "0")),
# if set, fastvideo will force attention to be computed in bfloat16
"FASTVIDEO_FORCE_ATTN_BF16":
lambda: bool(int(os.getenv("FASTVIDEO_FORCE_ATTN_BF16", "0"))),
# Backend for attention computation
# Available options:
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
+1 -106
View File
@@ -158,7 +158,6 @@ class FastVideoArgs:
"transformer": True,
"vae": True,
})
override_transformer_cls_name: str | None = None
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
@@ -397,12 +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",
)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -612,11 +605,6 @@ class TrainingArgs(FastVideoArgs):
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
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
@@ -639,7 +627,6 @@ class TrainingArgs(FastVideoArgs):
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
@@ -671,7 +658,6 @@ class TrainingArgs(FastVideoArgs):
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
@@ -692,30 +678,16 @@ class TrainingArgs(FastVideoArgs):
# distillation args
generator_update_interval: int = 5
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
intermediate_latents_visualization: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
independent_first_frame: bool = False
enable_gradient_masking: bool = True
gradient_mask_last_n_frames: int = 21
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
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -817,20 +789,6 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Directory to cache models")
# DMD model paths - separate paths for each network
parser.add_argument(
"--generator-model-path",
type=str,
help="Path to generator (student) model for DMD distillation")
parser.add_argument(
"--real-score-model-path",
type=str,
help="Path to real score (teacher) model for DMD distillation")
parser.add_argument(
"--fake-score-model-path",
type=str,
help="Path to fake score (critic) model for DMD distillation")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
@@ -901,10 +859,6 @@ 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")
@@ -1009,10 +963,6 @@ class TrainingArgs(FastVideoArgs):
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--betas",
type=str,
default=TrainingArgs.betas,
help="Betas for optimizer (format: 'beta1,beta2')")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
@@ -1063,13 +1013,6 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument(
"--dfake-gen-update-ratio",
type=int,
default=TrainingArgs.dfake_gen_update_ratio,
help=
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
)
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
@@ -1086,11 +1029,6 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.fake_score_learning_rate,
help="Learning rate for fake score transformer")
parser.add_argument(
"--fake-score-betas",
type=str,
default=TrainingArgs.fake_score_betas,
help="Betas for fake score optimizer (format: 'beta1,beta2')")
parser.add_argument(
"--fake-score-lr-scheduler",
type=str,
@@ -1103,49 +1041,6 @@ class TrainingArgs(FastVideoArgs):
"--simulate-generator-forward",
action=StoreBoolean,
help="Whether to simulate generator forward to match inference")
parser.add_argument(
"--warp-denoising-step",
action=StoreBoolean,
help=
"Whether to warp denoising step according to the scheduler time shift"
)
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=TrainingArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
parser.add_argument(
"--independent-first-frame",
action=StoreBoolean,
help="Whether the first frame is independent in causal generation")
parser.add_argument(
"--enable-gradient-masking",
action=StoreBoolean,
help="Whether to enable frame-level gradient masking")
parser.add_argument(
"--gradient-mask-last-n-frames",
type=int,
default=TrainingArgs.gradient_mask_last_n_frames,
help="Number of last frames to enable gradients for")
parser.add_argument(
"--validate-cache-structure",
action=StoreBoolean,
help="Whether to validate KV cache structure (debug flag)")
parser.add_argument(
"--same-step-across-blocks",
action=StoreBoolean,
help="Whether to use the same exit timestep for all blocks")
parser.add_argument(
"--last-step-only",
action=StoreBoolean,
help="Whether to only use the last timestep for training")
parser.add_argument(
"--context-noise",
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
return parser
@@ -1153,4 +1048,4 @@ class TrainingArgs(FastVideoArgs):
def parse_int_list(value: str) -> list[int]:
if not value:
return []
return [int(x.strip()) for x in value.split(",")]
return [int(x.strip()) for x in value.split(",")]
+4 -4
View File
@@ -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 + scale) + shift).flatten(1, 2)
(1.0 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1 + scale) + shift
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
@@ -267,11 +267,11 @@ class LayerNormScaleShift(nn.Module):
frame_seqlen = normalized.shape[1] // num_frames
output = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale) + shift).flatten(1, 2)
(1.0 + scale) + shift).flatten(1, 2)
else:
# scale.shape: [batch_size, 1, inner_dim]
# shift.shape: [batch_size, 1, inner_dim]
output = normalized * (1 + scale) + shift
output = normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
output = output.to(x.dtype)
+3 -5
View File
@@ -77,11 +77,9 @@ class BaseLayerWithLoRA(nn.Module):
lora_A = self.lora_A.to_local()
if not self.merged and not self.disable_lora:
lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(x, non_blocking=True))
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(x, non_blocking=True))
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
delta = x @ (
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
if self.lora_alpha != self.lora_rank:
delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore
+40 -68
View File
@@ -147,9 +147,6 @@ class CausalWanSelfAttention(nn.Module):
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
# kv_cache["k"] = kv_cache["k"].detach()
# kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
@@ -179,7 +176,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -212,7 +209,8 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -225,7 +223,8 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -250,34 +249,29 @@ class CausalWanTransformerBlock(nn.Module):
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb
e = self.scale_shift_table + temb.float()
# e.shape: [batch_size, num_frames, 6, inner_dim]
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
# assert shift_msa.dtype == torch.float32
# 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)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -291,6 +285,8 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -299,10 +295,13 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -365,7 +364,8 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -375,7 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 3
self.num_frame_per_block = 1
self.independent_first_frame = False
self.__post_init__()
@@ -487,16 +487,12 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -543,9 +539,14 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return torch.stack(output)
return output
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -556,8 +557,6 @@ class CausalWanTransformer3DModel(BaseDiT):
start_frame: int = 0,
**kwargs) -> torch.Tensor:
logger.info("timestep dtype: %s, timestep sum: %s", timestep.dtype, timestep.float().sum().item())
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
@@ -588,8 +587,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
@@ -602,12 +601,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -642,9 +637,14 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return torch.stack(output)
return output
def forward(
self,
@@ -652,34 +652,6 @@ class CausalWanTransformer3DModel(BaseDiT):
**kwargs
):
if kwargs.get('kv_cache', None) is not None:
noise_pred = self._forward_inference(*args, **kwargs)
return self._forward_inference(*args, **kwargs)
else:
noise_pred = self._forward_train(*args, **kwargs)
return noise_pred
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
return self._forward_train(*args, **kwargs)
+76 -93
View File
@@ -1,5 +1,3 @@
import torch
import torch.nn as nn
# SPDX-License-Identifier: Apache-2.0
import math
@@ -39,14 +37,16 @@ class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = nn.LayerNorm(in_features)
self.norm1 = FP32LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = nn.LayerNorm(out_features)
self.norm2 = FP32LayerNorm(out_features)
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
dtype = encoder_hidden_states_image.dtype
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.norm2(hidden_states).to(dtype)
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", freq_dtype=torch.float64)
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
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.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
q = self.norm_q(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.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k(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,16 +169,11 @@ class WanT2VCrossAttention(WanSelfAttention):
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
if envs.FASTVIDEO_FORCE_ATTN_BF16:
out_dtype = v.dtype
# compute attention
x = self.attn(q.to(torch.bfloat16), k.to(torch.bfloat16), v.to(torch.bfloat16)).to(out_dtype)
else:
# compute attention
x = self.attn(q, k, v)
# compute attention
x = self.attn(q, k, v)
# output
x = x.flatten(2)
@@ -218,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.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)
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)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
k_img = self.norm_added_k(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)
@@ -252,7 +247,7 @@ class WanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -283,29 +278,29 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_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)
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -324,11 +319,12 @@ 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
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
@@ -339,20 +335,22 @@ 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
e = self.scale_shift_table + temb.float()
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) * (1 + scale_msa) + shift_msa
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -364,12 +362,7 @@ class WanTransformerBlock(nn.Module):
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
if envs.FASTVIDEO_FORCE_ATTN_BF16:
out_dtype = value.dtype
attn_output, _ = self.attn1(query.to(torch.bfloat16), key.to(torch.bfloat16), value.to(torch.bfloat16))
attn_output = attn_output.to(out_dtype)
else:
attn_output, _ = self.attn1(query, key, value)
attn_output, _ = self.attn1(query, key, value)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
@@ -377,20 +370,26 @@ 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,
@@ -407,7 +406,7 @@ class WanTransformerBlock_VSA(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -439,7 +438,8 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -459,7 +459,8 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -479,22 +480,23 @@ 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
e = self.scale_shift_table + temb.float()
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) *
(1 + scale_msa) + shift_msa)
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
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.forward_native(query)
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -519,6 +521,8 @@ 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,
@@ -526,15 +530,17 @@ 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
@@ -592,7 +598,8 @@ class WanTransformer3DModel(CachableDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -652,12 +659,10 @@ 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,
freqs_sin) if freqs_cos is not None else None
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
@@ -667,8 +672,6 @@ 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:
@@ -725,35 +728,14 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return torch.stack(output)
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
return output
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
@@ -845,4 +827,5 @@ 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
+4 -24
View File
@@ -238,7 +238,6 @@ class TextEncoderLoader(ComponentLoader):
1]
target_device = get_local_torch_device()
logger.info("Loading text encoder in %s precision", encoder_precision)
# TODO(will): add support for other dtypes
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args, encoder_precision)
@@ -416,10 +415,6 @@ class TransformerLoader(ComponentLoader):
raise ValueError(
"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
@@ -435,23 +430,11 @@ class TransformerLoader(ComponentLoader):
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Check if we should use custom initialization weights
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
fastvideo_args.training_mode and
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
if use_custom_weights:
logger.info("Using custom initialization weights from: %s", custom_weights_path)
safetensors_list = [custom_weights_path]
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]
param_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_forward_precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
@@ -471,8 +454,7 @@ class TransformerLoader(ComponentLoader):
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
fsdp_inference=fastvideo_args.use_fsdp_inference,
# TODO(will): make these configurable
default_dtype=default_dtype,
param_dtype=param_dtype,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
training_mode=fastvideo_args.training_mode)
@@ -481,11 +463,9 @@ class TransformerLoader(ComponentLoader):
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
# Need to convert the model to the default_dtype
# Otherwise, the model master weights will be in param_dtype, and the gradients will also be in param_dtype
# This means the param update will be in lower precision, causing precision loss
logger.info("Converting model to dtype: %s", default_dtype)
model = model.to(default_dtype)
dtypes = set(param.dtype for param in model.parameters())
if len(dtypes) > 1:
model = model.to(default_dtype)
model = model.eval()
return model
+2 -3
View File
@@ -62,7 +62,6 @@ def maybe_load_fsdp_model(
device: torch.device,
hsdp_replicate_dim: int,
hsdp_shard_dim: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
@@ -88,7 +87,7 @@ def maybe_load_fsdp_model(
mp_policy=mp_policy,
)
with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(param_dtype), torch.device("meta"):
model = model_cls(**init_params)
# Check if we should use FSDP
@@ -126,7 +125,7 @@ def maybe_load_fsdp_model(
model,
weight_iterator,
device,
default_dtype,
param_dtype,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
@@ -635,31 +635,8 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
noise: torch.Tensor,
timestep: torch.IntTensor,
) -> torch.Tensor:
"""
Args:
clean_latent: the clean latent with shape [B, C, H, W],
where B is batch_size or batch_size * num_frames
noise: the noise with shape [B, C, H, W]
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
Returns:
the corrupted latent with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == clean_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(clean_latent.shape[0])
else:
assert timestep.numel() == clean_latent.shape[0]
else:
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
self.sigmas = self.sigmas.to(noise.device)
timestep = timestep.expand(clean_latent.shape[0])
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
@@ -22,10 +22,8 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
config_name = "scheduler_config.json"
order = 1
@register_to_config
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
@@ -64,15 +62,8 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
elif timestep.ndim == 0:
# handles the case where timestep is a scalar, this occurs when we
# use this scheduler for ODE trajectory
timestep = timestep.unsqueeze(0)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
timestep = timestep.to(model_output.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
+4 -26
View File
@@ -171,34 +171,12 @@ 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)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def pred_video_to_pred_noise(x0_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor, scheduler: Any) -> torch.Tensor:
"""
Convert x0 prediction to flow matching's prediction.
x0_pred: the x0 prediction with shape [B, C, H, W]
xt: the input noisy data with shape [B, C, H, W]
timestep: the timestep with shape [B]
pred = (x_t - x_0) / sigma_t
"""
# use higher precision for calculations
original_dtype = x0_pred.dtype
x0_pred, xt, sigmas, timesteps = map(
lambda x: x.double().to(x0_pred.device), [x0_pred, xt,
scheduler.sigmas,
scheduler.timesteps]
)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
flow_pred = (xt - x0_pred) / sigma_t
return flow_pred.to(original_dtype)
@@ -28,6 +28,10 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
+11 -31
View File
@@ -121,25 +121,14 @@ class ComposedPipelineBase(ABC):
model_path: str,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: PipelineConfig | None = None,
pipeline_config: str | PipelineConfig | None = None,
args: argparse.Namespace | None = None,
required_config_modules: list[str] | None = None,
loaded_modules: dict[str, torch.nn.Module]
| None = None,
**kwargs) -> "ComposedPipelineBase":
"""
Load a pipeline from a pretrained model.
Few different patterns are supported:
- Only provide model_path:
- This will load the pipeline in inference mode.
- The pipeline will be initialized with the default config.
- The pipeline will be initialized with the default modules.
- The pipeline will be initialized with the default stages.
- The pipeline will be initialized with the default stages.
- override the default config using pipeline_config or args or kwargs
- override the default modules using loaded_modules
- override the pipelineconfig
Load a pipeline from a pretrained model.
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
@@ -147,18 +136,9 @@ class ComposedPipelineBase(ABC):
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
if pipeline_config is not None:
fastvideo_args.pipeline_config = pipeline_config
if fastvideo_args.override_transformer_cls_name is not None:
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
fastvideo_args.pipeline_config = pipeline_config
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
if fastvideo_args.override_transformer_cls_name is not None:
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
fastvideo_args.pipeline_config = pipeline_config
logger.info("in 2 Overriding transformer cls name to %s", fastvideo_args.override_transformer_cls_name)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
for key, value in kwargs.items():
@@ -169,8 +149,7 @@ class ComposedPipelineBase(ABC):
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
fastvideo_args.pipeline_config.dit_precision = 'fp32'
# assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
@@ -258,19 +237,20 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# @TODO(Wei): Temporary hack
if "boundary_ratio" in model_index and model_index[
"boundary_ratio"] is not None:
logger.info(
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
)
self.required_config_modules.append("transformer_2")
logger.info("MoE pipeline detected. Setting boundary ratio to %s",
model_index["boundary_ratio"])
fastvideo_args.pipeline_config.dit_config.boundary_ratio = model_index[
"boundary_ratio"]
if fastvideo_args.boundary_ratio is None:
logger.info(
"MoE pipeline detected. Setting boundary ratio to %s",
model_index["boundary_ratio"])
fastvideo_args.boundary_ratio = model_index["boundary_ratio"]
model_index.pop("boundary_ratio", None)
# used by Wan2.2 ti2v
model_index.pop("expand_timesteps", None)
# some sanity checks
@@ -303,8 +283,8 @@ class ComposedPipelineBase(ABC):
architecture) in model_index.items():
if transformers_or_diffusers is None:
logger.warning(
"Module %s in model_index.json has null value, removing from required_config_modules",
module_name)
"Module in model_index.json has null value, removing from required_config_modules"
)
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
+1 -3
View File
@@ -129,7 +129,6 @@ class ForwardBatch:
timesteps: torch.Tensor | None = None
timestep: torch.Tensor | float | int | None = None
step_index: int | None = None
boundary_ratio: float | None = None
# Scheduler parameters
num_inference_steps: int = 50
@@ -151,7 +150,7 @@ class ForwardBatch:
output: torch.Tensor | None = None
return_trajectory_latents: bool = False
return_trajectory_decoded: bool = False
trajectory_timesteps: list[int] | None = None
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None
@@ -246,7 +245,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, torch.Tensor] = field(default_factory=dict)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
@@ -10,8 +10,6 @@ from torch.utils.data import DataLoader
from tqdm import tqdm
from fastvideo.dataset import getdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.preprocessing_datasets import PreprocessBatch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -19,6 +17,8 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import TextEncodingStage
from fastvideo.workflow.preprocess.parquet_io import (ParquetDatasetWriter,
records_to_table)
logger = init_logger(__name__)
@@ -423,3 +423,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
num_processed_samples = 0
def _final_flush_if_any(self):
if hasattr(self, 'dataset_writer'):
self.dataset_writer.flush()
@@ -15,17 +15,12 @@ from typing import Any
import numpy as np
import pyarrow as pa
import torch
from PIL import Image
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import (
ode_text_only_record_creator)
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.fastvideo_args import FastVideoArgs
@@ -41,11 +36,8 @@ from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
from fastvideo.forward_context import set_forward_context
from fastvideo.models.utils import pred_noise_to_pred_video, pred_video_to_pred_noise
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.distributed import get_local_torch_device
from fastvideo.workflow.preprocess.parquet_io import (ParquetDatasetWriter,
records_to_table)
logger = init_logger(__name__)
@@ -67,22 +59,26 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
logger.info('WTF flow_shift: %s',
fastvideo_args.pipeline_config.flow_shift)
assert fastvideo_args.pipeline_config.flow_shift == 5
# self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
# shift=fastvideo_args.pipeline_config.flow_shift)
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=48,
denoising_strength=1.0)
fastvideo_args.model_loaded["transformer"] = False
loader = TransformerLoader()
fastvideo_args.pipeline_config.dit_precision = "fp32" # Overwrite the precision to fp32 for transformer
fastvideo_args.pipeline_config.dit_forward_precision = "fp32"
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
self.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
# logger.info('WTF scheduler timesteps: %s',
# self.modules["scheduler"].timesteps)
# scheduler = FlowMatchScheduler(
# shift=8.0, sigma_min=0.0, extra_one_step=True)
# device = get_local_torch_device()
# # scheduler.num_train_timesteps = 100
# scheduler.set_timesteps(num_inference_steps=50, denoising_strength=1.0)
# scheduler.sigmas = scheduler.sigmas.to(device)
# self.modules["scheduler"] = scheduler
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
@@ -112,7 +108,6 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""Preprocess text-only data and generate trajectory information."""
for batch_idx, data in enumerate(self.pbar):
# logger.info("transformer weight sum: %s", sum(p.float().sum().item() for p in self.transformer.parameters()))
if data is None:
continue
@@ -143,7 +138,6 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
]
batch_captions = valid_data["text"]
self.prompt_encoding_stage.text_encoders[0] = self.prompt_encoding_stage.text_encoders[0].to(dtype=torch.bfloat16).to(dtype=torch.float32)
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
@@ -152,25 +146,22 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
logger.info("prompt_embeds sum: %s, prompt_embeds shape: %s, prompt_embeds dtype: %s", prompt_embeds.float().sum(), prompt_embeds.shape, prompt_embeds.dtype)
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
negative_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
negative_prompt,
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0]
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0]
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
@@ -195,131 +186,30 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
]
batch.num_inference_steps = 48
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 3.0
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
# result_batch = self.latent_preparation_stage(
# result_batch, fastvideo_args)
# result_batch = self.denoising_stage(result_batch,
# fastvideo_args)
noisy_input = []
# latents = result_batch.latents.permute(0, 2, 1, 3, 4)
latents = torch.randn(
[1, 21, 16, 60, 104], dtype=torch.float32, device=get_local_torch_device()
)
# logger.info("transformer weight sum: %s", sum(p.float().sum().item() for p in self.transformer.parameters()))
logger.info("latents sum: %s, latents shape: %s, latents dtype: %s", latents.float().sum(), latents.shape, latents.dtype)
logger.info("scheduler timesteps: %s", self.get_module("scheduler").timesteps)
for progress_id, t in enumerate(tqdm(self.get_module("scheduler").timesteps)):
timestep = t * \
torch.ones([1, 21], device=latents.device, dtype=torch.float32)
noisy_input.append(latents)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=None,
):
# logger.info("prompt_embed sum: %s, prompt_embed shape: %s, prompt_embed dtype: %s", prompt_embed.float().sum(), prompt_embed.shape, prompt_embed.dtype)
# logger.info("timestep: %s", timestep[:, 0])
# Run transformer
cond_pred_noise_btchw = self.transformer(
hidden_states=latents.permute(0, 2, 1, 3, 4),
encoder_hidden_states=prompt_embed,
timestep=timestep[:, 0]
).permute(0, 2, 1, 3, 4)
# logger.info("cond_pred_noise_btchw sum: %s, cond_pred_noise_btchw shape: %s, cond_pred_noise_btchw dtype: %s", cond_pred_noise_btchw.float().sum(), cond_pred_noise_btchw.shape, cond_pred_noise_btchw.dtype)
cond_pred_video_btchw = pred_noise_to_pred_video(
pred_noise=cond_pred_noise_btchw.flatten(0, 1),
noise_input_latent=latents.flatten(0, 1),
timestep=timestep.flatten(0, 1),
scheduler=self.get_module("scheduler")).unflatten(
0, cond_pred_noise_btchw.shape[:2])
# logger.info("cond_pred_video_btchw sum: %s, cond_pred_video_btchw shape: %s, cond_pred_video_btchw dtype: %s", cond_pred_video_btchw.float().sum(), cond_pred_video_btchw.shape, cond_pred_video_btchw.dtype)
with set_forward_context(
current_timestep=t,
attn_metadata=None,
forward_batch=result_batch,
):
# logger.info("latents sum: %s, latents shape: %s, latents dtype: %s", latents.float().sum(), latents.shape, latents.dtype)
# Run transformer
uncond_pred_noise_btchw = self.transformer(
latents.permute(0, 2, 1, 3, 4),
negative_prompt_embed,
timestep[:, 0]
).permute(0, 2, 1, 3, 4)
# logger.info("uncond_pred_noise_btchw sum: %s, uncond_pred_noise_btchw shape: %s, uncond_pred_noise_btchw dtype: %s", uncond_pred_noise_btchw.float().sum(), uncond_pred_noise_btchw.shape, uncond_pred_noise_btchw.dtype)
uncond_pred_video_btchw = pred_noise_to_pred_video(
pred_noise=uncond_pred_noise_btchw.flatten(0, 1),
noise_input_latent=latents.flatten(0, 1),
timestep=timestep.flatten(0, 1),
scheduler=self.get_module("scheduler")).unflatten(
0, uncond_pred_noise_btchw.shape[:2])
pred_video_btchw = uncond_pred_video_btchw + batch.guidance_scale * (
cond_pred_video_btchw - uncond_pred_video_btchw
)
# logger.info("pred_video_btchw sum: %s, pred_video_btchw shape: %s, pred_video_btchw dtype: %s", pred_video_btchw.float().sum(), pred_video_btchw.shape, pred_video_btchw.dtype)
pred_noise_btchw = pred_video_to_pred_noise(
x0_pred=pred_video_btchw.flatten(0, 1),
xt=latents.flatten(0, 1),
timestep=timestep.flatten(0, 1),
scheduler=self.get_module("scheduler")).unflatten(
0, pred_video_btchw.shape[:2])
# logger.info("pred_noise_btchw sum: %s, pred_noise_btchw shape: %s, pred_noise_btchw dtype: %s", pred_noise_btchw.float().sum(), pred_noise_btchw.shape, pred_noise_btchw.dtype)
latents = self.get_module("scheduler").step(
pred_noise_btchw.flatten(0, 1),
self.get_module("scheduler").timesteps[progress_id] * torch.ones(
[1, 21], device=latents.device, dtype=torch.long).flatten(0, 1),
latents.flatten(0, 1)
)[0].unflatten(dim=0, sizes=pred_noise_btchw.shape[:2])
# logger.info("latents sum: %s, latents shape: %s, latents dtype: %s", latents.float().sum(), latents.shape, latents.dtype)
noisy_input.append(latents)
noisy_inputs = torch.stack(noisy_input, dim=1)
noisy_inputs = noisy_inputs[:, [0, 12, 24, 36, -1]].half()
logger.info("noisy inputs sum: %s, noisy inputs shape: %s, noisy inputs dtype: %s", noisy_inputs.float().sum(), noisy_inputs.shape, noisy_inputs.dtype)
result_batch.trajectory_latents = noisy_inputs.permute(0, 1, 3, 2, 4, 5)
result_batch.trajectory_timesteps = torch.tensor([self.get_module("scheduler").timesteps[i] for i in [0, 12, 24, 36, -1]])
result_batch.latents = latents.permute(0, 2, 1, 3, 4)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
trajectory_latents = torch.stack(trajectory_latents, dim=0).squeeze(0)
# trajecotry_latents = trajectory_latents[:, [0, 12, 24, 36, -1]]
# Prepare extra features for text-only processing
extra_features = {
"trajectory_latents": trajectory_latents,
@@ -335,7 +225,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
@@ -364,16 +254,14 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (text-only ODE schema)
record: dict[str, Any] = ode_text_only_record_creator(
# Create record for Parquet dataset (without VAE latents for text-only)
record = self.create_text_only_record(
args,
video_name=video_name,
text_embedding=text_embedding,
caption=valid_data["text"][idx],
trajectory_latents=sample_extra_features[
"trajectory_latents"],
trajectory_timesteps=sample_extra_features[
"trajectory_timesteps"],
)
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
@@ -401,10 +289,78 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush(write_remainder=True)
written = self.dataset_writer.flush()
if written:
logger.info("Final flush wrote %s samples", written)
def create_text_only_record(
self,
args,
video_name: str,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{video_name}_{idx}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "text",
}
assert extra_features is not None, "extra_features is required"
assert "trajectory_latents" in extra_features, "trajectory_latents is required"
assert "trajectory_timesteps" in extra_features, "trajectory_timesteps is required"
# Add trajectory data if available
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features[
"trajectory_latents"][idx] if isinstance(
extra_features["trajectory_latents"],
list) else extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes":
trajectory_latents.tobytes(),
"trajectory_latents_shape":
list(trajectory_latents.shape),
"trajectory_latents_dtype":
str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features[
"trajectory_timesteps"][idx] if isinstance(
extra_features["trajectory_timesteps"],
list) else extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes":
trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape":
list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype":
str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
@@ -1,184 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Text-only Data Preprocessing pipeline implementation.
This module contains an implementation of the Text-only Data Preprocessing pipeline
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
"""
import os
from collections.abc import Iterator
from typing import Any
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import text_only_record_creator
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import TextEncodingStage
logger = init_logger(__name__)
class PreprocessPipeline_Text(BasePreprocessPipeline):
"""Text-only preprocessing pipeline implementation."""
_required_config_modules = ["text_encoder", "tokenizer"]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int = 0
def get_pyarrow_schema(self):
"""Return the PyArrow schema for text-only pipeline."""
return pyarrow_schema_text_only
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
def preprocess_text_only(self, fastvideo_args: FastVideoArgs, args):
"""Preprocess text-only data."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
logger.info("===== prompt_embeds: %s", prompt_embeds.shape)
logger.info("===== prompt_attention_masks: %s",
prompt_attention_masks.shape)
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, text_path in save_pbar:
text_name = os.path.basename(text_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Create record for Parquet dataset (text-only schema)
record = text_only_record_creator(
text_name=text_name,
text_embedding=text_embedding,
caption=valid_data["text"][idx],
)
batch_data.append(record)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
pyarrow_schema_text_only)
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self.num_processed_samples = 0
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush(write_remainder=True)
if written:
logger.info("Final flush wrote %s samples", written)
# Text-only record creation moved to fastvideo.dataset.dataloader.record_schema
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading text dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for text preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing text",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.text_data: dict[str, Any] = {} # Store text metadata and paths
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text
@@ -1,6 +1,5 @@
import argparse
import os
from typing import Any
from fastvideo import PipelineConfig
from fastvideo.configs.models.vaes import WanVAEConfig
@@ -14,8 +13,6 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -26,22 +23,13 @@ def main(args) -> None:
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs: dict[str, Any] = {}
if args.preprocess_task == "text_only":
kwargs = {
"text_encoder_cpu_offload": False,
}
else:
# Full config for video/image processing
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
"flow_shift": 5,
}
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
@@ -54,19 +42,13 @@ def main(args) -> None:
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "text_only":
PreprocessPipeline = PreprocessPipeline_Text
elif args.preprocess_task == "ode_trajectory":
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}")
logger.info("Preprocess task: %s using %s", args.preprocess_task,
PreprocessPipeline.__name__)
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -105,12 +87,7 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--flow_shift", type=float, default=None)
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
help="Type of preprocessing task to run")
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
@@ -78,8 +78,6 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
else:
assert False, "warp_denoising_step must be true"
timesteps = timesteps.to(get_local_torch_device())
logger.info("Using timesteps: %s", timesteps)
+12 -45
View File
@@ -140,12 +140,11 @@ class DenoisingStage(PipelineStage):
latents = latents[:, :, rank_in_sp_group, :, :, :]
batch.latents = latents
if batch.image_latent is not None:
if not fastvideo_args.pipeline_config.ti2v_task and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
image_latent = rearrange(batch.image_latent,
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
batch.image_latent = image_latent
image_latent = rearrange(batch.image_latent,
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
batch.image_latent = image_latent
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
@@ -205,14 +204,8 @@ class DenoisingStage(PipelineStage):
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_ratio is not None:
logger.info("Overriding boundary ratio from %s to %s",
boundary_ratio, batch.boundary_ratio)
boundary_ratio = batch.boundary_ratio
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
if fastvideo_args.boundary_ratio is not None:
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
@@ -254,7 +247,8 @@ class DenoisingStage(PipelineStage):
patch_size[2])
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
trajectory_timesteps: list[int] = []
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# Run denoising loop
@@ -284,27 +278,14 @@ class DenoisingStage(PipelineStage):
# Expand latents for I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
if batch.image_latent is not None:
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
elif batch.image_latent is not None and fastvideo_args.pipeline_config.t2v_as_i2v_task:
assert batch.image_latent is not None, "image latents should be provided for T2V to I2V task"
if rank_in_sp_group == 0:
logger.info("latent_model_input.shape: %s",
latent_model_input.shape)
latent_model_input = torch.cat([
batch.image_latent,
latent_model_input[:, :, 1:, :, :],
],
dim=2).to(target_dtype)
logger.info("latent_model_input.shape: %s",
latent_model_input.shape)
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
timestep = torch.stack([t]).to(get_local_torch_device())
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
@@ -319,13 +300,6 @@ class DenoisingStage(PipelineStage):
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
if rank_in_sp_group == 0:
latent_model_input = torch.cat([
batch.image_latent,
latent_model_input[:, :, 1:, :, :],
],
dim=2).to(target_dtype)
# Prepare inputs for transformer
guidance_expand = (
@@ -488,17 +462,10 @@ class DenoisingStage(PipelineStage):
trajectory_tensor = sequence_model_parallel_all_gather(
trajectory_tensor, dim=3)
if trajectory_tensor is not None:
batch.trajectory_timesteps = torch.tensor(trajectory_timesteps).cpu()
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
latents = torch.cat([
batch.image_latent,
latents[:, :, 1:, :, :],
],
dim=2)
# Update batch with final latents
batch.latents = latents
+48 -94
View File
@@ -105,81 +105,6 @@ class ImageVAEEncodingStage(PipelineStage):
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def encode_image(self,
image: PIL.Image.Image,
height: int,
width: int,
fastvideo_args: FastVideoArgs,
generator: torch.Generator | None = None) -> torch.Tensor:
"""
Encode image into latent space.
"""
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=height,
width=width).to(get_local_torch_device(), dtype=torch.float32)
# (B, C, H, W) -> (B, C, 1, H, W)
print(f"image.shape: {image.shape}")
image = image.unsqueeze(2)
print(f"after unsqueeze image.shape: {image.shape}")
return self.encode_tensor(image, fastvideo_args, generator)
def encode_tensor(self,
video_condition: torch.Tensor,
fastvideo_args: FastVideoArgs,
generator: torch.Generator | None = None) -> torch.Tensor:
"""
Encode frames into latent space.
"""
self.vae = self.vae.to(get_local_torch_device())
video_condition = video_condition.to(device=get_local_torch_device(),
dtype=torch.float32)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Encode Image
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
latent_condition = encoder_output.mean
else:
generator = generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latent_condition -= self.vae.shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.vae.scaling_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition = latent_condition * self.vae.scaling_factor
return latent_condition
def forward(
self,
batch: ForwardBatch,
@@ -232,28 +157,57 @@ class ImageVAEEncodingStage(PipelineStage):
# (B, C, H, W) -> (B, C, 1, H, W)
image = image.unsqueeze(2)
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
# repeat the image self.vae.temporal_compression_ratio times
video_condition = image.repeat(1, 1,
self.vae.temporal_compression_ratio,
1, 1)
# video_condition = image
logger.info("video_condition.shape: %s", video_condition.shape)
else:
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
image.shape[3], image.shape[4])
],
dim=2)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
image.shape[3], image.shape[4])
],
dim=2)
video_condition = video_condition.to(device=get_local_torch_device(),
dtype=torch.float32)
latent_condition = self.encode_tensor(video_condition, fastvideo_args,
batch.generator)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Encode Image
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
latent_condition = encoder_output.mean
else:
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latent_condition -= self.vae.shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.vae.scaling_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition = latent_condition * self.vae.scaling_factor
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
batch.image_latent = latent_condition
elif fastvideo_args.pipeline_config.t2v_as_i2v_task:
logger.info("latent_condition.shape: %s", latent_condition.shape)
batch.image_latent = latent_condition
else:
mask_lat_size = torch.ones(1, 1, num_frames, latent_height,
@@ -35,15 +35,9 @@ class InputValidationStage(PipelineStage):
"""Generate seeds for the inference"""
seed = batch.seed
num_videos_per_prompt = batch.num_videos_per_prompt
if isinstance(batch.prompt, list):
num_prompts = len(batch.prompt)
else:
num_prompts = 1
total_num_videos = num_prompts * num_videos_per_prompt
assert seed is not None
seeds = [seed + i for i in range(total_num_videos)]
seeds = [seed + i for i in range(num_videos_per_prompt)]
batch.seeds = seeds
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
batch.generator = [
@@ -3,7 +3,6 @@
Latent preparation stage for diffusion pipelines.
"""
import torch
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.distributed import get_local_torch_device
@@ -94,7 +93,7 @@ class LatentPreparationStage(PipelineStage):
# Generate or use provided latents
if latents is None:
latents = randn_tensor(shape,
generator=torch.Generator(device="cuda").manual_seed(1024),
generator=generator,
device=device,
dtype=dtype)
else:
+2 -2
View File
@@ -82,8 +82,8 @@ def rocm_platform_plugin() -> str | None:
logger.info("ROCm platform is available")
finally:
amdsmi.amdsmi_shut_down()
except Exception:
pass
except Exception as e:
logger.info("ROCm platform is unavailable: %s", e)
return "fastvideo.platforms.rocm.RocmPlatform" if is_rocm else None
@@ -1,123 +0,0 @@
import numpy as np
from fastvideo.dataset.dataloader.record_schema import (
basic_t2v_record_creator,
i2v_record_creator,
ode_text_only_record_creator,
text_only_record_creator,
)
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
def _mk_basic_batch(N: int) -> PreprocessBatch:
batch = PreprocessBatch(data_type="video")
batch.video_file_name = [f"vid_{i}" for i in range(N)]
batch.prompt = [f"caption_{i}" for i in range(N)]
batch.width = [640 for _ in range(N)]
batch.height = [360 for _ in range(N)]
batch.fps = [4 for _ in range(N)]
batch.num_frames = [2 for _ in range(N)]
# Latents: shape (N, C, T, H, W); per-record use latents[idx]
batch.latents = np.zeros((N, 4, 2, 8, 8), dtype=np.float32)
# Prompt embeds: list of per-record arrays [Seq, Dim]
batch.prompt_embeds = [np.ones((6, 16), dtype=np.float32) for _ in range(N)]
return batch
def test_basic_t2v_record_creator_fields():
N = 2
batch = _mk_basic_batch(N)
records = basic_t2v_record_creator(batch)
assert isinstance(records, list) and len(records) == N
for i, rec in enumerate(records):
assert rec["id"] == batch.video_file_name[i]
# Latents bytes/shape/dtype
assert isinstance(rec["vae_latent_bytes"], (bytes, bytearray))
assert rec["vae_latent_shape"] == list(batch.latents[i].shape)
assert rec["vae_latent_dtype"] == str(batch.latents[i].dtype)
# Text embedding
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(batch.prompt_embeds[i].shape)
assert rec["text_embedding_dtype"] == str(batch.prompt_embeds[i].dtype)
# Meta
assert rec["caption"] == batch.prompt[i]
assert rec["media_type"] == "video"
assert rec["width"] == int(batch.width[i])
assert rec["height"] == int(batch.height[i])
assert rec["num_frames"] == batch.latents[i].shape[1]
def test_i2v_record_creator_additional_fields():
N = 3
batch = _mk_basic_batch(N)
# image_embeds is a list of length 1, with an array of shape [N, D]
batch.image_embeds = [np.ones((N, 32), dtype=np.float32)]
# first frame latent per record
batch.image_latent = np.zeros((N, 4, 1, 8, 8), dtype=np.float32)
# pil image per record
batch.pil_image = np.zeros((N, 8, 8, 3), dtype=np.uint8)
records = i2v_record_creator(batch)
assert isinstance(records, list) and len(records) == N
for i, rec in enumerate(records):
# clip feature
assert isinstance(rec["clip_feature_bytes"], (bytes, bytearray))
assert rec["clip_feature_shape"] == list(batch.image_embeds[0][i].shape)
assert rec["clip_feature_dtype"] == str(batch.image_embeds[0][i].dtype)
# first frame latent
assert isinstance(rec["first_frame_latent_bytes"], (bytes, bytearray))
assert rec["first_frame_latent_shape"] == list(batch.image_latent[i].shape)
assert rec["first_frame_latent_dtype"] == str(batch.image_latent[i].dtype)
# pil image
assert isinstance(rec["pil_image_bytes"], (bytes, bytearray))
assert rec["pil_image_shape"] == list(batch.pil_image[i].shape)
assert rec["pil_image_dtype"] == str(batch.pil_image[i].dtype)
def test_ode_text_only_record_creator():
video_name = "ex"
caption = "a prompt"
text_embedding = np.ones((6, 16), dtype=np.float32)
traj = np.ones((5, 4, 2, 2), dtype=np.float32)
tsteps = np.arange(5, dtype=np.float32)
rec = ode_text_only_record_creator(
video_name=video_name,
text_embedding=text_embedding,
caption=caption,
trajectory_latents=traj,
trajectory_timesteps=tsteps,
)
assert rec["id"] == f"text_{video_name}"
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(text_embedding.shape)
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
assert rec["file_name"] == video_name
assert rec["caption"] == caption
assert rec["media_type"] == "text"
# Trajectory fields
assert isinstance(rec["trajectory_latents_bytes"], (bytes, bytearray))
assert rec["trajectory_latents_shape"] == list(traj.shape)
assert rec["trajectory_latents_dtype"] == str(traj.dtype)
assert isinstance(rec["trajectory_timesteps_bytes"], (bytes, bytearray))
assert rec["trajectory_timesteps_shape"] == list(tsteps.shape)
assert rec["trajectory_timesteps_dtype"] == str(tsteps.dtype)
def test_text_only_record_creator():
text_name = "note1"
caption = "a prompt"
text_embedding = np.ones((7, 16), dtype=np.float32)
rec = text_only_record_creator(
text_name=text_name,
text_embedding=text_embedding,
caption=caption,
)
assert rec["id"] == f"text_{text_name}"
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(text_embedding.shape)
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
assert rec["caption"] == caption
+20 -47
View File
@@ -8,9 +8,6 @@ from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
from fastvideo.wan.modules.tokenizers import HuggingfaceTokenizer
from fastvideo.wan.modules.t5 import umt5_xxl
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
@@ -41,22 +38,9 @@ def test_t5_encoder():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision_str = "fp32"
precision = PRECISION_TO_TYPE[precision_str]
# model1 = UMT5EncoderModel.from_pretrained(TEXT_ENCODER_PATH).to(
# precision).to(device).eval()
# tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
model1 = umt5_xxl(
encoder_only=True,
return_tokenizer=False,
dtype=torch.float32,
device=device,
).eval().requires_grad_(False)
model1.load_state_dict(
torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing-clean/wan_models/Wan2.1-T2V-1.3B/models_t5_umt5-xxl-enc-bf16.pth",
map_location='cpu', weights_only=False)
)
tokenizer1 = HuggingfaceTokenizer(
name="/mnt/weka/home/hao.zhang/wei/Self-Forcing-clean/wan_models/Wan2.1-T2V-1.3B/google/umt5-xxl/", seq_len=512, clean='whitespace')
model1 = UMT5EncoderModel.from_pretrained(TEXT_ENCODER_PATH).to(
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH,
@@ -65,9 +49,8 @@ def test_t5_encoder():
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(TEXT_ENCODER_PATH, args)
model2 = model2.to(dtype=torch.bfloat16).to(precision)
model2 = model2.to(precision)
model2.eval()
tokenizer2 = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
@@ -78,13 +61,8 @@ def test_t5_encoder():
logger.info("Model1 has %s parameters", len(params1))
logger.info("Model2 has %s parameters", len(params2))
model1_weight_sum = sum(p.float().sum().item() for p in model1.parameters())
model2_weight_sum = sum(p.float().sum().item() for p in model2.parameters())
logger.info("Model1 weight sum: %s", model1_weight_sum)
logger.info("Model2 weight sum: %s", model2_weight_sum)
# weight_diffs = []
# # check if embed_tokens are the same
weight_diffs = []
# check if embed_tokens are the same
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
@@ -92,21 +70,18 @@ def test_t5_encoder():
for idx in range(hf_config.num_hidden_layers):
for w in weights:
# name1 = w.format(idx)
name1 = w.format(idx)
name2 = w.format(idx)
# p1 = params1[name1]
p1 = params1[name1]
p2 = params2[name2]
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(device)
# assert_close(p1, p2, atol=1e-4, rtol=1e-4)
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
# Test with some sample prompts
# prompts = [
# "Once upon a time", "The quick brown fox jumps over",
# "In a galaxy far, far away"
# ]
prompts = [
"A vibrant scene of Kenyan golfers at a lush green golf course on a sunny day. The golfers, dressed in casual yet stylish attire, are teeing off with animated expressions, showcasing their enthusiasm for the game. Rolling hills and pristine greens stretch out behind them, creating a picturesque backdrop. In the foreground, a golf buggy and a caddy stand ready, adding to the serene atmosphere. The camera captures the action from a mid-shot angle, focusing on the golfers' dynamic motions as they swing their clubs."
"Once upon a time", "The quick brown fox jumps over",
"In a galaxy far, far away"
]
logger.info("Testing T5 encoder with sample prompts")
@@ -116,8 +91,7 @@ def test_t5_encoder():
logger.info("Testing prompt: %s", prompt)
# Tokenize the prompt
tokens1, mask = tokenizer1(prompt, return_mask=True, add_special_tokens=True)
tokens2 = tokenizer2(prompt,
tokens = tokenizer(prompt,
padding="max_length",
max_length=512,
truncation=True,
@@ -127,23 +101,22 @@ def test_t5_encoder():
# filter out padding input_ids
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
outputs1 = model1(tokens1.to(device),
mask.to(device))
outputs1 = model1(input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True).last_hidden_state
print("--------------------------------")
logger.info("Testing model2")
# Get outputs from our implementation
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs2 = model2(
input_ids=tokens2.input_ids,
attention_mask=tokens2.attention_mask,
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
).last_hidden_state
# Compare last hidden states
last_hidden_state1 = outputs1[mask == 1]
last_hidden_state2 = outputs2[tokens2.attention_mask == 1]
logger.info("last_hidden_state1 sum: %s", last_hidden_state1.float().sum())
logger.info("last_hidden_state2 sum: %s", last_hidden_state2.float().sum())
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
-4
View File
@@ -117,7 +117,3 @@ def run_inference_lora_tests():
@app.function(gpu="L40S:2", image=image, timeout=900)
def run_distill_dmd_tests():
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
@@ -1,292 +0,0 @@
# 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
@@ -1,144 +0,0 @@
# 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"
os.environ["FASTVIDEO_FORCE_ATTN_BF16"] = "1"
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.float32
precision_str = "fp32"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str, dit_forward_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args)
model1 = WanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing-clean/wan_models/Wan2.1-T2V-1.3B")
model1.eval()
model1 = model1.to(device).to(precision)
model1.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 = 120
seq_len = math.ceil((104 * 60) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
21,
16,
60,
104,
device=device,
generator=torch.Generator("cuda").manual_seed(1024),
dtype=precision)
logger.info("Hidden states sum: %s, Hidden states shape: %s, Hidden states dtype: %s", hidden_states.float().sum(), hidden_states.shape, hidden_states.dtype)
# Text embeddings [B, L, D] (including global token)
# encoder_hidden_states = torch.randn(batch_size,
# text_seq_len + 1,
# 4096,
# device=device,
# dtype=precision)
encoder_hidden_states = torch.load("../sf_cond_prompt_embeds.pt").to(device, dtype=precision)
logger.info("Encoder hidden states sum: %s, Encoder hidden states shape: %s, Encoder hidden states dtype: %s", encoder_hidden_states.float().sum(), encoder_hidden_states.shape, encoder_hidden_states.dtype)
# Timestep
timestep = torch.tensor([995.7627], device=device, dtype=precision)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states.permute(0, 2, 1, 3, 4),
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
).permute(0, 2, 1, 3, 4)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states.permute(0, 2, 1, 3, 4),
encoder_hidden_states=encoder_hidden_states,
timestep=timestep).permute(0, 2, 1, 3, 4)
# 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())
logger.info("Output 1 sum: %s", output1.float().sum())
logger.info("Output 2 sum: %s", output2.float().sum())
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()}"
@@ -1,144 +0,0 @@
# 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()}"
@@ -38,8 +38,7 @@ def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
data_type="video",
latents=torch.randn(B, 2),
prompt_embeds=[torch.randn(B, 1, 1)],
# Attention mask should be integer dtype in real pipelines
prompt_attention_mask=[torch.ones(B, 1, dtype=torch.int64)],
prompt_attention_mask=[torch.ones(B, 1)],
)
batch.video_file_name = [f"vid_{i}" for i in range(B)]
@@ -53,14 +52,16 @@ def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
out_dir = tmp_path / "saver_out"
saver.save_and_write_parquet_batch(batch, str(out_dir))
# First flush: should write one full file (3 rows), keep 2 in buffer
saver.flush_tables()
saver.flush_tables(str(out_dir))
files = sorted(out_dir.rglob("*.parquet"))
assert len(files) == 1
assert pq.read_table(str(files[0])).num_rows == 3
# Final flush: write remainder 2 rows
saver.flush_tables(write_remainder=True)
saver.flush_last(str(out_dir))
files2 = sorted(out_dir.rglob("*.parquet"))
assert len(files2) == 2
total = sum(pq.read_table(str(f)).num_rows for f in files2)
assert total == 5
+87 -412
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import json
import os
import time
from abc import abstractmethod
@@ -12,7 +11,6 @@ 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
@@ -38,11 +36,10 @@ from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
clip_grad_norm_while_handling_failing_dtensor_cases, count_trainable,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
from fastvideo.utils import is_vsa_available, set_random_seed
import wandb # isort: skip
@@ -72,11 +69,18 @@ class DistillationPipeline(TrainingPipeline):
current_trainstep: int
video_latent_shape: tuple[int, ...]
video_latent_shape_sp: tuple[int, ...]
real_score_transformer: torch.nn.Module
fake_score_transformer: torch.nn.Module
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_trainable(self) -> None:
super().set_trainable()
self.modules["real_score_transformer"].requires_grad_(False)
self.modules["vae"].requires_grad_(False)
def initialize_training_pipeline(self, training_args: TrainingArgs):
"""Initialize the distillation training pipeline with multiple models."""
logger.info("Initializing distillation pipeline...")
@@ -85,37 +89,14 @@ class DistillationPipeline(TrainingPipeline):
self.noise_scheduler = self.get_module("scheduler")
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
if training_args.real_score_model_path:
logger.info(
f"Loading real score transformer from: {training_args.real_score_model_path}"
)
self.real_score_transformer = self.load_module_from_path(
training_args.real_score_model_path, "transformer",
training_args)
else:
self.real_score_transformer = self.get_module(
"real_score_transformer")
if training_args.fake_score_model_path:
logger.info(
f"Loading fake score transformer from: {training_args.fake_score_model_path}"
)
self.fake_score_transformer = self.load_module_from_path(
training_args.fake_score_model_path, "transformer",
training_args)
else:
self.fake_score_transformer = self.get_module(
"fake_score_transformer")
self.real_score_transformer.requires_grad_(False)
# self.transformer is the generator model
self.real_score_transformer = self.get_module("real_score_transformer")
self.fake_score_transformer = self.get_module("fake_score_transformer")
self.real_score_transformer.eval()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
@@ -138,13 +119,10 @@ class DistillationPipeline(TrainingPipeline):
if fake_score_lr == 0.0:
fake_score_lr = training_args.learning_rate
betas_str = training_args.fake_score_betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=fake_score_lr,
betas=betas,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -172,19 +150,8 @@ class DistillationPipeline(TrainingPipeline):
self.training_args.pipeline_config.dmd_denoising_steps,
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()
self.denoising_step_list = timesteps[1000 -
self.denoising_step_list]
logger.info("Warping denoising_step_list")
self.denoising_step_list = self.denoising_step_list.to(
get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps: %s",
len(self.denoising_step_list), self.denoising_step_list)
logger.info("Distillation generator model to %s denoising steps",
len(self.denoising_step_list))
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.min_timestep = int(self.training_args.min_timestep_ratio *
@@ -194,82 +161,6 @@ class DistillationPipeline(TrainingPipeline):
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
self.generator_ema = 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(
f"Initialized generator EMA with decay={self.training_args.ema_decay}"
)
else:
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
def load_module_from_path(self, model_path: str, module_type: str,
training_args: "TrainingArgs"):
"""
Load a module from a specific path using the same loading logic as the pipeline.
Args:
model_path: Path to the model
module_type: Type of module to load (e.g., "transformer")
training_args: Training arguments
Returns:
The loaded module
"""
logger.info(f"Loading {module_type} from custom path: {model_path}")
# Set flag to prevent custom weight loading for teacher/critic models
training_args._loading_teacher_critic_model = True
try:
from fastvideo.models.loader.component_loader import (
PipelineComponentLoader)
# Download the model if it's a Hugging Face model ID
local_model_path = maybe_download_model(model_path)
logger.info(f"Model downloaded/found at: {local_model_path}")
config = verify_model_config_and_directory(local_model_path)
if module_type not in config:
if hasattr(self, '_extra_config_module_map'
) and module_type in self._extra_config_module_map:
extra_module = self._extra_config_module_map[module_type]
if extra_module in config:
module_type = extra_module
logger.info(f"Using {extra_module} for {module_type}")
else:
raise ValueError(
f"Module {module_type} not found in config at {local_model_path}"
)
else:
raise ValueError(
f"Module {module_type} not found in config at {local_model_path}"
)
module_info = config[module_type]
if module_info is None:
raise ValueError(
f"Module {module_type} has null value in config at {local_model_path}"
)
transformers_or_diffusers, architecture = module_info
component_path = os.path.join(local_model_path, module_type)
module = PipelineComponentLoader.load_module(
module_name=module_type,
component_model_path=component_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=training_args,
)
logger.info(
f"Successfully loaded {module_type} from {component_path}")
return module
finally:
# Always clean up the flag
if hasattr(training_args, '_loading_teacher_critic_model'):
delattr(training_args, '_loading_teacher_critic_model')
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
"""Initialize validation pipeline - must be implemented by subclasses."""
@@ -279,117 +170,11 @@ class DistillationPipeline(TrainingPipeline):
def _prepare_distillation(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Prepare training environment for distillation."""
self.transformer.requires_grad_(True)
self.transformer.train()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
return training_batch
def apply_ema_to_model(self, model):
"""Apply EMA weights to the model for validation or inference."""
if self.generator_ema is not None:
with self.generator_ema.apply_to_model(model):
return model
return model
def get_ema_model_copy(self):
"""Get a copy of the model with EMA weights applied."""
if self.generator_ema is not None:
ema_model = copy.deepcopy(self.transformer)
self.generator_ema.copy_to_unwrapped(ema_model)
return ema_model
return None
def is_ema_ready(self, current_step: int = None):
"""Check if EMA is ready for use (after ema_start_step)."""
if current_step is None:
current_step = getattr(self, 'current_trainstep', 0)
return (self.generator_ema is not None
and current_step >= self.training_args.ema_start_step)
def save_ema_weights(self, output_dir: str, step: int):
"""Save EMA weights separately for inference purposes."""
if self.generator_ema is None:
logger.warning("Cannot save EMA weights: EMA not initialized")
return
if not self.is_ema_ready():
logger.warning(
"Cannot save EMA weights: EMA not ready yet (step < ema_start_step)"
)
return
try:
ema_model = self.get_ema_model_copy()
if ema_model is None:
logger.warning("Failed to create EMA model copy")
return
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
os.makedirs(ema_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 = gather_state_dict_on_cpu_rank0(ema_model, device=None)
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)
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)
logger.info(f"EMA weights saved to {weight_path}")
del ema_model
except Exception as e:
logger.error(f"Failed to save EMA weights: {str(e)}")
def get_ema_stats(self):
"""Get EMA statistics for monitoring."""
if self.generator_ema is None:
return {
"ema_enabled": False,
"ema_decay": None,
"ema_start_step": self.training_args.ema_start_step,
"ema_ready": False,
"ema_step": self.current_trainstep,
}
return {
"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(),
"ema_step": self.current_trainstep,
}
def reset_ema(self):
"""Reset EMA to current model weights."""
if self.generator_ema is not None:
logger.info("Resetting EMA to current model weights")
self.generator_ema.update(self.transformer)
# Force update to current weights by setting decay to 0 temporarily
original_decay = self.generator_ema.decay
self.generator_ema.decay = 0.0
self.generator_ema.update(self.transformer)
self.generator_ema.decay = original_decay
logger.info("EMA reset completed")
else:
logger.warning("Cannot reset EMA: EMA not initialized")
def _build_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
@@ -546,7 +331,6 @@ class DistillationPipeline(TrainingPipeline):
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
original_latent = generator_pred_video
with torch.no_grad():
timestep = torch.randint(0,
self.num_train_timestep, [1],
@@ -571,7 +355,7 @@ class DistillationPipeline(TrainingPipeline):
noisy_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).detach().unflatten(0, (1, generator_pred_video.shape[1]))
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
@@ -620,24 +404,24 @@ class DistillationPipeline(TrainingPipeline):
pred_real_video_uncond) * self.real_score_guidance_scale
grad = (faker_score_pred_video - real_score_pred_video) / torch.abs(
original_latent - real_score_pred_video).mean()
generator_pred_video - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
original_latent.float(),
(original_latent.float() - grad.float()).detach())
generator_pred_video.float(),
(generator_pred_video.float() - grad.float()).detach())
training_batch.dmd_latent_vis_dict.update({
"training_batch_dmd_fwd_clean_latent":
training_batch.latents,
"generator_pred_video":
original_latent.detach(),
generator_pred_video,
"real_score_pred_video":
real_score_pred_video.detach(),
real_score_pred_video,
"faker_score_pred_video":
faker_score_pred_video.detach(),
faker_score_pred_video,
"dmd_timestep":
timestep.detach(),
timestep,
})
return dmd_loss
@@ -734,12 +518,12 @@ class DistillationPipeline(TrainingPipeline):
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
training_batch.unconditional_dict = unconditional_dict
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
training_batch.conditional_dict = conditional_dict
training_batch.unconditional_dict = unconditional_dict
training_batch.raw_latent_shape = training_batch.latents.shape
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
self.video_latent_shape = training_batch.latents.shape
@@ -802,15 +586,8 @@ class DistillationPipeline(TrainingPipeline):
(dmd_loss / gradient_accumulation_steps).backward()
total_dmd_loss += dmd_loss.detach().item()
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
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
avg_dmd_loss = torch.tensor(total_dmd_loss /
gradient_accumulation_steps,
device=self.device)
@@ -834,9 +611,6 @@ class DistillationPipeline(TrainingPipeline):
fake_score_latent_vis_dict.update(
batch_fake.fake_score_latent_vis_dict)
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
for param in self.fake_score_transformer.parameters():
# 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()
@@ -864,8 +638,7 @@ class DistillationPipeline(TrainingPipeline):
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)
self.fake_score_lr_scheduler, self.noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
@@ -896,14 +669,6 @@ class DistillationPipeline(TrainingPipeline):
sum(p.numel()
for p in self.fake_score_transformer.parameters()) / 1e9)
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")
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
training_args.inference_mode = True
@@ -935,18 +700,6 @@ class DistillationPipeline(TrainingPipeline):
transformer.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))
if use_ema_for_validation:
logger.info("Using EMA model for validation")
validation_transformer = self.transformer
ema_context = self.generator_ema.apply_to_model(
validation_transformer)
else:
validation_transformer = transformer
ema_context = None
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
@@ -962,98 +715,50 @@ class DistillationPipeline(TrainingPipeline):
step_videos: list[np.ndarray] = []
step_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)
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]
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",
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)
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
# 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,
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)
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Log validation results for this step
world_group = get_world_group()
@@ -1130,16 +835,16 @@ class DistillationPipeline(TrainingPipeline):
latents.dtype)
else:
latents += self.vae.shift_factor
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, latents
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[latent_key] = wandb.Video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, latents
# Process DMD training data if available - use decode_stage instead of self.vae.decode
if 'generator_pred_video' in dmd_latents_vis_dict:
@@ -1191,6 +896,14 @@ class DistillationPipeline(TrainingPipeline):
else:
set_random_seed(seed + self.global_rank)
# Check trainable params
num_trainable_generator = round(
count_trainable(self.transformer) / 1e9, 3)
num_trainable_critic = round(
count_trainable(self.fake_score_transformer) / 1e9, 3)
logger.info(
"rank: %s: # of trainable params in generator: %sB, # of trainable params in critic: %sB",
self.global_rank, num_trainable_generator, num_trainable_critic)
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
@@ -1200,10 +913,6 @@ class DistillationPipeline(TrainingPipeline):
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", seed)
# Initialize current_trainstep for EMA ready checks
#TODO: check if needed
self.current_trainstep = self.init_steps
# Resume from checkpoint if specified (this will restore random states)
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -1247,14 +956,6 @@ class DistillationPipeline(TrainingPipeline):
self.current_trainstep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
if (step >= self.training_args.ema_start_step) and \
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
self.generator_ema = EMA_FSDP(
self.transformer, decay=self.training_args.ema_decay)
logger.info(
f"Created generator EMA at step {step} with decay={self.training_args.ema_decay}"
)
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
@@ -1268,19 +969,11 @@ class DistillationPipeline(TrainingPipeline):
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"total_loss":
f"{total_loss:.4f}",
"generator_loss":
f"{generator_loss:.4f}",
"fake_score_loss":
f"{fake_score_loss:.4f}",
"step_time":
f"{step_time:.2f}s",
"grad_norm":
grad_norm,
"ema":
"✓" if (self.generator_ema is not None and self.is_ema_ready())
else "✗",
"total_loss": f"{total_loss:.4f}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
@@ -1308,15 +1001,6 @@ class DistillationPipeline(TrainingPipeline):
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
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
ema_stats = self.get_ema_stats()
log_data.update(ema_stats)
if training_batch.dmd_latent_vis_dict:
dmd_additional_logs = {
"generator_timestep":
@@ -1348,8 +1032,7 @@ class DistillationPipeline(TrainingPipeline):
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)
self.fake_score_lr_scheduler, self.noise_random_generator)
if self.transformer:
self.transformer.train()
@@ -1366,11 +1049,7 @@ class DistillationPipeline(TrainingPipeline):
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)
only_save_generator_weight=True)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
if self.training_args.log_visualization:
@@ -1390,11 +1069,7 @@ class DistillationPipeline(TrainingPipeline):
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)
self.noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
cleanup_dist_env_and_memory()
-492
View File
@@ -1,492 +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(f"ARG dmd_denoising_steps: {training_args.pipeline_config.dmd_denoising_steps}")
logger.info(
f"ARG dmd_denoising_steps: {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())
# 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(f"timesteps: {timesteps}")
self.dmd_denoising_steps = timesteps[1000 -
self.dmd_denoising_steps]
logger.info(
f"warped self.dmd_denoising_steps: {self.dmd_denoising_steps}")
# assert False, "warp_denoising_step must be false"
else:
assert False, "warp_denoising_step must be true"
logger.info("not warped")
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
get_local_torch_device())
logger.info(f"denoising_step_list: {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.min_timestep = int(self.training_args.min_timestep_ratio *
# self.num_train_timestep)
# self.max_timestep = int(self.training_args.max_timestep_ratio *
# self.num_train_timestep)
# self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
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,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
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): # type: ignore[override]
# 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
path = "/mnt/weka/home/hao.zhang/wei/FastVideo/data/ode_vidprom_1_fv/00000.pt"
b = torch.load(path)
for k, v in b.items():
logger.info(f"b[{k}]: {type(v)}")
if isinstance(v, torch.Tensor):
logger.info(f"b[{k}]: {v.shape}")
else:
logger.info(f"b[{k}]: {v}")
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(f"trajectory_latents: {trajectory_latents.shape}")
logger.info(
f"encoder_hidden_states: {training_batch.encoder_hidden_states.shape}"
)
return training_batch, trajectory_latents.to(device,
dtype=torch.bfloat16), None
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 = {}
device = get_local_torch_device()
target_latent = traj_latents[:, -1]
# logger.info(f"traj_latents: {traj_latents.shape}")
# logger.info(f"traj_timesteps: {traj_timesteps.shape}")
# 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:
# Use the first sample's trajectory timesteps; assumed identical across batches
# s_steps = traj_timesteps[0].to(torch.long) # [S]
# dmd = cast(torch.Tensor, self.dmd_denoising_steps).to(s_steps.device) # [K]
# distances_ks: [K, S] = |s_steps - dmd|
# distances_ks = (s_steps.unsqueeze(0) - dmd.unsqueeze(1)).abs()
# self._cached_closest_idx_per_dmd = distances_ks.argmin(dim=1).to(torch.long).cpu() # [K]
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36], dtype=torch.long).cpu()
logger.info(
f"self._cached_closest_idx_per_dmd: {self._cached_closest_idx_per_dmd}"
)
logger.info(
f"corresponding timesteps: {self.noise_scheduler.timesteps[self._cached_closest_idx_per_dmd]}"
)
# logger.info(f"traj_latents: {traj_latents.shape}")
# 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]
# relevant_traj_latents = torch.index_select(
# traj_latents,
# dim=1,
# index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
relevant_traj_latents = traj_latents
logger.info(f"relevant_traj_latents sum: {relevant_traj_latents.float().sum().item()}, relevant_traj_latents: {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(f"indexes: {indexes.shape}")
logger.info(f"indexes: {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)
# noisy_input = noisy_input.unsqueeze(0)
# # Sample a single DMD step for the whole batch and fetch its cached nearest S-index
# K = len(self.dmd_denoising_steps)
# dmd_idx = torch.randint(0, K, (1,), device=device)
# logger.info(f"dmd_idx: {dmd_idx}")
# assert self._cached_closest_idx_per_dmd is not None
# nearest_s_idx = int(self._cached_closest_idx_per_dmd[int(dmd_idx.item())])
# nearest_idx = torch.full((B,), nearest_s_idx, device=device, dtype=torch.long)
# batch_indices = torch.arange(B, device=device)
# noisy_input = traj_latents[batch_indices, nearest_idx] # [B, C, T, H, W]
# target_latent = traj_latents[batch_indices, -1] # [B, C, T, H, W]
# t = traj_timesteps[batch_indices, nearest_idx] # [B]
# Scale model input as in inference for consistency with stored trajectories
# noisy_input = self.modules["scheduler"].scale_model_input(noisy_input, t)
# logger.info(f"indexes: {indexes.shape}")
# logger.info(f"indexes: {indexes}")
timestep = self.dmd_denoising_steps[indexes]
# logger.info(f"timestep: {timestep.shape}")
# logger.info(f"timestep: {timestep}")
# 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()
logger.info("noisy input sum: %s, noisy input shape: %s, noisy input dtype: %s", noisy_input.float().sum().item(), noisy_input.shape, noisy_input.dtype)
logger.info("encoder hidden states sum: %s, encoder hidden states shape: %s, encoder hidden states dtype: %s", encoder_hidden_states.float().sum().item(), encoder_hidden_states.shape, encoder_hidden_states.dtype)
logger.info("timestep sum: %s, timestep shape: %s, timestep dtype: %s", timestep.float().sum().item(), timestep.shape, timestep.dtype)
logger.info("model dtype set: %s. model weight sum: %s", set(p.dtype for p in self.transformer.parameters()), sum(p.float().sum().item() for p in self.transformer.parameters()))
# model_dtype = next(self.transformer.parameters()).dtype
model_dtype = torch.bfloat16
input_kwargs = {
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep,
"return_dict": False,
"scheduler": self.modules["scheduler"]
}
# Predict noise and step the scheduler to obtain next latent
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
pred_video = self.transformer(**input_kwargs)
# logger.info(f"noise_pred: {noise_pred.shape}")
# logger.info("noise pred sum: %s, noise pred shape: %s, noise pred dtype: %s", noise_pred.float().sum().item(), noise_pred.shape, noise_pred.dtype)
# if isinstance(noise_pred, (tuple, list)):
# noise_pred = noise_pred[0]
# 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=model_dtype).flatten(0, 1),
# scheduler=self.modules["scheduler"]).unflatten(
# 0, noise_pred.shape[:2])
logger.info("pred video sum: %s, pred video shape: %s, pred video dtype: %s", pred_video.float().sum().item(), pred_video.shape, pred_video.dtype)
latent_vis_dict["pred_video"] = pred_video.permute(
0, 2, 1, 3, 4).detach().clone().cpu()
# noisy_input = pred_noise_to_pred_video(noise_pred, noisy_input, t, self.modules["scheduler"])
# next_latent_pred = self.modules["scheduler"].step(
# noise_pred, t, current_latents, return_dict=False)[0]
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]
B, S = traj_latents.shape[0], traj_latents.shape[1]
if S < 2:
raise ValueError("Trajectory must contain at least 2 steps")
# Sample per-sample current step i in [0, S-2]
# idx = torch.randint(low=0, high=S - 1, size=(B, ),
# device=traj_latents.device)
# Gather current latents and next latents
# batch_indices = torch.arange(B, device=traj_latents.device)
# current_latents = traj_latents[batch_indices, idx] # [B, C, T,H,W]
# current_latent = traj_timesteps[:, -1, :, :, :, :]
# target_latents = traj_latents[:, -1, :, :, :, :]
# Corresponding timesteps t (long) -> cast per sample
# t = traj_timesteps[:, -1, :, :, :, :]
# if t.dtype != torch.long:
# t = t.long()
# 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, training_batch.current_timestep, 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()
logger.info("Loss sum: %s, Loss dtype: %s", loss.float().sum().item(), loss.dtype)
assert self.transformer.blocks[0].to_q.weight.dtype == torch.float32
logger.info("blocks[0].to_q param sum before backprop: %s", self.transformer.blocks[0].to_q.weight.float().sum().item())
logger.info("blocks[0].to_q param grad sum: %s", self.transformer.blocks[0].to_q.weight.grad.float().sum().item())
logger.info("Transformer grad dtype: %s", set(p.grad.dtype for p in self.transformer.parameters()))
logger.info("Transformer param dtype before backprop: %s", set(p.dtype for p in self.transformer.parameters()))
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)
grad_norm = None
self.optimizer.step()
logger.info("Transformer param dtype after backprop: %s", set(p.dtype for p in self.transformer.parameters()))
logger.info("blocks[0].to_q param param after backprop sum: %s", self.transformer.blocks[0].to_q.weight.float().sum().item())
assert False
# 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)
# dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
# fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
# fake_score_log_keys = ['generator_pred_video']
# dmd_log_keys = ['faker_score_pred_video', 'real_score_pred_video']
def main(args) -> None:
logger.info("Starting ODE-init training pipeline...")
logger.info(f"ARG dmd_denoising_steps: {args.dmd_denoising_steps}")
pipeline = ODEInitTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("ODE-init training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
File diff suppressed because it is too large Load Diff
+28 -48
View File
@@ -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,26 +39,20 @@ 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, get_scheduler, get_sigmas,
load_checkpoint, normalize_dit_input, save_checkpoint,
compute_density_for_timestep_sampling, count_trainable, 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,
# set_random_seed, shallow_asdict)
from fastvideo.utils import (is_vsa_available,
from fastvideo.utils import (is_vmoba_available, 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.
@@ -118,18 +112,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
# Set grads for proper modules based on the training mode (Distill, LoRA, etc.)
self.set_trainable()
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
# Parse betas from string format "beta1,beta2"
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -281,20 +272,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
@@ -319,8 +310,7 @@ 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":
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
@@ -441,7 +431,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
local_main_process_only=False)
if not self.post_init_called:
self.post_init()
num_trainable_params = _get_trainable_params(self.transformer)
num_trainable_params = count_trainable(self.transformer)
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
@@ -486,9 +476,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
@@ -530,14 +520,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(
_get_trainable_params(self.transformer) / 1e9, 3)
count_trainable(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
@@ -573,7 +559,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(_get_trainable_params(self.transformer) / 1e9, 3))
round(count_trainable(self.transformer) / 1e9, 3))
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
@@ -641,7 +627,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
@@ -735,8 +720,3 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
transformer.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")
+4 -173
View File
@@ -202,7 +202,6 @@ def save_distillation_checkpoint(generator_transformer,
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.
@@ -234,8 +233,6 @@ def save_distillation_checkpoint(generator_transformer,
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -349,14 +346,10 @@ def load_checkpoint(transformer,
"""
if not os.path.exists(checkpoint_path):
logger.warning("Checkpoint path %s does not exist", checkpoint_path)
assert False
return 0
# Extract step number from checkpoint path
try:
step = int(os.path.basename(checkpoint_path).split('-')[-1])
except:
step = 1
step = int(os.path.basename(checkpoint_path).split('-')[-1])
if rank == 0:
logger.info("Loading checkpoint from step %s", step)
@@ -409,8 +402,7 @@ def load_distillation_checkpoint(generator_transformer,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None) -> int:
noise_generator=None) -> int:
"""
Load distillation checkpoint with both generator and fake_score models.
Returns the step number from which training should resume.
@@ -464,18 +456,6 @@ def load_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
if generator_ema is not None:
try:
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully", rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank, str(e))
# Load critic distributed checkpoint
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
"critic")
@@ -1300,154 +1280,5 @@ def get_scheduler(
last_epoch=last_epoch)
class EMA_FSDP:
"""
FSDP2-friendly EMA with two modes:
- mode="local_shard" (default): maintain float32 CPU EMA of local parameter shards on every rank.
Provides a context manager to temporarily swap EMA weights into the live model for teacher forward.
- mode="rank0_full": maintain a consolidated float32 CPU EMA of full parameters on rank 0 only
using gather_state_dict_on_cpu_rank0(). Useful for checkpoint export; not for teacher forward.
Usage (local_shard for CM teacher):
ema = EMA_FSDP(model, decay=0.999, mode="local_shard")
for step in ...:
ema.update(model)
with ema.apply_to_model(model):
with torch.no_grad():
y_teacher = model(...)
Usage (rank0_full for export):
ema = EMA_FSDP(model, decay=0.999, mode="rank0_full")
ema.update(model)
ema.state_dict() # on rank 0
"""
def __init__(self, module, decay: float = 0.999, mode: str = "local_shard"):
self.decay = float(decay)
self.mode = mode
self.shadow: dict[str, torch.Tensor] = {}
self.rank = dist.get_rank() if dist.is_initialized() else 0
if self.mode not in {"local_shard", "rank0_full"}:
raise ValueError(f"Unsupported EMA_FSDP mode: {self.mode}")
self._init_shadow(module)
@staticmethod
def _to_local_tensor(t: torch.Tensor) -> torch.Tensor:
# DTensor-aware to_local fetch; fall back to raw tensor
try:
from torch.distributed.tensor import DTensor # type: ignore
if isinstance(t, DTensor):
return t.to_local()
except Exception:
pass
return t
@torch.no_grad()
def _init_shadow(self, module):
if self.mode == "rank0_full":
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
if self.rank == 0:
self.shadow = {k: v.detach().clone().float().cpu() for k, v in cpu_state.items()}
else:
self.shadow = {}
return
# local_shard: maintain EMA of local shards for requires_grad params
self.shadow = {}
for name, p in module.named_parameters():
if not p.requires_grad:
continue
local = self._to_local_tensor(p.detach())
self.shadow[name] = local.clone().float().cpu()
@torch.no_grad()
def update(self, module):
d = self.decay
if self.mode == "rank0_full":
if self.rank != 0:
return
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
for n, v in cpu_state.items():
v_cpu = v.detach().float().cpu()
if n not in self.shadow:
self.shadow[n] = v_cpu.clone()
else:
self.shadow[n].mul_(d).add_(v_cpu, alpha=1.0 - d)
return
# local_shard: update local shard EMA on every rank
for name, p in module.named_parameters():
if not p.requires_grad:
continue
local = self._to_local_tensor(p.detach())
v_cpu = local.float().cpu()
if name not in self.shadow:
self.shadow[name] = v_cpu.clone()
else:
self.shadow[name].mul_(d).add_(v_cpu, alpha=1.0 - d)
def state_dict(self) -> dict[str, torch.Tensor]:
if self.mode == "rank0_full":
return {k: v.clone() for k, v in self.shadow.items()} if self.rank == 0 else {}
return {k: v.clone() for k, v in self.shadow.items()}
def load_state_dict(self, sd: dict[str, torch.Tensor]):
self.shadow = {k: v.clone() for k, v in sd.items()}
@torch.no_grad()
def copy_to_unwrapped(self, module) -> None:
"""
Copy EMA weights into a non-sharded (unwrapped) module. Intended for export/eval.
For mode="rank0_full", only rank 0 has the full EMA state.
"""
if self.mode == "rank0_full" and self.rank != 0:
return
name_to_param = dict(module.named_parameters())
for n, w in self.shadow.items():
if n in name_to_param:
p = name_to_param[n]
p.data.copy_(w.to(dtype=p.dtype, device=p.device))
class _ApplyEMACtx:
def __init__(self, ema: "EMA_FSDP", module):
self.ema = ema
self.module = module
self.saved: dict[str, torch.Tensor] = {}
def __enter__(self):
if self.ema.mode != "local_shard":
raise RuntimeError("EMA apply_to_model is only supported for mode='local_shard'")
with torch.no_grad():
for name, p in self.module.named_parameters():
if not p.requires_grad:
continue
# Save local shard
p_local = EMA_FSDP._to_local_tensor(p.detach())
if p_local.numel() == 0:
# Nothing to swap on this rank for this param
continue
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
if name in self.ema.shadow:
ema_cpu = self.ema.shadow[name]
if ema_cpu.numel() != p_local.numel():
# Shard shape mismatch (e.g., empty shard here), skip
continue
# Copy EMA shard into local param shard
p_local.copy_(ema_cpu.to(dtype=p_local.dtype, device=p_local.device))
return self.module
def __exit__(self, exc_type, exc, tb):
with torch.no_grad():
for name, p in self.module.named_parameters():
if name in self.saved:
p_local = EMA_FSDP._to_local_tensor(p.detach())
if p_local.numel() == 0:
continue
saved_local = self.saved[name]
if saved_local.numel() != p_local.numel():
continue
p_local.copy_(saved_local)
self.saved.clear()
return False
def apply_to_model(self, module):
return EMA_FSDP._ApplyEMACtx(self, module)
def count_trainable(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
@@ -1,72 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import WanCausalDMDPipeline
from fastvideo.training.self_forcing_distillation_pipeline import SelfForcingDistillationPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
"""
A self-forcing distillation pipeline for Wan that uses the self-forcing methodology
with DMD for video generation.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
validation_pipeline = WanCausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
self.validation_pipeline = validation_pipeline
def main(args) -> None:
logger.info("Starting Wan self-forcing distillation pipeline...")
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Wan self-forcing distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
main(args)
@@ -1,211 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.basic.wan.wan_i2v_pipeline import (
WanImageToVideoPipeline)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanT2VI2VTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_t2v
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.dit_cpu_offload = True
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
pipeline_config = PipelineConfig.from_pretrained(
training_args.model_path)
pipeline_config.vae_config.load_encoder = True
self.validation_pipeline = WanImageToVideoPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
pipeline_config=pipeline_config,
loaded_modules={
"transformer": self.get_module("transformer"),
},
required_config_modules=[
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
],
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
dit_cpu_offload=True,
)
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
batch = next(self.train_loader_iter)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
# clip_features = batch['clip_feature']
# image_latents = batch['first_frame_latent']
# image_latents = image_latents[:, :, :self.training_args.num_latent_t]
# pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
# training_batch.preprocessed_image = pil_image.to(
# get_local_torch_device())
# training_batch.image_embeds = clip_features.to(get_local_torch_device())
# training_batch.image_latents = image_latents.to(
# get_local_torch_device())
training_batch.infos = infos
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
latents = training_batch.latents
logger.info("latents.shape: %s", latents.shape)
first_frame_latent = latents[:, :, 0, :, :]
logger.info("first_frame_latent.shape: %s", first_frame_latent.shape)
logger.info("training_batch.noisy_model_input.shape: %s",
training_batch.noisy_model_input.shape)
training_batch.noisy_model_input = torch.cat([
first_frame_latent.unsqueeze(2),
training_batch.noisy_model_input[:, :, 1:, :, :]
],
dim=2)
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
# Image Embeds for conditioning
# image_embeds = training_batch.image_embeds
# assert torch.isnan(image_embeds).sum() == 0
# image_embeds = image_embeds.to(get_local_torch_device(),
# dtype=torch.bfloat16)
# encoder_hidden_states_image = image_embeds
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
# "encoder_hidden_states_image":
# encoder_hidden_states_image,
"return_dict":
False,
}
return training_batch
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.image_path = validation_batch['video_path']
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanT2VI2VTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("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)
-4
View File
@@ -2,10 +2,6 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
import argparse
from einops import rearrange
import torchvision
import numpy as np
import imageio
import ctypes
import hashlib
import importlib
-2
View File
@@ -1,2 +0,0 @@
Code in this folder is modified from https://github.com/Wan-Video/Wan2.1
Apache-2.0 License
-3
View File
@@ -1,3 +0,0 @@
from . import configs, distributed, modules
from .image2video import WanI2V
from .text2video import WanT2V
-42
View File
@@ -1,42 +0,0 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from .wan_t2v_14B import t2v_14B
from .wan_t2v_1_3B import t2v_1_3B
from .wan_i2v_14B import i2v_14B
import copy
import os
os.environ['TOKENIZERS_PARALLELISM'] = 'false'
# the config of t2i_14B is the same as t2v_14B
t2i_14B = copy.deepcopy(t2v_14B)
t2i_14B.__name__ = 'Config: Wan T2I 14B'
WAN_CONFIGS = {
't2v-14B': t2v_14B,
't2v-1.3B': t2v_1_3B,
'i2v-14B': i2v_14B,
't2i-14B': t2i_14B,
}
SIZE_CONFIGS = {
'720*1280': (720, 1280),
'1280*720': (1280, 720),
'480*832': (480, 832),
'832*480': (832, 480),
'1024*1024': (1024, 1024),
}
MAX_AREA_CONFIGS = {
'720*1280': 720 * 1280,
'1280*720': 1280 * 720,
'480*832': 480 * 832,
'832*480': 832 * 480,
}
SUPPORTED_SIZES = {
't2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
't2v-1.3B': ('480*832', '832*480'),
'i2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
't2i-14B': tuple(SIZE_CONFIGS.keys()),
}
-19
View File
@@ -1,19 +0,0 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
# ------------------------ Wan shared config ------------------------#
wan_shared_cfg = EasyDict()
# t5
wan_shared_cfg.t5_model = 'umt5_xxl'
wan_shared_cfg.t5_dtype = torch.bfloat16
wan_shared_cfg.text_len = 512
# transformer
wan_shared_cfg.param_dtype = torch.bfloat16
# inference
wan_shared_cfg.num_train_timesteps = 1000
wan_shared_cfg.sample_fps = 16
wan_shared_cfg.sample_neg_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
-35
View File
@@ -1,35 +0,0 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
from .shared_config import wan_shared_cfg
# ------------------------ Wan I2V 14B ------------------------#
i2v_14B = EasyDict(__name__='Config: Wan I2V 14B')
i2v_14B.update(wan_shared_cfg)
i2v_14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
i2v_14B.t5_tokenizer = 'google/umt5-xxl'
# clip
i2v_14B.clip_model = 'clip_xlm_roberta_vit_h_14'
i2v_14B.clip_dtype = torch.float16
i2v_14B.clip_checkpoint = 'models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth'
i2v_14B.clip_tokenizer = 'xlm-roberta-large'
# vae
i2v_14B.vae_checkpoint = 'Wan2.1_VAE.pth'
i2v_14B.vae_stride = (4, 8, 8)
# transformer
i2v_14B.patch_size = (1, 2, 2)
i2v_14B.dim = 5120
i2v_14B.ffn_dim = 13824
i2v_14B.freq_dim = 256
i2v_14B.num_heads = 40
i2v_14B.num_layers = 40
i2v_14B.window_size = (-1, -1)
i2v_14B.qk_norm = True
i2v_14B.cross_attn_norm = True
i2v_14B.eps = 1e-6
-29
View File
@@ -1,29 +0,0 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from easydict import EasyDict
from .shared_config import wan_shared_cfg
# ------------------------ Wan T2V 14B ------------------------#
t2v_14B = EasyDict(__name__='Config: Wan T2V 14B')
t2v_14B.update(wan_shared_cfg)
# t5
t2v_14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
t2v_14B.t5_tokenizer = 'google/umt5-xxl'
# vae
t2v_14B.vae_checkpoint = 'Wan2.1_VAE.pth'
t2v_14B.vae_stride = (4, 8, 8)
# transformer
t2v_14B.patch_size = (1, 2, 2)
t2v_14B.dim = 5120
t2v_14B.ffn_dim = 13824
t2v_14B.freq_dim = 256
t2v_14B.num_heads = 40
t2v_14B.num_layers = 40
t2v_14B.window_size = (-1, -1)
t2v_14B.qk_norm = True
t2v_14B.cross_attn_norm = True
t2v_14B.eps = 1e-6
-29
View File
@@ -1,29 +0,0 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from easydict import EasyDict
from .shared_config import wan_shared_cfg
# ------------------------ Wan T2V 1.3B ------------------------#
t2v_1_3B = EasyDict(__name__='Config: Wan T2V 1.3B')
t2v_1_3B.update(wan_shared_cfg)
# t5
t2v_1_3B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
t2v_1_3B.t5_tokenizer = 'google/umt5-xxl'
# vae
t2v_1_3B.vae_checkpoint = 'Wan2.1_VAE.pth'
t2v_1_3B.vae_stride = (4, 8, 8)
# transformer
t2v_1_3B.patch_size = (1, 2, 2)
t2v_1_3B.dim = 1536
t2v_1_3B.ffn_dim = 8960
t2v_1_3B.freq_dim = 256
t2v_1_3B.num_heads = 12
t2v_1_3B.num_layers = 30
t2v_1_3B.window_size = (-1, -1)
t2v_1_3B.qk_norm = True
t2v_1_3B.cross_attn_norm = True
t2v_1_3B.eps = 1e-6
-33
View File
@@ -1,33 +0,0 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from functools import partial
import torch
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
def shard_model(
model,
device_id,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32,
process_group=None,
sharding_strategy=ShardingStrategy.FULL_SHARD,
sync_module_states=True,
):
model = FSDP(
module=model,
process_group=process_group,
sharding_strategy=sharding_strategy,
auto_wrap_policy=partial(
lambda_auto_wrap_policy, lambda_fn=lambda m: m in model.blocks),
mixed_precision=MixedPrecision(
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
buffer_dtype=buffer_dtype),
device_id=device_id,
use_orig_params=True,
sync_module_states=sync_module_states)
return model

Some files were not shown because too many files have changed in this diff Show More