Compare commits

..
Author SHA1 Message Date
SolitaryThinker d42fdc56f2 chekcpoint 2025-08-26 22:54:58 +00:00
William Lin 3ef04f1654 [misc] [docs] Various fixes for logging and docs (#758) 2025-08-23 21:13:50 -07:00
Jinzhe Pan 0eced76a41 [Feat][Preprocess] support multi-gpus (#753) 2025-08-23 11:34:42 +08:00
Jinzhe Pan 3ab6470d1a [Feat][Preprocess] support merged dataset (#752) 2025-08-22 15:29:33 -07:00
Wenxuan Tan 989a03532c Optionally use unmerged weights for inference (#745) 2025-08-22 15:20:31 -07:00
William Lin fa15369a02 [bugfix] Check that model_index.json module is in required_modules list before removing (#756) 2025-08-22 14:36:44 -07:00
Zhang Peiyuan 78a9cb88d8 [Fix] fix seed in dmd denoising loop (#736) 2025-08-21 18:06:16 -07:00
Peng Xiaoand肖鹏 a0bff12746 [bugfix] [dmd] Align backward simulation with dmd2 sample back (#744)
Co-authored-by: 肖鹏 <xiaopeng1@aishi.ai>
2025-08-20 22:25:33 -07:00
William Lin 98f2af94e5 [bugfix] Missing Docker file for cuda12.9 (#750) 2025-08-20 15:34:31 -07:00
William Lin 46f7b6d574 [Docker] add 12.9 docker image and also fix py3.10 and py3.11 dockerfile (#749) 2025-08-20 15:31:15 -07:00
Jinzhe Pan 911a6a6a35 [Feat][Preprocessing] i2v preprocessing workflow (#737) 2025-08-14 20:47:25 -07:00
Zhang Peiyuan 38c7949d5c Update WeChat group link (#739) 2025-08-14 15:03:35 -07:00
Jinzhe Pan 7e7a0dba9d feat: preprocess validation dataset only when exist (#734) 2025-08-12 02:16:31 -07:00
Zhang Peiyuan f62e210ae6 Fix vsa backward gQ (#735) 2025-08-11 21:43:13 -07:00
William Lin 6ceb4942a0 [bugfix] [dmd] Fix backward simulation and also naming in wan_i2v_dmd_pipeline (#731) 2025-08-10 21:13:30 -07:00
William LinandRandNMR73 8cae5e4708 [feature] add Gradio live serving demo code (#727)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-08-10 15:34:03 -07:00
William Lin 2a773fa34e [bugfix] [distill] remove i2v validation schema import in distill (#728) 2025-08-09 20:47:42 -07:00
Wenxuan Tan 5357f63327 Fix LoRA load from training checkpoint (#719) 2025-08-09 20:46:00 -05:00
William Lin 60f61c8101 [bugfix] fix pyproject install and VSA precision test (#726) 2025-08-08 18:45:03 -07:00
Jiali Chen 3d75ba8251 update version selection for VSA workflow (#725) 2025-08-08 13:05:16 -07:00
Wenxuan Tan 6c6bcd914d Remove all empty_cache (#713) 2025-08-07 22:50:38 -07:00
Jiali Chen f79b08de81 add cicd workflow for publishing VSA kernel (#723) 2025-08-07 18:53:05 -07:00
Jinzhe Pan f2bc037fff [Fix] training pipeline pin_cpu_memory issue (#692) 2025-08-07 02:31:20 -07:00
Jinzhe Pan 86604a684b [3/3][Preprocess] add preprocessing workflows (#645) 2025-08-07 01:49:07 -07:00
Zhang Peiyuan 47bd1e0178 [Misc] change installation logic of vsa (#721) 2025-08-06 21:54:09 -07:00
Wei ZhouandSolitaryThinker c41305ad18 [Feat] Add Wan2.2 14B MoE (#688)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-06 20:31:03 -07:00
Zhang Peiyuan 98ce9034f0 [Chore] Include our demo in the readme. (#720) 2025-08-06 19:29:40 -07:00
William Lin 0ceff110da [chore] Release 0.1.5 (#717) 2025-08-06 13:07:52 -07:00
Yongqi Chen 1d018acb3e [Feature]Add Data-free distillation readme (#710) 2025-08-05 14:27:39 -04:00
Yongqi Chen 7d8cf38dbe Fix typo (#709) 2025-08-04 20:21:21 -07:00
Yongqi Chen 8d483fe4aa [Bugfix] Fix neg_prompt bug when training from local cp (#708) 2025-08-04 15:54:06 -07:00
Zhang Peiyuan c1191250bf Add WeChat group link (#707) 2025-08-04 15:19:01 -07:00
Wenxuan Tan 4b7266349a [misc] Remove allow_tf32 in scripts (#705) 2025-08-04 15:37:56 -05:00
Yongqi Chen 22f9b7681f [Feature]Update Wan2.2+DMD doc example (#706) 2025-08-04 16:14:22 -04:00
Yongqi Chen 589d32cc39 [Feature] Update Readme and scripts (#703) 2025-08-04 15:02:32 -04:00
Hao Zhang 89199837db Update readme pre-release (#704) 2025-08-04 11:54:31 -07:00
131 changed files with 4225 additions and 779 deletions
+1
View File
@@ -163,6 +163,7 @@ 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"
+15
View File
@@ -18,6 +18,12 @@ on:
required: false
default: false
type: boolean
python_3_12_cuda_12_9:
description: 'Build Python 3.12 image Cuda 12.9'
required: false
default: false
type: boolean
permissions:
contents: read
@@ -49,4 +55,13 @@ jobs:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12
tag_suffix: py3.12
secrets: inherit
build-python-3-12-cuda-12-9:
if: ${{ github.event.inputs.python_3_12_cuda_12_9 == 'true' }}
uses: ./.github/workflows/build-image-template.yml
with:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12.cuda12.9.1
tag_suffix: py3.12-cuda12.9.1
secrets: inherit
+257
View File
@@ -0,0 +1,257 @@
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/
+1
View File
@@ -22,6 +22,7 @@ exclude: |
examples/.*|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml|
.github/workflows/vsa-publish.yml|
.github/workflows/build-image-template.yml|
docs/source/inference/support_matrix.md
)
+10 -8
View File
@@ -1,5 +1,5 @@
<div align="center">
<img src=assets/logo.png width="30%"/>
<img src=assets/logos/logo.svg width="30%"/>
</div>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
@@ -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://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> |
| 🕹️ <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/wZPZTLKg" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -64,12 +64,15 @@ See below for recipes and datasets:
## Inference
### Generating Your First Video
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
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:
```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",
@@ -151,11 +154,10 @@ If you find FastVideo useful, please considering citing our work:
year = {2024},
}
@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},
@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},
year={2025}
}
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 46 KiB

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

Binary file not shown.

Before

Width:  |  Height:  |  Size: 31 KiB

+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_block_sparse.py # test VSA
python tests/test_vsa.py # test VSA
```
### Benchmark
```bash
-2
View File
@@ -86,8 +86,6 @@ def benchmark_attention(configurations):
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
torch.cuda.empty_cache()
return results
+10 -13
View File
@@ -51,19 +51,16 @@ for k in kernels:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
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'])
]
ext_modules = [
CUDAExtension('vsa_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
]
+23 -26
View File
@@ -13,9 +13,9 @@ BLOCK_M = 64
BLOCK_N = 64
def pytorch_test(Q, K, V, block_sparse_mask, dO):
q_ = Q.clone().requires_grad_()
k_ = K.clone().requires_grad_()
v_ = V.clone().requires_grad_()
q_ = Q.clone().float().requires_grad_()
k_ = K.clone().float().requires_grad_()
v_ = V.clone().float().requires_grad_()
QK = torch.matmul(q_, k_.transpose(-2, -1))
QK /= (q_.size(-1) ** 0.5)
@@ -35,9 +35,9 @@ def pytorch_test(Q, K, V, block_sparse_mask, dO):
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
Q = Q.clone().requires_grad_()
K = K.clone().requires_grad_()
V = V.clone().requires_grad_()
Q = Q.detach().requires_grad_()
K = K.detach().requires_grad_()
V = V.detach().requires_grad_()
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
@@ -60,11 +60,9 @@ def get_non_pad_index(
return index_pad[index_mask]
def generate_tensor(shape, mean, std, dtype, device):
def generate_tensor(shape, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
return tensor
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
@@ -75,7 +73,7 @@ def vsa_pad(x, non_pad_index, num_blocks, block_size):
padded_x[:, :, non_pad_index, :] = x
return padded_x
def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_mode='all'):
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
results = {
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
@@ -91,10 +89,10 @@ def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_m
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
K = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
V = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
# dO_padded = torch.zeros_like(dO_padded)
# dO_padded[:, :, non_pad_index, :] = dO
@@ -107,7 +105,8 @@ def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_m
abs_diff = torch.abs(diff)
results[name]['sum_diff'] += torch.sum(abs_diff).item()
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
results[name]['max_diff'] = max(results[name]['max_diff'], torch.max(abs_diff).item())
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
if torch.cuda.is_available():
torch.cuda.empty_cache()
@@ -119,27 +118,27 @@ def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_m
return results
def generate_error_graphs(h, d, mean, std, error_mode='all'):
def generate_error_graphs(h, d, error_mode='all'):
test_configs = [
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
]
print(f"\nError Analysis for h={h}, d={d}, mean={mean}, std={std}, mode={error_mode}")
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
print("=" * 150)
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
f"{'gQ Avg':<12} {'gQ Max':<12} "
f"{'gK Avg':<12} {'gK Max':<12} "
f"{'gV Avg':<12} {'gV Max':<12} "
f"{'gO Avg':<12} {'gO Max':<12}")
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
f"{'gK Avg':<12} {'Rel gK Max':<12} "
f"{'gV Avg':<12} {'Rel gV Max':<12} "
f"{'gO Avg':<12} {'Rel gO Max':<12}")
print("-" * 150)
for config in test_configs:
num_blocks = config["num_blocks"]
k = config["k"]
description = config["description"]
results = check_correctness(h, d, num_blocks, k, mean, std, error_mode=error_mode)
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
print(f"{description:<20} {num_blocks:<8} {k:<4} "
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
@@ -150,10 +149,8 @@ def generate_error_graphs(h, d, mean, std, error_mode='all'):
if __name__ == "__main__":
h, d = 16, 128
mean = 0.0
std = 1
print("Block Sparse Attention with Variable Block Sizes Analysis")
print("=" * 60)
for mode in ['backward']:
generate_error_graphs(h, d, mean, std, error_mode=mode)
generate_error_graphs(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
+4 -3
View File
@@ -1,12 +1,13 @@
import torch
from typing import Tuple
block_sparse_attn=None
try:
import torch
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
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
except ImportError:
else:
from vsa.block_sparse_wrapper import block_sparse_attn_triton
block_sparse_fwd = None
block_sparse_bwd = None
+4 -4
View File
@@ -568,7 +568,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
__syncthreads(); // wait for sd_smem shared memory write
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
warpgroup::mma_commit_group();
tma::store_async_wait();
warpgroup::mma_async_wait();
// store qg to shared memory
warpgroup::store(qg_smem, qg_reg);
@@ -578,6 +577,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if (threadIdx.x / 32 == 0) {
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
tma::store_add_async(g.qg, qg_smem, tile_idx);
tma::store_async_wait();
}
}
@@ -624,7 +624,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
__syncthreads(); // wait for sd_smem shared memory write
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
warpgroup::mma_commit_group();
tma::store_async_wait();
warpgroup::mma_async_wait();
// store qg to shared memory
warpgroup::store(qg_smem, qg_reg);
@@ -634,13 +633,14 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if (threadIdx.x / 32 == 0) {
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
tma::store_add_async(g.qg, qg_smem, tile_idx);
tma::store_async_wait();
}
}
// store kq and vq
// ! the following two line seems unnecessary.
tma::store_async_wait(); // ensure qg is finished
// tma::store_async_wait(); // ensure qg is finished
__syncthreads();
warpgroup::store(kg_smem[0], kg_reg);
@@ -1174,4 +1174,4 @@ block_sparse_attention_backward(torch::Tensor q,
return {qg, kg, vg};
//cudadevicesynchronize();
}
}
+93 -90
View File
@@ -11,95 +11,6 @@ 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,
@@ -179,4 +90,96 @@ 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)
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)
+44 -20
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.10.0 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.10 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_sta.py install
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
EXPOSE 22
+44 -20
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.11.11 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.11 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_sta.py install
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
EXPOSE 22
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
+72
View File
@@ -0,0 +1,72 @@
FROM nvidia/cuda:12.9.1-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.9
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.12 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_sta.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
EXPOSE 22
+1 -2
View File
@@ -96,8 +96,7 @@ copybutton_prompt_is_regexp = True
#
html_title = project
html_theme = 'sphinx_book_theme'
html_logo = '../../assets/logo.jpg'
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
html_logo = '../../assets/logos/icon_simple.svg'
html_theme_options = {
'path_to_docs': 'docs/source',
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
@@ -3,12 +3,12 @@
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
**Images:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
## Starting the container
```bash
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
```
This will:
@@ -6,7 +6,7 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
## Creating a new pod
Choose a GPU that supports CUDA 12.4
Choose a GPU that supports CUDA 12.8
Pick 1 or 2 L40S GPU(s)
+24 -3
View File
@@ -22,10 +22,20 @@ source ~/.bashrc
Create and activate a Conda environment for FastVideo:
```
conda create -n fastvideo python=3.10 -y
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Install `uv` (optional, but recommended):
From instructions on [uv](https://astral.sh/uv/):
```
curl -LsSf https://astral.sh/uv/install.sh | sh
# or
wget -qO- https://astral.sh/uv/install.sh | sh
```
Clone the FastVideo repository and go to the FastVideo directory:
```
@@ -36,10 +46,10 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
```bash
pip install -e .[dev]
uv pip install -e .[dev]
# Can also install flash-attn (optional)
pip install flash-attn==2.7.4.post1 --no-build-isolation
uv pip install flash-attn --no-build-isolation
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
@@ -50,3 +60,14 @@ pre-commit run --all-files
# Unit tests
pytest tests/
```
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
```
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
# make sure you have ninja installed
uv pip install ninja
python setup.py install
```
+22 -4
View File
@@ -6,8 +6,9 @@ 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 **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
- **[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
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).
@@ -34,7 +35,7 @@ python scripts/huggingface/download_hf.py \
## 🚀 Training Scripts
### 1.3B Model Sparse-Distill
### Wan2.1 1.3B Model Sparse-Distill
For the 1.3B model, we use **4 nodes with 32 H200 GPUs** (8 GPUs per node):
@@ -50,7 +51,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)
### 14B Model Sparse-Distill
### Wan2.1 14B Model Sparse-Distill
For the 14B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
@@ -67,3 +68,20 @@ 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
@@ -6,7 +6,7 @@ Instructions to install FastVideo for NVIDIA CUDA GPUs.
- **OS: Linux or Windows WSL**
- **Python: 3.10-3.12**
- **CUDA 12.4**
- **CUDA 12.8**
- **At least 1 NVIDIA GPU**
## Set up using Python
@@ -38,6 +38,7 @@ conda activate fastvideo
:::{tip}
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
Note that you can also use `uv` to install FastVideo in a Conda environment.
:::
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
@@ -60,7 +61,7 @@ uv pip install fastvideo
Also optionally install flash-attn:
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
pip install flash-attn --no-build-isolation
```
### Installation from Source
@@ -87,7 +88,7 @@ uv pip install -e .
#### Flash Attention
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
pip install flash-attn --no-build-isolation
```
## Set up using Docker
@@ -102,7 +103,7 @@ If you're planning to contribute to FastVideo please see the following page:
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA 12.4 support
- NVIDIA GPU with CUDA 12.8 support
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
@@ -39,6 +39,7 @@ conda activate fastvideo
:::{tip}
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
Note that you can also use `uv` to install FastVideo in a Conda environment.
:::
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
+2 -3
View File
@@ -1,6 +1,6 @@
# Welcome to FastVideo
:::{figure} ../../assets/logo.png
:::{figure} ../../assets/logos/logo.svg
:align: center
:alt: FastVideo
:class: no-scaled-link
@@ -9,7 +9,7 @@
:::{raw} html
<p style="text-align:center">
<strong>FastVideo is An unified inference and post-training framework for accelerated video generation.
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.
</strong>
</p>
@@ -101,7 +101,6 @@ sliding_tile_attention/demo
:maxdepth: 1
video_sparse_attention/installation
video_sparse_attention/demo
:::
:::{toctree}
@@ -5,7 +5,7 @@ This page contains step-by-step instructions to get you quickly started with vid
## Requirements
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.4
- **CUDA**: 12.8
- **GPU**: At least one NVIDIA GPU
## Installation
+50 -6
View File
@@ -6,6 +6,7 @@ The symbols used have the following meanings:
- ✅ = Full compatibility
- ❌ = No compatibility
- ⭕ = Does not apply to this model
## Models x Optimization
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
@@ -37,51 +38,94 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
* TeaCache
* Sliding Tile Attn
* Sage Attn
* Video Sparse Attention (VSA)
- * FastWan2.1 T2V 1.3B
* `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
* 480P
* ⭕
* ⭕
* ⭕
* ✅
- * FastWan2.2 TI2V 5B Full Attn*
* `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers`
* 720P
* ⭕
* ⭕
* ⭕
* ✅
- * Wan2.2 TI2V 5B
* `Wan-AI/Wan2.2-TI2V-5B-Diffusers`
* 720P
* ⭕
* ⭕
* ✅
* ⭕
- * Wan2.2 T2V A14B
* `Wan-AI/Wan2.2-T2V-A14B-Diffusers`
* 480P<br>720P
* ❌
* ❌
* ✅
* ⭕
- * Wan2.2 I2V A14B
* `Wan-AI/Wan2.2-I2V-A14B-Diffusers`
* 480P<br>720P
* ❌
* ❌
* ✅
* ⭕
- * HunyuanVideo
* `hunyuanvideo-community/HunyuanVideo`
* 720px1280p<br>544px960p
* ❌
* ✅
* ✅
* ⭕
- * FastHunyuan
* `FastVideo/FastHunyuan-diffusers`
* 720px1280p<br>544px960p
* ❌
* ✅
* ✅
- * Wan T2V 1.3B
* ⭕
- * Wan2.1 T2V 1.3B
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan T2V 14B
* ⭕
- * Wan2.1 T2V 14B
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
* 480P, 720P
* ✅
* ✅*
* ✅
- * Wan I2V 480P
* ⭕
- * Wan2.1 I2V 480P
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan I2V 720P
* ⭕
- * Wan2.1 I2V 720P
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
* 720P
* ✅
* ✅*
* ✅
* ✅
* ⭕
- * StepVideo T2V
* `FastVideo/stepvideo-t2v-diffusers`
* 768px768px204f<br>544px992px204f<br>544px992px136f
* ❌
* ❌
* ✅
* ⭕
:::
**Note**: there are some known quality issues with Wan2.1 + Sliding Tile Attn. We are working on fixing this issue.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
## Special requirements
@@ -53,3 +53,9 @@ 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
```
@@ -1,3 +0,0 @@
(vsa-demo)=
# 🎬 Demo
@@ -23,10 +23,10 @@ sudo apt update
sudo apt install clang-11
```
Set up CUDA environment (if using CUDA 12.4):
Set up CUDA environment (if using CUDA 12.8):
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
@@ -59,3 +59,9 @@ 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
```
@@ -86,7 +86,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "1.0" # not used for dmd inference
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
@@ -102,7 +102,6 @@ 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 "1.0" # not used for dmd inference
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
@@ -102,7 +102,6 @@ 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 "1.0" # not used for dmd inference
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
@@ -102,7 +102,6 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -0,0 +1,13 @@
# 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 "1.0" # not used for dmd inference
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
@@ -108,7 +108,6 @@ 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 "1.0" # not used for dmd inference
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
@@ -108,7 +108,6 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
@@ -9,4 +9,14 @@ git submodule update --init --recursive
python setup_vsa.py install
```
### TODO
### 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/
```
@@ -63,7 +63,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "1.0" # not used for dmd inference
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
@@ -77,7 +77,6 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
+1 -2
View File
@@ -16,8 +16,7 @@ def main():
dit_cpu_offload=False,
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=False,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# 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=False,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
+48
View File
@@ -0,0 +1,48 @@
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()
@@ -0,0 +1,30 @@
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()
-59
View File
@@ -1,59 +0,0 @@
# FastVideo Gradio Demo
This is a Gradio-based web interface for generating videos using the FastVideo framework. The demo allows users to create videos from text prompts with various customization options.
## Overview
The demo uses the FastVideo framework to generate videos based on text prompts. It provides a simple web interface built with Gradio that allows users to:
- Enter text prompts to generate videos
- Customize video parameters (dimensions, number of frames, etc.)
- Use negative prompts to guide the generation process
- Set or randomize seeds for reproducibility
---
## Usage
Run the demo with:
```bash
python examples/inference/gradio/gradio_demo.py
```
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
---
## Model Initialization
This demo initializes a `VideoGenerator` with the minimum required arguments for inference. Users can seamlessly adjust inference options between generations, including prompts, resolution, video length, or even the number of inference steps, *without ever needing to reload the model*.
## Video Generation
The core functionality is in the `generate_video` function, which:
1. Processes user inputs
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
3. Returns an output path that Gradio uses to display the generated video
## Gradio Interface
The interface is built with several components:
- A text input for the prompt
- A video display for the result
- Inference options in a collapsible accordion:
- Height and width sliders
- Number of frames slider
- Guidance scale slider
- Inference steps slider
- Negative prompt options
- Seed controls
### Inference Options
- **Height/Width**: Control the resolution of the generated video
- **Number of Frames**: Set how many frames to generate
- **Guidance Scale**: Control how closely the generation follows the prompt
- **Inference Steps**: More steps can improve quality but take longer
- **Negative Prompt**: Specify what you don't want to see in the video
- **Seed**: Control randomness for reproducible results
-169
View File
@@ -1,169 +0,0 @@
import argparse
import os
from copy import deepcopy
import gradio as gr
import torch
from fastvideo import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Gradio Demo")
parser.add_argument("--model_path",
type=str,
default="FastVideo/FastHunyuan-diffusers",
help="Path to the model")
parser.add_argument("--num_gpus",
type=int,
default=1,
help="Number of GPUs to use")
parser.add_argument("--output_path",
type=str,
default="outputs",
help="Path to save generated videos")
parsed_args = parser.parse_args()
# args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=parsed_args.model_path, num_gpus=parsed_args.num_gpus)
default_params = SamplingParam.from_pretrained(parsed_args.model_path)
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
params = deepcopy(default_params)
params.prompt = prompt
params.negative_prompt = negative_prompt
params.seed = seed
params.guidance_scale = guidance_scale
params.num_frames = num_frames
params.height = height
params.width = width
params.num_inference_steps = num_inference_steps
if randomize_seed:
params.seed = torch.randint(0, 1000000, (1, )).item()
if not use_negative_prompt:
params.negative_prompt = None
generator.generate_video(prompt=prompt, sampling_param=params)
output_path = os.path.join(parsed_args.output_path,
f"{params.prompt[:100]}.mp4")
return output_path, params.seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
with gr.Blocks() as demo:
gr.Markdown("# FastVideo Inference Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=default_params.height,
)
width = gr.Slider(label="Width",
minimum=256,
maximum=1024,
step=32,
value=default_params.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=default_params.num_frames,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=default_params.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=default_params.num_inference_steps,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=default_params.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=default_params.negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
@@ -0,0 +1,730 @@
import argparse
import os
import requests
import base64
import time
import gradio as gr
from fastvideo.configs.sample.base import SamplingParam
MODEL_PATH_MAPPING = {
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
}
class RayServeClient:
def __init__(self, backend_url: str):
self.backend_url = backend_url
self.session = requests.Session()
def check_health(self) -> bool:
try:
response = self.session.get(f"{self.backend_url}/health", timeout=5)
return response.status_code == 200
except requests.exceptions.RequestException:
return False
def generate_video(self, request_data: dict) -> dict:
start_time = time.time()
try:
headers = {"Content-Type": "application/json"}
response = self.session.post(
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=300
)
round_trip_time = time.time() - start_time
if response.status_code == 200:
result = response.json()
backend_total = result.get("total_time", 0)
network_time = round_trip_time - backend_total
result["network_time"] = network_time
return result
else:
return {"success": False, "error_message": f"HTTP {response.status_code}: {response.text}"}
except requests.exceptions.RequestException as e:
return {"success": False, "error_message": f"Request failed: {str(e)}"}
def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str:
if not video_data:
return None
try:
if video_data.startswith('data:video/'):
video_data = video_data.split(',')[1]
video_bytes = base64.b64decode(video_data)
safe_prompt = prompt[:50].replace(' ', '_').replace('/', '_').replace('\\', '_')
video_filename = f"{safe_prompt}.mp4"
video_path = os.path.join(output_dir, video_filename)
os.makedirs(output_dir, exist_ok=True)
with open(video_path, 'wb') as f:
f.write(video_bytes)
return video_path
except Exception as e:
print(f"Failed to save video: {e}")
return None
def create_timing_display(inference_time, encoding_time, network_time, total_time, stage_execution_times, num_frames):
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
timing_html = f"""
<div style="margin: 10px 0;">
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
<div class="timing-card timing-card-highlight">
<div style="font-size: 20px;">🚀</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🧠</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🎬</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
<div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🌐</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
<div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">📊</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
</div>
</div>"""
if inference_time > 0:
fps = num_frames / inference_time
timing_html += f"""
<div class="performance-card" style="margin-top: 15px;">
<span style="font-weight: bold;">Generation Speed: </span>
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
</div>"""
return timing_html + "</div>"
def load_example_prompts():
def contains_chinese(text):
return any('\u4e00' <= char <= '\u9fff' for char in text)
def load_from_file(filepath):
prompts, labels = [], []
try:
with open(filepath, "r", encoding='utf-8') as f:
for line in f:
line = line.strip()
if line and not contains_chinese(line):
label = line[:100] + "..." if len(line) > 100 else line
labels.append(label)
prompts.append(line)
except Exception as e:
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("prompts/prompts_final.txt")
if not examples:
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
example_labels = ["Crowded rooftop bar at night"]
return examples, example_labels
def create_gradio_interface(backend_url: str, default_params: dict[str, SamplingParam]):
client = RayServeClient(backend_url)
def generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
):
if not client.check_health():
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
# Validate dimensions
max_pixels = 720 * 1280
if height * width > max_pixels:
return None, f"Video dimensions too large. Maximum: 720x1280 pixels", ""
if progress:
progress(0.1, desc="Checking backend health...")
request_data = {
"prompt": prompt,
"negative_prompt": negative_prompt,
"use_negative_prompt": use_negative_prompt,
"seed": seed,
"guidance_scale": guidance_scale,
"num_frames": num_frames,
"height": height,
"width": width,
"randomize_seed": randomize_seed,
"return_frames": False,
"image_path": None,
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
}
if progress:
progress(0.4, desc="Generating video...")
response = client.generate_video(request_data)
if progress:
progress(0.8, desc="Processing response...")
if response.get("success", False):
video_data = response.get("video_data", "")
used_seed = response.get("seed", seed)
inference_time = response.get("inference_time", 0.0)
encoding_time = response.get("encoding_time", 0.0)
total_time = response.get("total_time", 0.0)
network_time = response.get("network_time", 0.0)
stage_execution_times = response.get("stage_execution_times", [])
timing_details = create_timing_display(
inference_time, encoding_time, network_time, total_time,
stage_execution_times, num_frames
)
if video_data:
if progress:
progress(0.9, desc="Saving video...")
video_path = save_video_from_base64(video_data, "outputs", prompt)
if progress:
progress(1.0, desc="Generation complete!")
if video_path and os.path.exists(video_path):
return video_path, used_seed, timing_details
else:
return None, "Failed to save video", ""
else:
return None, "No video data received from backend", ""
else:
error_msg = response.get("error_message", "Unknown error occurred")
return None, f"Generation failed: {error_msg}", ""
examples, example_labels = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
button_primary_background_fill_hover="#1d4ed8",
button_primary_text_color="white",
slider_color="#2563eb",
checkbox_background_color_selected="#2563eb",
)
def get_default_values(model_name):
model_path = MODEL_PATH_MAPPING.get(model_name)
if model_path and model_path in default_params:
params = default_params[model_path]
return {
'height': params.height,
'width': params.width,
'num_frames': params.num_frames,
'guidance_scale': params.guidance_scale,
'seed': params.seed,
}
return {
'height': 448,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
with gr.Blocks(title="FastWan", theme=theme) as demo:
gr.Image("assets/logos/logo.svg", show_label=False, container=False, height=80)
gr.HTML("""
<div style="text-align: center; margin-bottom: 10px;">
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
with gr.Accordion("🎥 What Is FastVideo?", open=False):
gr.HTML("""
<div style="padding: 20px; line-height: 1.6;">
<p style="font-size: 16px; margin-bottom: 15px;">
FastVideo is an inference and post-training framework for diffusion models. It 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>
</div>
""")
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
label="Select Model",
interactive=True
)
with gr.Row():
example_dropdown = gr.Dropdown(
choices=example_labels,
label="Example Prompts",
value=None,
interactive=True,
allow_custom_value=False
)
with gr.Row():
with gr.Column(scale=6):
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=3,
placeholder="Describe your scene...",
container=False,
lines=3,
autofocus=True,
)
with gr.Column(scale=1, min_width=120, elem_classes="center-button"):
run_button = gr.Button("Run", variant="primary", size="lg")
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", visible=False)
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
with gr.Row(equal_height=True, elem_classes="main-content-row"):
with gr.Column(scale=1, elem_classes="advanced-options-column"):
with gr.Group():
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Advanced Options</div>")
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=initial_values['guidance_scale'],
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
with gr.Column(scale=1, elem_classes="video-column"):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
container=True,
elem_classes="video-component"
)
gr.HTML("""
<style>
.center-button {
display: flex !important;
justify-content: center !important;
height: 100% !important;
padding-top: 1.4em !important;
}
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
}
.main {
max-width: 1200px !important;
margin: 0 auto !important;
}
.gr-form, .gr-box, .gr-group {
max-width: 1200px !important;
}
.gr-video {
max-width: 500px !important;
margin: 0 auto !important;
}
.main-content-row {
display: flex !important;
align-items: flex-start !important;
min-height: 500px !important;
gap: 20px !important;
}
.advanced-options-column,
.video-column {
display: flex !important;
flex-direction: column !important;
flex: 1 !important;
min-height: 400px !important;
align-items: stretch !important;
}
.video-column > * {
margin-top: 0 !important;
}
.video-column .gr-video,
.video-component {
margin-top: 0 !important;
padding-top: 0 !important;
}
.video-column .gr-video .gr-form {
margin-top: 0 !important;
}
.advanced-options-column .gr-group,
.video-column .gr-video {
margin-top: 0 !important;
vertical-align: top !important;
}
.advanced-options-column > *:last-child,
.video-column > *:last-child {
flex-grow: 0 !important;
}
@media (max-width: 1400px) {
.main-content-row {
min-height: 600px !important;
}
.advanced-options-column,
.video-column {
min-height: 600px !important;
}
}
@media (max-width: 1200px) {
.main-content-row {
flex-direction: column !important;
align-items: stretch !important;
}
.advanced-options-column,
.video-column {
min-height: auto !important;
width: 100% !important;
}
}
.timing-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 8px;
text-align: center;
min-height: 80px;
display: flex;
flex-direction: column;
justify-content: center;
}
.timing-card-highlight {
background: var(--background-fill-primary) !important;
border: 2px solid var(--color-accent) !important;
}
.performance-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
.gr-number input[readonly] {
background-color: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color-subdued) !important;
cursor: default !important;
text-align: center !important;
font-weight: 500 !important;
}
</style>
""")
def on_example_select(example_label):
if example_label and example_label in example_labels:
index = example_labels.index(example_label)
return examples[index]
return ""
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
)
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected. We are also rate-limiting users to 3 requests per minute.</p>
</div>
""")
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
def on_model_selection_change(selected_model):
if not selected_model:
selected_model = "FastWan2.1-T2V-1.3B"
model_path = MODEL_PATH_MAPPING.get(selected_model)
if model_path and model_path in default_params:
params = default_params[model_path]
return (
gr.update(value=params.height),
gr.update(value=params.width),
gr.update(value=params.num_frames),
gr.update(value=params.guidance_scale),
gr.update(value=params.seed),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=3.0),
gr.update(value=1024),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
result_path, seed_or_error, timing_details = generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
)
if result_path and os.path.exists(result_path):
return (
result_path,
seed_or_error,
gr.update(visible=False),
gr.update(visible=True, value=timing_details),
)
else:
return (
None,
seed_or_error,
gr.update(visible=True, value=seed_or_error),
gr.update(visible=False),
)
run_button.click(
fn=handle_generation,
inputs=[
model_selection,
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
randomize_seed,
],
outputs=[result, seed_output, error_output, timing_display],
concurrency_limit=20,
)
return demo
def main():
parser = argparse.ArgumentParser(description="FastVideo Gradio Frontend")
parser.add_argument("--backend_url", type=str, default="http://localhost:8000",
help="URL of the Ray Serve backend")
parser.add_argument("--t2v_model_paths", type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--host", type=str, default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port", type=int, default=7860,
help="Port to bind to")
args = parser.parse_args()
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(args.backend_url, default_params)
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
print(f"Backend URL: {args.backend_url}")
print(f"T2V Models: {args.t2v_model_paths}")
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
import uvicorn
app = FastAPI()
@app.get("/logo.svg")
def get_logo():
return FileResponse(
"assets/logos/logo.svg",
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
@app.get("/favicon.ico")
def get_favicon():
favicon_path = "assets/logos/icon_simple.svg"
if os.path.exists(favicon_path):
return FileResponse(
favicon_path,
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
else:
raise HTTPException(status_code=404, detail="Favicon not found")
@app.get("/", response_class=HTMLResponse)
def index(request: Request):
base_url = str(request.base_url).rstrip('/')
return f"""
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:description" content="Make video generation go blurrrrrrr">
<meta property="og:image" content="{base_url}/logo.svg">
<meta property="og:image:width" content="1200">
<meta property="og:image:height" content="630">
<meta property="og:site_name" content="FastWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
<meta property="twitter:image" content="{base_url}/logo.svg">
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
<link rel="icon" type="image/png" sizes="16x16" href="/favicon.ico">
<link rel="apple-touch-icon" href="/favicon.ico">
<style>
body, html {{
margin: 0;
padding: 0;
height: 100%;
overflow: hidden;
}}
iframe {{
width: 100%;
height: 100vh;
border: none;
}}
</style>
</head>
<body>
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
</body>
</html>
"""
app = gr.mount_gradio_app(
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
)
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()
@@ -0,0 +1,379 @@
import time
import os
import torch
import base64
import io
from copy import deepcopy
from typing import Dict, Any, Optional, List
import signal
import sys
import ray
from ray import serve
from fastapi import FastAPI, Request, Response
from pydantic import BaseModel
import numpy as np
from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.util import get_remote_address
from slowapi.errors import RateLimitExceeded
import imageio
from ray.serve.handle import DeploymentHandle
from prometheus_client import Counter, Histogram, generate_latest
NUM_GPUS = 16
DEFAULT_FPS = 16
SEED_RANGE_MAX = 1_000_000
SUPPORTED_MODELS = [
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
]
MODEL_CONFIGS = {
"1.3B": {
"num_cpus": 2,
"text_encoder_cpu_offload": False,
"dit_cpu_offload": False,
"vae_cpu_offload": False,
"VSA_sparsity": 0.8,
},
"14B": {
"num_cpus": 16,
"text_encoder_cpu_offload": True,
"dit_cpu_offload": True,
"vae_cpu_offload": False,
"VSA_sparsity": 0.9,
}
}
class VideoGenerationRequest(BaseModel):
prompt: str
negative_prompt: Optional[str] = None
use_negative_prompt: bool = False
seed: int = 42
guidance_scale: float = 7.5
num_frames: int = 21
height: int = 448
width: int = 832
randomize_seed: bool = False
return_frames: bool = False
model_path: Optional[str] = None
class VideoGenerationResponse(BaseModel):
video_data: Optional[str] = None
seed: int
success: bool
error_message: Optional[str] = None
generation_time: Optional[float] = None
model_load_time: Optional[float] = None
inference_time: Optional[float] = None
encoding_time: Optional[float] = None
total_time: Optional[float] = None
stage_names: Optional[List[str]] = None
stage_execution_times: Optional[List[float]] = None
def encode_video_to_base64(frames: List[np.ndarray], fps: int = DEFAULT_FPS) -> str:
if not frames:
return ""
try:
buffer = io.BytesIO()
imageio.mimsave(buffer, frames, fps=fps, format="mp4")
buffer.seek(0)
video_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
return f"data:video/mp4;base64,{video_base64}"
except Exception as e:
print(f"Warning: Failed to encode video: {e}")
return ""
def setup_model_environment(model_path: str) -> None:
if "fullattn" in model_path.lower():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
def process_generation_result(result: Any) -> tuple[List[np.ndarray], float, List[str], List[float]]:
frames = result if isinstance(result, list) else result.get("frames", [])
generation_time = result.get("generation_time", 0.0) if isinstance(result, dict) else 0.0
logging_info = result.get("logging_info", None)
if logging_info:
stage_names = logging_info.get_execution_order()
stage_execution_times = [
logging_info.get_stage_info(stage_name).get("execution_time", 0.0)
for stage_name in stage_names
]
else:
stage_names = []
stage_execution_times = []
return frames, generation_time, stage_names, stage_execution_times
def prepare_sampling_params(video_request: VideoGenerationRequest, default_params: Any) -> Any:
params = deepcopy(default_params)
params.prompt = video_request.prompt
if video_request.use_negative_prompt:
params.negative_prompt = video_request.negative_prompt
params.seed = (video_request.seed if not video_request.randomize_seed
else torch.randint(0, SEED_RANGE_MAX, (1,)).item())
params.randomize_seed = video_request.randomize_seed
params.guidance_scale = video_request.guidance_scale
params.num_frames = video_request.num_frames
params.height = video_request.height
params.width = video_request.width
params.save_video = False
params.return_frames = False
return params
class BaseModelDeployment:
def __init__(self, model_path: str, output_path: str = "outputs"):
self.model_path = model_path
self.output_path = output_path
self.generator = None
self.default_params = None
os.makedirs(self.output_path, exist_ok=True)
setup_model_environment(self.model_path)
def _initialize_generator(self, config: Dict[str, Any]) -> None:
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
print(f"Initializing model: {self.model_path}")
self.generator = VideoGenerator.from_pretrained(
model_path=self.model_path,
num_gpus=1,
use_fsdp_inference=True,
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
dit_cpu_offload=config["dit_cpu_offload"],
vae_cpu_offload=config["vae_cpu_offload"],
VSA_sparsity=config["VSA_sparsity"],
enable_stage_verification=False,
)
self.default_params = SamplingParam.from_pretrained(self.model_path)
def generate_video(self, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
total_start_time = time.time()
params = prepare_sampling_params(video_request, self.default_params)
inference_start_time = time.time()
result = self.generator.generate_video(
prompt=video_request.prompt,
sampling_param=params,
save_video=False,
return_frames=False,
)
inference_time = time.time() - inference_start_time
frames, generation_time, stage_names, stage_execution_times = process_generation_result(result)
encoding_start_time = time.time()
video_data = encode_video_to_base64(frames, fps=DEFAULT_FPS)
encoding_time = time.time() - encoding_start_time
total_time = time.time() - total_start_time
return VideoGenerationResponse(
video_data=video_data,
seed=params.seed,
success=True,
generation_time=generation_time,
inference_time=inference_time,
encoding_time=encoding_time,
total_time=total_time,
stage_names=stage_names,
stage_execution_times=stage_execution_times,
)
@serve.deployment(
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
)
class T2VModelDeployment(BaseModelDeployment):
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
super().__init__(t2v_model_path, output_path)
self._initialize_generator(MODEL_CONFIGS["1.3B"])
print("✅ T2V model initialized successfully")
@serve.deployment(
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
)
class T2V14BModelDeployment(BaseModelDeployment):
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
super().__init__(t2v_14b_model_path, output_path)
# Override environment for 14B model
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
self._initialize_generator(MODEL_CONFIGS["14B"])
print("✅ T2V 14B model initialized successfully")
app = FastAPI()
limiter = Limiter(key_func=get_remote_address)
app.state.limiter = limiter
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
@serve.deployment(num_replicas=50, ray_actor_options={"num_cpus": 2})
@serve.ingress(app)
class FastVideoAPI:
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
self.t2v_deployments = t2v_deployments
# Initialize Prometheus metrics
self.request_count = Counter('fastvideo_requests_total', 'Total FastVideo requests', ['model_type', 'status'])
self.request_duration = Histogram('fastvideo_request_duration_seconds', 'FastVideo request duration', ['model_type'])
self.video_generation_time = Histogram('fastvideo_video_generation_seconds', 'Video generation time', ['model_type'])
def _get_model_name(self, model_path: Optional[str]) -> str:
return model_path.split('/')[-1] if model_path else "unknown"
def _record_metrics(self, model_name: str, status: str, duration: float, response: Optional[VideoGenerationResponse] = None) -> None:
self.request_count.labels(model_type=model_name, status=status).inc()
self.request_duration.labels(model_type=model_name).observe(duration)
if response and hasattr(response, 'generation_time') and response.generation_time:
self.video_generation_time.labels(model_type=model_name).observe(response.generation_time)
@app.post("/generate_video", response_model=VideoGenerationResponse)
@limiter.limit("10/minute")
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
"""Route the request to the appropriate model deployment based on model_path."""
start_time = time.time()
model_name = self._get_model_name(video_request.model_path)
try:
if video_request.model_path not in self.t2v_deployments:
raise ValueError(f"Model {video_request.model_path} not found")
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
response = await response_ref
self._record_metrics(model_name, "success", time.time() - start_time, response)
return response
except Exception as e:
self._record_metrics(model_name, "error", time.time() - start_time)
return VideoGenerationResponse(
video_data=None,
seed=video_request.seed,
success=False,
error_message=str(e),
generation_time=0,
inference_time=0,
encoding_time=0,
total_time=0,
)
@app.get("/health")
@limiter.limit("10/minute")
async def health_check(self, request: Request) -> Dict[str, str]:
return {"status": "healthy"}
@app.get("/metrics")
async def metrics(self) -> Response:
return Response(generate_latest(), media_type="text/plain")
def validate_configuration(model_paths: List[str], replicas: List[int]) -> None:
assert len(model_paths) == len(replicas), "Number of models and replicas must match"
assert sum(replicas) <= NUM_GPUS, f"Total replicas ({sum(replicas)}) must be <= {NUM_GPUS}"
for model, replica_count in zip(model_paths, replicas):
assert model in SUPPORTED_MODELS, f"Model {model} not supported"
assert replica_count > 0, f"Replicas must be greater than 0"
def start_ray_serve(
*,
t2v_model_paths: str,
t2v_model_replicas: str,
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
) -> None:
if not ray.is_initialized():
ray.init()
model_paths = t2v_model_paths.split(",")
replicas = [int(r) for r in t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
t2v_deps = {}
for model_path, replica_count in zip(model_paths, replicas):
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
api = FastVideoAPI.bind(t2v_deps)
serve.run(api, route_prefix="/", name="fast_video")
print(f"Ray Serve backend started at http://{host}:{port}")
for model_path, replica_count in zip(model_paths, replicas):
print(f"T2V Model: {model_path} | Replicas: {replica_count}")
print(f"Health check: http://{host}:{port}/health")
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
def setup_signal_handlers() -> None:
signal.signal(signal.SIGINT, lambda *_: sys.exit(0))
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0))
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
parser.add_argument("--t2v_model_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--output_path",
type=str,
default="outputs",
help="Path to save generated videos")
parser.add_argument("--host",
type=str,
default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port",
type=int,
default=8000,
help="Port to bind to")
args = parser.parse_args()
model_paths = args.t2v_model_paths.split(",")
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
start_ray_serve(
t2v_model_paths=args.t2v_model_paths,
t2v_model_replicas=args.t2v_model_replicas,
output_path=args.output_path,
host=args.host,
port=args.port,
)
setup_signal_handlers()
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
while True:
time.sleep(3600)
+3
View File
@@ -0,0 +1,3 @@
python examples/inference/gradio/start_ray_serve_app.py \
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" \
--t2v_model_replicas "4,4"
@@ -0,0 +1,257 @@
"""
Startup script for FastVideo with Ray Serve backend and Gradio frontend.
This script starts both the backend and frontend services.
"""
import argparse
import os
import subprocess
import sys
import time
import threading
import signal
import requests
from pathlib import Path
from typing import Dict, Any, Optional
DEFAULT_BACKEND_HOST = "0.0.0.0"
DEFAULT_BACKEND_PORT = 8000
DEFAULT_FRONTEND_HOST = "0.0.0.0"
DEFAULT_FRONTEND_PORT = 7860
DEFAULT_OUTPUT_PATH = "outputs"
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_T2V_REPLICAS = "4,4"
HEALTH_CHECK_TIMEOUT = 5
HEALTH_CHECK_MAX_RETRIES = 100
HEALTH_CHECK_INTERVAL = 2
PROCESS_SHUTDOWN_TIMEOUT = 5
PROCESS_MONITOR_INTERVAL = 1
PROJECT_ROOT = Path(__file__).parent.parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))
class ServiceManager:
def __init__(self, args: argparse.Namespace):
self.args = args
self.backend_process: Optional[subprocess.Popen] = None
self.frontend_process: Optional[subprocess.Popen] = None
self.backend_url = f"http://{args.backend_host}:{args.backend_port}"
def check_backend_health(self, max_retries: int = HEALTH_CHECK_MAX_RETRIES) -> bool:
health_url = f"{self.backend_url}/health"
for attempt in range(max_retries):
try:
response = requests.get(health_url, timeout=HEALTH_CHECK_TIMEOUT)
if response.status_code == 200:
print(f"✅ Backend is healthy at {self.backend_url}")
return True
except requests.exceptions.RequestException:
pass
if attempt < max_retries - 1:
print(f"⏳ Waiting for backend to start... ({attempt + 1}/{max_retries})")
time.sleep(HEALTH_CHECK_INTERVAL)
print(f"❌ Backend failed to start within {max_retries * HEALTH_CHECK_INTERVAL} seconds")
return False
def _create_monitor_thread(self, process: subprocess.Popen, service_name: str) -> threading.Thread:
def monitor():
if process.stdout:
for line in process.stdout:
print(f"[{service_name}] {line.rstrip()}")
thread = threading.Thread(target=monitor, daemon=True)
thread.start()
return thread
def _start_service(self, script_name: str, args_dict: Dict[str, Any], service_name: str) -> subprocess.Popen:
script_path = Path(__file__).parent / script_name
cmd = [sys.executable, str(script_path)]
for key, value in args_dict.items():
cmd.extend([f"--{key}", str(value)])
print(f"🚀 Starting {service_name}...")
print(f"Command: {' '.join(cmd)}")
process = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
universal_newlines=True,
bufsize=1
)
self._create_monitor_thread(process, service_name.upper())
return process
def start_backend(self) -> subprocess.Popen:
backend_args = {
"t2v_model_paths": self.args.t2v_model_paths,
"t2v_model_replicas": self.args.t2v_model_replicas,
"output_path": self.args.output_path,
"host": self.args.backend_host,
"port": self.args.backend_port
}
self.backend_process = self._start_service("ray_serve_backend.py", backend_args, "backend")
return self.backend_process
def start_frontend(self) -> subprocess.Popen:
frontend_args = {
"backend_url": self.backend_url,
"t2v_model_paths": self.args.t2v_model_paths,
"host": self.args.frontend_host,
"port": self.args.frontend_port
}
self.frontend_process = self._start_service("gradio_frontend.py", frontend_args, "frontend")
return self.frontend_process
def shutdown_services(self) -> None:
print("\n🛑 Shutting down services...")
processes = []
if self.frontend_process:
self.frontend_process.terminate()
processes.append(("frontend", self.frontend_process))
if self.backend_process:
self.backend_process.terminate()
processes.append(("backend", self.backend_process))
for name, process in processes:
try:
process.wait(timeout=PROCESS_SHUTDOWN_TIMEOUT)
print(f"✅ {name.capitalize()} stopped gracefully")
except subprocess.TimeoutExpired:
print(f"⚠️ Force killing {name} process...")
process.kill()
print("✅ All services stopped")
def monitor_processes(self) -> None:
if not self.backend_process or not self.frontend_process:
print("❌ Processes not properly initialized")
return
try:
while True:
if self.frontend_process.poll() is not None:
print("❌ Frontend process died unexpectedly")
break
if self.backend_process.poll() is not None:
print("❌ Backend process died unexpectedly")
break
time.sleep(PROCESS_MONITOR_INTERVAL)
except KeyboardInterrupt:
pass
self.shutdown_services()
def setup_signal_handlers(service_manager: ServiceManager) -> None:
def signal_handler(signum: int, frame: Any) -> None:
service_manager.shutdown_services()
sys.exit(0)
signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)
def print_startup_info(args: argparse.Namespace) -> None:
print("🎬 FastVideo Ray Serve App")
print("=" * 50)
print(f"T2V Models: {args.t2v_model_paths}")
print(f"T2V Model Replicas: {args.t2v_model_replicas}")
print(f"Output: {args.output_path}")
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
print("=" * 50)
def parse_arguments() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="FastVideo Ray Serve App")
parser.add_argument("--t2v_model_paths",
type=str,
default=DEFAULT_T2V_MODELS,
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default=DEFAULT_T2V_REPLICAS,
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--output_path",
type=str,
default=DEFAULT_OUTPUT_PATH,
help="Path to save generated videos")
parser.add_argument("--backend_host",
type=str,
default=DEFAULT_BACKEND_HOST,
help="Backend host to bind to")
parser.add_argument("--backend_port",
type=int,
default=DEFAULT_BACKEND_PORT,
help="Backend port to bind to")
parser.add_argument("--frontend_host",
type=str,
default=DEFAULT_FRONTEND_HOST,
help="Frontend host to bind to")
parser.add_argument("--frontend_port",
type=int,
default=DEFAULT_FRONTEND_PORT,
help="Frontend port to bind to")
parser.add_argument("--skip_backend_check",
action="store_true",
help="Skip backend health check")
return parser.parse_args()
def main() -> None:
args = parse_arguments()
os.makedirs(args.output_path, exist_ok=True)
print_startup_info(args)
service_manager = ServiceManager(args)
setup_signal_handlers(service_manager)
try:
service_manager.start_backend()
if not args.skip_backend_check:
if not service_manager.check_backend_health():
print("❌ Backend failed to start. Terminating...")
service_manager.shutdown_services()
sys.exit(1)
service_manager.start_frontend()
print("\n🎉 Both services are starting up!")
print(f"📺 Frontend will be available at: http://{args.frontend_host}:{args.frontend_port}")
print(f"🔧 Backend API will be available at: http://{args.backend_host}:{args.backend_port}")
print("\nPress Ctrl+C to stop both services...")
service_manager.monitor_processes()
except Exception as e:
print(f"❌ Unexpected error: {e}")
service_manager.shutdown_services()
sys.exit(1)
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=False,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
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": 5.0,
"guidance_scale": 6.0,
"num_inference_steps": 32,
"seed": 42,
}
@@ -11,22 +11,23 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
dit_cpu_offload=False,
vae_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1250/transformer",
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-160/transformer",
lora_nickname="crush_smol"
)
generator.unmerge_lora_weights()
kwargs = {
"height": 480,
"width": 832,
"num_frames": 77,
"guidance_scale": 5.0,
"guidance_scale": 6.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 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."
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."
video = generator.generate_video(
prompt,
@@ -34,6 +35,12 @@ 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=False,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
)
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=False,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
@@ -54,7 +54,7 @@ validation_args=(
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
@@ -69,7 +69,6 @@ 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 "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
@@ -103,7 +103,6 @@ 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,7 +101,6 @@ optimizer_args=(
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--dit_precision "fp32"
@@ -101,7 +101,6 @@ 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 "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
@@ -69,7 +69,6 @@ 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 "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
@@ -103,7 +103,6 @@ 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 "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
@@ -69,7 +69,6 @@ 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"
@@ -0,0 +1,24 @@
#!/bin/bash
GPU_NUM=2 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
torchrun --nproc_per_node=$GPU_NUM \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type i2v \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--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
@@ -52,9 +52,9 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 50
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
@@ -69,7 +69,6 @@ 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"
@@ -78,6 +77,7 @@ 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 "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 500
--checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -100,7 +100,6 @@ 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"
@@ -7,7 +7,7 @@ export WANDB_MODE=online
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=2
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
@@ -52,16 +52,16 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 50
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 500
--checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -69,7 +69,6 @@ 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,6 +76,7 @@ miscellaneous_args=(
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--resume_from_checkpoint "checkpoints/wan_t2v_finetune_lora/checkpoint-160"
)
torchrun \
@@ -0,0 +1,24 @@
#!/bin/bash
GPU_NUM=2 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--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
+45
View File
@@ -1,9 +1,36 @@
import dataclasses
from enum import Enum
from typing import Any, Optional
from fastvideo.configs.utils import update_config_from_args
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
class DatasetType(str, Enum):
"""
Enumeration for different dataset types.
"""
HF = "hf"
MERGED = "merged"
@classmethod
def from_string(cls, value: str) -> "DatasetType":
"""Convert string to DatasetType enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid dataset type: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [dataset_type.value for dataset_type in cls]
@dataclasses.dataclass
class PreprocessConfig:
@@ -12,6 +39,7 @@ class PreprocessConfig:
# Model and dataset configuration
model_path: str = ""
dataset_path: str = ""
dataset_type: DatasetType = DatasetType.HF
dataset_output_dir: str = "./output"
# Dataloader configuration
@@ -35,6 +63,9 @@ class PreprocessConfig:
# Model configuration
training_cfg_rate: float = 0.0
# framework configuration
seed: int = 42
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "preprocess") -> FlexibleArgumentParser:
@@ -52,6 +83,12 @@ class PreprocessConfig:
type=str,
default=PreprocessConfig.dataset_path,
help="Path to the dataset directory for preprocessing")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-type",
type=str,
choices=DatasetType.choices(),
default=PreprocessConfig.dataset_type.value,
help="Type of the dataset")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-output-dir",
type=str,
@@ -123,6 +160,10 @@ class PreprocessConfig:
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
type=int,
default=PreprocessConfig.seed,
help="Seed for random number generator")
return parser
@@ -130,6 +171,10 @@ class PreprocessConfig:
def from_kwargs(cls, kwargs: dict[str,
Any]) -> Optional["PreprocessConfig"]:
"""Create PreprocessConfig from keyword arguments."""
if 'dataset_type' in kwargs and isinstance(kwargs['dataset_type'], str):
kwargs['dataset_type'] = DatasetType.from_string(
kwargs['dataset_type'])
preprocess_config = cls()
if not update_config_from_args(
preprocess_config, kwargs, prefix="preprocess", pop_args=True):
+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.in_channels if self.added_kv_proj_dim is None else self.out_channels
self.num_channels_latents = self.out_channels
@dataclass
+3 -1
View File
@@ -31,7 +31,9 @@ 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-TI2V-5B-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
# 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 | 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"
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
+15 -10
View File
@@ -6,12 +6,19 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
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)
# 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,
)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -29,10 +36,8 @@ 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,
"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
}
+8 -2
View File
@@ -129,9 +129,15 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
pass
guidance_scale: float = 4.0
guidance_scale_2: float = 3.0
num_inference_steps: int = 40
fps: int = 16
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
pass
guidance_scale: float = 3.5
guidance_scale_2: float = 3.5
num_inference_steps: int = 40
fps: int = 16
+19 -2
View File
@@ -9,6 +9,7 @@ diffusion models.
import math
import os
import time
from copy import deepcopy
from typing import Any
import imageio
@@ -202,6 +203,8 @@ class VideoGenerator:
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
else:
sampling_param = deepcopy(sampling_param)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
@@ -275,6 +278,7 @@ class VideoGenerator:
width: {target_width}
video_length: {sampling_param.num_frames}
prompt: {prompt}
image_path: {sampling_param.image_path}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
infer_steps: {sampling_param.num_inference_steps}
@@ -304,7 +308,8 @@ class VideoGenerator:
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
@@ -334,9 +339,11 @@ class VideoGenerator:
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time
"generation_time": gen_time,
"logging_info": logging_info,
}
def set_lora_adapter(self,
@@ -344,6 +351,16 @@ class VideoGenerator:
lora_path: str | None = None) -> None:
self.executor.set_lora_adapter(lora_nickname, lora_path)
def unmerge_lora_weights(self) -> None:
"""
Use unmerged weights for inference to produce videos that align with
validation videos generated during training.
"""
self.executor.unmerge_lora_weights()
def merge_lora_weights(self) -> None:
self.executor.merge_lora_weights()
def shutdown(self):
"""
Shutdown the video generator.
+3
View File
@@ -158,6 +158,9 @@ class FastVideoArgs:
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
# MoE parameters used by Wan2.2
boundary_ratio: float | None = None
@property
def training_mode(self) -> bool:
return not self.inference_mode
+4 -4
View File
@@ -76,7 +76,7 @@ class BaseLayerWithLoRA(nn.Module):
lora_B = self.lora_B.to_local()
lora_A = self.lora_A.to_local()
if (self.training_mode or not self.merged) and not self.disable_lora:
if not self.merged and not self.disable_lora:
delta = x @ (
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
@@ -101,15 +101,15 @@ class BaseLayerWithLoRA(nn.Module):
B: torch.Tensor,
training_mode: bool = False,
lora_path: str | None = None) -> None:
self.lora_A = A # share storage with weights in the pipeline
self.lora_B = B
self.lora_A = torch.nn.Parameter(
A) # share storage with weights in the pipeline
self.lora_B = torch.nn.Parameter(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 -2
View File
@@ -120,14 +120,14 @@ def _info(logger: Logger,
if not _warned_local_main_process and local_main_process_only:
logger.warning(
'%s is_local_main_process is set to True, logging only from the local main process.%s',
'%s By default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s',
GREEN,
RESET,
)
_warned_local_main_process = True
if not _warned_main_process and main_process_only:
logger.warning(
'%s is_main_process_only is set to True, logging only from the main process.%s',
'%s is_main_process_only is set to True, logging only from the main (RANK==0) process.%s',
GREEN,
RESET,
)
-2
View File
@@ -672,11 +672,9 @@ 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,8 +72,7 @@ class ComponentLoader(ABC):
module_loaders = {
"scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"),
"real_score_transformer": (TransformerLoader, "diffusers"),
"fake_score_transformer": (TransformerLoader, "diffusers"),
"transformer_2": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
@@ -452,6 +451,7 @@ 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,
+2 -1
View File
@@ -8,10 +8,11 @@ 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']
required_attributes = ['timesteps', 'order', 'num_train_timesteps']
for attr in required_attributes:
if not hasattr(self, attr):
@@ -134,6 +134,7 @@ 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
@@ -145,6 +146,8 @@ 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,6 +117,7 @@ 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
@@ -132,6 +133,8 @@ 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,6 +273,7 @@ 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,12 +12,10 @@ 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,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
@@ -67,7 +65,7 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
@@ -12,12 +12,10 @@ 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,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
@@ -48,11 +46,14 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
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="conditioning_stage",
stage=ConditioningStage())
@@ -67,11 +68,12 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
stage=ImageVAEEncodingStage(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,6 +61,7 @@ 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))
+39 -12
View File
@@ -37,6 +37,7 @@ 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] = {}
@@ -148,9 +149,7 @@ class ComposedPipelineBase(ABC):
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
# fastvideo_args.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)
@@ -239,6 +238,18 @@ 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)
@@ -248,12 +259,20 @@ 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:
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]
logger.warning(
"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']
"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()}"
)
# all the component models used by the pipeline
required_modules = self.required_config_modules
@@ -263,6 +282,11 @@ class ComposedPipelineBase(ABC):
for module_name, (transformers_or_diffusers,
architecture) in model_index.items():
if transformers_or_diffusers is None:
logger.warning(
"Module in model_index.json has null value, removing from required_config_modules"
)
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
@@ -271,14 +295,17 @@ class ComposedPipelineBase(ABC):
logger.info("Using module %s already provided", module_name)
modules[module_name] = loaded_modules[module_name]
continue
if 'transformer' in module_name:
loading_module_name = module_name.split("_")[-1]
# 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]
else:
loading_module_name = module_name
load_module_name = module_name
component_model_path = os.path.join(self.model_path,
loading_module_name)
load_module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
module_name=load_module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=fastvideo_args,
+8
View File
@@ -217,3 +217,11 @@ class LoRAPipeline(ComposedPipelineBase):
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
def merge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.merge_lora_weights()
def unmerge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.unmerge_lora_weights()
+62 -9
View File
@@ -9,15 +9,55 @@ in a functional manner, reducing the need for explicit parameter passing.
import pprint
from dataclasses import asdict, dataclass, field
from typing import Any
from typing import TYPE_CHECKING, Any
import PIL.Image
import torch
if TYPE_CHECKING:
from torchcodec.decoders import VideoDecoder
import time
from collections import OrderedDict
from fastvideo.attention import AttentionMetadata
from fastvideo.configs.sample.teacache import TeaCacheParams, WanTeaCacheParams
class PipelineLoggingInfo:
"""Simple approach using OrderedDict to track stage metrics."""
def __init__(self):
# OrderedDict preserves insertion order and allows easy access
self.stages: OrderedDict[str, dict[str, Any]] = OrderedDict()
def add_stage_execution_time(self, stage_name: str, execution_time: float):
"""Add execution time for a stage."""
if stage_name not in self.stages:
self.stages[stage_name] = {}
self.stages[stage_name]['execution_time'] = execution_time
self.stages[stage_name]['timestamp'] = time.time()
def add_stage_metric(self, stage_name: str, metric_name: str, value: Any):
"""Add any metric for a stage."""
if stage_name not in self.stages:
self.stages[stage_name] = {}
self.stages[stage_name][metric_name] = value
def get_stage_info(self, stage_name: str) -> dict[str, Any]:
"""Get all info for a specific stage."""
return self.stages.get(stage_name, {})
def get_execution_order(self) -> list[str]:
"""Get stages in execution order."""
return list(self.stages.keys())
def get_total_execution_time(self) -> float:
"""Get total pipeline execution time."""
return sum(
stage.get('execution_time', 0) for stage in self.stages.values())
@dataclass
class ForwardBatch:
"""
@@ -37,7 +77,7 @@ class ForwardBatch:
# Image inputs
image_path: str | None = None
image_embeds: list[torch.Tensor] = field(default_factory=list)
pil_image: PIL.Image.Image | None = None
pil_image: torch.Tensor | PIL.Image.Image | None = None
preprocessed_image: torch.Tensor | None = None
# Text inputs
@@ -75,15 +115,15 @@ class ForwardBatch:
image_latent: torch.Tensor | None = None
# Latent dimensions
height_latents: int | None = None
width_latents: int | None = None
num_frames: int = 1 # Default for image models
height_latents: list[int] | int | None = None
width_latents: list[int] | int | None = None
num_frames: list[int] | 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: int | None = None
width: int | None = None
fps: int | None = None
height: list[int] | int | None = None
width: list[int] | int | None = None
fps: list[int] | int | None = None
# Timesteps
timesteps: torch.Tensor | None = None
@@ -93,6 +133,7 @@ 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
@@ -128,6 +169,10 @@ class ForwardBatch:
# VSA parameters
VSA_sparsity: float = 0.0
# Logging info
logging_info: PipelineLoggingInfo = field(
default_factory=PipelineLoggingInfo)
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
@@ -136,6 +181,8 @@ 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)
@@ -189,4 +236,10 @@ class TrainingBatch:
fake_score_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)
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)
+10 -9
View File
@@ -25,6 +25,11 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"HunyuanVideoPipeline": "hunyuan",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
WorkloadType.I2V: "PreprocessPipelineI2V",
WorkloadType.T2V: "PreprocessPipelineT2V",
}
class PipelineType(str, Enum):
"""
@@ -65,15 +70,11 @@ class _PipelineRegistry:
arch = _PIPELINE_NAME_TO_ARCHITECTURE_NAME[pipeline_name_in_config]
return set(self.pipelines[pipeline_type.value][arch].keys())
def _load_preprocessing_pipeline_cls(
def _load_preprocess_pipeline_cls(
self, workload_type: WorkloadType,
arch: str) -> type[ComposedPipelineBase] | None:
if workload_type == WorkloadType.I2V:
pipeline_name = "I2VPreprocessPipeline"
elif workload_type == WorkloadType.T2V:
pipeline_name = "T2VPreprocessPipeline"
else:
raise ValueError(f"Invalid workload type: {workload_type.value}")
pipeline_name = _PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME[
workload_type]
return self.pipelines[
PipelineType.PREPROCESS.value][arch][pipeline_name]
@@ -90,7 +91,7 @@ class _PipelineRegistry:
return None
if pipeline_type == PipelineType.PREPROCESS:
return self._load_preprocessing_pipeline_cls(workload_type, arch)
return self._load_preprocess_pipeline_cls(workload_type, arch)
elif pipeline_type == PipelineType.BASIC:
return self.pipelines[
pipeline_type.value][arch][pipeline_name_in_config]
@@ -131,7 +132,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, preprocessing, training).
pipeline_types: The pipeline types to load (basic, preprocess, training).
If None, loads all types.
Returns:
@@ -0,0 +1,107 @@
import random
from collections.abc import Callable
from typing import cast
import numpy as np
import torch
from einops import rearrange
from torchvision import transforms
from fastvideo.dataset.transform import (CenterCropResizeVideo,
TemporalRandomCrop)
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.pipelines.pipeline_batch_info import (ForwardBatch,
PreprocessBatch)
from fastvideo.pipelines.stages.base import PipelineStage
class VideoTransformStage(PipelineStage):
"""
Crop a video in temporal dimension.
"""
def __init__(self, train_fps: int, num_frames: int, max_height: int,
max_width: int, do_temporal_sample: bool) -> None:
self.train_fps = train_fps
self.num_frames = num_frames
if do_temporal_sample:
self.temporal_sample_fn: Callable | None = TemporalRandomCrop(
num_frames)
else:
self.temporal_sample_fn = None
self.video_transform = transforms.Compose([
CenterCropResizeVideo((max_height, max_width)),
])
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.fps, list)
assert isinstance(batch.num_frames, list)
if batch.data_type != "video":
return batch
if len(batch.video_loader) == 0:
raise ValueError("Video loader is not set")
video_pixel_batch = []
for i in range(len(batch.video_loader)):
frame_interval = batch.fps[i] / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch.num_frames[i],
frame_interval).astype(int)
if len(frame_indices) > self.num_frames:
if self.temporal_sample_fn is not None:
begin_index, end_index = self.temporal_sample_fn(
len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
else:
frame_indices = frame_indices[:self.num_frames]
video = batch.video_loader[i].get_frames_at(frame_indices).data
video = self.video_transform(video)
video_pixel_batch.append(video)
video_pixel_values = torch.stack(video_pixel_batch)
video_pixel_values = rearrange(video_pixel_values,
"b t c h w -> b c t h w")
video_pixel_values = video_pixel_values.to(torch.uint8)
if fastvideo_args.workload_type == WorkloadType.I2V:
batch.pil_image = video_pixel_values[:, :, 0, :, :]
video_pixel_values = video_pixel_values.float() / 255.0
batch.latents = video_pixel_values
batch.num_frames = [video_pixel_values.shape[2]] * len(
batch.video_loader)
batch.height = [video_pixel_values.shape[3]] * len(batch.video_loader)
batch.width = [video_pixel_values.shape[4]] * len(batch.video_loader)
return cast(ForwardBatch, batch)
class TextTransformStage(PipelineStage):
"""
Process text data according to the cfg rate.
"""
def __init__(self, cfg_uncondition_drop_rate: float, seed: int) -> None:
self.cfg_rate = cfg_uncondition_drop_rate
self.rng = random.Random(seed)
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
prompts = []
for prompt in batch.prompt:
if not isinstance(prompt, list):
prompt = [prompt]
prompt = self.rng.choice(prompt)
prompt = prompt if self.rng.random() > self.cfg_rate else ""
prompts.append(prompt)
batch.prompt = prompts
return cast(ForwardBatch, batch)
@@ -0,0 +1,23 @@
from fastvideo.distributed import (
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser
from fastvideo.workflow.workflow_base import WorkflowBase
logger = init_logger(__name__)
def main(fastvideo_args: FastVideoArgs) -> None:
maybe_init_distributed_environment_and_model_parallel(1, 1)
preprocess_workflow_cls = WorkflowBase.get_workflow_cls(fastvideo_args)
preprocess_workflow = preprocess_workflow_cls(fastvideo_args)
preprocess_workflow.run()
if __name__ == "__main__":
parser = FlexibleArgumentParser()
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
fastvideo_args = FastVideoArgs.from_cli_args(args)
main(fastvideo_args)
@@ -0,0 +1,83 @@
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.preprocess.preprocess_stages import (
TextTransformStage, VideoTransformStage)
from fastvideo.pipelines.stages import (EncodingStage, ImageEncodingStage,
TextEncodingStage)
from fastvideo.pipelines.stages.image_encoding import ImageVAEEncodingStage
class PreprocessPipelineI2V(ComposedPipelineBase):
_required_config_modules = [
"image_encoder", "image_processor", "text_encoder", "tokenizer", "vae"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
assert fastvideo_args.preprocess_config is not None
self.add_stage(stage_name="text_transform_stage",
stage=TextTransformStage(
cfg_uncondition_drop_rate=fastvideo_args.
preprocess_config.training_cfg_rate,
seed=fastvideo_args.preprocess_config.seed,
))
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(
stage_name="video_transform_stage",
stage=VideoTransformStage(
train_fps=fastvideo_args.preprocess_config.train_fps,
num_frames=fastvideo_args.preprocess_config.num_frames,
max_height=fastvideo_args.preprocess_config.max_height,
max_width=fastvideo_args.preprocess_config.max_width,
do_temporal_sample=fastvideo_args.preprocess_config.
do_temporal_sample,
))
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_vae_encoding_stage",
stage=ImageVAEEncodingStage(
vae=self.get_module("vae"), ))
self.add_stage(stage_name="video_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae"), ))
class PreprocessPipelineT2V(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
assert fastvideo_args.preprocess_config is not None
self.add_stage(stage_name="text_transform_stage",
stage=TextTransformStage(
cfg_uncondition_drop_rate=fastvideo_args.
preprocess_config.training_cfg_rate,
seed=fastvideo_args.preprocess_config.seed,
))
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(
stage_name="video_transform_stage",
stage=VideoTransformStage(
train_fps=fastvideo_args.preprocess_config.train_fps,
num_frames=fastvideo_args.preprocess_config.num_frames,
max_height=fastvideo_args.preprocess_config.max_height,
max_width=fastvideo_args.preprocess_config.max_width,
do_temporal_sample=fastvideo_args.preprocess_config.
do_temporal_sample,
))
self.add_stage(stage_name="video_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae"), ))
EntryClass = [PreprocessPipelineI2V, PreprocessPipelineT2V]
+3 -1
View File
@@ -12,7 +12,8 @@ from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import ImageEncodingStage
from fastvideo.pipelines.stages.image_encoding import (ImageEncodingStage,
ImageVAEEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
from fastvideo.pipelines.stages.stepvideo_encoding import (
@@ -32,6 +33,7 @@ __all__ = [
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"ImageVAEEncodingStage",
"TextEncodingStage",
"StepvideoPromptEncodingStage",
]
+2
View File
@@ -155,6 +155,8 @@ class PipelineStage(ABC):
execution_time = time.perf_counter() - start_time
logger.info("[%s] Execution completed in %s ms", stage_name,
execution_time * 1000)
batch.logging_info.add_stage_execution_time(
stage_name, execution_time)
except Exception as e:
execution_time = time.perf_counter() - start_time
logger.error("[%s] Error during execution after %s ms: %s",
-3
View File
@@ -3,7 +3,6 @@
Decoding stage for diffusion pipelines.
"""
import gc
import weakref
import torch
@@ -140,8 +139,6 @@ class DecodingStage(PipelineStage):
del self.vae
if pipeline is not None and "vae" in pipeline.modules:
del pipeline.modules["vae"]
gc.collect()
torch.mps.empty_cache()
fastvideo_args.model_loaded["vae"] = False
return batch
+37 -17
View File
@@ -3,7 +3,6 @@
Denoising stage for diffusion pipelines.
"""
import gc
import inspect
import weakref
from collections.abc import Iterable
@@ -57,9 +56,14 @@ class DenoisingStage(PipelineStage):
the initial noise into the final output.
"""
def __init__(self, transformer, scheduler, pipeline=None) -> None:
def __init__(self,
transformer,
scheduler,
pipeline=None,
transformer_2=None) -> None:
super().__init__()
self.transformer = transformer
self.transformer_2 = transformer_2
self.scheduler = scheduler
self.pipeline = weakref.ref(pipeline) if pipeline else None
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
@@ -184,6 +188,12 @@ class DenoisingStage(PipelineStage):
assert neg_prompt_embeds is not None
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
if fastvideo_args.boundary_ratio is not None:
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -191,6 +201,23 @@ class DenoisingStage(PipelineStage):
if hasattr(self, 'interrupt') and self.interrupt:
continue
if boundary_timestep is None or t >= boundary_timestep:
if (fastvideo_args.dit_cpu_offload
and self.transformer_2 is not None and next(
self.transformer_2.parameters()).device.type
== 'cuda'):
self.transformer_2.to('cpu')
current_model = self.transformer
current_guidance_scale = batch.guidance_scale
else:
# low-noise stage in wan2.2
if fastvideo_args.dit_cpu_offload and next(
self.transformer.parameters(
)).device.type == 'cuda':
self.transformer.to('cpu')
current_model = self.transformer_2
current_guidance_scale = batch.guidance_scale_2
# Expand latents for I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
@@ -257,7 +284,7 @@ class DenoisingStage(PipelineStage):
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred = self.transformer(
noise_pred = current_model(
latent_model_input,
prompt_embeds,
t_expand,
@@ -276,7 +303,7 @@ class DenoisingStage(PipelineStage):
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred_uncond = self.transformer(
noise_pred_uncond = current_model(
latent_model_input,
neg_prompt_embeds,
t_expand,
@@ -285,7 +312,7 @@ class DenoisingStage(PipelineStage):
**neg_cond_kwargs,
)
noise_pred_text = noise_pred
noise_pred = noise_pred_uncond + batch.guidance_scale * (
noise_pred = noise_pred_uncond + current_guidance_scale * (
noise_pred_text - noise_pred_uncond)
# Apply guidance rescale if needed
@@ -296,14 +323,12 @@ class DenoisingStage(PipelineStage):
noise_pred_text,
guidance_rescale=batch.guidance_rescale,
)
# Compute the previous noisy sample
latents = self.scheduler.step(noise_pred,
t,
latents,
**extra_step_kwargs,
return_dict=False)[0]
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
@@ -329,8 +354,6 @@ class DenoisingStage(PipelineStage):
del self.transformer
if pipeline is not None and "transformer" in pipeline.modules:
del pipeline.modules["transformer"]
gc.collect()
torch.mps.empty_cache()
fastvideo_args.model_loaded["transformer"] = False
logger.info("Memory after deallocating transformer: %s",
torch.mps.current_allocated_memory())
@@ -650,12 +673,8 @@ class DmdDenoisingStage(DenoisingStage):
# Get latents and embeddings
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
# TODO(yongqi) hard code prepare latents
latents = torch.randn(
latents.permute(0, 2, 1, 3, 4).shape,
dtype=torch.bfloat16,
device="cuda",
generator=torch.Generator(device="cuda").manual_seed(42))
latents = latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
@@ -772,8 +791,9 @@ class DmdDenoisingStage(DenoisingStage):
next_timestep = timesteps[i + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video.device)
noise = torch.randn(video_raw_latent_shape,
device=self.device,
dtype=pred_video.dtype)
dtype=pred_video.dtype,
generator=batch.generator[0]).to(
self.device)
if sp_group:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",

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