Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
632d56f61b | ||
|
|
ef281dc698 | ||
|
|
32d92ce38d | ||
|
|
d8b04897c7 | ||
|
|
700de312e4 | ||
|
|
85f0a112c1 | ||
|
|
77e7d995f3 | ||
|
|
8a93b834c1 | ||
|
|
8b2f2b556c | ||
|
|
242e38e98d | ||
|
|
21fa907a9d |
+18
-19
@@ -117,11 +117,11 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -133,10 +133,10 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -147,10 +147,10 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
@@ -161,12 +161,11 @@ steps:
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/tests/test_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
|
||||
@@ -18,12 +18,6 @@ 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
|
||||
@@ -55,13 +49,4 @@ 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
|
||||
@@ -104,17 +104,16 @@ jobs:
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
- 'csrc/attn/sliding_tile_attn/tk/**'
|
||||
- 'csrc/attn/sliding_tile_attn/setup.py'
|
||||
- 'csrc/attn/sliding_tile_attn/config_sta.py'
|
||||
- 'csrc/attn/sliding_tile_attn/st_attn.cpp'
|
||||
- 'csrc/attn/st_attn/**'
|
||||
- 'csrc/attn/setup_sta.py'
|
||||
- 'csrc/attn/config_sta.py'
|
||||
- 'csrc/attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/video_sparse_attn/**'
|
||||
- 'csrc/attn/video_sparse_attn/tk/**'
|
||||
- 'csrc/attn/video_sparse_attn/setup.py'
|
||||
- 'csrc/attn/video_sparse_attn/config_vsa.py'
|
||||
- 'csrc/attn/video_sparse_attn/vsa.cpp'
|
||||
- 'csrc/attn/vsa/**'
|
||||
- 'csrc/attn/tk/**'
|
||||
- 'csrc/attn/setup_vsa.py'
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
@@ -235,7 +234,7 @@ jobs:
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
@@ -373,4 +372,4 @@ jobs:
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
cd csrc/attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -144,13 +144,13 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
cd csrc/attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -165,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/sliding_tile_attn/dist/*.whl
|
||||
path: csrc/attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -239,11 +239,11 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/sliding_tile_attn/dist/
|
||||
packages-dir: csrc/attn/dist/
|
||||
|
||||
@@ -1,257 +0,0 @@
|
||||
name: Publish Video Sparse Attention Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/video_sparse_attn/setup.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/video_sparse_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.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/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_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/video_sparse_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/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/video_sparse_attn/dist/
|
||||
+2
-6
@@ -1,7 +1,3 @@
|
||||
[submodule "csrc/attn/video_sparse_attn/tk"]
|
||||
path = csrc/attn/video_sparse_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
[submodule "csrc/attn/sliding_tile_attn/tk"]
|
||||
path = csrc/attn/sliding_tile_attn/tk
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -22,7 +22,6 @@ 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,21 +1,21 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
<img src=assets/logo.jpg width="30%"/>
|
||||
</div>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
**FastVideo is a unified framework for accelerated video generation.**
|
||||
|
||||
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.
|
||||
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/rG0QpZdw" target="_blank"> <b> WeChat </b> </a> |
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/fastwan.png width="90%"/>
|
||||
<img src=assets/perf.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
@@ -23,19 +23,20 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
- End-to-end post-training support:
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
|
||||
- Data preprocessing pipeline for video data
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Diverse hardware and OS support
|
||||
- Support H100, A100, 4090
|
||||
- Support Linux, Windows, MacOS
|
||||
- Cutting edge models
|
||||
- Wan2.1 T2V, I2V
|
||||
- HunyuanVideo
|
||||
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- StepVideo T2V
|
||||
- Distillation support
|
||||
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
## Getting Started
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
@@ -51,31 +52,17 @@ pip install fastvideo
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
|
||||
|
||||
## Sparse Distillation
|
||||
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
See below for recipes and datasets:
|
||||
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
|:-------------------------------------------------------------------------------------------: |:---------------------------------------------------------------------------------------------------------------: |:--------------------------------------------------------------------------------------------------------: |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
|
||||
## Inference
|
||||
### Generating Your First Video
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Create a file called `example.py` with the following code:
|
||||
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
@@ -109,41 +96,41 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
More FastWan Models Coming Soon!
|
||||
- [ ] Add FastWan2.1-T2V-14B
|
||||
- [ ] Add FastWan2.2-T2V-14B
|
||||
- [ ] Add FastWan2.2-I2V-14B
|
||||
<!-- - Optimization features
|
||||
- Code updates -->
|
||||
- More models support
|
||||
<!-- - [ ] Add CogvideoX model -->
|
||||
- [x] Add StepVideo to V1
|
||||
- Optimization features
|
||||
- [x] Teacache in V1
|
||||
- [x] SageAttention in V1
|
||||
- Code updates
|
||||
- [x] V1 Configuration API
|
||||
- [ ] Support Training in V1
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
|
||||
- [Triton](https://github.com/triton-lang/triton)
|
||||
- [DMD2](https://github.com/tianweiy/DMD2)
|
||||
- [PCM](https://github.com/G-U-N/Phased-Consistency-Model)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
|
||||
We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you find FastVideo useful, please considering citing our work:
|
||||
If you use FastVideo for your research, please cite our work:
|
||||
|
||||
```bibtex
|
||||
@software{fastvideo2024,
|
||||
@@ -154,17 +141,31 @@ If you find FastVideo useful, please considering citing our work:
|
||||
year = {2024},
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2505.13389},
|
||||
year={2025}
|
||||
@misc{zhang2025vsafastervideodiffusion,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2505.13389},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2505.13389},
|
||||
}
|
||||
|
||||
@article{zhang2025fast,
|
||||
title={Fast video generation with sliding tile attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2502.04507},
|
||||
year={2025}
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.04507},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.04507},
|
||||
}
|
||||
@misc{ding2025efficientvditefficientvideodiffusion,
|
||||
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
|
||||
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.06155},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.06155},
|
||||
}
|
||||
```
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 194 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 149 KiB |
@@ -1,6 +0,0 @@
|
||||
<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>
|
||||
|
Before Width: | Height: | Size: 691 B |
Binary file not shown.
|
After Width: | Height: | Size: 31 KiB |
@@ -1,2 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
include config.py
|
||||
+3
-3
@@ -25,9 +25,9 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
(If you use CUDA12.4)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
@@ -84,7 +84,7 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
### Test
|
||||
```bash
|
||||
python tests/test_sta.py # test STA
|
||||
python tests/test_vsa.py # test VSA
|
||||
python tests/test_block_sparse.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
|
||||
@@ -86,6 +86,8 @@ def benchmark_attention(configurations):
|
||||
# print(f"Average TFLOPS: {tflops_bwd}")
|
||||
# print("=" * 60)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return results
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
off_hz = tl.program_id(2)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config_sta import kernels, sources, target
|
||||
from csrc.attn.config_sta import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -9,7 +9,7 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.6"
|
||||
VERSION = "0.0.4"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
@@ -9,10 +9,10 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.3"
|
||||
VERSION = "0.0.1"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
@@ -51,16 +51,19 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
|
||||
ext_modules = [
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
]
|
||||
ext_modules = []
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
ext_modules = [
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
recursive-include tk *
|
||||
include config_sta.py
|
||||
@@ -1,87 +0,0 @@
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python ../benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
Submodule csrc/attn/sliding_tile_attn/tk deleted from 6c27e28c81
+26
-23
@@ -13,9 +13,9 @@ BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
def pytorch_test(Q, K, V, block_sparse_mask, dO):
|
||||
q_ = Q.clone().float().requires_grad_()
|
||||
k_ = K.clone().float().requires_grad_()
|
||||
v_ = V.clone().float().requires_grad_()
|
||||
q_ = Q.clone().requires_grad_()
|
||||
k_ = K.clone().requires_grad_()
|
||||
v_ = V.clone().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.detach().requires_grad_()
|
||||
K = K.detach().requires_grad_()
|
||||
V = V.detach().requires_grad_()
|
||||
Q = Q.clone().requires_grad_()
|
||||
K = K.clone().requires_grad_()
|
||||
V = V.clone().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,9 +60,11 @@ def get_non_pad_index(
|
||||
|
||||
return index_pad[index_mask]
|
||||
|
||||
def generate_tensor(shape, dtype, device):
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
return tensor
|
||||
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()
|
||||
|
||||
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)
|
||||
@@ -73,7 +75,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, num_iterations=20, error_mode='all'):
|
||||
def check_correctness(h, d, num_blocks, k, mean, std, 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},
|
||||
@@ -89,10 +91,10 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
|
||||
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), 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)
|
||||
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)
|
||||
|
||||
# dO_padded = torch.zeros_like(dO_padded)
|
||||
# dO_padded[:, :, non_pad_index, :] = dO
|
||||
@@ -105,8 +107,7 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
|
||||
abs_diff = torch.abs(diff)
|
||||
results[name]['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).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())
|
||||
results[name]['max_diff'] = max(results[name]['max_diff'], torch.max(abs_diff).item())
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -118,27 +119,27 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_graphs(h, d, error_mode='all'):
|
||||
def generate_error_graphs(h, d, mean, std, 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}, mode={error_mode}")
|
||||
print(f"\nError Analysis for h={h}, d={d}, mean={mean}, std={std}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
|
||||
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}")
|
||||
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}")
|
||||
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, error_mode=error_mode)
|
||||
results = check_correctness(h, d, num_blocks, k, mean, std, 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} "
|
||||
@@ -149,8 +150,10 @@ def generate_error_graphs(h, d, 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, error_mode=mode)
|
||||
generate_error_graphs(h, d, mean, std, error_mode=mode)
|
||||
print("\nAnalysis completed for all modes.")
|
||||
|
||||
Submodule
+1
Submodule csrc/attn/tk added at 1719fb7264
@@ -1,61 +0,0 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python ../tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python ../benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
Submodule csrc/attn/video_sparse_attn/tk deleted from 6c27e28c81
@@ -1,185 +0,0 @@
|
||||
import torch
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
|
||||
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
|
||||
from vsa.index import map_to_index
|
||||
from typing import Tuple, Optional
|
||||
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.int()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_triton(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: 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()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
return dq, dk, dv
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
|
||||
def _block_sparse_attn_backward_triton_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: 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()
|
||||
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_triton(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def setup_context_triton(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, M = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
|
||||
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
|
||||
|
||||
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_SM90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
variable_block_sizes = variable_block_sizes.int()
|
||||
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
|
||||
def _block_sparse_attn_SM90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
|
||||
B, H, S, D = q_padded.shape
|
||||
o_padded = torch.empty_like(q_padded)
|
||||
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_SM90(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
|
||||
)
|
||||
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
|
||||
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
|
||||
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
|
||||
return grad_q_padded, grad_k_padded, grad_v_padded
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
|
||||
def _block_sparse_attn_backward_SM90_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
torch._check(grad_output_padded.dtype == torch.bfloat16)
|
||||
torch._check(lse_padded.dtype == torch.float32)
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_SM90(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
def setup_context_SM90(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, lse_padded = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
|
||||
@@ -1,13 +1,12 @@
|
||||
import torch
|
||||
from typing import Tuple
|
||||
block_sparse_attn=None
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
|
||||
block_sparse_attn = block_sparse_attn_SM90
|
||||
else:
|
||||
except ImportError:
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_triton
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
+4
-4
@@ -568,6 +568,7 @@ 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);
|
||||
@@ -577,7 +578,6 @@ 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,6 +624,7 @@ 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);
|
||||
@@ -633,14 +634,13 @@ 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();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,182 @@
|
||||
import torch
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
|
||||
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
|
||||
from vsa.index import map_to_index
|
||||
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,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.int()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_triton(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: 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()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
return dq, dk, dv
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
|
||||
def _block_sparse_attn_backward_triton_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: 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()
|
||||
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_triton(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def setup_context_triton(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
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)
|
||||
@@ -1,9 +1,7 @@
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -11,25 +9,17 @@ 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/*
|
||||
|
||||
# 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
|
||||
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 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
|
||||
ENV PATH=/opt/conda/bin:$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
|
||||
RUN conda create --name fastvideo-dev python=3.10.0 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
@@ -37,36 +27,22 @@ 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.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
|
||||
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
|
||||
|
||||
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
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.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
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
@@ -1,9 +1,7 @@
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -11,25 +9,17 @@ 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/*
|
||||
|
||||
# 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
|
||||
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 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
|
||||
ENV PATH=/opt/conda/bin:$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
|
||||
RUN conda create --name fastvideo-dev python=3.11.11 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
@@ -37,36 +27,22 @@ 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.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
|
||||
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
|
||||
|
||||
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
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.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
|
||||
|
||||
EXPOSE 22
|
||||
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.3 --no-build-isolation
|
||||
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
python setup_sta.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
python setup_vsa.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
@@ -1,72 +0,0 @@
|
||||
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/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 194 KiB |
+2
-1
@@ -96,7 +96,8 @@ copybutton_prompt_is_regexp = True
|
||||
#
|
||||
html_title = project
|
||||
html_theme = 'sphinx_book_theme'
|
||||
html_logo = '../../assets/logos/icon_simple.svg'
|
||||
html_logo = '../../assets/logo.jpg'
|
||||
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
|
||||
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:
|
||||
|
||||
**Images:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev: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:py3.12-latest
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev: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.8
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||
|
||||
@@ -22,20 +22,10 @@ source ~/.bashrc
|
||||
Create and activate a Conda environment for FastVideo:
|
||||
|
||||
```
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda create -n fastvideo python=3.10 -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:
|
||||
|
||||
```
|
||||
@@ -46,10 +36,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
|
||||
uv pip install -e .[dev]
|
||||
pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
uv pip install flash-attn --no-build-isolation
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
pre-commit install --hook-type pre-commit --hook-type commit-msg
|
||||
@@ -60,14 +50,3 @@ 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,19 +6,11 @@ We introduce a new finetuning strategy - **Sparse-distill**, which jointly integ
|
||||
|
||||
We provide two distilled models:
|
||||
|
||||
- **[FastWan2.1-T2V-1.3B-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers)**: 3-step inference, up to **16 FPS** on H100 GPU
|
||||
- **[FastWan2.1-T2V-14B-480P-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-480P-Diffusers)**: 3-step inference, up to **60x speed up** at 480P, **90x speed up** at 720P for denoising loop
|
||||
- **[FastWan2.2-TI2V-5B-FullAttn-Diffusers](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers)**: 3-step inference, up to **50x speed up** at 720P for denoising loop
|
||||
- **[FastWan2.1-T2V-1.3B-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers)**: 3-step inference, up to **20 FPS** on H100 GPU
|
||||
- **[FastWan2.1-T2V-14B-480P-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-480P-Diffusers)**: 3-step inference, up to **50x speed up** at 480P, **70x speed up** at 720P for denoising loop
|
||||
|
||||
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
|
||||
|
||||
## ⚙️ Inference
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
```
|
||||
|
||||
## 🗂️ Dataset
|
||||
|
||||
We use the **FastVideo 480P Synthetic Wan dataset** ([FastVideo/Wan-Syn_77x448x832_600k](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k)) for distillation, which contains 600k synthetic latents.
|
||||
@@ -35,7 +27,7 @@ python scripts/huggingface/download_hf.py \
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
### Wan2.1 1.3B Model Sparse-Distill
|
||||
### 1.3B Model Sparse-Distill
|
||||
|
||||
For the 1.3B model, we use **4 nodes with 32 H200 GPUs** (8 GPUs per node):
|
||||
|
||||
@@ -51,7 +43,7 @@ sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_1.3B.sl
|
||||
- VSA attention sparsity: 0.8
|
||||
- Training steps: 4000 (~12 hours)
|
||||
|
||||
### Wan2.1 14B Model Sparse-Distill
|
||||
### 14B Model Sparse-Distill
|
||||
|
||||
For the 14B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
|
||||
|
||||
@@ -69,19 +61,9 @@ sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_14B.slu
|
||||
- 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):
|
||||
## ⚙️ Inference
|
||||
Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
# Multi-node training (8 nodes, 64 GPUs total)
|
||||
sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
|
||||
bash scripts/inference/v1_inference_wan_dmd.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.8**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Set up using Python
|
||||
@@ -38,7 +38,6 @@ 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:
|
||||
@@ -61,7 +60,7 @@ uv pip install fastvideo
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -88,7 +87,7 @@ uv pip install -e .
|
||||
#### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
@@ -103,7 +102,7 @@ If you're planning to contribute to FastVideo please see the following page:
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA 12.8 support
|
||||
- NVIDIA GPU with CUDA 12.4 support
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
|
||||
@@ -39,7 +39,6 @@ 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:
|
||||
|
||||
+17
-9
@@ -1,6 +1,6 @@
|
||||
# Welcome to FastVideo
|
||||
|
||||
:::{figure} ../../assets/logos/logo.svg
|
||||
:::{figure} ../../assets/logo.jpg
|
||||
:align: center
|
||||
:alt: FastVideo
|
||||
:class: no-scaled-link
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
:::{raw} html
|
||||
<p style="text-align:center">
|
||||
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.
|
||||
<strong>FastVideo is a unified framework for accelerated video generation.
|
||||
</strong>
|
||||
</p>
|
||||
|
||||
@@ -21,10 +21,11 @@
|
||||
</p>
|
||||
:::
|
||||
|
||||
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.
|
||||
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<img src=_static/images/fastwan.png width="100%"/>
|
||||
<img src=_static/images/perf.png width="100%"/>
|
||||
</div>
|
||||
|
||||
## Key Features
|
||||
@@ -32,13 +33,19 @@ FastVideo is an inference and post-training framework for diffusion models. It f
|
||||
FastVideo has the following features:
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- E2E post-training support
|
||||
- Data preprocessing pipeline for video data.
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Cutting edge models
|
||||
- Wan2.1 T2V, I2V
|
||||
- HunyuanVideo
|
||||
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- StepVideo T2V
|
||||
- Distillation support
|
||||
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
## Documentation
|
||||
|
||||
@@ -101,6 +108,7 @@ 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.8
|
||||
- **CUDA**: 12.4
|
||||
- **GPU**: At least one NVIDIA GPU
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -6,7 +6,6 @@ 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.
|
||||
@@ -38,94 +37,51 @@ 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
|
||||
* ❌
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 T2V 1.3B
|
||||
- * Wan T2V 1.3B
|
||||
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 T2V 14B
|
||||
- * Wan T2V 14B
|
||||
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
|
||||
* 480P, 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 I2V 480P
|
||||
- * Wan I2V 480P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 I2V 720P
|
||||
- * Wan I2V 720P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
|
||||
* 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * StepVideo T2V
|
||||
* `FastVideo/stepvideo-t2v-diffusers`
|
||||
* 768px768px204f<br>544px992px204f<br>544px992px136f
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
:::
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
**Note**: there are some known quality issues with Wan2.1 + Sliding Tile Attn. We are working on fixing this issue.
|
||||
|
||||
## Special requirements
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
You can install the Sliding Tile Attention package using
|
||||
|
||||
```
|
||||
pip install st_attn
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
@@ -12,6 +12,7 @@ We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have impleme
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
cd csrc/sliding_tile_attention/
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
@@ -21,20 +22,14 @@ sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
Set up CUDA environment (if using CUDA 12.4):
|
||||
Install STA:
|
||||
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
Install STA:
|
||||
|
||||
```bash
|
||||
cd csrc/attn/sliding_tile_attn/
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
# 🧪 Test
|
||||
@@ -58,9 +53,3 @@ out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
|
||||
# 🚀Inference
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
(vsa-demo)=
|
||||
|
||||
# 🎬 Demo
|
||||
@@ -4,7 +4,8 @@
|
||||
You can install the Video Sparse Attention package using
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
@@ -22,10 +23,10 @@ sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
Set up CUDA environment (if using CUDA 12.8):
|
||||
Set up CUDA environment (if using CUDA 12.4):
|
||||
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
@@ -33,9 +34,9 @@ export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
Install VSA:
|
||||
|
||||
```bash
|
||||
cd csrc/attn/video_sparse_attn/
|
||||
cd csrc/attn/
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
# 🧪 Test
|
||||
@@ -58,9 +59,3 @@ from vsa import video_sparse_attn
|
||||
output = video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight)
|
||||
|
||||
```
|
||||
|
||||
# 🚀Inference
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_VSA.sh
|
||||
```
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
A person reading a book with words that float off the pages and form pictures.
|
||||
A person diving into a pool of liquid crystal, creating ripples of light.
|
||||
A handheld shot chasing after a group of friends laughing and playing on the beach at sunset.
|
||||
A mysterious ancient temple hidden in the jungle.
|
||||
A high-speed train navigating a steep descent.
|
||||
a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a winter storm
|
||||
A cheetah accelerating to full speed while chasing its prey.
|
||||
A serene orchard is in full bloom, with trees heavy with blossoms and bees buzzing around, darting from flower to flower in a display of natural harmony.
|
||||
A little child let out a big yawn
|
||||
Subtle reflections of a woman on the window of a train moving at hyper-speed in a Japanese city.
|
||||
A truck left along the edge of a cliff, revealing the stunning coastal landscape below with waves crashing against the rocks.
|
||||
A red bird transforms into a flag
|
||||
A zoom-out from a single leaf on a tree to reveal the entire forest, showcasing the vastness and diversity of the woodland.
|
||||
A slow-motion video of a liquid droplet bouncing on a water-repellent surface.
|
||||
Static camera shot. A dinasour running near some lions and chasing them away.
|
||||
an adorable kangaroo wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a beautiful sunset
|
||||
A zoom-in on an artist's brush touching the canvas, highlighting the texture of the paint and the strokes being made.
|
||||
an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a colorful festival
|
||||
A woman is ascending to the sky from the ground
|
||||
View out a window of a giant strange creature walking in rundown city at night, one single street lamp dimly lighting the area.
|
||||
An arc shot around a lone tree in a vast, foggy field at dawn, revealing the changing light and shadows.
|
||||
A person sculpting a statue out of a waterfall, the water solidifying under their touch.
|
||||
The person's forehead creased with concentration as she worked on a challenging puzzle.
|
||||
The person's cheeks flushed with pleasure as she savored a delicious meal.
|
||||
Hand-drawn simple line art, a young kid looking up into space with a wondrous expression on his face.
|
||||
A crab made of different jewlery is walking on the beach. As it walks, it drops different jewelry pieces like diamonds, pearls, etc
|
||||
Gold coins are falling out when elevator door opens
|
||||
the scene transitions from huge waves into a snowy mountain at sunset
|
||||
a giant cathedral is completely filled with cats. there are cats everywhere you look. a man enters the cathedral and bows before the giant cat king sitting on a throne.
|
||||
A mother dog gently picks up a piece of meat and carefully places it in her puppy's bowl, her eyes filled with warmth and care as she watches her little one eat.
|
||||
A soap bubble floating in the air, displaying iridescent colors that shift and change as it moves through different angles of light.
|
||||
A truck left alongside a train moving through the countryside, matching its speed and revealing the changing landscape.
|
||||
An astronaut walking between stone buildings.
|
||||
A close-up shot of the person's face reveals his fear and desperation as he navigates the ship through the storm.
|
||||
A frozen lake slowly cracking and thawing as spring arrives, with sheets of ice breaking apart and drifting across the surface.
|
||||
A FPV shot zooming through a tunnel into a vibrant underwater space.
|
||||
a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a colorful festival
|
||||
A person sips on a smoothie, the cool and fruity flavors refreshing her mouth.
|
||||
In a vibrant theater, a magician in dazzling attire stands center stage, pulling a comically oversized rubber chicken from an ornate, old-fashioned box. His costume shimmers under the stage lights, adding to the spectacle. The crowd erupts in laughter and applause, their faces filled with joy and amazement. The magician's expression hints at mischievous delight as he holds up the rubber chicken, his performance bringing cheer to the audience.
|
||||
A hamster running on a spinning wheel.
|
||||
A quaint village nestled in a valley is surrounded by blooming cherry blossoms, with petals drifting through the air as villagers go about their daily activities, adding life to the scene.
|
||||
In a tranquil forest clearing, a sparkling waterfall cascades down into a clear pool, surrounded by lush greenery and flowers, with occasional birds fluttering by.
|
||||
A woman beamed with pride as she watched her child perform on stage.
|
||||
an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a winter storm
|
||||
A man is eating salad
|
||||
An Asian girl wearing a bright yellow T-shirt and white pants is Hip-Hop dancing
|
||||
nighttime footage of a hermit crab using an incandescent lightbulb as its shell
|
||||
a toy robot wearing a green dress and a sun hat taking a pleasant stroll in Antarctica during a beautiful sunset
|
||||
A goat operating a food truck, serving gourmet grilled cheese sandwiches to a line of animals.
|
||||
Macro shot. Man in an antique scuba helmet with dark glass walking out of a flower
|
||||
A bustling train station in the heart of a vibrant city.
|
||||
Light filtering through a canopy of autumn leaves, casting warm, dappled patterns of yellow, orange, and red onto the ground.
|
||||
Chimneys in the setting sun
|
||||
A longboarder accelerating downhill, carving through turns.
|
||||
A couple runs through a sudden downpour, laughing and splashing in puddles as they try to find shelter.
|
||||
A glass of iced coffee condensing water on the outside, with droplets forming and sliding down the glass in slow motion.
|
||||
macro shot of a leaf showing tiny trains moving through its veins
|
||||
A corgi wearing sunglasses walks on the beach of a tropical island
|
||||
Borneo wildlife on the Kinabatangan River
|
||||
A beautiful silhouette animation shows a wolf howling at the moon, feeling lonely, until it finds its pack.
|
||||
an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a colorful festival
|
||||
A green monster made of plants walks through an airport.
|
||||
A close up view of a glass sphere that has a zen garden within it. There is a small dwarf in the sphere who is raking the zen garden and creating patterns in the sand.
|
||||
A person on a scooter colliding with a park bench, the scooter tipping over.
|
||||
A tilt-up from a city street, ascending to show the skyline with its mix of modern and historic architecture.
|
||||
A chef tossing a pancake into the air and catching it.
|
||||
A woman whispering a secret into a friend's ear.
|
||||
A vulture circling high in the sky.
|
||||
A medieval castle overlooking a bustling renaissance fair.
|
||||
a toy robot wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a beautiful sunset
|
||||
A man standing in front of a burning building giving the 'thumbs up' sign.
|
||||
The person's cheeks flushed with embarrassment as he told a funny story.
|
||||
Llamas and Emus are playing chess
|
||||
A woman sipping a steaming cup of tea.
|
||||
A tree root bursting through the seat of an ancient, weathered bench, intertwining with the wood.
|
||||
Smoke rises from the chimney of a cozy log cabin nestled in the woods, with soft light glowing from the windows, suggesting a warm and inviting atmosphere.
|
||||
A close-up of sparkling water being poured into a glass, capturing the detailed flow and bubbles.
|
||||
a woman wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a beautiful sunset
|
||||
The Glenfinnan Viaduct is a historic railway bridge in Scotland, UK, that crosses over the west highland line between the towns of Mallaig and Fort William. It is a stunning sight as a steam train leaves the bridge, traveling over the arch-covered viaduct. The landscape is dotted with lush greenery and rocky mountains, creating a picturesque backdrop for the train journey. The sky is blue and the sun is shining, making for a beautiful day to explore this majestic spot.
|
||||
A piece of elastic fabric being pulled and stretched, then returning to its original size when the tension is released.
|
||||
a woman wearing a green dress and a sun hat taking a pleasant stroll in Antarctica during a beautiful sunset
|
||||
A video of a water jet cutting through metal, showing the powerful and precise movement of water.
|
||||
Car mirrors and sunsets
|
||||
Giant Pandas are eating hot noodles in a Chinese restaurant
|
||||
A rally car taking a fast turn on a track
|
||||
a toy robot wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a colorful festival
|
||||
A crystal-clear icicle slowly dripping as it melts in the warmth of the midday sun, each drop sparkling as it falls.
|
||||
A tilt-down from a chandelier in a grand hall, revealing the ornate decor and people mingling below.
|
||||
A man is playing the drums under the water
|
||||
A person playing an electric guitar made of lightning, with thunderous sound waves.
|
||||
A person floating in a bubble, drifting over a bustling cityscape.
|
||||
A tilt-down from a starry night sky, revealing a quiet forest clearing bathed in moonlight.
|
||||
A pan right through a dense jungle, moving past lush vegetation and exotic wildlife.
|
||||
Close-up of a man eating an apple.
|
||||
A low-angle shot of a dancer leaping gracefully into the air, making their movement appear even more dynamic and powerful.
|
||||
A woman is search her bag trying to find something.
|
||||
A bulldozer clears debris from a demolished building, making way for new construction.
|
||||
A man sighed in relief as the doctor delivered the good news.
|
||||
A tsunami coming through an alley in Bulgaria, dynamic movement.
|
||||
Blooming Flowers
|
||||
A push-in through a dense crowd at a festival, moving towards a performer on stage who is captivating the audience.
|
||||
A truck right through a tranquil garden, moving past blooming flowers, trees, and a small fountain.
|
||||
The person's eyes sparkled with excitement as he greeted a friend.
|
||||
A person playing chess with a robot on a floating platform above the ocean.
|
||||
A gentle breeze rustles the leaves as someone walks down a serene forest path, sunlight filtering through the trees and shifting patterns on the ground as branches sway.
|
||||
A rollercoaster ride from a city to a desert and then to an ice world
|
||||
A pan left across an ancient library, moving from shelf to shelf, showcasing rows of leather-bound books.
|
||||
A mother otter floating on her back in a river, cradling her pup on her stomach to keep it safe and warm in the gentle current.
|
||||
an adorable kangaroo wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a colorful festival
|
||||
a woman wearing a green dress and a sun hat taking a pleasant stroll in Mumbai India during a colorful festival
|
||||
A delicate layer of morning frost melting off a flower petal, the tiny droplets glistening like diamonds in the light.
|
||||
A panda is cooking for her child, her child is next to her.
|
||||
Macro shot of a man wearing an antique diving helmet with dark glass and a jetpack walking on the veins of a leaf. Realistic style
|
||||
an old man wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a beautiful sunset
|
||||
A girl is unfolding a birthday gift.
|
||||
A pencil drawing an architectural plan.
|
||||
A handheld camera following a dog running through a park, bouncing and tilting as it captures the dog's joyful exploration.
|
||||
A pan left across a serene beach at sunrise, moving from the darkened shore to the brightening horizon.
|
||||
A group of people are clapping to celebrate
|
||||
Vendors set up stalls at a bustling farmer’s market, displaying fresh fruits and vegetables, while people stroll through, selecting produce and enjoying the lively atmosphere.
|
||||
A police helicopter hovers above a high-speed chase, guiding officers on the ground to apprehend a suspect.
|
||||
A paper origami dragon riding a boat in waves. Realistic style.
|
||||
A close-up of a droplet of dew forming on a leaf, capturing the detailed surface tension.
|
||||
a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a beautiful sunset
|
||||
A dry rainbow rose is coming back to life.
|
||||
A glass falling off a table and shattering on the floor.
|
||||
A marathon runner crossing the finish line after a grueling race.
|
||||
A zoom-in on a drop of morning dew on a leaf, showing the reflection of the surrounding world within it.
|
||||
A child blowing on hot cocoa to cool it down.
|
||||
A squad of futsal players showcasing their skills on an indoor court.
|
||||
A princess is brushing her long golden hair in the garden.
|
||||
A close-up of a pair of eyes, revealing the subtle emotions and reflections within them.
|
||||
A tracking shot of a group of cyclists racing through a forest trail, with trees and foliage rushing by.
|
||||
A woman yawning widely at the end of a long day.
|
||||
an old man wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a colorful festival
|
||||
Hidden within a garden, an ancient fountain trickles with water, surrounded by vibrant flowers and lush greenery that seem to whisper secrets of the past.
|
||||
A Chinese man sits at a table and eats noodles with chopsticks
|
||||
A pink pig running fast toward the camera in an alley in Tokyo.
|
||||
Strange creatures move through a mysterious, foggy marsh, their silhouettes barely visible through the dense mist as they navigate the eerie, otherworldly landscape.
|
||||
Tour of an art gallery with many beautiful works of art in different styles.
|
||||
FPV flying through a colorful coral lined streets of an underwater suburban neighborhood.
|
||||
Aerial view of Santorini during the blue hour, showcasing the stunning architecture of white Cycladic buildings with blue domes. The caldera views are breathtaking, and the lighting creates a beautiful, serene atmosphere.
|
||||
Camera zoom out. A couple walking along the beach as the sun sets over the ocean.
|
||||
an extreme close up shot of a woman's eye, with her iris appearing as earth
|
||||
a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a colorful festival
|
||||
an old man wearing a green dress and a sun hat taking a pleasant stroll in Mumbai India during a winter storm
|
||||
an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a winter storm
|
||||
A martial artist breaking a board with a powerful punch.
|
||||
People gather on a peaceful beach at sunset, a bonfire crackling as they sit around, enjoying the warmth and the sight of the sun dipping below the horizon.
|
||||
A close-up of a waterfall, showing the detailed movement of water as it crashes down.
|
||||
A child is blowing bubbles
|
||||
a woman wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a winter storm
|
||||
A wide-angle perspective of a serene lake surrounded by mountains, reflecting the sky and creating a sense of infinite space.
|
||||
The person's eyebrows arched in skepticism as she listened to a dubious claim.
|
||||
an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a beautiful sunset
|
||||
a woman wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a colorful festival
|
||||
a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Antarctica during a colorful festival
|
||||
a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a colorful festival
|
||||
A chef flips a pancake and puts cream on it.
|
||||
An astronaut runs on the surface of the moon, the low angle shot shows the vast background of the moon, the movement is smooth and appears lightweight
|
||||
A man's face lit up with happiness as he received a heartfelt compliment.
|
||||
A futuristic spaceport hums with activity as ships of various shapes and sizes take off and land on multiple platforms, their engines glowing with vibrant colors.
|
||||
A person knitting a scarf using beams of light instead of yarn.
|
||||
A pedestal up from the edge of a canyon, gradually revealing the expansive landscape and river below.
|
||||
a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a colorful festival
|
||||
an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a colorful festival
|
||||
A person walking up a staircase made of clouds leading to a floating castle.
|
||||
Monks meditate in a serene mountaintop temple, sitting in quiet reflection as the wind gently moves through the surrounding trees, creating a sense of peace and tranquility.
|
||||
An aerial shot of a bustling city intersection at rush hour, capturing the organized chaos of cars and pedestrians.
|
||||
A pair of hands skillfully knitting a colorful scarf, the yarn winding through their fingers with each stitch.
|
||||
Close-up, a Chinese child is eating dumplings
|
||||
A kite losing wind and falling to the ground.
|
||||
Bioluminescent waves gently wash ashore on a deserted beach, illuminating the sand with each cresting wave as a figure walks along the water's edge, leaving glowing footprints.
|
||||
A red panda taking a bite of a pizza
|
||||
A close-up shot of a young woman driving a car, looking thoughtful, blurred green forest visible through the rainy car window.
|
||||
A high-speed video of a splash created by a stone thrown into a pond.
|
||||
A metal rod being bent slightly by a force and then springing back to its original straight shape when the force is removed.
|
||||
A hedgehog in a knight's armor, riding a toy horse into a medieval castle.
|
||||
A bird made of fresh oranges rushes out of the orange
|
||||
A low altitude first person perspective camera tracking shot of a soccer player's feet dribbling the ball on the groud in a soccer field, Sports Videography, Motion Tracking camera shot
|
||||
A tranquil island retreat features swaying palm trees and hammocks strung between them, inviting guests to relax and enjoy the serene beauty of the surroundings.
|
||||
a spooky haunted mansion, with friendly jack o lanterns and ghost characters welcoming trick or treaters to the entrance, tilt shift photography
|
||||
A coconut tree made of dollar bills at sunset, with bills falling off like leaves.
|
||||
A motocross bike accelerating out of a tight turn on a dirt track.
|
||||
A tranquil Zen garden with a gently flowing stream and koi fish.
|
||||
A green monster made of leaves walks through the airport, carrying a suitcase.
|
||||
A time-lapse of a frost-covered leaf gradually thawing in the morning sunlight, with tiny water droplets forming and trickling down.
|
||||
A woman practicing her archery skills at a range.
|
||||
A slow-motion video of ink being injected into a tank of water, creating intricate and beautiful patterns.
|
||||
a woman wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a winter storm
|
||||
The person's forehead creased with worry as he listened to bad news.
|
||||
An arc shot around a grand piano being played in an empty concert hall, the motion revealing the intricate details of the instrument.
|
||||
A person conducting a symphony of animals in a forest clearing.
|
||||
A truck right alongside a flowing river, capturing the movement of the water and the surrounding forest.
|
||||
A rocket blasting off from the launch pad, accelerating rapidly into the sky.
|
||||
Workers move through a picturesque vineyard during the harvest season, carefully picking grapes and placing them into baskets as the sun bathes the vines in a warm glow.
|
||||
A person is eating an ice cream.
|
||||
An over-the-shoulder perspective of a chef meticulously plating a dish in a bustling kitchen.
|
||||
A man looked away in shame when confronted with his wrongdoing.
|
||||
A person is savoring a slice of pizza at a pizzeria.
|
||||
@@ -1,76 +0,0 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,13 +0,0 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,4 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 720P dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "FastVideo/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
|
||||
@@ -4,7 +4,9 @@ These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -86,7 +86,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -102,6 +102,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -86,7 +86,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -102,6 +102,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -86,7 +86,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -102,6 +102,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
# Wan2.2-5B Distill Example
|
||||
These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA methods.
|
||||
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### 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.
|
||||
@@ -1,144 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_Wan2.2/t2v_g2e5_f1e5_%j.out
|
||||
#SBATCH --error=dmd_Wan2.2/t2v_g2e5_f1e5_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
export WANDB_API_KEY=your_wandb_api_key
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name Wan_distillation
|
||||
--output_dir "your_output_dir"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 31
|
||||
--num_height 704
|
||||
--num_width 1280
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DIR"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--lr_scheduler "cosine_with_min_lr"
|
||||
--min_lr_ratio 0.5
|
||||
--lr_warmup_steps 100
|
||||
--fake_score_learning_rate 1e-5
|
||||
--fake_score_lr_scheduler "cosine_with_min_lr"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 200
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3
|
||||
--simulate_generator_forward
|
||||
--log_visualization # disable if oom
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -1,145 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_Wan2.2/t2v_g2e5_f1e5_%j.out
|
||||
#SBATCH --error=dmd_Wan2.2/t2v_g2e5_f1e5_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
export WANDB_API_KEY=your_wandb_api_key
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name Wan_distillation
|
||||
--output_dir "your_output_dir"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 31
|
||||
--num_height 704
|
||||
--num_width 1280
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DIR"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--lr_scheduler "cosine_with_min_lr"
|
||||
--min_lr_ratio 0.5
|
||||
--lr_warmup_steps 100
|
||||
--fake_score_learning_rate 1e-5
|
||||
--fake_score_lr_scheduler "cosine_with_min_lr"
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 200
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3
|
||||
--simulate_generator_forward
|
||||
--log_visualization # disable if oom
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -4,17 +4,9 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
```bash
|
||||
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/download_dataset.sh
|
||||
```
|
||||
|
||||
### 2. Configure and run distillation:
|
||||
|
||||
#### For DMD-only distillation:
|
||||
```bash
|
||||
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/
|
||||
```
|
||||
### TODO
|
||||
@@ -63,7 +63,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -77,6 +77,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
import os
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -8,29 +12,41 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "FastVideo/FastWan2.1-T2V-14B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
text_encoder_cpu_offload=False,
|
||||
# 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"
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
# sampling_param.image_path = "test.jpg"
|
||||
# sampling_param.num_inference_steps = 0
|
||||
# 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."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
i2v_prompt = "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
i2v_prompt = "A little girl is packing a suitcase and the contents starts flying out of the suitcase everywhere."
|
||||
prompt = i2v_prompt
|
||||
# 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."
|
||||
# )
|
||||
results = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
stage_names = results["stage_names"]
|
||||
stage_execution_times = results["stage_execution_times"]
|
||||
# print(logging_info)
|
||||
print(stage_names)
|
||||
print(stage_execution_times)
|
||||
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
return
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
@@ -17,7 +17,7 @@ def main():
|
||||
use_fsdp_inference=True,
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
@@ -31,6 +31,7 @@ def main():
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. The puddles reflect glowing signs in kanji, advertising ramen, karaoke, and VR arcades. A woman in a translucent raincoat walks briskly with an LED umbrella. Steam rises from a street food cart, and a cat darts across the screen. Raindrops are visible on the camera lens, creating a cinematic bokeh effect."
|
||||
)
|
||||
prompt = "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."
|
||||
start_time = time.perf_counter()
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
end_time = time.perf_counter()
|
||||
@@ -45,7 +46,7 @@ def main():
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
start_time = time.perf_counter()
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=False)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
|
||||
@@ -1,31 +0,0 @@
|
||||
import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_causal"
|
||||
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.
|
||||
model_name = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,48 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,30 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
|
||||
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, height=832, width=480, num_frames=81)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,41 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
|
||||
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.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# I2V is triggered just by passing in an image_path argument
|
||||
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, output_path=OUTPUT_PATH, save_video=True, image_path=image_path)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
# T2V mode
|
||||
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.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,38 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
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.1-I2V-14B-480P-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-I2V-14B-480P-Diffusers")
|
||||
sampling_param.num_frames = 61
|
||||
sampling_param.num_inference_steps = 40
|
||||
sampling_param.guidance_scale = 5.0
|
||||
sampling_param.height = 448
|
||||
sampling_param.width = 832
|
||||
sampling_param.seed = 1024
|
||||
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 = (
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
)
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -261,7 +261,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
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.Image("fastvideo-logos/main/svg/full.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>
|
||||
@@ -637,10 +637,10 @@ def main():
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/logo.svg")
|
||||
@app.get("/logo.png")
|
||||
def get_logo():
|
||||
return FileResponse(
|
||||
"assets/logos/logo.svg",
|
||||
"fastvideo-logos/main/svg/full.svg",
|
||||
media_type="image/svg+xml",
|
||||
headers={
|
||||
"Cache-Control": "public, max-age=3600",
|
||||
@@ -650,7 +650,7 @@ def main():
|
||||
|
||||
@app.get("/favicon.ico")
|
||||
def get_favicon():
|
||||
favicon_path = "assets/logos/icon_simple.svg"
|
||||
favicon_path = "fastvideo-logos/main/svg/icon-simple.svg"
|
||||
|
||||
if os.path.exists(favicon_path):
|
||||
return FileResponse(
|
||||
@@ -683,7 +683,7 @@ def main():
|
||||
<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" content="{base_url}/logo.png">
|
||||
<meta property="og:image:width" content="1200">
|
||||
<meta property="og:image:height" content="630">
|
||||
<meta property="og:site_name" content="FastWan">
|
||||
@@ -692,7 +692,7 @@ def main():
|
||||
<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">
|
||||
<meta property="twitter:image" content="{base_url}/logo.png">
|
||||
<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">
|
||||
|
||||
@@ -10,7 +10,7 @@ def main():
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=False,
|
||||
lora_path="benjamin-paine/steamboat-willie-1.3b",
|
||||
lora_nickname="steamboat"
|
||||
)
|
||||
@@ -18,7 +18,7 @@ def main():
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"guidance_scale": 6.0,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"seed": 42,
|
||||
}
|
||||
|
||||
@@ -11,23 +11,22 @@ def main():
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-160/transformer",
|
||||
pin_cpu_memory=False,
|
||||
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1250/transformer",
|
||||
lora_nickname="crush_smol"
|
||||
)
|
||||
generator.unmerge_lora_weights()
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77,
|
||||
"guidance_scale": 6.0,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"seed": 42,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press."
|
||||
prompt = "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table."
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
@@ -35,12 +34,6 @@ def main():
|
||||
save_video=True,
|
||||
**kwargs
|
||||
)
|
||||
prompt = "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press."
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -14,7 +14,7 @@ def main():
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
load_time = time.perf_counter() - start_time
|
||||
print(f"Model loading time: {load_time:.2f} seconds")
|
||||
|
||||
@@ -12,7 +12,7 @@ def main():
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
load_time = time.perf_counter() - start_time
|
||||
print(f"Model loading time: {load_time:.2f} seconds")
|
||||
|
||||
@@ -54,7 +54,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "6.0"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -69,6 +69,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -88,7 +88,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "6.0"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -103,6 +103,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -101,6 +101,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -7,7 +7,9 @@ These are e2e example scripts for finetuning Wan2.1 T2V with VSA to accelerate i
|
||||
## Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
### Download the synthetic dataset:
|
||||
|
||||
@@ -101,6 +101,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--dit_precision "fp32"
|
||||
|
||||
+3
@@ -2,3 +2,6 @@
|
||||
|
||||
# 480P dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
|
||||
|
||||
# 720P dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "FastVideo/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -54,7 +54,7 @@ validation_args=(
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "6.0"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
@@ -69,6 +69,7 @@ optimizer_args=(
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user