Compare commits
36
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d42fdc56f2 | ||
|
|
3ef04f1654 | ||
|
|
0eced76a41 | ||
|
|
3ab6470d1a | ||
|
|
989a03532c | ||
|
|
fa15369a02 | ||
|
|
78a9cb88d8 | ||
|
|
a0bff12746 | ||
|
|
98f2af94e5 | ||
|
|
46f7b6d574 | ||
|
|
911a6a6a35 | ||
|
|
38c7949d5c | ||
|
|
7e7a0dba9d | ||
|
|
f62e210ae6 | ||
|
|
6ceb4942a0 | ||
|
|
8cae5e4708 | ||
|
|
2a773fa34e | ||
|
|
5357f63327 | ||
|
|
60f61c8101 | ||
|
|
3d75ba8251 | ||
|
|
6c6bcd914d | ||
|
|
f79b08de81 | ||
|
|
f2bc037fff | ||
|
|
86604a684b | ||
|
|
47bd1e0178 | ||
|
|
c41305ad18 | ||
|
|
98ce9034f0 | ||
|
|
0ceff110da | ||
|
|
1d018acb3e | ||
|
|
7d8cf38dbe | ||
|
|
8d483fe4aa | ||
|
|
c1191250bf | ||
|
|
4b7266349a | ||
|
|
22f9b7681f | ||
|
|
589d32cc39 | ||
|
|
89199837db |
@@ -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"
|
||||
|
||||
@@ -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
|
||||
@@ -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/
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 46 KiB |
@@ -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 |
@@ -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 |
@@ -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
@@ -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
|
||||
|
||||
@@ -86,8 +86,6 @@ def benchmark_attention(configurations):
|
||||
# print(f"Average TFLOPS: {tflops_bwd}")
|
||||
# print("=" * 60)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return results
|
||||
|
||||
|
||||
|
||||
+10
-13
@@ -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
@@ -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.")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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 . .
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user