Compare commits

..
Author SHA1 Message Date
“BrianChen1129” dbd38332c9 update 2025-08-23 21:47:06 +00:00
“BrianChen1129” be27c17a7a update 2025-08-14 03:51:06 +00:00
“BrianChen1129” f1c718ef85 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-12 22:28:17 +00:00
“BrianChen1129” 28df93fddf Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-12 02:32:12 +00:00
“BrianChen1129” 741dfb2a23 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-11 03:44:32 +00:00
“BrianChen1129” 690362e0b6 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-04 22:28:44 +00:00
“BrianChen1129” 9e1bd28dd3 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-04 19:27:20 +00:00
“BrianChen1129” 793ab20451 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-04 17:10:09 +00:00
“BrianChen1129” 7331b8f59a Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-04 15:51:33 +00:00
“BrianChen1129” c4b239a62a Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-04 02:55:47 +00:00
“BrianChen1129” 084d2dba64 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-08-02 03:12:26 +00:00
“BrianChen1129” 3187698cc2 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-31 02:34:39 +00:00
“BrianChen1129” 5ae466be21 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-30 22:38:16 +00:00
“BrianChen1129” f70cdc4c2e Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-30 20:22:03 +00:00
“BrianChen1129” 62fe15b64f Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-30 19:46:52 +00:00
“BrianChen1129” 6b43af5274 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-30 10:30:02 +00:00
“BrianChen1129” 3efe6da6b0 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-30 03:06:02 +00:00
“BrianChen1129” 3fc247ddf4 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-30 00:42:45 +00:00
“BrianChen1129” 76006d24cf Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-29 05:41:17 +00:00
“BrianChen1129” 708fe18bb8 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-29 03:17:18 +00:00
“BrianChen1129” a0522cd7f0 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-29 00:20:14 +00:00
“BrianChen1129” df6afbb8b0 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-28 07:28:42 +00:00
“BrianChen1129” 6d446c6eb9 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-27 22:46:52 +00:00
“BrianChen1129” ebae1321cf Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-27 20:02:15 +00:00
“BrianChen1129” 5c94cbad94 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-27 07:41:01 +00:00
“BrianChen1129” 862ac8a6a1 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-27 03:31:20 +00:00
“BrianChen1129” e1b2c40857 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-25 19:33:45 +00:00
“BrianChen1129” 93521970f6 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-23 20:11:45 +00:00
“BrianChen1129” e7c363d313 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-22 03:03:17 +00:00
“BrianChen1129” fc63becc92 sy 2025-07-15 05:26:10 +00:00
“BrianChen1129” 3973469f70 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-07-01 02:34:09 +00:00
“BrianChen1129” 9c135a3845 Merge branch 'main' of github.com:hao-ai-lab/FastVideo 2025-06-30 08:15:57 +00:00
“BrianChen1129” bc8b24ca9b update 2025-06-30 04:10:50 +00:00
160 changed files with 1853 additions and 6796 deletions
+18 -41
View File
@@ -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,12 @@ steps:
agents:
queue: "default"
- path:
- "csrc/attn/video_sparse_attn/**"
- "csrc/attn/video_sparse_attn/tk/**"
- "csrc/attn/vsa/**"
- "csrc/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/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -176,26 +176,3 @@ steps:
- TEST_TYPE=precision_vsa
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VMoBA"
env:
- TEST_TYPE=precision_vmoba
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/vmoba/**"
- "fastvideo/attention/backends/vmoba.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests VMoBA"
env:
- TEST_TYPE=inference_vmoba
agents:
queue: "default"
-9
View File
@@ -109,15 +109,6 @@ case "$TEST_TYPE" in
log "Running distillation DMD tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
;;
# run_inference_tests_vmoba
"inference_vmoba")
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
;;
"precision_vmoba")
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
-15
View File
@@ -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
+11 -12
View File
@@ -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
+11 -11
View File
@@ -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/
+11 -11
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/attn/video_sparse_attn/setup.py"
- "csrc/attn/setup_vsa.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/attn/video_sparse_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_vsa.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_vsa.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -152,13 +152,13 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn/video_sparse_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_vsa.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/attn/video_sparse_attn
cd csrc/attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
@@ -173,7 +173,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/attn/video_sparse_attn/dist/*.whl
path: csrc/attn/dist/*.whl
retention-days: 90
publish_package:
@@ -247,11 +247,11 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn/video_sparse_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_vsa.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/attn/video_sparse_attn/dist/
packages-dir: csrc/attn/dist/
+2 -6
View File
@@ -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
-1
View File
@@ -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
)
+2 -2
View File
@@ -1,5 +1,5 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
<img src=assets/logo.png width="30%"/>
</div>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
@@ -7,7 +7,7 @@
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://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-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/qqPzbrw" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 46 KiB

-6
View File
@@ -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

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

Before

Width:  |  Height:  |  Size: 5.7 KiB

-6
View File
@@ -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
+2 -2
View File
@@ -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
```
+4
View File
@@ -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/')))
-2
View File
@@ -1,2 +0,0 @@
recursive-include tk *
include config_sta.py
-87
View File
@@ -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.
+26 -23
View File
@@ -13,9 +13,9 @@ BLOCK_M = 64
BLOCK_N = 64
def pytorch_test(Q, K, V, block_sparse_mask, dO):
q_ = Q.clone().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
-61
View File
@@ -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.
-32
View File
@@ -1,32 +0,0 @@
# Attention Kernel Used in FastVideo
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
### Installation
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
### Usage
You can use `moba_attn_varlen` in the following ways:
**Install from source:**
```bash
python setup.py install
```
**Import after installation:**
```python
from vmoba import moba_attn_varlen
```
**Or import directly from the project root:**
```python
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
```
### Verify if you have successfully installed
```bash
python csrc/attn/vmoba_attn/vmoba/vmoba.py
```
-24
View File
@@ -1,24 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from setuptools import find_packages, setup
PACKAGE_NAME = "vmoba"
VERSION = "0.0.0"
AUTHOR = "JianzongWu"
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
URL = "https://github.com/KwaiVGI/VMoBA"
setup(
name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
classifiers=[
"Programming Language :: Python :: 3",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[]
)
@@ -1,97 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import pytest
import random
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
"""
Generates random data for testing the variable-length attention function.
"""
torch.manual_seed(42)
random.seed(42)
torch.cuda.manual_seed_all(42)
# Generate sequence lengths for each item in the batch
if batch_size > 1:
# Ensure sequence lengths are reasonably distributed
avg_seqlen = total_seqlen // batch_size
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
remaining_len = total_seqlen - sum(seqlens)
if remaining_len > 0:
seqlens.append(remaining_len)
else: # Adjust if sum exceeds total_seqlen
seqlens.append(avg_seqlen)
current_sum = sum(seqlens)
seqlens[-1] -= (current_sum - total_seqlen)
# Ensure all lengths are positive
seqlens = [max(1, s) for s in seqlens]
# Final adjustment to match total_seqlen
seqlens[-1] += total_seqlen - sum(seqlens)
else:
seqlens = [total_seqlen]
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
max_seqlen = max(seqlens) if seqlens else 0
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
return q, k, v, cu_seqlens, max_seqlen
@pytest.mark.parametrize("batch_size", [1, 2])
@pytest.mark.parametrize("total_seqlen", [512, 1024])
@pytest.mark.parametrize("num_heads", [8])
@pytest.mark.parametrize("head_dim", [64])
@pytest.mark.parametrize("moba_chunk_size", [64])
@pytest.mark.parametrize("moba_topk", [2, 4])
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
def test_moba_attn_varlen_forward(
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
):
"""
Tests the forward pass of moba_attn_varlen for basic correctness.
It checks output shape, dtype, and for the presence of NaNs/Infs.
"""
if dtype == torch.float32:
pytest.skip("float32 is not supported in flash attention")
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
batch_size, total_seqlen, num_heads, head_dim, dtype
)
# Ensure chunk size is not larger than the smallest sequence length
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
if moba_chunk_size > min_seqlen:
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
try:
output = moba_attn_varlen(
q=q,
k=k,
v=v,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=moba_chunk_size,
moba_topk=moba_topk,
select_mode=select_mode,
threshold_type=threshold_type,
simsum_threshold=0.5, # A reasonable default for threshold mode
)
except Exception as e:
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
# 1. Check output shape
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
# 2. Check output dtype
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
# 3. Check for NaNs or Infs in the output
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
-2
View File
@@ -1,2 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
-860
View File
@@ -1,860 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
import random
import time
import os
import torch
from typing import Tuple
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
from functools import lru_cache
from einops import rearrange
@lru_cache(maxsize=16)
def calc_chunks(cu_seqlen, moba_chunk_size):
"""
Calculate chunk boundaries.
For vision tasks we include all chunks (even the last one which might be shorter)
so that every chunk can be selected.
"""
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
cu_num_chunk = torch.ones(
batch_num_chunk.numel() + 1,
device=cu_seqlen.device,
dtype=batch_num_chunk.dtype,
)
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
num_chunk = cu_num_chunk[-1]
chunk_sizes = torch.full(
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
)
chunk_sizes[0] = 0
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
chunk_to_batch = torch.zeros(
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
)
chunk_to_batch[cu_num_chunk[1:-1]] = 1
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
# Do not filter out any chunk
filtered_chunk_indices = torch.arange(
num_chunk, device=cu_seqlen.device, dtype=torch.int32
)
num_filtered_chunk = num_chunk
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
# --- Threshold Selection Helper Functions ---
def _select_threshold_query_head(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects chunks for each <query, head> pair based on threshold.
Normalization and sorting happen along the chunk dimension (dim=0).
"""
C, H, S = gate.shape
eps = 1e-6
# LSE‐style normalization per <head, query> (across chunks)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
row_min = gate_min_val.amin(dim=0) # (H, S)
row_max = gate_masked.amax(dim=0) # (H, S)
denom = row_max - row_min
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
# 2) compute how much more normalized weight we need beyond self
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
# 3) zero out the self‐chunk in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0
# 4) sort the other chunks by descending norm, per <head,seq>
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
# 5) cumulative‑sum the sorted norms per <head,seq>
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
any_cond = cond.any(dim=0) # (H, S)
# Find the index of the first True value along dim 0. If none, use C-1.
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
# 7) build a mask in sorted order up to that cutoff
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
# 8) scatter it back to original chunk order
others_mask = torch.zeros_like(gate, dtype=torch.bool)
others_mask.scatter_(0, sorted_idx, sorted_mask)
# 9) finally, include every self‐chunk plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_block(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <query, head> pairs for each block based on threshold.
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
"""
C, H, S = gate.shape
HS = H * S
eps = 1e-6
# LSE‐style normalization per block (across heads and queries)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_denom = block_max - block_min
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
gate_norm = (gate - block_min) / block_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights *per block*
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort the other <head, seq> pairs by descending norm, per block
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
# 5) cumulative‑sum the sorted norms per block
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
any_cond = cond_flat.any(dim=1) # (C,)
# Find the index of the first True value along dim 1. If none, use HS-1.
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
# 7) build a mask in sorted order up to that cutoff per block
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
# 8) scatter it back to original <head, seq> order per block
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_overall(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query, head> triplets globally based on threshold.
Normalization and sorting happen across all valid entries.
"""
C, H, S = gate.shape
CHS = C * H * S
eps = 1e-6
# LSE‐style normalization globally across all valid entries
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
overall_max = gate_masked.max() # scalar
overall_min = gate_min_val.min() # scalar
overall_denom = overall_max - overall_min
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights globally
self_norm_sum_overall = self_norm_entries.sum() # scalar
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
total_norm_sum_overall = gate_norm.sum() # scalar
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort all other entries by descending norm, globally
others_flat = others_norm.flatten() # (C*H*S,)
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
# Only sort the valid 'other' entries
valid_others_indices = torch.where(valid_others_mask_flat)[0]
valid_others_values = others_flat[valid_others_indices]
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
# 5) cumulative‑sum the sorted valid 'other' norms globally
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
any_cond = cond_values.any() # scalar
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
cutoff_idx_in_sorted = torch.where(
any_cond,
cond_values.float().argmax(dim=0),
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
)
# 7) build a mask selecting the top-k others based on the cutoff
# Select the original indices corresponding to the top entries in the sorted list
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
# 8) create the mask in the original flat shape
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
others_mask_flat[selected_other_indices] = True
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_head_global(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query> globally for each head based on threshold.
"""
C, H, S = gate.shape
eps = 1e-6
# 1) LSE‐style normalization per head (across chunks and sequence dims)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
denom = max_per_head - min_per_head
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
gate_norm = (gate - min_per_head) / denom
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 2) sum normalized self‐chunk contributions per head
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
# 3) total normalized sum per head
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
# 4) how much more normalized weight needed per head
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
remain_ratio = torch.clamp(remain_ratio, min=0.0)
# 5) zero out self‐chunk entries to focus on "others"
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
# 6) flatten chunk and sequence dims, per head
CS = C * S
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
# 7) vectorized selection of “others” per head
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
has_cutoff = cond.any(dim=1) # (H,)
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
# 8) reshape selection mask back to (C, H, S)
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
# 9) include self‐chunks plus selected others, and obey valid mask
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
return final_gate_mask
class MixedAttention(torch.autograd.Function):
@staticmethod
def forward(
ctx,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
):
ctx.max_seqlen = max_seqlen
ctx.moba_chunk_size = moba_chunk_size
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
# Non-causal self-attention branch
# return out, softmax_lse, S_dmask, rng_state
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=q,
k=k,
v=v,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
# MOBA attention branch (non-causal)
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
output_2d = output.view(-1, q.shape[2])
max_lse_1d = self_attn_lse_sh.view(-1)
max_lse_1d = max_lse_1d.index_reduce(
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
)
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
moba_attn_lse = (
moba_attn_lse.view(-1)
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
.reshape_as(moba_attn_lse)
)
mixed_attn_se_sh = self_attn_lse_sh.exp()
moba_attn_se = moba_attn_lse.exp()
mixed_attn_se_sh.view(-1).index_add_(
0, moba_q_sh_indices, moba_attn_se.view(-1)
)
mixed_attn_lse_sh = mixed_attn_se_sh.log()
# Combine self-attention output
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
output_2d += self_attn_out_sh.reshape_as(output_2d)
# Combine MOBA attention output
mixed_attn_lse = (
mixed_attn_lse_sh.view(-1)
.index_select(0, moba_q_sh_indices)
.view_as(moba_attn_lse)
)
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
output = output.to(q.dtype)
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
ctx.save_for_backward(
output,
mixed_attn_lse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
)
return output
@staticmethod
def backward(ctx, d_output):
max_seqlen = ctx.max_seqlen
moba_chunk_size = ctx.moba_chunk_size
softmax_scale = ctx.softmax_scale
(
output,
mixed_attn_vlse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
) = ctx.saved_tensors
d_output = d_output.contiguous()
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
_ = _flash_attn_varlen_backward(
dout=d_output,
q=q,
k=k,
v=v,
out=output,
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
dq=dq,
dk=dk,
dv=dv,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
headdim = q.shape[-1]
d_moba_output = (
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
moba_output = (
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
mixed_attn_vlse = (
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
)
dmq = torch.empty_like(moba_q)
dmkv = torch.empty_like(moba_kv)
_ = _flash_attn_varlen_backward(
dout=d_moba_output,
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
out=moba_output,
softmax_lse=mixed_attn_vlse,
dq=dmq,
dk=dmkv[:,0],
dv=dmkv[:,1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
def moba_attn_varlen(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens: torch.Tensor,
max_seqlen: int,
moba_chunk_size: int,
moba_topk: int,
select_mode: str = 'threshold', # "topk" or "threshold"
simsum_threshold: float = 0.25,
threshold_type: str = 'query_head',
) -> torch.Tensor:
"""
Accelerated MOBA attention for vision tasks with proper LSE normalization.
This version:
- Splits KV into chunks.
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
by amplifying the diagonal (self-chunk) logits.
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
reduction so that attending to each query over the selected chunks is equivalent
to the original algorithm.
"""
# Stack keys and values.
kv = torch.stack((k, v), dim=1)
seqlen, num_head, head_dim = q.shape
# Compute chunk boundaries.
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
cu_seqlens, moba_chunk_size
)
self_attn_cu_seqlen = cu_chunk
# Update top-k selection to include the self chunk.
moba_topk = min(moba_topk, num_filtered_chunk)
# --- Build filtered KV from chunks ---
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
max_chunk_len = int(chunk_lengths.max().item())
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
indices = chunk_starts.unsqueeze(1) + range_tensor
indices = torch.clamp(indices, max=kv.shape[0] - 1)
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
# Compute key_gate_weight over valid tokens.
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
key_sum = (key_values * valid_mask_exp).sum(dim=1)
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
# Compute gate logits between key_gate_weight and queries.
q_float = q.float()
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
# Amplify the diagonal (self chunk) contributions.
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
amplification_factor = 1e9 # Example factor; adjust as needed.
origin_gate = gate.clone()
gate = gate.clone()
if select_mode == "topk":
gate[gate_self_chunk_mask] += amplification_factor
# Exclude positions that are outside the valid batch boundaries.
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
if select_mode == 'topk':
# We amplify self‐chunk in gate already, so self entries will rank highest.
valid_gate_mask = gate != -float("inf")
if threshold_type == 'query_head':
# === per‐<head,seq> top-k across chunks (original behavior) ===
# gate: (C, H, S)
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
gate_idx_mask.scatter_(0, gate_topk_idx, True)
gate_mask = valid_gate_mask & gate_idx_mask
elif threshold_type == 'overall':
# === global top-k across all (chunk, head, seq) entries ===
C, H, S = gate.shape
flat_gate = gate.flatten()
flat_mask = valid_gate_mask.flatten()
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
# pick topk global entries
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
others_mask_flat[idx] = True
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
elif threshold_type == 'head_global':
# per-head top-k across all chunks and sequence positions
C, H, S = gate.shape
CS = C * S
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
# pick top-k indices per head
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
gate_idx_flat.scatter_(1, topk_idx, True)
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
else:
raise ValueError(
f"Invalid threshold_type for topk: {threshold_type}. "
"Choose 'query_head', 'block', or 'overall'."
)
elif select_mode == 'threshold':
# Delegate to the specific thresholding function
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
if threshold_type == 'query_head':
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'block':
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'overall':
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'head_global':
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
else:
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
else:
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
# eliminate self_chunk in MoBA branch
gate_mask = gate_mask & ~gate_self_chunk_mask
# if gate_mask is all false, perform flash_attn instead
if gate_mask.sum() == 0:
return flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
)
# Determine which query positions are selected.
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
# Build cumulative sequence lengths for the selected queries.
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
q_zero_mask = moba_seqlen_q == 0
valid_expert_mask = ~q_zero_mask
if q_zero_mask.sum() > 0:
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
moba_cu_seqlen_q = torch.cat(
(
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
moba_seqlen_q.cumsum(dim=0),
),
dim=0,
).to(torch.int32)
# Rearrange gathered KV for the MOBA branch.
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
if q_zero_mask.sum() > 0:
experts_tensor = experts_tensor[valid_expert_mask]
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
mask = seq_range < valid_expert_lengths.unsqueeze(1)
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
moba_cu_seqlen_kv = torch.cat(
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
valid_expert_lengths.cumsum(dim=0)],
dim=0,
).to(torch.int32)
assert (
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
return MixedAttention.apply(
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
)
def process_moba_input(
x,
patch_resolution,
chunk_size,
):
"""
Process inputs for the attention function.
Args:
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
Returns:
torch.Tensor: Processed input tensor.
"""
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
else:
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
if len(chunk_size) == 2:
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
elif len(chunk_size) == 3:
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
else:
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
return x, moba_chunk_size
def process_moba_output(
x,
patch_resolution,
chunk_size,
):
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
pass
elif len(chunk_size) == 2:
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
elif len(chunk_size) == 3:
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
return x
# TEST
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
random.seed(0)
torch.manual_seed(0)
torch.cuda.manual_seed(0)
device = torch.cuda.current_device()
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
max_seqlen = q.shape[1]
q = rearrange(q, "b s ... -> (b s) ...")
k = rearrange(k, "b s ... -> (b s) ...")
v = rearrange(v, "b s ... -> (b s) ...")
return q, k, v, cu_seqlens, max_seqlen
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
"""Speed test comparing flash_attn vs moba_attention"""
# Get data
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
vo_grad = torch.randn_like(q)
# Warmup
warmup_iters = 3
perf_test_iters = 10
# Warmup
for _ in range(warmup_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
start_flash = time.perf_counter()
for _ in range(perf_test_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
# Warmup
for _ in range(warmup_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
start_moba = time.perf_counter()
for _ in range(perf_test_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
print(f"Speedup: {time_flash / time_moba:.2f}x")
if __name__ == "__main__":
"""
CUDA_VISIBLE_DEVICES=1 \
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
"""
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
@@ -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);
@@ -624,6 +625,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);
+21 -45
View File
@@ -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
+21 -45
View File
@@ -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
+6 -6
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.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
-72
View File
@@ -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
+2 -1
View File
@@ -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)
+3 -24
View File
@@ -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,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:
+2 -1
View File
@@ -1,6 +1,6 @@
# Welcome to FastVideo
:::{figure} ../../assets/logos/logo.svg
:::{figure} ../../assets/logo.png
:align: center
:alt: FastVideo
:class: no-scaled-link
@@ -101,6 +101,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 -50
View File
@@ -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
@@ -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
@@ -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:
@@ -4,7 +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
```
### Data-free Distillation
@@ -4,7 +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:
@@ -98,7 +98,6 @@ dmd_args=(
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port $MASTER_PORT \
fastvideo/training/wan_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
@@ -1,112 +0,0 @@
#!/bin/bash
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export MASTER_PORT=29501
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# Configs
NUM_GPUS=1
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir="checkpoints/wan_t2v_finetune"
--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"
--training_state_checkpointing_steps=500
--weight_only_checkpointing_steps=500
--lora_rank 32
--lora_training True
)
# Parallel arguments
parallel_args=(
--num_gpus 1
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-4
--mixed_precision="bf16"
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
)
# 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.5
--VSA_sparsity 0.8
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port $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,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,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()
+6 -6
View File
@@ -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">
@@ -14,10 +14,9 @@ def main():
vae_cpu_offload=True,
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",
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1000/transformer",
lora_nickname="crush_smol"
)
generator.unmerge_lora_weights()
kwargs = {
"height": 480,
"width": 832,
@@ -1,47 +0,0 @@
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
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.
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.
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
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.
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
The video shows a large 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.
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
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.
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
@@ -1,93 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 77
--warp_denoising_step
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,24 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory"
@@ -1,40 +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": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -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:
@@ -1,24 +0,0 @@
#!/bin/bash
GPU_NUM=2 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
torchrun --nproc_per_node=$GPU_NUM \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type i2v \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 77 \
--preprocess.train_fps 16 \
--preprocess.samples_per_file 8 \
--preprocess.flush_frequency 8 \
--preprocess.video_length_tolerance_range 5
@@ -7,7 +7,7 @@ export WANDB_MODE=online
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
NUM_GPUS=2
# export CUDA_VISIBLE_DEVICES=4,5
@@ -76,7 +76,6 @@ miscellaneous_args=(
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--resume_from_checkpoint "checkpoints/wan_t2v_finetune_lora/checkpoint-160"
)
torchrun \
@@ -1,20 +1,19 @@
#!/bin/bash
GPU_NUM=2 # 2,4,8
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATASET_PATH="data/crush-smol/"
MODEL_TYPE="wan"
DATASET_PATH="data/crush-smol-test/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
fastvideo/pipelines/preprocess/v1_preprocessing_new.py \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.video_loader_type torchvision \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.preprocess_video_batch_size 4 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
@@ -28,4 +28,4 @@
"num_frames": 77
}
]
}
}
-214
View File
@@ -1,214 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import re
from dataclasses import dataclass
import torch
from einops import rearrange
from flash_attn.bert_padding import pad_input
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class VMOBAAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_name() -> str:
return "VMOBA_ATTN"
@staticmethod
def get_impl_cls() -> type["VMOBAAttentionImpl"]:
return VMOBAAttentionImpl
@staticmethod
def get_metadata_cls() -> type["VideoMobaAttentionMetadata"]:
return VideoMobaAttentionMetadata
@staticmethod
def get_builder_cls() -> type["VideoMobaAttentionMetadataBuilder"]:
return VideoMobaAttentionMetadataBuilder
@dataclass
class VideoMobaAttentionMetadata(AttentionMetadata):
current_timestep: int
temporal_chunk_size: int
temporal_topk: int
spatial_chunk_size: tuple[int, int]
spatial_topk: int
st_chunk_size: tuple[int, int, int]
st_topk: int
moba_select_mode: str
moba_threshold: float
moba_threshold_type: str
patch_resolution: list[int]
first_full_step: int = 12
first_full_layer: int = 0
# temporal_layer -> spatial_layer -> st_layer
temporal_layer: int = 1
spatial_layer: int = 1
st_layer: int = 1
class VideoMobaAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
temporal_chunk_size: int,
temporal_topk: int,
spatial_chunk_size: tuple[int, int],
spatial_topk: int,
st_chunk_size: tuple[int, int, int],
st_topk: int,
moba_select_mode: str = 'threshold',
moba_threshold: float = 0.25,
moba_threshold_type: str = 'query_head',
device: torch.device = None,
first_full_layer: int = 0,
first_full_step: int = 12,
temporal_layer: int = 1,
spatial_layer: int = 1,
st_layer: int = 1,
**kwargs,
) -> VideoMobaAttentionMetadata:
if device is None:
device = torch.device("cpu")
assert raw_latent_shape[0] % patch_size[0] == 0 and raw_latent_shape[
1] % patch_size[1] == 0 and raw_latent_shape[2] % patch_size[
2] == 0, f"spatial patch_resolution {raw_latent_shape} should be divisible by patch_size {patch_size}"
patch_resolution = [
t // pt for t, pt in zip(raw_latent_shape, patch_size, strict=False)
]
return VideoMobaAttentionMetadata(
current_timestep=current_timestep,
temporal_chunk_size=temporal_chunk_size,
temporal_topk=temporal_topk,
spatial_chunk_size=spatial_chunk_size,
spatial_topk=spatial_topk,
st_chunk_size=st_chunk_size,
st_topk=st_topk,
moba_select_mode=moba_select_mode,
moba_threshold=moba_threshold,
moba_threshold_type=moba_threshold_type,
patch_resolution=patch_resolution,
first_full_layer=first_full_layer,
first_full_step=first_full_step,
temporal_layer=temporal_layer,
spatial_layer=spatial_layer,
st_layer=st_layer,
)
class VMOBAAttentionImpl(AttentionImpl):
def __init__(self,
num_heads,
head_size,
softmax_scale,
causal=False,
num_kv_heads=None,
prefix="",
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
if not match:
raise ValueError(f"Invalid prefix: {prefix}")
return int(match.group(1))
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
"""
query: [B, L, H, D]
key: [B, L, H, D]
value: [B, L, H, D]
attn_metadata: AttentionMetadata
"""
batch_size, sequence_length, num_heads, head_dim = query.shape
# select chunk type according to layer idx:
loop_layer_num = attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer
moba_layer = self.layer_idx - attn_metadata.first_full_layer
if moba_layer % loop_layer_num < attn_metadata.temporal_layer:
moba_chunk_size = attn_metadata.temporal_chunk_size
moba_topk = attn_metadata.temporal_topk
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer:
moba_chunk_size = attn_metadata.spatial_chunk_size
moba_topk = attn_metadata.spatial_topk
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer:
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
# torch.distributed.breakpoint()
query, chunk_size = process_moba_input(query,
attn_metadata.patch_resolution,
moba_chunk_size)
key, chunk_size = process_moba_input(key,
attn_metadata.patch_resolution,
moba_chunk_size)
value, chunk_size = process_moba_input(value,
attn_metadata.patch_resolution,
moba_chunk_size)
max_seqlen = query.shape[1]
indices_q = torch.arange(0,
query.shape[0] * query.shape[1],
device=query.device)
cu_seqlens = torch.arange(0,
query.shape[0] * query.shape[1] + 1,
query.shape[1],
dtype=torch.int32,
device=query.device)
query = rearrange(query, "b s ... -> (b s) ...")
key = rearrange(key, "b s ... -> (b s) ...")
value = rearrange(value, "b s ... -> (b s) ...")
# current_timestep=attn_metadata.current_timestep
hidden_states = moba_attn_varlen(
query,
key,
value,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=chunk_size,
moba_topk=moba_topk,
select_mode=attn_metadata.moba_select_mode,
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = process_moba_output(hidden_states,
attn_metadata.patch_resolution,
moba_chunk_size)
return hidden_states
@@ -1,16 +0,0 @@
{
"temporal_chunk_size": 2,
"temporal_topk": 2,
"spatial_chunk_size": [4, 13],
"spatial_topk": 6,
"st_chunk_size": [4, 4, 13],
"st_topk": 18,
"moba_select_mode": "topk",
"moba_threshold": 0.25,
"moba_threshold_type": "query_head",
"first_full_layer": 0,
"first_full_step": 12,
"temporal_layer": 1,
"spatial_layer": 1,
"st_layer": 1
}
@@ -1,16 +0,0 @@
{
"temporal_chunk_size": 2,
"temporal_topk": 3,
"spatial_chunk_size": [3, 4],
"spatial_topk": 20,
"st_chunk_size": [4, 6, 4],
"st_topk": 15,
"moba_select_mode": "threshold",
"moba_threshold": 0.25,
"moba_threshold_type": "query_head",
"first_full_layer": 0,
"first_full_step": 12,
"temporal_layer": 1,
"spatial_layer": 1,
"st_layer": 1
}
-79
View File
@@ -1,59 +1,9 @@
import dataclasses
from enum import Enum
from typing import Any, Optional
from fastvideo.configs.utils import update_config_from_args
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
class DatasetType(str, Enum):
"""
Enumeration for different dataset types.
"""
HF = "hf"
MERGED = "merged"
@classmethod
def from_string(cls, value: str) -> "DatasetType":
"""Convert string to DatasetType enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid dataset type: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [dataset_type.value for dataset_type in cls]
class VideoLoaderType(str, Enum):
"""
Enumeration for different video loaders.
"""
TORCHCODEC = "torchcodec"
TORCHVISION = "torchvision"
@classmethod
def from_string(cls, value: str) -> "VideoLoaderType":
"""Convert string to VideoLoader enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid video loader: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [video_loader.value for video_loader in cls]
@dataclasses.dataclass
class PreprocessConfig:
@@ -62,7 +12,6 @@ class PreprocessConfig:
# Model and dataset configuration
model_path: str = ""
dataset_path: str = ""
dataset_type: DatasetType = DatasetType.HF
dataset_output_dir: str = "./output"
# Dataloader configuration
@@ -74,7 +23,6 @@ class PreprocessConfig:
flush_frequency: int = 256
# Video processing parameters
video_loader_type: VideoLoaderType = VideoLoaderType.TORCHCODEC
max_height: int = 480
max_width: int = 848
num_frames: int = 163
@@ -87,9 +35,6 @@ class PreprocessConfig:
# Model configuration
training_cfg_rate: float = 0.0
# framework configuration
seed: int = 42
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "preprocess") -> FlexibleArgumentParser:
@@ -107,12 +52,6 @@ class PreprocessConfig:
type=str,
default=PreprocessConfig.dataset_path,
help="Path to the dataset directory for preprocessing")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-type",
type=str,
choices=DatasetType.choices(),
default=PreprocessConfig.dataset_type.value,
help="Type of the dataset")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-output-dir",
type=str,
@@ -144,12 +83,6 @@ class PreprocessConfig:
help="How often to save to parquet files")
# Video processing parameters
preprocess_args.add_argument(
f"--{prefix_with_dot}video-loader-type",
type=str,
choices=VideoLoaderType.choices(),
default=PreprocessConfig.video_loader_type.value,
help="Type of the video loader")
preprocess_args.add_argument(f"--{prefix_with_dot}max-height",
type=int,
default=PreprocessConfig.max_height,
@@ -190,10 +123,6 @@ class PreprocessConfig:
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
type=int,
default=PreprocessConfig.seed,
help="Seed for random number generator")
return parser
@@ -201,14 +130,6 @@ class PreprocessConfig:
def from_kwargs(cls, kwargs: dict[str,
Any]) -> Optional["PreprocessConfig"]:
"""Create PreprocessConfig from keyword arguments."""
if 'dataset_type' in kwargs and isinstance(kwargs['dataset_type'], str):
kwargs['dataset_type'] = DatasetType.from_string(
kwargs['dataset_type'])
if 'video_loader_type' in kwargs and isinstance(
kwargs['video_loader_type'], str):
kwargs['video_loader_type'] = VideoLoaderType.from_string(
kwargs['video_loader_type'])
preprocess_config = cls()
if not update_config_from_args(
preprocess_config, kwargs, prefix="preprocess", pop_args=True):
+3 -7
View File
@@ -15,13 +15,9 @@ class DiTArchConfig(ArchConfig):
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
)
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
hidden_size: int = 0
num_attention_heads: int = 0
@@ -92,12 +92,6 @@ class WanVideoArchConfig(DiTArchConfig):
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
+2 -3
View File
@@ -4,13 +4,12 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
"get_pipeline_config_cls_from_name"
]
-3
View File
@@ -85,9 +85,6 @@ class PipelineConfig:
# DMD parameters
dmd_denoising_steps: list[int] | None = field(default=None)
# Wan2.2 TI2V parameters
ti2v_task: bool = False
# Compilation
# enable_torch_compile: bool = False
+7 -14
View File
@@ -7,14 +7,10 @@ from collections.abc import Callable
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
# isort: off
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig)
# isort: on
from fastvideo.configs.pipelines.wan import (FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -35,10 +31,9 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": FastWan2_2_TI2V_5B_Config,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
# Add other specific weight variants
}
@@ -48,7 +43,6 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -61,7 +55,6 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
+11 -23
View File
@@ -12,13 +12,13 @@ from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.Tensor = outputs.last_hidden_state
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
mask: torch.tensor = outputs.attention_mask
hidden_state: torch.tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
prompt_embeds_tensor: torch.Tensor = torch.stack([
prompt_embeds_tensor: torch.tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in prompt_embeds
],
@@ -39,12 +39,12 @@ class WanT2V480PConfig(PipelineConfig):
vae_sp: bool = False
# Denoising stage
flow_shift: float | None = 3.0
flow_shift: int = 3
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
@@ -68,7 +68,7 @@ class WanT2V720PConfig(WanT2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: float | None = 5.0
flow_shift: int = 5
@dataclass
@@ -94,7 +94,7 @@ class WanI2V720PConfig(WanI2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: float | None = 5.0
flow_shift: int = 5
@dataclass
@@ -104,7 +104,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: float | None = 8.0
flow_shift: int = 8
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
@@ -115,7 +115,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
@dataclass
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
flow_shift: float | None = 5.0
flow_shift: int = 5
ti2v_task: bool = True
def __post_init__(self) -> None:
@@ -125,7 +125,7 @@ class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
@dataclass
class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
flow_shift: float | None = 5.0
flow_shift: int = 5
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
@@ -138,15 +138,3 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
@dataclass
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
pass
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
is_causal: bool = True
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
-21
View File
@@ -47,8 +47,6 @@ class SamplingParam:
# Misc
save_video: bool = True
return_frames: bool = False
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
@@ -193,25 +191,6 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--moba-config-path",
type=str,
default=None,
help=
"Path to a JSON file containing V-MoBA specific configurations.",
)
parser.add_argument(
"--return-trajectory-latents",
action="store_true",
default=SamplingParam.return_trajectory_latents,
help="Whether to return the trajectory",
)
parser.add_argument(
"--return-trajectory-decoded",
action="store_true",
default=SamplingParam.return_trajectory_decoded,
help="Whether to return the decoded trajectory",
)
return parser
+3 -20
View File
@@ -10,7 +10,6 @@ from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
FastWanT2V480PConfig,
Wan2_1_Fun_1_3B_InP_SamplingParam,
Wan2_2_I2V_A14B_SamplingParam,
Wan2_2_T2V_A14B_SamplingParam,
Wan2_2_TI2V_5B_SamplingParam,
@@ -18,7 +17,6 @@ from fastvideo.configs.sample.wan import (
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
SelfForcingWanT2V480PConfig,
)
# isort: on
from fastvideo.logger import init_logger
@@ -30,31 +28,16 @@ logger = init_logger(__name__)
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
# Wan2.2
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
# Add other specific weight variants
}
-23
View File
@@ -107,21 +107,6 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
fps: int = 16
# =============================================
# ============= Wan2.1 Fun Models =============
# =============================================
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
# =============================================
# ============= Wan2.2 TI2V Models =============
# =============================================
@@ -156,11 +141,3 @@ class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale_2: float = 3.5
num_inference_steps: int = 40
fps: int = 16
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
pass
+2 -8
View File
@@ -4,7 +4,7 @@ from torchvision.transforms import Lambda
from fastvideo.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.validation_dataset import ValidationDataset
@@ -37,15 +37,9 @@ def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
seed=args.seed)
def gettextdataset(args) -> TextDataset:
return TextDataset(data_merge_path=args.data_merge_path,
args=args,
seed=args.seed)
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
"VideoCaptionMergedDataset"
]
-264
View File
@@ -1,264 +0,0 @@
"""
Utilities for converting preprocessing records (dicts) into Arrow tables and
writing Parquet datasets in fixed-size chunks.
This module centralizes table construction and Parquet file writing so
pipelines only need to define their PyArrow schema and produce per-sample
record dictionaries.
Key APIs:
- records_to_table(records, schema): Safely convert a list of dictionaries into
a pa.Table, casting to the provided schema.
- ParquetDatasetWriter: Buffer tables and flush to a directory as multiple
Parquet files with a fixed number of rows per file. Uses temporary files and
atomic rename to avoid partially written outputs.
"""
from __future__ import annotations
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any
import pyarrow as pa
import pyarrow.parquet as pq
def records_to_table(records: list[dict[str, Any]], schema: pa.Schema) -> pa.Table:
"""Build a PyArrow table from Python record dicts using an explicit schema.
Arrow will cast values to the target schema when possible (e.g., promoting
Python ints/floats to pa.int64/pa.float64), eliminating hand-written per-
field array construction.
Args:
records: List of dictionaries, each representing one row. Keys must
match schema field names.
schema: Target PyArrow schema. Controls field names and types.
Returns:
pa.Table: In-memory table matching the provided schema. If ``records``
is empty, returns an empty table with the given schema.
"""
if not records:
return pa.table({}, schema=schema)
return pa.Table.from_pylist(records, schema=schema)
class ParquetDatasetWriter:
"""Accumulate tables and flush them to a Parquet directory in fixed-size chunks.
Behavior:
- Writes files under worker-specific subdirectories for parallelism.
- Uses temporary files and atomic rename to avoid partial files being left
behind on failure.
- Only full chunks of ``samples_per_file`` rows are written on each flush;
any remainder rows are re-buffered for the next flush.
Note:
- Instances are not meant to be shared across processes. Create one writer
per process if using multiprocessing.
"""
def __init__(self, out_dir: str, samples_per_file: int, compression: str = "zstd") -> None:
"""Initialize the dataset writer.
Args:
out_dir: Output directory where Parquet files will be written.
samples_per_file: Fixed number of rows per Parquet file.
compression: Compression codec passed to ``pyarrow.parquet.write_table``
(e.g., ``"zstd"``, ``"snappy"``, ``"gzip"``).
"""
self.out_dir = out_dir
self.samples_per_file = max(int(samples_per_file), 1)
self.compression = compression
os.makedirs(self.out_dir, exist_ok=True)
self._tables: list[pa.Table] = []
def append_table(self, table: pa.Table) -> None:
"""Append a non-empty table to the internal buffer.
Args:
table: A ``pa.Table`` to buffer. Empty or ``None`` tables are ignored.
"""
if table is None or len(table) == 0:
return
self._tables.append(table)
def _combine(self) -> pa.Table | None:
"""Combine all buffered tables into a single table, if any.
Returns:
A concatenated table, a single table if only one was buffered, or
``None`` if no tables are buffered.
"""
if not self._tables:
return None
if len(self._tables) == 1:
return self._tables[0]
return pa.concat_tables(self._tables, promote_options='none')
def flush(self, num_workers: int | None = None, write_remainder: bool = False) -> int:
"""Write accumulated tables to disk and clear the written portion.
Only complete chunks of size ``samples_per_file`` are written. Any
remainder rows are kept buffered for the next flush.
Args:
num_workers: Optional override for the number of parallel workers
used to write chunks. Defaults to ``min(cpu_count, chunks)``.
write_remainder: If True, also write any leftover rows (< samples_per_file)
as a final small Parquet file (useful for the last flush at the
end of preprocessing).
Returns:
int: Number of rows successfully written in this flush call.
"""
combined = self._combine()
self._tables = []
if combined is None or len(combined) == 0:
return 0
num_samples = len(combined)
total_chunks = num_samples // self.samples_per_file
if total_chunks == 0:
if not write_remainder:
# Not enough to form a full chunk; keep buffered for next round
# Re-buffer and return 0 written
self._tables = [combined]
return 0
# Last flush: write the small remainder as a final file in worker_0
worker_dir = os.path.join(self.out_dir, "worker_0")
os.makedirs(worker_dir, exist_ok=True)
# Determine next index
num_parquets = 0
for _, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
chunk_path = os.path.join(worker_dir, f"data_chunk_{num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
pq.write_table(combined, temp_path, compression=self.compression)
if os.path.exists(chunk_path):
os.remove(chunk_path)
os.rename(temp_path, chunk_path)
return num_samples
# Only write full chunks; keep remainder for next flush
written_rows = total_chunks * self.samples_per_file
remainder = num_samples - written_rows
table_to_write = combined.slice(0, written_rows)
remainder_table = combined.slice(written_rows, remainder) if remainder > 0 else None
if remainder_table is not None and len(remainder_table) > 0:
if write_remainder:
# Write the remainder as a final small file (worker_0)
worker_dir = os.path.join(self.out_dir, "worker_0")
os.makedirs(worker_dir, exist_ok=True)
num_parquets = 0
for _, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
remainder_path = os.path.join(worker_dir,
f"data_chunk_{num_parquets}.parquet")
temp_path = remainder_path + '.tmp'
pq.write_table(remainder_table,
temp_path,
compression=self.compression)
if os.path.exists(remainder_path):
os.remove(remainder_path)
os.rename(temp_path, remainder_path)
else:
self._tables = [remainder_table]
# Parallel write by chunk ranges
if num_workers is None:
num_workers = min(multiprocessing.cpu_count(), max(total_chunks, 1))
num_workers = max(int(num_workers), 1)
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
work_ranges: list[tuple[int, int, pa.Table, int, str, int, str]] = []
for worker_id in range(num_workers):
start_chunk = worker_id * chunks_per_worker
end_chunk = min((worker_id + 1) * chunks_per_worker, total_chunks)
if start_chunk < end_chunk:
work_ranges.append(
(
start_chunk,
end_chunk,
table_to_write,
worker_id,
self.out_dir,
self.samples_per_file,
self.compression,
)
)
written_total = 0
if len(work_ranges) == 1:
written_total += _process_chunk_range(work_ranges[0])
return written_total
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = [executor.submit(_process_chunk_range, args) for args in work_ranges]
for f in futures:
written_total += f.result()
return written_total + (len(remainder_table) if write_remainder and remainder_table is not None else 0)
def _process_chunk_range(args: Any) -> int:
"""Worker function to write a contiguous range of chunk files.
Args:
args: Tuple containing
- start_chunk (int): inclusive start chunk index
- end_chunk (int): exclusive end chunk index
- table (pa.Table): concatenated table containing all rows to write
- worker_id (int): numeric worker identifier
- output_dir (str): base output directory
- samples_per_file (int): rows per chunk file
- compression (str): compression codec for Parquet
Returns:
int: Total number of rows written by this worker.
"""
start_chunk, end_chunk, table, worker_id, output_dir, samples_per_file, compression = args
total_written = 0
num_samples = len(table)
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Offset to continue numbering if files exist
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_chunk, end_chunk):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
if end_sample <= start_sample:
continue
chunk = table.slice(start_sample, end_sample - start_sample)
chunk_path = os.path.join(worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
pq.write_table(chunk, temp_path, compression=compression)
if os.path.exists(chunk_path):
os.remove(chunk_path)
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception:
if os.path.exists(temp_path):
os.remove(temp_path)
raise
return total_written
-24
View File
@@ -50,7 +50,6 @@ pyarrow_schema_i2v = pa.schema([
pa.field("fps", pa.float64()),
])
pyarrow_schema_t2v = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
@@ -79,26 +78,3 @@ pyarrow_schema_t2v = pa.schema([
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# --- ODE Trajectory ---
pa.field("trajectory_latents_bytes", pa.binary()),
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
pa.field("trajectory_latents_dtype", pa.string()),
pa.field("trajectory_timesteps_bytes", pa.binary()),
pa.field("trajectory_timesteps_shape", pa.list_(pa.int64())),
pa.field("trajectory_timesteps_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # Always 'text' for text-only
])
-131
View File
@@ -628,134 +628,3 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
class TextDataset(torch.utils.data.IterableDataset,
torch.distributed.checkpoint.stateful.Stateful):
"""
Text-only dataset for processing prompts from a simple text file.
Assumes that data_merge_path is a text file with one prompt per line:
A cat playing with a ball
A dog running in the park
A person cooking dinner
...
This dataset processes text data through text encoding stages only.
"""
def __init__(self,
data_merge_path: str,
args,
start_idx: int = 0,
seed: int = 42):
self.data_merge_path = data_merge_path
self.start_idx = start_idx
self.args = args
self.seed = seed
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
# Initialize text encoding stage
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=getattr(args, 'training_cfg_rate', 0.0),
seed=self.seed)
# Process text data
self.processed_batches = self._process_text_data()
def _load_text_data(self) -> list[str]:
"""Load text prompts from file."""
prompts = []
with open(self.data_merge_path, 'r', encoding='utf-8') as f:
for line in f:
line = line.strip()
if line: # Skip empty lines
prompts.append(line)
logger.info(f"Loaded {len(prompts)} text prompts from {self.data_merge_path}")
return prompts
def _process_text_data(self) -> list[PreprocessBatch]:
"""Process the text prompts through text encoding stage."""
raw_prompts = self._load_text_data()
processed_batches = []
for idx, prompt in enumerate(raw_prompts):
# Create a text-only batch with dummy path
batch = PreprocessBatch(
path=f"text_prompt_{idx}",
cap=[prompt], # TextEncodingStage expects a list
resolution=None,
fps=None,
duration=None,
num_frames=0,
sample_frame_index=None,
sample_num_frames=0
)
processed_batches.append(batch)
logger.info(f"Processed {len(processed_batches)} text batches")
return processed_batches
def __iter__(self):
"""Iterator for the dataset."""
# Set up distributed sampling if needed
if torch.distributed.is_available() and torch.distributed.is_initialized():
rank = torch.distributed.get_rank()
world_size = torch.distributed.get_world_size()
else:
rank = 0
world_size = 1
# Calculate chunk for this rank
total_items = len(self.processed_batches)
items_per_rank = math.ceil(total_items / world_size)
start_idx = rank * items_per_rank + self.start_idx
end_idx = min(start_idx + items_per_rank, total_items)
# Yield items for this rank
for idx in range(start_idx, end_idx):
if idx < len(self.processed_batches):
yield self._get_item(idx)
def _get_item(self, idx: int) -> dict:
"""Get a single processed text item."""
batch = self.processed_batches[idx]
# Apply text encoding stage
batch = self.text_encoding_stage.process(batch)
# Build result dictionary for text-only processing with required schema fields
result = {
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
# Required schema fields for ODE trajectory processing
"id": f"text_{idx}",
"file_name": batch.path,
"caption": batch.text,
"media_type": "text",
"width": 1,
"height": 1,
"num_frames": 0,
"duration_sec": 0.0,
"fps": 0.0,
}
return result
def state_dict(self) -> dict[str, Any]:
"""Return state dict for checkpointing."""
return {"processed_batches": self.processed_batches}
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
+1 -1
View File
@@ -5,7 +5,7 @@ import numpy as np
import torch
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
-13
View File
@@ -344,9 +344,6 @@ class VideoGenerator:
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
def set_lora_adapter(self,
@@ -354,16 +351,6 @@ class VideoGenerator:
lora_path: str | None = None) -> None:
self.executor.set_lora_adapter(lora_nickname, lora_path)
def unmerge_lora_weights(self) -> None:
"""
Use unmerged weights for inference to produce videos that align with
validation videos generated during training.
"""
self.executor.unmerge_lora_weights()
def merge_lora_weights(self) -> None:
self.executor.merge_lora_weights()
def shutdown(self):
"""
Shutdown the video generator.
+2 -25
View File
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py
"""The arguments of FastVideo Inference."""
import argparse
import dataclasses
import json
from contextlib import contextmanager
from dataclasses import field
from enum import Enum
@@ -119,7 +119,7 @@ class FastVideoArgs:
output_type: str = "pil"
# CPU offload parameters
dit_cpu_offload: bool = True
dit_cpu_offload: bool = False
use_fsdp_inference: bool = True
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
@@ -139,10 +139,6 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
# V-MoBA parameters
moba_config_path: str | None = None
moba_config: dict[str, Any] = field(default_factory=dict)
# Master port for distributed training/inference
master_port: int | None = None
@@ -170,16 +166,6 @@ class FastVideoArgs:
return not self.inference_mode
def __post_init__(self):
if self.moba_config_path:
try:
with open(self.moba_config_path) as f:
self.moba_config = json.load(f)
logger.info("Loaded V-MoBA config from %s",
self.moba_config_path)
except (FileNotFoundError, json.JSONDecodeError) as e:
logger.error("Failed to load V-MoBA config from %s: %s",
self.moba_config_path, e)
raise
self.check_fastvideo_args()
@staticmethod
@@ -999,15 +985,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
# V-MoBA parameters
parser.add_argument(
"--moba-config-path",
type=str,
default=None,
help=
"Path to a JSON file containing V-MoBA specific configurations.",
)
# Distillation arguments
parser.add_argument("--generator-update-interval",
type=int,
+6 -60
View File
@@ -100,16 +100,7 @@ class ScaleResidual(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
"""Apply gated residual connection."""
# x.shape: [batch_size, seq_len, inner_dim]
if gate.dim() == 4:
# gate.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
return residual + (x.unflatten(
dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
else:
# gate.shape: [batch_size, 1, inner_dim]
return residual + x * gate
return residual + x * gate
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
@@ -168,7 +159,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor | int, shift: torch.Tensor,
gate: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
@@ -180,41 +171,12 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
- residual value (value after residual connection
but before normalization)
"""
# x.shape: [batch_size, seq_len, inner_dim]
# Apply residual connection with gating
if isinstance(gate, int):
# used by cross-attention, should be 1
assert gate == 1
residual_output = residual + x
elif isinstance(gate, torch.Tensor):
if gate.dim() == 4:
# gate.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
residual_output = residual + (
x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
gate).flatten(1, 2)
else:
# used by bidirectional self attention
# gate.shape: [batch_size, 1, inner_dim]
residual_output = residual + x * gate
else:
raise ValueError(f"Gate type {type(gate)} not supported")
# residual_output.shape: [batch_size, seq_len, inner_dim]
residual_output = residual + x * gate
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
# shift.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
modulated = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1.0 + scale) + shift
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
@@ -256,24 +218,8 @@ class LayerNormScaleShift(nn.Module):
def forward(self, x: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
# x.shape: [batch_size, seq_len, inner_dim]
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
normalized = normalized.float()
if scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
output = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
else:
# scale.shape: [batch_size, 1, inner_dim]
# shift.shape: [batch_size, 1, inner_dim]
output = normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
output = output.to(x.dtype)
return output
return normalized * (1.0 + scale) + shift
+2 -2
View File
@@ -63,7 +63,7 @@ class BaseLayerWithLoRA(nn.Module):
device=self.base_layer.weight.device,
dtype=self.base_layer.weight.dtype))
torch.nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
torch.nn.init.zeros_(self.lora_B)
torch.nn.init.kaiming_uniform_(self.lora_B, a=math.sqrt(5))
else:
self.lora_A = None
self.lora_B = None
@@ -76,7 +76,7 @@ class BaseLayerWithLoRA(nn.Module):
lora_B = self.lora_B.to_local()
lora_A = self.lora_A.to_local()
if not self.merged and not self.disable_lora:
if (self.training_mode or not self.merged) and not self.disable_lora:
delta = x @ (
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
-9
View File
@@ -29,9 +29,6 @@ import torch
from fastvideo.distributed.parallel_state import get_sp_group
from fastvideo.layers.custom_op import CustomOp
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
@@ -270,7 +267,6 @@ def get_nd_rotary_pos_embed(
sp_rank: int = 0,
sp_world_size: int = 1,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
@@ -296,9 +292,6 @@ def get_nd_rotary_pos_embed(
full_grid = get_meshgrid_nd(
start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
if start_frame > 0:
full_grid[0] += start_frame
# Shard the grid if using sequence parallelism (sp_world_size > 1)
assert shard_dim < len(
rope_dim_list
@@ -377,7 +370,6 @@ def get_rotary_pos_embed(
interpolation_factor=1.0,
shard_dim: int = 0,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate rotary positional embeddings for the given sizes.
@@ -421,7 +413,6 @@ def get_rotary_pos_embed(
sp_rank=sp_rank,
sp_world_size=sp_world_size,
dtype=dtype,
start_frame=start_frame,
)
return freqs_cos, freqs_sin
+1 -5
View File
@@ -86,16 +86,12 @@ class TimestepEmbedder(nn.Module):
dtype=dtype)
self.freq_dtype = freq_dtype
def forward(self,
t: torch.Tensor,
timestep_seq_len: int | None = None) -> torch.Tensor:
def forward(self, t: torch.Tensor) -> torch.Tensor:
t_freq = timestep_embedding(t,
self.frequency_embedding_size,
self.max_period,
dtype=self.freq_dtype).to(
self.mlp.fc_in.weight.dtype)
if timestep_seq_len is not None:
t_freq = t_freq.unflatten(0, (1, timestep_seq_len))
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb

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