Compare commits

..
Author SHA1 Message Date
sfc-gh-hazhang 4bfcdbd0ec lint 2025-08-04 11:53:02 -07:00
135 changed files with 472 additions and 4831 deletions
-1
View File
@@ -163,7 +163,6 @@ steps:
- path:
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/tests/test_vsa.py"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
-257
View File
@@ -1,257 +0,0 @@
name: Publish Video Sparse Attention Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/attn/setup_vsa.py"
workflow_dispatch:
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
cd csrc/attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_vsa.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup_vsa.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build_wheels:
name: Build Wheel
needs: check-version-change
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12', '3.13']
# For version reference https://pytorch.org/get-started/previous-versions/
torch-cuda:
- torch-version: '2.5.1'
cuda-version: '12.4.1'
torch-cuda-short: 'cu124'
- torch-version: '2.6.0'
cuda-version: '12.6.3'
torch-cuda-short: 'cu126'
- torch-version: '2.7.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
steps:
- name: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.torch-cuda.cuda-version }}
linux-local-args: '["--toolkit"]'
method: 'network'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_vsa.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
# Get the correct version format
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
# Rename with version information
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
- name: Upload wheel artifact
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/attn/dist/*.whl
retention-days: 90
publish_package:
name: Publish package
needs: [build_wheels, check-version-change]
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Install CUDA 12.4.1
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: 12.4.1
linux-local-args: '["--toolkit"]'
method: 'network'
sub-packages: '["nvcc"]'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-12.4.1
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch 2.5.1+cu12.4.1
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_vsa.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/attn/dist/
+7 -9
View File
@@ -7,7 +7,7 @@
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/qqPzbrw" target="_blank"> <b> WeChat </b> </a> |
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers" target="_blank"><b>FastWan2.1</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers" target="_blank"><b>FastWan2.2</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> |
</p>
<div align="center">
@@ -64,15 +64,12 @@ See below for recipes and datasets:
## Inference
### Generating Your First Video
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Create a file called `example.py` with the following code:
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
```python
import os
from fastvideo import VideoGenerator
def main():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
@@ -154,10 +151,11 @@ If you find FastVideo useful, please considering citing our work:
year = {2024},
}
@article{zhang2025vsa,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
@article{zhang2025faster,
title={Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric P and Zhang, Hao},
journal={arXiv e-prints},
pages={arXiv--2505},
year={2025}
}
+1 -1
View File
@@ -84,7 +84,7 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
### Test
```bash
python tests/test_sta.py # test STA
python tests/test_vsa.py # test VSA
python tests/test_block_sparse.py # test VSA
```
### Benchmark
```bash
+2
View File
@@ -86,6 +86,8 @@ def benchmark_attention(configurations):
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
torch.cuda.empty_cache()
return results
+13 -10
View File
@@ -51,16 +51,19 @@ for k in kernels:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
ext_modules = [
CUDAExtension('vsa_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
]
ext_modules = []
import torch
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
ext_modules = [
CUDAExtension('vsa_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
]
+3 -4
View File
@@ -1,13 +1,12 @@
import torch
from typing import Tuple
block_sparse_attn=None
import torch
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
try:
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
block_sparse_attn = block_sparse_attn_SM90
else:
except ImportError:
from vsa.block_sparse_wrapper import block_sparse_attn_triton
block_sparse_fwd = None
block_sparse_bwd = None
+90 -93
View File
@@ -11,6 +11,95 @@ from typing import Tuple, Optional
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_SM90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor]:
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
variable_block_sizes = variable_block_sizes.int()
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
return o_padded, lse_padded
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
def _block_sparse_attn_SM90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
B, H, S, D = q_padded.shape
o_padded = torch.empty_like(q_padded)
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
return o_padded, lse_padded
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_backward_SM90(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
)
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
return grad_q_padded, grad_k_padded, grad_v_padded
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
def _block_sparse_attn_backward_SM90_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch._check(grad_output_padded.dtype == torch.bfloat16)
torch._check(lse_padded.dtype == torch.float32)
grad_output_padded = grad_output_padded.contiguous()
dq = torch.empty_like(grad_output_padded)
dk = torch.empty_like(grad_output_padded)
dv = torch.empty_like(grad_output_padded)
return dq, dk, dv
def backward_SM90(ctx, grad_output1, grad_output2):
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def setup_context_SM90(ctx, inputs, output):
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
o_padded, lse_padded = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
def block_sparse_attn_triton(
q: torch.Tensor,
@@ -90,96 +179,4 @@ def setup_context_triton(ctx, inputs, output):
o_padded, M = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_SM90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor]:
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
variable_block_sizes = variable_block_sizes.int()
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
return o_padded, lse_padded
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
def _block_sparse_attn_SM90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
B, H, S, D = q_padded.shape
o_padded = torch.empty_like(q_padded)
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
return o_padded, lse_padded
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_backward_SM90(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
)
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
return grad_q_padded, grad_k_padded, grad_v_padded
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
def _block_sparse_attn_backward_SM90_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch._check(grad_output_padded.dtype == torch.bfloat16)
torch._check(lse_padded.dtype == torch.float32)
grad_output_padded = grad_output_padded.contiguous()
dq = torch.empty_like(grad_output_padded)
dk = torch.empty_like(grad_output_padded)
dv = torch.empty_like(grad_output_padded)
return dq, dk, dv
def backward_SM90(ctx, grad_output1, grad_output2):
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def setup_context_SM90(ctx, inputs, output):
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
o_padded, lse_padded = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
+4 -22
View File
@@ -6,9 +6,8 @@ We introduce a new finetuning strategy - **Sparse-distill**, which jointly integ
We provide two distilled models:
- **[FastWan2.1-T2V-1.3B-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers)**: 3-step inference, up to **16 FPS** on H100 GPU
- **[FastWan2.1-T2V-14B-480P-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-480P-Diffusers)**: 3-step inference, up to **60x speed up** at 480P, **90x speed up** at 720P for denoising loop
- **[FastWan2.2-TI2V-5B-FullAttn-Diffusers](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers)**: 3-step inference, up to **50x speed up** at 720P for denoising loop
- **[FastWan2.1-T2V-1.3B-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers)**: 3-step inference, up to **20 FPS** on H100 GPU
- **[FastWan2.1-T2V-14B-480P-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-480P-Diffusers)**: 3-step inference, up to **50x speed up** at 480P, **70x speed up** at 720P for denoising loop
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
@@ -35,7 +34,7 @@ python scripts/huggingface/download_hf.py \
## 🚀 Training Scripts
### Wan2.1 1.3B Model Sparse-Distill
### 1.3B Model Sparse-Distill
For the 1.3B model, we use **4 nodes with 32 H200 GPUs** (8 GPUs per node):
@@ -51,7 +50,7 @@ sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_1.3B.sl
- VSA attention sparsity: 0.8
- Training steps: 4000 (~12 hours)
### Wan2.1 14B Model Sparse-Distill
### 14B Model Sparse-Distill
For the 14B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
@@ -68,20 +67,3 @@ sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_14B.slu
- VSA attention sparsity: 0.9
- Training steps: 3000 (~52 hours)
- HSDP shard dim: 8
### Wan2.2 5B Model Sparse-Distill
For the 5B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
```bash
# Multi-node training (8 nodes, 64 GPUs total)
sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
```
**Key Configuration:**
- Global batch size: 64
- Sequence parallel size: 1
- Gradient accumulation steps: 1
- Learning rate: 2e-5
- Training steps: 3000 (~12 hours)
- HSDP shard dim: 1
+1 -1
View File
@@ -9,7 +9,7 @@
:::{raw} html
<p style="text-align:center">
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.
<strong>FastVideo is An unified inference and post-training framework for accelerated video generation.
</strong>
</p>
@@ -53,9 +53,3 @@ out = sliding_tile_attention(q, k, v, window_size, text_length)
out = sliding_tile_attention(q, k, v, window_size, 0, False)
```
# 🚀Inference
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
@@ -59,9 +59,3 @@ from vsa import video_sparse_attn
output = video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight)
```
# 🚀Inference
```bash
bash scripts/inference/v1_inference_wan_VSA.sh
```
@@ -1,16 +0,0 @@
# Wan2.1-I2V-1.3B-InP Crush-Smol Example
These are e2e example scripts for finetuning Wan2.1 T2V 1.3B InP on the crush-smol dataset.
## Execute the following commands from `FastVideo/` to run training:
### Download crush-smol dataset:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
### Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
@@ -1,111 +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="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp_single/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
NUM_GPUS=8
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_distill"
--output_dir "checkpoints/wan_i2v_distill"
--wandb_run_name "wan_i2v_distill"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 8
--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 8
)
# 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
# --log_visualization
--log_visualization
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--lr_scheduler "constant"
--fake_score_learning_rate 8e-7
--fake_score_lr_scheduler "constant"
--mixed_precision "bf16"
--training_state_checkpointing_steps 2000
--weight_only_checkpointing_steps 2000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
# --enable_gradient_checkpointing_type "full"
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--simulate_generator_forward
--real_score_guidance_scale 5
--VSA_sparsity 0.8
)
# 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/wan_i2v_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_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,111 +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-I2V-14B-480P-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
NUM_GPUS=8
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
export HF_HUB_ENABLE_HF_TRANSFER=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_distill"
--output_dir "checkpoints/wan_i2v_distill"
--wandb_run_name "wan_i2v_distill"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 8
--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 8
)
# 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
# --log_visualization
--log_visualization
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--lr_scheduler "constant"
--fake_score_learning_rate 8e-7
--fake_score_lr_scheduler "constant"
--mixed_precision "bf16"
--checkpointing_steps 2000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
# --enable_gradient_checkpointing_type "full"
)
dmd_args=(
--dmd_denoising_steps '1000,757,522'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--simulate_generator_forward
--real_score_guidance_scale 3.5
--VSA_sparsity 0.8
)
# 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/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -1,129 +0,0 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# Training arguments
training_args=(
--tracker_project_name wan_i2v_finetune
--output_dir "checkpoints/wan_i2v_finetune"
--max_train_steps 2000
--train_batch_size 2
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
--hsdp_shard_dim $NUM_GPUS
)
# 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 10
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--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"
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_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="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_i2v_1_3b_inp/"
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 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 "i2v"
@@ -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,31 +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
}
]
}
@@ -1,16 +0,0 @@
# Wan2.1-I2V-1.3B-InP Crush-Smol Example
These are e2e example scripts for finetuning Wan2.1 T2V 1.3B InP on the crush-smol dataset.
## Execute the following commands from `FastVideo/` to run training:
### Download crush-smol dataset:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
### Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
@@ -1,122 +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-I2V-14B-480P-Diffusers"
# MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_DIR="data/crush-smol_processed_wan21_i2v_14b/combined_parquet_dataset/"
# DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-I2V/crush_smol/validation_orig.json"
NUM_GPUS=8
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_distill"
--output_dir "checkpoints/wan_i2v_distill"
--wandb_run_name "14b_i2v_cm"
--max_train_steps 1500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--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 8
)
# 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
# --log_visualization
--log_visualization
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "3"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--lr_scheduler "constant"
# --min_lr_ratio 0.5
# --lr_warmup_steps 50
--fake_score_learning_rate 7e-7
--fake_score_lr_scheduler "constant"
--mixed_precision "bf16"
--training_state_checkpointing_steps 2000
--weight_only_checkpointing_steps 2000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
# --enable_gradient_checkpointing_type "full"
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--warp_denoising_step True
--min_timestep_ratio 0.02
--max_timestep_ratio 0.96
--generator_update_interval 5
--simulate_generator_forward
--simulate_forward_interval 1
--real_score_guidance_scale 5
--VSA_sparsity 0.8
--regression_loss_weight 0.01
--use_regression_loss False
--cm_loss_weight 1
--ema_decay 0.999
--cm_use_ema_teacher
)
# 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/wan_i2v_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_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,111 +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-I2V-14B-480P-Diffusers"
DATA_DIR="data/crush-smol_processed_wan21_i2v_14b/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-I2V/crush_smol/validation.json"
NUM_GPUS=8
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
export HF_HUB_ENABLE_HF_TRANSFER=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_distill"
--output_dir "checkpoints/wan_i2v_distill"
--wandb_run_name "wan21_i2v_14b_ode_50steps"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 8
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 8
)
# 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
# --log_visualization
--log_visualization
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--lr_scheduler "constant"
--fake_score_learning_rate 8e-7
--fake_score_lr_scheduler "constant"
--mixed_precision "bf16"
--checkpointing_steps 2000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
# --enable_gradient_checkpointing_type "full"
)
dmd_args=(
--dmd_denoising_steps '1000,757,522'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--simulate_generator_forward
--real_score_guidance_scale 3.5
--VSA_sparsity 0.8
)
# 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/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -1,129 +0,0 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# Training arguments
training_args=(
--tracker_project_name wan_i2v_finetune
--output_dir "checkpoints/wan_i2v_finetune"
--max_train_steps 2000
--train_batch_size 2
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
--hsdp_shard_dim $NUM_GPUS
)
# 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 10
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--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"
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,26 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
MODEL_TYPE="wan"
# DATA_MERGE_PATH="data/single/merge.txt"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_wan21_t2v_14b/"
# OUTPUT_DIR="data/single_processed_i2v_14b/"
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 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"
@@ -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,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
}
]
}
@@ -86,7 +86,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
@@ -102,6 +102,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -86,7 +86,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
@@ -102,6 +102,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -86,7 +86,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
@@ -102,6 +102,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -1,13 +0,0 @@
# Wan2.2-5B Distill Example
These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA methods.
### 0. Make sure you have installed VSA
```bash
cd csrc/attn
git submodule update --init --recursive
python setup_vsa.py install
```
### Data-free Distillation
When `--simulate_generator_forward` is enabled, distillation becomes data-free by simulating intermediate steps through forward inference of the generator. This helps avoid training–inference mismatch. See Section 4.5 of [DMD2](https://arxiv.org/pdf/2405.14867) for details.
@@ -87,7 +87,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DIR"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
@@ -108,6 +108,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -87,7 +87,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DIR"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
@@ -108,6 +108,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -9,14 +9,4 @@ git submodule update --init --recursive
python setup_vsa.py install
```
### 1. Download dataset:
```bash
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/download_dataset.sh
```
### 2. Configure and run distillation:
#### For DMD-only distillation:
```bash
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/
```
### TODO
@@ -63,7 +63,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
@@ -77,6 +77,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -1,118 +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-14B-Diffusers"
# MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_DIR="data/crush-smol_processed_wan21_t2v_14b/combined_parquet_dataset/"
# DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.1-I2V/crush_smol/validation_orig.json"
NUM_GPUS=8
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_distill"
--output_dir "checkpoints/wan_i2v_distill"
--wandb_run_name "cm_t2v"
--max_train_steps 1500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--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 8
)
# 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
# --log_visualization
--log_visualization
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "3"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 4e-6
--lr_scheduler "constant"
# --min_lr_ratio 0.5
# --lr_warmup_steps 50
--mixed_precision "bf16"
--training_state_checkpointing_steps 2000
--weight_only_checkpointing_steps 2000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
# --enable_gradient_checkpointing_type "full"
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--warp_denoising_step True
--min_timestep_ratio 0.02
--max_timestep_ratio 0.96
--real_score_guidance_scale 5
--VSA_sparsity 0.8
--regression_loss_weight 0.01
--use_regression_loss False
--cm_loss_weight 1
--ema_decay 0.999
--cm_weighing_function "constant"
# --cm_use_ema_teacher
)
# 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/wan_cm_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
+2 -1
View File
@@ -16,7 +16,8 @@ def main():
dit_cpu_offload=False,
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"
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=False,
# image_encoder_cpu_offload=False,
)
+1 -1
View File
@@ -17,7 +17,7 @@ def main():
use_fsdp_inference=True,
# Adjust these offload parameters if you have < 32GB of VRAM
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
-48
View File
@@ -1,48 +0,0 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
if __name__ == "__main__":
main()
@@ -1,30 +0,0 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, height=832, width=480, num_frames=81)
if __name__ == "__main__":
main()
@@ -10,7 +10,7 @@ def main():
dit_cpu_offload=False,
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"
pin_cpu_memory=False,
lora_path="benjamin-paine/steamboat-willie-1.3b",
lora_nickname="steamboat"
)
@@ -18,7 +18,7 @@ def main():
"height": 480,
"width": 832,
"num_frames": 81,
"guidance_scale": 6.0,
"guidance_scale": 5.0,
"num_inference_steps": 32,
"seed": 42,
}
@@ -11,22 +11,22 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
dit_cpu_offload=False,
vae_cpu_offload=True,
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"
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1000/transformer",
pin_cpu_memory=False,
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1250/transformer",
lora_nickname="crush_smol"
)
kwargs = {
"height": 480,
"width": 832,
"num_frames": 77,
"guidance_scale": 6.0,
"guidance_scale": 5.0,
"num_inference_steps": 50,
"seed": 42,
}
# Generate video with LoRA style
prompt = "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press."
prompt = "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."
video = generator.generate_video(
prompt,
@@ -34,12 +34,6 @@ def main():
save_video=True,
**kwargs
)
prompt = "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press."
video = generator.generate_video(
prompt,
output_path=OUTPUT_PATH,
save_video=True,
**kwargs
)
if __name__ == "__main__":
main()
@@ -14,7 +14,7 @@ def main():
dit_cpu_offload=False,
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"
pin_cpu_memory=False,
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
@@ -12,7 +12,7 @@ def main():
dit_cpu_offload=False,
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"
pin_cpu_memory=False,
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
@@ -9,19 +9,17 @@ MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
NUM_GPUS=8
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "outputs/wan_i2v_finetune"
--wandb_run_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 2000
--train_batch_size 1
--train_batch_size 4
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
@@ -32,10 +30,10 @@ training_args=(
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 8
--sp_size 4
--tp_size 4
--hsdp_replicate_dim 2
--hsdp_shard_dim 4
)
# Model arguments
@@ -56,12 +54,12 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-5
--mixed_precision "bf16"
--checkpointing_steps 2000
--weight_decay 1e-4
@@ -71,6 +69,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -88,7 +88,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
@@ -103,6 +103,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -101,6 +101,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--dit_precision "fp32"
@@ -101,6 +101,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--dit_precision "fp32"
@@ -54,7 +54,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
@@ -69,6 +69,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -88,7 +88,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
@@ -103,6 +103,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -54,7 +54,7 @@ validation_args=(
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
@@ -69,6 +69,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -52,9 +52,9 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
@@ -69,6 +69,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -77,7 +78,6 @@ miscellaneous_args=(
--num_euler_timesteps 50
--ema_start_step 0
--enable_gradient_checkpointing_type "full"
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
)
torchrun \
@@ -85,14 +85,14 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 400
--checkpointing_steps 500
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -100,6 +100,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -52,16 +52,16 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 400
--checkpointing_steps 500
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -69,6 +69,7 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
@@ -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"
DATASET_PATH="data/crush-smol-test/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocessing_new.py \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 4 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 77 \
--preprocess.train_fps 16 \
--preprocess.samples_per_file 8 \
--preprocess.flush_frequency 8 \
--preprocess.video_length_tolerance_range 5
+1 -1
View File
@@ -96,7 +96,7 @@ class WanVideoArchConfig(DiTArchConfig):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
self.num_channels_latents = self.in_channels if self.added_kv_proj_dim is None else self.out_channels
@dataclass
+1 -4
View File
@@ -31,10 +31,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": FastWan2_2_TI2V_5B_Config,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig
# Add other specific weight variants
}
+1 -1
View File
@@ -20,7 +20,7 @@ class SamplingParam:
# Text inputs
prompt: str | list[str] | None = None
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
negative_prompt: str | None = None
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
+10 -18
View File
@@ -6,20 +6,12 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
FastWanT2V480PConfig,
Wan2_2_I2V_A14B_SamplingParam,
Wan2_2_T2V_A14B_SamplingParam,
Wan2_2_TI2V_5B_SamplingParam,
WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
Wan2_1_Fun_1_3B_InP_SamplingParam,
)
# isort: on
from fastvideo.configs.sample.wan import (FastWanT2V480PConfig,
Wan2_2_TI2V_5B_SamplingParam,
WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam)
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -37,10 +29,10 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
# "Wan-AI/Wan2.2-T2V-A14B-Diffusers":
# Wan2_2_T2V_A14B_SamplingParam,
# "Wan-AI/Wan2.2-I2V-A14B-Diffusers":
# Wan2_2_I2V_A14B_SamplingParam,
# Add other specific weight variants
}
+3 -24
View File
@@ -107,28 +107,13 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
fps: int = 16
# =============================================
# ============= Wan2.1 Fun Models =============
# =============================================
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
negative_prompt: str = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
# =============================================
# ============= Wan2.2 TI2V Models =============
# =============================================
@dataclass
class Wan2_2_Base_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.2 TI2V 5B model."""
negative_prompt: str = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
@dataclass
@@ -144,15 +129,9 @@ 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
guidance_scale_2: float = 3.0
num_inference_steps: int = 40
fps: int = 16
pass
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 3.5
guidance_scale_2: float = 3.5
num_inference_steps: int = 40
fps: int = 16
pass
+10 -59
View File
@@ -158,12 +158,6 @@ class FastVideoArgs:
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
# MoE parameters used by Wan2.2
boundary_ratio: float | None = None
# scheduler parameters
warp_denoising_step: bool = False
@property
def training_mode(self) -> bool:
return not self.inference_mode
@@ -385,15 +379,6 @@ class FastVideoArgs:
default=FastVideoArgs.enable_stage_verification,
help="Enable input/output verification for pipeline stages",
)
# scheduler parameters
parser.add_argument(
"--warp-denoising-step",
action=StoreBoolean,
default=FastVideoArgs.warp_denoising_step,
help="Warp denoising step for scheduler",
)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -604,6 +589,8 @@ class TrainingArgs(FastVideoArgs):
dit_model_name_or_path: str = ""
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
training_cfg_rate: float = 0.0
precondition_outputs: bool = False
@@ -684,16 +671,6 @@ class TrainingArgs(FastVideoArgs):
log_visualization: bool = False
# simulate generator forward to match inference
simulate_generator_forward: bool = False
simulate_forward_interval: int = 1
regression_loss_weight: float = 0.0
use_regression_loss: bool = False
# consistency model parameters
cm_loss_weight: float = 0.0
cm_use_ema_teacher: bool = False
ema_decay: float = 0.999
ema_start_step: int = 0
cm_weighing_function: str = "constant"
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -796,6 +773,14 @@ class TrainingArgs(FastVideoArgs):
help="Directory to cache models")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
default=0.999,
help="EMA decay rate")
parser.add_argument("--ema-start-step",
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--training-cfg-rate",
type=float,
help="Classifier-free guidance scale")
@@ -1030,40 +1015,6 @@ class TrainingArgs(FastVideoArgs):
"--simulate-generator-forward",
action=StoreBoolean,
help="Whether to simulate generator forward to match inference")
parser.add_argument("--simulate-forward-interval",
type=int,
default=TrainingArgs.simulate_forward_interval,
help="Ratio of steps to simulate generator forward")
parser.add_argument("--regression-loss-weight",
type=float,
default=TrainingArgs.regression_loss_weight,
help="Weight for regression loss")
parser.add_argument("--use-regression-loss",
action=StoreBoolean,
help="Whether to use regression loss")
# Consistency model arguments
parser.add_argument("--ema-decay",
type=float,
default=0.999,
help="EMA decay rate")
parser.add_argument("--ema-start-step",
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--cm-loss-weight",
type=float,
default=TrainingArgs.cm_loss_weight,
help="Weight for consistency loss")
parser.add_argument(
"--cm-use-ema-teacher",
action=StoreBoolean,
help="Whether to use EMA teacher for consistency loss")
parser.add_argument("--cm-weighing-function",
type=str,
choices=["constant", "sigma_sqrt"],
default=TrainingArgs.cm_weighing_function,
help="Weighting function for consistency loss")
return parser
+3 -3
View File
@@ -101,15 +101,15 @@ class BaseLayerWithLoRA(nn.Module):
B: torch.Tensor,
training_mode: bool = False,
lora_path: str | None = None) -> None:
self.lora_A = torch.nn.Parameter(
A) # share storage with weights in the pipeline
self.lora_B = torch.nn.Parameter(B)
self.lora_A = A # share storage with weights in the pipeline
self.lora_B = B
self.disable_lora = False
if not training_mode:
self.merge_lora_weights()
self.lora_path = lora_path
@torch.no_grad()
# @torch.compile()
def merge_lora_weights(self) -> None:
if self.disable_lora:
return
+2
View File
@@ -672,9 +672,11 @@ class WanTransformer3DModel(CachableDiT):
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
+2 -2
View File
@@ -72,7 +72,8 @@ class ComponentLoader(ABC):
module_loaders = {
"scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"),
"transformer_2": (TransformerLoader, "diffusers"),
"real_score_transformer": (TransformerLoader, "diffusers"),
"fake_score_transformer": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
@@ -451,7 +452,6 @@ class TransformerLoader(ComponentLoader):
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
cpu_offload=fastvideo_args.dit_cpu_offload,
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
fsdp_inference=fastvideo_args.use_fsdp_inference,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
+1 -2
View File
@@ -8,11 +8,10 @@ import torch
class BaseScheduler(ABC):
timesteps: torch.Tensor
order: int
num_train_timesteps: int
def __init__(self, *args, **kwargs) -> None:
# Check if subclass has defined all required properties
required_attributes = ['timesteps', 'order', 'num_train_timesteps']
required_attributes = ['timesteps', 'order']
for attr in required_attributes:
if not hasattr(self, attr):
@@ -134,7 +134,6 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.timesteps = sigmas * num_train_timesteps
self.num_train_timesteps = num_train_timesteps
self._step_index: int | None = None
self._begin_index: int | None = None
@@ -146,8 +145,6 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
BaseScheduler.__init__(self)
@property
def shift(self) -> float:
"""
@@ -117,7 +117,6 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.sigmas = sigmas
self.timesteps = sigmas * num_train_timesteps
self.num_train_timesteps = num_train_timesteps
self.model_outputs = [None] * solver_order
self.timestep_list: list[Any | None] = [None] * solver_order
@@ -133,8 +132,6 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
BaseScheduler.__init__(self)
@property
def step_index(self):
"""
@@ -273,7 +273,6 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.num_inference_steps = None
timesteps = np.linspace(0, num_train_timesteps - 1, num_train_timesteps, dtype=np.float32)[::-1].copy()
self.timesteps = torch.from_numpy(timesteps)
self.num_train_timesteps = num_train_timesteps
self.model_outputs = [None] * solver_order
self.timestep_list = [None] * solver_order
self.lower_order_nums = 0
@@ -12,10 +12,12 @@ from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
from fastvideo.pipelines.stages import (ImageEncodingStage, ConditioningStage,
DecodingStage, DmdDenoisingStage,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
@@ -65,7 +67,7 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
@@ -12,10 +12,12 @@ from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
from fastvideo.pipelines.stages import (ImageEncodingStage, ConditioningStage,
DecodingStage, DenoisingStage,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
@@ -46,14 +48,11 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
tokenizers=[self.get_module("tokenizer")],
))
if (self.get_module("image_encoder") is not None
and self.get_module("image_processor") is not None):
self.add_stage(
stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
@@ -68,12 +67,11 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
@@ -61,7 +61,6 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self))
+12 -35
View File
@@ -37,7 +37,6 @@ class ComposedPipelineBase(ABC):
is_video_pipeline: bool = False # To be overridden by video pipelines
_required_config_modules: list[str] = []
_extra_config_module_map: dict[str, str] = {}
training_args: TrainingArgs | None = None
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
modules: dict[str, torch.nn.Module] = {}
@@ -149,7 +148,9 @@ 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.precision = fastvideo_args.master_weight_type
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
@@ -238,18 +239,6 @@ class ComposedPipelineBase(ABC):
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")
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)
model_index.pop("expand_timesteps", None)
@@ -259,20 +248,12 @@ class ComposedPipelineBase(ABC):
) > 1, "model_index.json must contain at least one pipeline module"
for module_name in self.required_config_modules:
if module_name not in model_index and module_name in self._extra_config_module_map:
extra_module_value = self._extra_config_module_map[module_name]
if module_name not in model_index:
logger.warning(
"model_index.json does not contain a %s module, but found {%s: %s} in _extra_config_module_map, adding to model_index.",
module_name, module_name, extra_module_value)
if extra_module_value in model_index:
logger.info("Using module %s for %s", extra_module_value,
module_name)
model_index[module_name] = model_index[extra_module_value]
continue
else:
raise ValueError(
f"Required module key: {module_name} value: {model_index.get(module_name)} was not found in loaded modules {model_index.keys()}"
)
"model_index.json does not contain a %s module, adding %s to model_index",
module_name, module_name)
if 'transformer' in module_name:
model_index[module_name] = model_index['transformer']
# all the component models used by the pipeline
required_modules = self.required_config_modules
@@ -282,7 +263,6 @@ class ComposedPipelineBase(ABC):
for module_name, (transformers_or_diffusers,
architecture) in model_index.items():
if transformers_or_diffusers is None:
self.required_config_modules.remove(module_name)
continue
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
@@ -291,17 +271,14 @@ class ComposedPipelineBase(ABC):
logger.info("Using module %s already provided", module_name)
modules[module_name] = loaded_modules[module_name]
continue
# we load the module from the extra config module map if it exists
if module_name in self._extra_config_module_map:
load_module_name = self._extra_config_module_map[module_name]
if 'transformer' in module_name:
loading_module_name = module_name.split("_")[-1]
else:
load_module_name = module_name
loading_module_name = module_name
component_model_path = os.path.join(self.model_path,
load_module_name)
loading_module_name)
module = PipelineComponentLoader.load_module(
module_name=load_module_name,
module_name=module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=fastvideo_args,
+8 -24
View File
@@ -9,14 +9,11 @@ in a functional manner, reducing the need for explicit parameter passing.
import pprint
from dataclasses import asdict, dataclass, field
from typing import TYPE_CHECKING, Any
from typing import Any
import PIL.Image
import torch
if TYPE_CHECKING:
from torchcodec.decoders import VideoDecoder
from fastvideo.attention import AttentionMetadata
from fastvideo.configs.sample.teacache import TeaCacheParams, WanTeaCacheParams
@@ -78,15 +75,15 @@ class ForwardBatch:
image_latent: torch.Tensor | None = None
# Latent dimensions
height_latents: list[int] | int | None = None
width_latents: list[int] | int | None = None
num_frames: list[int] | int = 1 # Default for image models
height_latents: int | None = None
width_latents: int | None = None
num_frames: int = 1 # Default for image models
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
# Original dimensions (before VAE scaling)
height: list[int] | int | None = None
width: list[int] | int | None = None
fps: list[int] | int | None = None
height: int | None = None
width: int | None = None
fps: int | None = None
# Timesteps
timesteps: torch.Tensor | None = None
@@ -96,7 +93,6 @@ class ForwardBatch:
# Scheduler parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
guidance_rescale: float = 0.0
eta: float = 0.0
sigmas: list[float] | None = None
@@ -118,8 +114,6 @@ class ForwardBatch:
# Misc
save_video: bool = True
return_frames: bool = False
return_trajectory_latents: bool = False
trajectory_latents: list[torch.Tensor] = field(default_factory=list)
# TeaCache parameters
enable_teacache: bool = False
@@ -142,8 +136,6 @@ class ForwardBatch:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
if self.guidance_scale_2 is None:
self.guidance_scale_2 = self.guidance_scale
def __str__(self):
return pprint.pformat(asdict(self), indent=2, width=120)
@@ -195,14 +187,6 @@ class TrainingBatch:
# Distillation losses
generator_loss: float = 0.0
fake_score_loss: float = 0.0
regression_loss: float = 0.0
consistency_loss: float = 0.0
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
@dataclass
class PreprocessBatch(ForwardBatch):
video_loader: list["VideoDecoder"] = field(default_factory=list)
video_file_name: list[str] = field(default_factory=list)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
+3 -3
View File
@@ -65,7 +65,7 @@ class _PipelineRegistry:
arch = _PIPELINE_NAME_TO_ARCHITECTURE_NAME[pipeline_name_in_config]
return set(self.pipelines[pipeline_type.value][arch].keys())
def _load_preprocess_pipeline_cls(
def _load_preprocessing_pipeline_cls(
self, workload_type: WorkloadType,
arch: str) -> type[ComposedPipelineBase] | None:
if workload_type == WorkloadType.I2V:
@@ -90,7 +90,7 @@ class _PipelineRegistry:
return None
if pipeline_type == PipelineType.PREPROCESS:
return self._load_preprocess_pipeline_cls(workload_type, arch)
return self._load_preprocessing_pipeline_cls(workload_type, arch)
elif pipeline_type == PipelineType.BASIC:
return self.pipelines[
pipeline_type.value][arch][pipeline_name_in_config]
@@ -131,7 +131,7 @@ def import_pipeline_classes(
Import pipeline classes based on the pipeline type and workload type.
Args:
pipeline_types: The pipeline types to load (basic, preprocess, training).
pipeline_types: The pipeline types to load (basic, preprocessing, training).
If None, loads all types.
Returns:

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