Compare commits

...
Author SHA1 Message Date
RandNMR73 da5ca94091 preprocessing text 2025-09-06 05:56:13 +00:00
William Lin d3ceb67e66 [misc] Update Slack invite link (#786) 2025-09-05 12:16:18 -07:00
Zhang Peiyuan 7ac153a5ca Update WeChat Link 2025-09-05 11:40:47 -07:00
William Lin d1e7aa0abd [CI] Add ssim test for causal inference (#784) 2025-09-05 01:23:01 -07:00
William Lin 2d846c55a1 [misc] Improve text encoding stage (#774) 2025-09-04 17:51:27 -07:00
Jinzhe Pan b318063c0a [Preprocess][Fix] video quality issue (#773) 2025-09-03 20:47:33 -07:00
Jinzhe Pan 4aa307be55 [Preprocess][Feat] support torchvision to load video in new preprocessing (#761) 2025-09-01 23:37:01 -07:00
William Lin 055e52e5ea [misc] [VSA] [STA] fix tk_root in setup.py for VSA and STA (#772) 2025-08-29 01:13:37 -07:00
William Lin 7d2069596b [bugfix] [VSA] [STA] Fix MANIFEST.in for VSA and STA; Move tk into both directories (#771) 2025-08-29 00:51:05 -07:00
William Lin c45009c9a4 [bugfix] fix STA install setup.py import (#770) 2025-08-28 23:02:53 -07:00
William LinandPeiyuan Zhang b91020b407 [VSA] [STA] Fix directory structure for pypi publishing (#769)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-08-28 22:34:03 -07:00
William Lin 2dcc5ea4f6 [chore] Release 0.1.6 (#768) 2025-08-28 20:56:21 -07:00
Wei ZhouandSolitaryThinker 359151d9a0 [Feature] Add wan2.2 5b i2v (#760)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-28 18:15:59 -07:00
Wei ZhouandSolitaryThinker ce67cd3729 [Feat] Support Self-Forcing's Causal Inference for Wan2.1 T2V 1.3B (#766)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-28 16:47:49 -07:00
Zhang Peiyuan 7c554e5da8 Update Community Link (#765) 2025-08-27 16:12:47 -07:00
William Lin 663ea33ff1 [bugfix] Fix wrong HF model string for FastWan2.2 5B (#763) 2025-08-26 22:05:40 -07:00
William Lin 3ef04f1654 [misc] [docs] Various fixes for logging and docs (#758) 2025-08-23 21:13:50 -07:00
Jinzhe Pan 0eced76a41 [Feat][Preprocess] support multi-gpus (#753) 2025-08-23 11:34:42 +08:00
Jinzhe Pan 3ab6470d1a [Feat][Preprocess] support merged dataset (#752) 2025-08-22 15:29:33 -07:00
Wenxuan Tan 989a03532c Optionally use unmerged weights for inference (#745) 2025-08-22 15:20:31 -07:00
William Lin fa15369a02 [bugfix] Check that model_index.json module is in required_modules list before removing (#756) 2025-08-22 14:36:44 -07:00
Zhang Peiyuan 78a9cb88d8 [Fix] fix seed in dmd denoising loop (#736) 2025-08-21 18:06:16 -07:00
Peng Xiaoand肖鹏 a0bff12746 [bugfix] [dmd] Align backward simulation with dmd2 sample back (#744)
Co-authored-by: 肖鹏 <xiaopeng1@aishi.ai>
2025-08-20 22:25:33 -07:00
William Lin 98f2af94e5 [bugfix] Missing Docker file for cuda12.9 (#750) 2025-08-20 15:34:31 -07:00
William Lin 46f7b6d574 [Docker] add 12.9 docker image and also fix py3.10 and py3.11 dockerfile (#749) 2025-08-20 15:31:15 -07:00
Jinzhe Pan 911a6a6a35 [Feat][Preprocessing] i2v preprocessing workflow (#737) 2025-08-14 20:47:25 -07:00
Zhang Peiyuan 38c7949d5c Update WeChat group link (#739) 2025-08-14 15:03:35 -07:00
Jinzhe Pan 7e7a0dba9d feat: preprocess validation dataset only when exist (#734) 2025-08-12 02:16:31 -07:00
Zhang Peiyuan f62e210ae6 Fix vsa backward gQ (#735) 2025-08-11 21:43:13 -07:00
William Lin 6ceb4942a0 [bugfix] [dmd] Fix backward simulation and also naming in wan_i2v_dmd_pipeline (#731) 2025-08-10 21:13:30 -07:00
William LinandRandNMR73 8cae5e4708 [feature] add Gradio live serving demo code (#727)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-08-10 15:34:03 -07:00
134 changed files with 6329 additions and 813 deletions
+18 -18
View File
@@ -117,11 +117,11 @@ steps:
queue: "default"
- path:
- "fastvideo/**"
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "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"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -133,10 +133,10 @@ steps:
queue: "default"
- path:
- "fastvideo/**"
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "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"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -147,10 +147,10 @@ steps:
agents:
queue: "default"
- path:
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "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"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -161,12 +161,12 @@ steps:
agents:
queue: "default"
- path:
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/video_sparse_attn/**"
- "csrc/attn/video_sparse_attn/tk/**"
- "csrc/attn/tests/test_vsa.py"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "csrc/attn/video_sparse_attn/setup.py"
- "csrc/attn/video_sparse_attn/config_vsa.py"
- "csrc/attn/video_sparse_attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
+15
View File
@@ -18,6 +18,12 @@ on:
required: false
default: false
type: boolean
python_3_12_cuda_12_9:
description: 'Build Python 3.12 image Cuda 12.9'
required: false
default: false
type: boolean
permissions:
contents: read
@@ -49,4 +55,13 @@ jobs:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12
tag_suffix: py3.12
secrets: inherit
build-python-3-12-cuda-12-9:
if: ${{ github.event.inputs.python_3_12_cuda_12_9 == 'true' }}
uses: ./.github/workflows/build-image-template.yml
with:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12.cuda12.9.1
tag_suffix: py3.12-cuda12.9.1
secrets: inherit
+12 -11
View File
@@ -104,16 +104,17 @@ jobs:
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/st_attn/**'
- 'csrc/attn/setup_sta.py'
- 'csrc/attn/config_sta.py'
- 'csrc/attn/st_attn.cpp'
- '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'
vsa-kernel-paths: &vsa-kernel-paths
- 'csrc/attn/vsa/**'
- 'csrc/attn/tk/**'
- 'csrc/attn/setup_vsa.py'
- 'csrc/attn/config_vsa.py'
- 'csrc/attn/vsa.cpp'
- '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'
vsa-paths: &vsa-paths
- 'fastvideo/**'
- *common-paths
@@ -234,7 +235,7 @@ jobs:
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
training-test:
needs: change-filter
if: >-
@@ -372,4 +373,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/setup_sta.py"
- "csrc/attn/sliding_tile_attn/setup.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/attn
cd csrc/attn/sliding_tile_attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -144,13 +144,13 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
cd csrc/attn/sliding_tile_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py bdist_wheel --dist-dir=dist
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/attn
cd csrc/attn/sliding_tile_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/dist/*.whl
path: csrc/attn/sliding_tile_attn/dist/*.whl
retention-days: 90
publish_package:
@@ -239,11 +239,11 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
cd csrc/attn/sliding_tile_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py sdist --dist-dir=dist
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/attn/dist/
packages-dir: csrc/attn/sliding_tile_attn/dist/
+11 -11
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/attn/setup_vsa.py"
- "csrc/attn/video_sparse_attn/setup.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/attn
cd csrc/attn/video_sparse_attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_vsa.py)
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup_vsa.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -152,13 +152,13 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
cd csrc/attn/video_sparse_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_vsa.py bdist_wheel --dist-dir=dist
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/attn
cd csrc/attn/video_sparse_attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
@@ -173,7 +173,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/attn/dist/*.whl
path: csrc/attn/video_sparse_attn/dist/*.whl
retention-days: 90
publish_package:
@@ -247,11 +247,11 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
cd csrc/attn/video_sparse_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_vsa.py sdist --dist-dir=dist
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/attn/dist/
packages-dir: csrc/attn/video_sparse_attn/dist/
+6 -2
View File
@@ -1,3 +1,7 @@
[submodule "csrc/attn/tk"]
path = csrc/attn/tk
[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
url = https://github.com/HazyResearch/ThunderKittens.git
+1
View File
@@ -22,6 +22,7 @@ exclude: |
examples/.*|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml|
.github/workflows/vsa-publish.yml|
.github/workflows/build-image-template.yml|
docs/source/inference/support_matrix.md
)
+2 -2
View File
@@ -1,5 +1,5 @@
<div align="center">
<img src=assets/logo.png width="30%"/>
<img src=assets/logos/logo.svg width="30%"/>
</div>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
@@ -7,7 +7,7 @@
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/qqPzbrw" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://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> |
</p>
<div align="center">
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 46 KiB

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

After

Width:  |  Height:  |  Size: 691 B

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

After

Width:  |  Height:  |  Size: 5.7 KiB

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

After

Width:  |  Height:  |  Size: 691 B

Binary file not shown.

Before

Width:  |  Height:  |  Size: 31 KiB

+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.4)
(If you use CUDA12.8)
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
-4
View File
@@ -1,4 +0,0 @@
off_hz = tl.program_id(2)
b = off_hz // H
h = off_hz % H
meta_base = ((b * H + h) * q_tiles + q_blk)
@@ -1,2 +1,2 @@
recursive-include tk *
include config.py
include config_sta.py
+87
View File
@@ -0,0 +1,87 @@
# 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.
@@ -1,7 +1,7 @@
import os
import subprocess
from csrc.attn.config_sta import kernels, sources, target
from 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.4"
VERSION = "0.0.6"
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"
+23 -26
View File
@@ -13,9 +13,9 @@ BLOCK_M = 64
BLOCK_N = 64
def pytorch_test(Q, K, V, block_sparse_mask, dO):
q_ = Q.clone().requires_grad_()
k_ = K.clone().requires_grad_()
v_ = V.clone().requires_grad_()
q_ = Q.clone().float().requires_grad_()
k_ = K.clone().float().requires_grad_()
v_ = V.clone().float().requires_grad_()
QK = torch.matmul(q_, k_.transpose(-2, -1))
QK /= (q_.size(-1) ** 0.5)
@@ -35,9 +35,9 @@ def pytorch_test(Q, K, V, block_sparse_mask, dO):
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
Q = Q.clone().requires_grad_()
K = K.clone().requires_grad_()
V = V.clone().requires_grad_()
Q = Q.detach().requires_grad_()
K = K.detach().requires_grad_()
V = V.detach().requires_grad_()
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
@@ -60,11 +60,9 @@ def get_non_pad_index(
return index_pad[index_mask]
def generate_tensor(shape, mean, std, dtype, device):
def generate_tensor(shape, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
return tensor
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
@@ -75,7 +73,7 @@ def vsa_pad(x, non_pad_index, num_blocks, block_size):
padded_x[:, :, non_pad_index, :] = x
return padded_x
def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_mode='all'):
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
results = {
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
@@ -91,10 +89,10 @@ def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_m
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
K = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
V = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
# dO_padded = torch.zeros_like(dO_padded)
# dO_padded[:, :, non_pad_index, :] = dO
@@ -107,7 +105,8 @@ def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_m
abs_diff = torch.abs(diff)
results[name]['sum_diff'] += torch.sum(abs_diff).item()
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
results[name]['max_diff'] = max(results[name]['max_diff'], torch.max(abs_diff).item())
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
if torch.cuda.is_available():
torch.cuda.empty_cache()
@@ -119,27 +118,27 @@ def check_correctness(h, d, num_blocks, k, mean, std, num_iterations=20, error_m
return results
def generate_error_graphs(h, d, mean, std, error_mode='all'):
def generate_error_graphs(h, d, error_mode='all'):
test_configs = [
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
]
print(f"\nError Analysis for h={h}, d={d}, mean={mean}, std={std}, mode={error_mode}")
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
print("=" * 150)
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
f"{'gQ Avg':<12} {'gQ Max':<12} "
f"{'gK Avg':<12} {'gK Max':<12} "
f"{'gV Avg':<12} {'gV Max':<12} "
f"{'gO Avg':<12} {'gO Max':<12}")
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
f"{'gK Avg':<12} {'Rel gK Max':<12} "
f"{'gV Avg':<12} {'Rel gV Max':<12} "
f"{'gO Avg':<12} {'Rel gO Max':<12}")
print("-" * 150)
for config in test_configs:
num_blocks = config["num_blocks"]
k = config["k"]
description = config["description"]
results = check_correctness(h, d, num_blocks, k, mean, std, error_mode=error_mode)
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
print(f"{description:<20} {num_blocks:<8} {k:<4} "
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
@@ -150,10 +149,8 @@ def generate_error_graphs(h, d, mean, std, error_mode='all'):
if __name__ == "__main__":
h, d = 16, 128
mean = 0.0
std = 1
print("Block Sparse Attention with Variable Block Sizes Analysis")
print("=" * 60)
for mode in ['backward']:
generate_error_graphs(h, d, mean, std, error_mode=mode)
generate_error_graphs(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
Submodule csrc/attn/tk deleted from 1719fb7264
+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config_vsa.py
+61
View File
@@ -0,0 +1,61 @@
# 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.
@@ -9,10 +9,10 @@ target = target.lower()
# Package metadata
PACKAGE_NAME = "vsa"
VERSION = "0.0.1"
VERSION = "0.0.3"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
# Set environment variables
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
@@ -568,7 +568,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
__syncthreads(); // wait for sd_smem shared memory write
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
warpgroup::mma_commit_group();
tma::store_async_wait();
warpgroup::mma_async_wait();
// store qg to shared memory
warpgroup::store(qg_smem, qg_reg);
@@ -578,6 +577,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if (threadIdx.x / 32 == 0) {
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
tma::store_add_async(g.qg, qg_smem, tile_idx);
tma::store_async_wait();
}
}
@@ -624,7 +624,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
__syncthreads(); // wait for sd_smem shared memory write
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
warpgroup::mma_commit_group();
tma::store_async_wait();
warpgroup::mma_async_wait();
// store qg to shared memory
warpgroup::store(qg_smem, qg_reg);
@@ -634,13 +633,14 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if (threadIdx.x / 32 == 0) {
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
tma::store_add_async(g.qg, qg_smem, tile_idx);
tma::store_async_wait();
}
}
// store kq and vq
// ! the following two line seems unnecessary.
tma::store_async_wait(); // ensure qg is finished
// tma::store_async_wait(); // ensure qg is finished
__syncthreads();
warpgroup::store(kg_smem[0], kg_reg);
@@ -1174,4 +1174,4 @@ block_sparse_attention_backward(torch::Tensor q,
return {qg, kg, vg};
//cudadevicesynchronize();
}
}
+45 -21
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.10.0 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.10 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/sliding_tile_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
# 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
EXPOSE 22
+45 -21
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.11.11 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.11 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/sliding_tile_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
# 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
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.0.post2 --no-build-isolation
uv pip install --no-cache-dir flash-attn==2.8.3 --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 && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup_sta.py install
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
python setup.py install
EXPOSE 22
EXPOSE 22
+72
View File
@@ -0,0 +1,72 @@
FROM nvidia/cuda:12.9.1-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.9
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.12 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/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
+1 -2
View File
@@ -96,8 +96,7 @@ copybutton_prompt_is_regexp = True
#
html_title = project
html_theme = 'sphinx_book_theme'
html_logo = '../../assets/logo.jpg'
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
html_logo = '../../assets/logos/icon_simple.svg'
html_theme_options = {
'path_to_docs': 'docs/source',
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
@@ -3,12 +3,12 @@
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
**Images:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
## Starting the container
```bash
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
```
This will:
@@ -6,7 +6,7 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
## Creating a new pod
Choose a GPU that supports CUDA 12.4
Choose a GPU that supports CUDA 12.8
Pick 1 or 2 L40S GPU(s)
+24 -3
View File
@@ -22,10 +22,20 @@ source ~/.bashrc
Create and activate a Conda environment for FastVideo:
```
conda create -n fastvideo python=3.10 -y
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Install `uv` (optional, but recommended):
From instructions on [uv](https://astral.sh/uv/):
```
curl -LsSf https://astral.sh/uv/install.sh | sh
# or
wget -qO- https://astral.sh/uv/install.sh | sh
```
Clone the FastVideo repository and go to the FastVideo directory:
```
@@ -36,10 +46,10 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
```bash
pip install -e .[dev]
uv pip install -e .[dev]
# Can also install flash-attn (optional)
pip install flash-attn==2.7.4.post1 --no-build-isolation
uv pip install flash-attn --no-build-isolation
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
@@ -50,3 +60,14 @@ pre-commit run --all-files
# Unit tests
pytest tests/
```
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
```
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
# make sure you have ninja installed
uv pip install ninja
python setup.py install
```
@@ -6,7 +6,7 @@ Instructions to install FastVideo for NVIDIA CUDA GPUs.
- **OS: Linux or Windows WSL**
- **Python: 3.10-3.12**
- **CUDA 12.4**
- **CUDA 12.8**
- **At least 1 NVIDIA GPU**
## Set up using Python
@@ -38,6 +38,7 @@ conda activate fastvideo
:::{tip}
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
Note that you can also use `uv` to install FastVideo in a Conda environment.
:::
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
@@ -60,7 +61,7 @@ uv pip install fastvideo
Also optionally install flash-attn:
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
pip install flash-attn --no-build-isolation
```
### Installation from Source
@@ -87,7 +88,7 @@ uv pip install -e .
#### Flash Attention
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
pip install flash-attn --no-build-isolation
```
## Set up using Docker
@@ -102,7 +103,7 @@ If you're planning to contribute to FastVideo please see the following page:
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA 12.4 support
- NVIDIA GPU with CUDA 12.8 support
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
@@ -39,6 +39,7 @@ conda activate fastvideo
:::{tip}
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
Note that you can also use `uv` to install FastVideo in a Conda environment.
:::
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
+1 -2
View File
@@ -1,6 +1,6 @@
# Welcome to FastVideo
:::{figure} ../../assets/logo.png
:::{figure} ../../assets/logos/logo.svg
:align: center
:alt: FastVideo
:class: no-scaled-link
@@ -101,7 +101,6 @@ sliding_tile_attention/demo
:maxdepth: 1
video_sparse_attention/installation
video_sparse_attention/demo
:::
:::{toctree}
@@ -5,7 +5,7 @@ This page contains step-by-step instructions to get you quickly started with vid
## Requirements
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.4
- **CUDA**: 12.8
- **GPU**: At least one NVIDIA GPU
## Installation
+50 -6
View File
@@ -6,6 +6,7 @@ The symbols used have the following meanings:
- ✅ = Full compatibility
- ❌ = No compatibility
- ⭕ = Does not apply to this model
## Models x Optimization
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
@@ -37,51 +38,94 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
* TeaCache
* Sliding Tile Attn
* Sage Attn
* Video Sparse Attention (VSA)
- * FastWan2.1 T2V 1.3B
* `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
* 480P
* ⭕
* ⭕
* ⭕
* ✅
- * FastWan2.2 TI2V 5B Full Attn*
* `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers`
* 720P
* ⭕
* ⭕
* ⭕
* ✅
- * Wan2.2 TI2V 5B
* `Wan-AI/Wan2.2-TI2V-5B-Diffusers`
* 720P
* ⭕
* ⭕
* ✅
* ⭕
- * Wan2.2 T2V A14B
* `Wan-AI/Wan2.2-T2V-A14B-Diffusers`
* 480P<br>720P
* ❌
* ❌
* ✅
* ⭕
- * Wan2.2 I2V A14B
* `Wan-AI/Wan2.2-I2V-A14B-Diffusers`
* 480P<br>720P
* ❌
* ❌
* ✅
* ⭕
- * HunyuanVideo
* `hunyuanvideo-community/HunyuanVideo`
* 720px1280p<br>544px960p
* ❌
* ✅
* ✅
* ⭕
- * FastHunyuan
* `FastVideo/FastHunyuan-diffusers`
* 720px1280p<br>544px960p
* ❌
* ✅
* ✅
- * Wan T2V 1.3B
* ⭕
- * Wan2.1 T2V 1.3B
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan T2V 14B
* ⭕
- * Wan2.1 T2V 14B
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
* 480P, 720P
* ✅
* ✅*
* ✅
- * Wan I2V 480P
* ⭕
- * Wan2.1 I2V 480P
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan I2V 720P
* ⭕
- * Wan2.1 I2V 720P
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
* 720P
* ✅
* ✅*
* ✅
* ✅
* ⭕
- * StepVideo T2V
* `FastVideo/stepvideo-t2v-diffusers`
* 768px768px204f<br>544px992px204f<br>544px992px136f
* ❌
* ❌
* ✅
* ⭕
:::
**Note**: there are some known quality issues with Wan2.1 + Sliding Tile Attn. We are working on fixing this issue.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
## Special requirements
@@ -4,7 +4,7 @@
You can install the Sliding Tile Attention package using
```
pip install st_attn==0.0.4
pip install st_attn
```
# Building from Source
@@ -12,7 +12,6 @@ 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
@@ -22,14 +21,20 @@ sudo apt update
sudo apt install clang-11
```
Install STA:
Set up CUDA environment (if using CUDA 12.4):
```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_sta.py install
python setup.py install
```
# 🧪 Test
@@ -1,3 +0,0 @@
(vsa-demo)=
# 🎬 Demo
@@ -4,8 +4,7 @@
You can install the Video Sparse Attention package using
```bash
git submodule update --init --recursive
python setup_vsa.py install
pip install vsa
```
# Building from Source
@@ -23,10 +22,10 @@ sudo apt update
sudo apt install clang-11
```
Set up CUDA environment (if using CUDA 12.4):
Set up CUDA environment (if using CUDA 12.8):
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
@@ -34,9 +33,9 @@ export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
Install VSA:
```bash
cd csrc/attn/
cd csrc/attn/video_sparse_attn/
git submodule update --init --recursive
python setup_vsa.py install
python setup.py install
```
# 🧪 Test
@@ -4,9 +4,7 @@ These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using
### 0. Make sure you have installed VSA
```bash
cd csrc/attn
git submodule update --init --recursive
python setup_vsa.py install
pip install vsa
```
### 1. Download dataset:
@@ -4,9 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
### 0. Make sure you have installed VSA
```bash
cd csrc/attn
git submodule update --init --recursive
python setup_vsa.py install
pip install vsa
```
### Data-free Distillation
@@ -4,9 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
### 0. Make sure you have installed VSA
```bash
cd csrc/attn
git submodule update --init --recursive
python setup_vsa.py install
pip install vsa
```
### 1. Download dataset:
@@ -0,0 +1,31 @@
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()
@@ -0,0 +1,41 @@
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()
-59
View File
@@ -1,59 +0,0 @@
# FastVideo Gradio Demo
This is a Gradio-based web interface for generating videos using the FastVideo framework. The demo allows users to create videos from text prompts with various customization options.
## Overview
The demo uses the FastVideo framework to generate videos based on text prompts. It provides a simple web interface built with Gradio that allows users to:
- Enter text prompts to generate videos
- Customize video parameters (dimensions, number of frames, etc.)
- Use negative prompts to guide the generation process
- Set or randomize seeds for reproducibility
---
## Usage
Run the demo with:
```bash
python examples/inference/gradio/gradio_demo.py
```
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
---
## Model Initialization
This demo initializes a `VideoGenerator` with the minimum required arguments for inference. Users can seamlessly adjust inference options between generations, including prompts, resolution, video length, or even the number of inference steps, *without ever needing to reload the model*.
## Video Generation
The core functionality is in the `generate_video` function, which:
1. Processes user inputs
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
3. Returns an output path that Gradio uses to display the generated video
## Gradio Interface
The interface is built with several components:
- A text input for the prompt
- A video display for the result
- Inference options in a collapsible accordion:
- Height and width sliders
- Number of frames slider
- Guidance scale slider
- Inference steps slider
- Negative prompt options
- Seed controls
### Inference Options
- **Height/Width**: Control the resolution of the generated video
- **Number of Frames**: Set how many frames to generate
- **Guidance Scale**: Control how closely the generation follows the prompt
- **Inference Steps**: More steps can improve quality but take longer
- **Negative Prompt**: Specify what you don't want to see in the video
- **Seed**: Control randomness for reproducible results
-169
View File
@@ -1,169 +0,0 @@
import argparse
import os
from copy import deepcopy
import gradio as gr
import torch
from fastvideo import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Gradio Demo")
parser.add_argument("--model_path",
type=str,
default="FastVideo/FastHunyuan-diffusers",
help="Path to the model")
parser.add_argument("--num_gpus",
type=int,
default=1,
help="Number of GPUs to use")
parser.add_argument("--output_path",
type=str,
default="outputs",
help="Path to save generated videos")
parsed_args = parser.parse_args()
# args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=parsed_args.model_path, num_gpus=parsed_args.num_gpus)
default_params = SamplingParam.from_pretrained(parsed_args.model_path)
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
params = deepcopy(default_params)
params.prompt = prompt
params.negative_prompt = negative_prompt
params.seed = seed
params.guidance_scale = guidance_scale
params.num_frames = num_frames
params.height = height
params.width = width
params.num_inference_steps = num_inference_steps
if randomize_seed:
params.seed = torch.randint(0, 1000000, (1, )).item()
if not use_negative_prompt:
params.negative_prompt = None
generator.generate_video(prompt=prompt, sampling_param=params)
output_path = os.path.join(parsed_args.output_path,
f"{params.prompt[:100]}.mp4")
return output_path, params.seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
with gr.Blocks() as demo:
gr.Markdown("# FastVideo Inference Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=default_params.height,
)
width = gr.Slider(label="Width",
minimum=256,
maximum=1024,
step=32,
value=default_params.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=default_params.num_frames,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=default_params.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=default_params.num_inference_steps,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=default_params.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=default_params.negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
@@ -0,0 +1,730 @@
import argparse
import os
import requests
import base64
import time
import gradio as gr
from fastvideo.configs.sample.base import SamplingParam
MODEL_PATH_MAPPING = {
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
}
class RayServeClient:
def __init__(self, backend_url: str):
self.backend_url = backend_url
self.session = requests.Session()
def check_health(self) -> bool:
try:
response = self.session.get(f"{self.backend_url}/health", timeout=5)
return response.status_code == 200
except requests.exceptions.RequestException:
return False
def generate_video(self, request_data: dict) -> dict:
start_time = time.time()
try:
headers = {"Content-Type": "application/json"}
response = self.session.post(
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=300
)
round_trip_time = time.time() - start_time
if response.status_code == 200:
result = response.json()
backend_total = result.get("total_time", 0)
network_time = round_trip_time - backend_total
result["network_time"] = network_time
return result
else:
return {"success": False, "error_message": f"HTTP {response.status_code}: {response.text}"}
except requests.exceptions.RequestException as e:
return {"success": False, "error_message": f"Request failed: {str(e)}"}
def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str:
if not video_data:
return None
try:
if video_data.startswith('data:video/'):
video_data = video_data.split(',')[1]
video_bytes = base64.b64decode(video_data)
safe_prompt = prompt[:50].replace(' ', '_').replace('/', '_').replace('\\', '_')
video_filename = f"{safe_prompt}.mp4"
video_path = os.path.join(output_dir, video_filename)
os.makedirs(output_dir, exist_ok=True)
with open(video_path, 'wb') as f:
f.write(video_bytes)
return video_path
except Exception as e:
print(f"Failed to save video: {e}")
return None
def create_timing_display(inference_time, encoding_time, network_time, total_time, stage_execution_times, num_frames):
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
timing_html = f"""
<div style="margin: 10px 0;">
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
<div class="timing-card timing-card-highlight">
<div style="font-size: 20px;">🚀</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🧠</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🎬</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
<div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🌐</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
<div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">📊</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
</div>
</div>"""
if inference_time > 0:
fps = num_frames / inference_time
timing_html += f"""
<div class="performance-card" style="margin-top: 15px;">
<span style="font-weight: bold;">Generation Speed: </span>
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
</div>"""
return timing_html + "</div>"
def load_example_prompts():
def contains_chinese(text):
return any('\u4e00' <= char <= '\u9fff' for char in text)
def load_from_file(filepath):
prompts, labels = [], []
try:
with open(filepath, "r", encoding='utf-8') as f:
for line in f:
line = line.strip()
if line and not contains_chinese(line):
label = line[:100] + "..." if len(line) > 100 else line
labels.append(label)
prompts.append(line)
except Exception as e:
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("prompts/prompts_final.txt")
if not examples:
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
example_labels = ["Crowded rooftop bar at night"]
return examples, example_labels
def create_gradio_interface(backend_url: str, default_params: dict[str, SamplingParam]):
client = RayServeClient(backend_url)
def generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
):
if not client.check_health():
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
# Validate dimensions
max_pixels = 720 * 1280
if height * width > max_pixels:
return None, f"Video dimensions too large. Maximum: 720x1280 pixels", ""
if progress:
progress(0.1, desc="Checking backend health...")
request_data = {
"prompt": prompt,
"negative_prompt": negative_prompt,
"use_negative_prompt": use_negative_prompt,
"seed": seed,
"guidance_scale": guidance_scale,
"num_frames": num_frames,
"height": height,
"width": width,
"randomize_seed": randomize_seed,
"return_frames": False,
"image_path": None,
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
}
if progress:
progress(0.4, desc="Generating video...")
response = client.generate_video(request_data)
if progress:
progress(0.8, desc="Processing response...")
if response.get("success", False):
video_data = response.get("video_data", "")
used_seed = response.get("seed", seed)
inference_time = response.get("inference_time", 0.0)
encoding_time = response.get("encoding_time", 0.0)
total_time = response.get("total_time", 0.0)
network_time = response.get("network_time", 0.0)
stage_execution_times = response.get("stage_execution_times", [])
timing_details = create_timing_display(
inference_time, encoding_time, network_time, total_time,
stage_execution_times, num_frames
)
if video_data:
if progress:
progress(0.9, desc="Saving video...")
video_path = save_video_from_base64(video_data, "outputs", prompt)
if progress:
progress(1.0, desc="Generation complete!")
if video_path and os.path.exists(video_path):
return video_path, used_seed, timing_details
else:
return None, "Failed to save video", ""
else:
return None, "No video data received from backend", ""
else:
error_msg = response.get("error_message", "Unknown error occurred")
return None, f"Generation failed: {error_msg}", ""
examples, example_labels = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
button_primary_background_fill_hover="#1d4ed8",
button_primary_text_color="white",
slider_color="#2563eb",
checkbox_background_color_selected="#2563eb",
)
def get_default_values(model_name):
model_path = MODEL_PATH_MAPPING.get(model_name)
if model_path and model_path in default_params:
params = default_params[model_path]
return {
'height': params.height,
'width': params.width,
'num_frames': params.num_frames,
'guidance_scale': params.guidance_scale,
'seed': params.seed,
}
return {
'height': 448,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
with gr.Blocks(title="FastWan", theme=theme) as demo:
gr.Image("assets/logos/logo.svg", show_label=False, container=False, height=80)
gr.HTML("""
<div style="text-align: center; margin-bottom: 10px;">
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
with gr.Accordion("🎥 What Is FastVideo?", open=False):
gr.HTML("""
<div style="padding: 20px; line-height: 1.6;">
<p style="font-size: 16px; margin-bottom: 15px;">
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
</p>
</div>
""")
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
label="Select Model",
interactive=True
)
with gr.Row():
example_dropdown = gr.Dropdown(
choices=example_labels,
label="Example Prompts",
value=None,
interactive=True,
allow_custom_value=False
)
with gr.Row():
with gr.Column(scale=6):
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=3,
placeholder="Describe your scene...",
container=False,
lines=3,
autofocus=True,
)
with gr.Column(scale=1, min_width=120, elem_classes="center-button"):
run_button = gr.Button("Run", variant="primary", size="lg")
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", visible=False)
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
with gr.Row(equal_height=True, elem_classes="main-content-row"):
with gr.Column(scale=1, elem_classes="advanced-options-column"):
with gr.Group():
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Advanced Options</div>")
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=initial_values['guidance_scale'],
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
with gr.Column(scale=1, elem_classes="video-column"):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
container=True,
elem_classes="video-component"
)
gr.HTML("""
<style>
.center-button {
display: flex !important;
justify-content: center !important;
height: 100% !important;
padding-top: 1.4em !important;
}
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
}
.main {
max-width: 1200px !important;
margin: 0 auto !important;
}
.gr-form, .gr-box, .gr-group {
max-width: 1200px !important;
}
.gr-video {
max-width: 500px !important;
margin: 0 auto !important;
}
.main-content-row {
display: flex !important;
align-items: flex-start !important;
min-height: 500px !important;
gap: 20px !important;
}
.advanced-options-column,
.video-column {
display: flex !important;
flex-direction: column !important;
flex: 1 !important;
min-height: 400px !important;
align-items: stretch !important;
}
.video-column > * {
margin-top: 0 !important;
}
.video-column .gr-video,
.video-component {
margin-top: 0 !important;
padding-top: 0 !important;
}
.video-column .gr-video .gr-form {
margin-top: 0 !important;
}
.advanced-options-column .gr-group,
.video-column .gr-video {
margin-top: 0 !important;
vertical-align: top !important;
}
.advanced-options-column > *:last-child,
.video-column > *:last-child {
flex-grow: 0 !important;
}
@media (max-width: 1400px) {
.main-content-row {
min-height: 600px !important;
}
.advanced-options-column,
.video-column {
min-height: 600px !important;
}
}
@media (max-width: 1200px) {
.main-content-row {
flex-direction: column !important;
align-items: stretch !important;
}
.advanced-options-column,
.video-column {
min-height: auto !important;
width: 100% !important;
}
}
.timing-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 8px;
text-align: center;
min-height: 80px;
display: flex;
flex-direction: column;
justify-content: center;
}
.timing-card-highlight {
background: var(--background-fill-primary) !important;
border: 2px solid var(--color-accent) !important;
}
.performance-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
.gr-number input[readonly] {
background-color: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color-subdued) !important;
cursor: default !important;
text-align: center !important;
font-weight: 500 !important;
}
</style>
""")
def on_example_select(example_label):
if example_label and example_label in example_labels:
index = example_labels.index(example_label)
return examples[index]
return ""
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
)
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected. We are also rate-limiting users to 3 requests per minute.</p>
</div>
""")
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
def on_model_selection_change(selected_model):
if not selected_model:
selected_model = "FastWan2.1-T2V-1.3B"
model_path = MODEL_PATH_MAPPING.get(selected_model)
if model_path and model_path in default_params:
params = default_params[model_path]
return (
gr.update(value=params.height),
gr.update(value=params.width),
gr.update(value=params.num_frames),
gr.update(value=params.guidance_scale),
gr.update(value=params.seed),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=3.0),
gr.update(value=1024),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
result_path, seed_or_error, timing_details = generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
)
if result_path and os.path.exists(result_path):
return (
result_path,
seed_or_error,
gr.update(visible=False),
gr.update(visible=True, value=timing_details),
)
else:
return (
None,
seed_or_error,
gr.update(visible=True, value=seed_or_error),
gr.update(visible=False),
)
run_button.click(
fn=handle_generation,
inputs=[
model_selection,
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
randomize_seed,
],
outputs=[result, seed_output, error_output, timing_display],
concurrency_limit=20,
)
return demo
def main():
parser = argparse.ArgumentParser(description="FastVideo Gradio Frontend")
parser.add_argument("--backend_url", type=str, default="http://localhost:8000",
help="URL of the Ray Serve backend")
parser.add_argument("--t2v_model_paths", type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--host", type=str, default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port", type=int, default=7860,
help="Port to bind to")
args = parser.parse_args()
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(args.backend_url, default_params)
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
print(f"Backend URL: {args.backend_url}")
print(f"T2V Models: {args.t2v_model_paths}")
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
import uvicorn
app = FastAPI()
@app.get("/logo.svg")
def get_logo():
return FileResponse(
"assets/logos/logo.svg",
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
@app.get("/favicon.ico")
def get_favicon():
favicon_path = "assets/logos/icon_simple.svg"
if os.path.exists(favicon_path):
return FileResponse(
favicon_path,
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
else:
raise HTTPException(status_code=404, detail="Favicon not found")
@app.get("/", response_class=HTMLResponse)
def index(request: Request):
base_url = str(request.base_url).rstrip('/')
return f"""
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:description" content="Make video generation go blurrrrrrr">
<meta property="og:image" content="{base_url}/logo.svg">
<meta property="og:image:width" content="1200">
<meta property="og:image:height" content="630">
<meta property="og:site_name" content="FastWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
<meta property="twitter:image" content="{base_url}/logo.svg">
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
<link rel="icon" type="image/png" sizes="16x16" href="/favicon.ico">
<link rel="apple-touch-icon" href="/favicon.ico">
<style>
body, html {{
margin: 0;
padding: 0;
height: 100%;
overflow: hidden;
}}
iframe {{
width: 100%;
height: 100vh;
border: none;
}}
</style>
</head>
<body>
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
</body>
</html>
"""
app = gr.mount_gradio_app(
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
)
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()
@@ -0,0 +1,379 @@
import time
import os
import torch
import base64
import io
from copy import deepcopy
from typing import Dict, Any, Optional, List
import signal
import sys
import ray
from ray import serve
from fastapi import FastAPI, Request, Response
from pydantic import BaseModel
import numpy as np
from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.util import get_remote_address
from slowapi.errors import RateLimitExceeded
import imageio
from ray.serve.handle import DeploymentHandle
from prometheus_client import Counter, Histogram, generate_latest
NUM_GPUS = 16
DEFAULT_FPS = 16
SEED_RANGE_MAX = 1_000_000
SUPPORTED_MODELS = [
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
]
MODEL_CONFIGS = {
"1.3B": {
"num_cpus": 2,
"text_encoder_cpu_offload": False,
"dit_cpu_offload": False,
"vae_cpu_offload": False,
"VSA_sparsity": 0.8,
},
"14B": {
"num_cpus": 16,
"text_encoder_cpu_offload": True,
"dit_cpu_offload": True,
"vae_cpu_offload": False,
"VSA_sparsity": 0.9,
}
}
class VideoGenerationRequest(BaseModel):
prompt: str
negative_prompt: Optional[str] = None
use_negative_prompt: bool = False
seed: int = 42
guidance_scale: float = 7.5
num_frames: int = 21
height: int = 448
width: int = 832
randomize_seed: bool = False
return_frames: bool = False
model_path: Optional[str] = None
class VideoGenerationResponse(BaseModel):
video_data: Optional[str] = None
seed: int
success: bool
error_message: Optional[str] = None
generation_time: Optional[float] = None
model_load_time: Optional[float] = None
inference_time: Optional[float] = None
encoding_time: Optional[float] = None
total_time: Optional[float] = None
stage_names: Optional[List[str]] = None
stage_execution_times: Optional[List[float]] = None
def encode_video_to_base64(frames: List[np.ndarray], fps: int = DEFAULT_FPS) -> str:
if not frames:
return ""
try:
buffer = io.BytesIO()
imageio.mimsave(buffer, frames, fps=fps, format="mp4")
buffer.seek(0)
video_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
return f"data:video/mp4;base64,{video_base64}"
except Exception as e:
print(f"Warning: Failed to encode video: {e}")
return ""
def setup_model_environment(model_path: str) -> None:
if "fullattn" in model_path.lower():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
def process_generation_result(result: Any) -> tuple[List[np.ndarray], float, List[str], List[float]]:
frames = result if isinstance(result, list) else result.get("frames", [])
generation_time = result.get("generation_time", 0.0) if isinstance(result, dict) else 0.0
logging_info = result.get("logging_info", None)
if logging_info:
stage_names = logging_info.get_execution_order()
stage_execution_times = [
logging_info.get_stage_info(stage_name).get("execution_time", 0.0)
for stage_name in stage_names
]
else:
stage_names = []
stage_execution_times = []
return frames, generation_time, stage_names, stage_execution_times
def prepare_sampling_params(video_request: VideoGenerationRequest, default_params: Any) -> Any:
params = deepcopy(default_params)
params.prompt = video_request.prompt
if video_request.use_negative_prompt:
params.negative_prompt = video_request.negative_prompt
params.seed = (video_request.seed if not video_request.randomize_seed
else torch.randint(0, SEED_RANGE_MAX, (1,)).item())
params.randomize_seed = video_request.randomize_seed
params.guidance_scale = video_request.guidance_scale
params.num_frames = video_request.num_frames
params.height = video_request.height
params.width = video_request.width
params.save_video = False
params.return_frames = False
return params
class BaseModelDeployment:
def __init__(self, model_path: str, output_path: str = "outputs"):
self.model_path = model_path
self.output_path = output_path
self.generator = None
self.default_params = None
os.makedirs(self.output_path, exist_ok=True)
setup_model_environment(self.model_path)
def _initialize_generator(self, config: Dict[str, Any]) -> None:
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
print(f"Initializing model: {self.model_path}")
self.generator = VideoGenerator.from_pretrained(
model_path=self.model_path,
num_gpus=1,
use_fsdp_inference=True,
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
dit_cpu_offload=config["dit_cpu_offload"],
vae_cpu_offload=config["vae_cpu_offload"],
VSA_sparsity=config["VSA_sparsity"],
enable_stage_verification=False,
)
self.default_params = SamplingParam.from_pretrained(self.model_path)
def generate_video(self, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
total_start_time = time.time()
params = prepare_sampling_params(video_request, self.default_params)
inference_start_time = time.time()
result = self.generator.generate_video(
prompt=video_request.prompt,
sampling_param=params,
save_video=False,
return_frames=False,
)
inference_time = time.time() - inference_start_time
frames, generation_time, stage_names, stage_execution_times = process_generation_result(result)
encoding_start_time = time.time()
video_data = encode_video_to_base64(frames, fps=DEFAULT_FPS)
encoding_time = time.time() - encoding_start_time
total_time = time.time() - total_start_time
return VideoGenerationResponse(
video_data=video_data,
seed=params.seed,
success=True,
generation_time=generation_time,
inference_time=inference_time,
encoding_time=encoding_time,
total_time=total_time,
stage_names=stage_names,
stage_execution_times=stage_execution_times,
)
@serve.deployment(
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
)
class T2VModelDeployment(BaseModelDeployment):
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
super().__init__(t2v_model_path, output_path)
self._initialize_generator(MODEL_CONFIGS["1.3B"])
print("✅ T2V model initialized successfully")
@serve.deployment(
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
)
class T2V14BModelDeployment(BaseModelDeployment):
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
super().__init__(t2v_14b_model_path, output_path)
# Override environment for 14B model
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
self._initialize_generator(MODEL_CONFIGS["14B"])
print("✅ T2V 14B model initialized successfully")
app = FastAPI()
limiter = Limiter(key_func=get_remote_address)
app.state.limiter = limiter
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
@serve.deployment(num_replicas=50, ray_actor_options={"num_cpus": 2})
@serve.ingress(app)
class FastVideoAPI:
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
self.t2v_deployments = t2v_deployments
# Initialize Prometheus metrics
self.request_count = Counter('fastvideo_requests_total', 'Total FastVideo requests', ['model_type', 'status'])
self.request_duration = Histogram('fastvideo_request_duration_seconds', 'FastVideo request duration', ['model_type'])
self.video_generation_time = Histogram('fastvideo_video_generation_seconds', 'Video generation time', ['model_type'])
def _get_model_name(self, model_path: Optional[str]) -> str:
return model_path.split('/')[-1] if model_path else "unknown"
def _record_metrics(self, model_name: str, status: str, duration: float, response: Optional[VideoGenerationResponse] = None) -> None:
self.request_count.labels(model_type=model_name, status=status).inc()
self.request_duration.labels(model_type=model_name).observe(duration)
if response and hasattr(response, 'generation_time') and response.generation_time:
self.video_generation_time.labels(model_type=model_name).observe(response.generation_time)
@app.post("/generate_video", response_model=VideoGenerationResponse)
@limiter.limit("10/minute")
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
"""Route the request to the appropriate model deployment based on model_path."""
start_time = time.time()
model_name = self._get_model_name(video_request.model_path)
try:
if video_request.model_path not in self.t2v_deployments:
raise ValueError(f"Model {video_request.model_path} not found")
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
response = await response_ref
self._record_metrics(model_name, "success", time.time() - start_time, response)
return response
except Exception as e:
self._record_metrics(model_name, "error", time.time() - start_time)
return VideoGenerationResponse(
video_data=None,
seed=video_request.seed,
success=False,
error_message=str(e),
generation_time=0,
inference_time=0,
encoding_time=0,
total_time=0,
)
@app.get("/health")
@limiter.limit("10/minute")
async def health_check(self, request: Request) -> Dict[str, str]:
return {"status": "healthy"}
@app.get("/metrics")
async def metrics(self) -> Response:
return Response(generate_latest(), media_type="text/plain")
def validate_configuration(model_paths: List[str], replicas: List[int]) -> None:
assert len(model_paths) == len(replicas), "Number of models and replicas must match"
assert sum(replicas) <= NUM_GPUS, f"Total replicas ({sum(replicas)}) must be <= {NUM_GPUS}"
for model, replica_count in zip(model_paths, replicas):
assert model in SUPPORTED_MODELS, f"Model {model} not supported"
assert replica_count > 0, f"Replicas must be greater than 0"
def start_ray_serve(
*,
t2v_model_paths: str,
t2v_model_replicas: str,
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
) -> None:
if not ray.is_initialized():
ray.init()
model_paths = t2v_model_paths.split(",")
replicas = [int(r) for r in t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
t2v_deps = {}
for model_path, replica_count in zip(model_paths, replicas):
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
api = FastVideoAPI.bind(t2v_deps)
serve.run(api, route_prefix="/", name="fast_video")
print(f"Ray Serve backend started at http://{host}:{port}")
for model_path, replica_count in zip(model_paths, replicas):
print(f"T2V Model: {model_path} | Replicas: {replica_count}")
print(f"Health check: http://{host}:{port}/health")
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
def setup_signal_handlers() -> None:
signal.signal(signal.SIGINT, lambda *_: sys.exit(0))
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0))
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
parser.add_argument("--t2v_model_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--output_path",
type=str,
default="outputs",
help="Path to save generated videos")
parser.add_argument("--host",
type=str,
default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port",
type=int,
default=8000,
help="Port to bind to")
args = parser.parse_args()
model_paths = args.t2v_model_paths.split(",")
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
start_ray_serve(
t2v_model_paths=args.t2v_model_paths,
t2v_model_replicas=args.t2v_model_replicas,
output_path=args.output_path,
host=args.host,
port=args.port,
)
setup_signal_handlers()
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
while True:
time.sleep(3600)
+3
View File
@@ -0,0 +1,3 @@
python examples/inference/gradio/start_ray_serve_app.py \
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" \
--t2v_model_replicas "4,4"
@@ -0,0 +1,257 @@
"""
Startup script for FastVideo with Ray Serve backend and Gradio frontend.
This script starts both the backend and frontend services.
"""
import argparse
import os
import subprocess
import sys
import time
import threading
import signal
import requests
from pathlib import Path
from typing import Dict, Any, Optional
DEFAULT_BACKEND_HOST = "0.0.0.0"
DEFAULT_BACKEND_PORT = 8000
DEFAULT_FRONTEND_HOST = "0.0.0.0"
DEFAULT_FRONTEND_PORT = 7860
DEFAULT_OUTPUT_PATH = "outputs"
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_T2V_REPLICAS = "4,4"
HEALTH_CHECK_TIMEOUT = 5
HEALTH_CHECK_MAX_RETRIES = 100
HEALTH_CHECK_INTERVAL = 2
PROCESS_SHUTDOWN_TIMEOUT = 5
PROCESS_MONITOR_INTERVAL = 1
PROJECT_ROOT = Path(__file__).parent.parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))
class ServiceManager:
def __init__(self, args: argparse.Namespace):
self.args = args
self.backend_process: Optional[subprocess.Popen] = None
self.frontend_process: Optional[subprocess.Popen] = None
self.backend_url = f"http://{args.backend_host}:{args.backend_port}"
def check_backend_health(self, max_retries: int = HEALTH_CHECK_MAX_RETRIES) -> bool:
health_url = f"{self.backend_url}/health"
for attempt in range(max_retries):
try:
response = requests.get(health_url, timeout=HEALTH_CHECK_TIMEOUT)
if response.status_code == 200:
print(f"✅ Backend is healthy at {self.backend_url}")
return True
except requests.exceptions.RequestException:
pass
if attempt < max_retries - 1:
print(f"⏳ Waiting for backend to start... ({attempt + 1}/{max_retries})")
time.sleep(HEALTH_CHECK_INTERVAL)
print(f"❌ Backend failed to start within {max_retries * HEALTH_CHECK_INTERVAL} seconds")
return False
def _create_monitor_thread(self, process: subprocess.Popen, service_name: str) -> threading.Thread:
def monitor():
if process.stdout:
for line in process.stdout:
print(f"[{service_name}] {line.rstrip()}")
thread = threading.Thread(target=monitor, daemon=True)
thread.start()
return thread
def _start_service(self, script_name: str, args_dict: Dict[str, Any], service_name: str) -> subprocess.Popen:
script_path = Path(__file__).parent / script_name
cmd = [sys.executable, str(script_path)]
for key, value in args_dict.items():
cmd.extend([f"--{key}", str(value)])
print(f"🚀 Starting {service_name}...")
print(f"Command: {' '.join(cmd)}")
process = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
universal_newlines=True,
bufsize=1
)
self._create_monitor_thread(process, service_name.upper())
return process
def start_backend(self) -> subprocess.Popen:
backend_args = {
"t2v_model_paths": self.args.t2v_model_paths,
"t2v_model_replicas": self.args.t2v_model_replicas,
"output_path": self.args.output_path,
"host": self.args.backend_host,
"port": self.args.backend_port
}
self.backend_process = self._start_service("ray_serve_backend.py", backend_args, "backend")
return self.backend_process
def start_frontend(self) -> subprocess.Popen:
frontend_args = {
"backend_url": self.backend_url,
"t2v_model_paths": self.args.t2v_model_paths,
"host": self.args.frontend_host,
"port": self.args.frontend_port
}
self.frontend_process = self._start_service("gradio_frontend.py", frontend_args, "frontend")
return self.frontend_process
def shutdown_services(self) -> None:
print("\n🛑 Shutting down services...")
processes = []
if self.frontend_process:
self.frontend_process.terminate()
processes.append(("frontend", self.frontend_process))
if self.backend_process:
self.backend_process.terminate()
processes.append(("backend", self.backend_process))
for name, process in processes:
try:
process.wait(timeout=PROCESS_SHUTDOWN_TIMEOUT)
print(f"✅ {name.capitalize()} stopped gracefully")
except subprocess.TimeoutExpired:
print(f"⚠️ Force killing {name} process...")
process.kill()
print("✅ All services stopped")
def monitor_processes(self) -> None:
if not self.backend_process or not self.frontend_process:
print("❌ Processes not properly initialized")
return
try:
while True:
if self.frontend_process.poll() is not None:
print("❌ Frontend process died unexpectedly")
break
if self.backend_process.poll() is not None:
print("❌ Backend process died unexpectedly")
break
time.sleep(PROCESS_MONITOR_INTERVAL)
except KeyboardInterrupt:
pass
self.shutdown_services()
def setup_signal_handlers(service_manager: ServiceManager) -> None:
def signal_handler(signum: int, frame: Any) -> None:
service_manager.shutdown_services()
sys.exit(0)
signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)
def print_startup_info(args: argparse.Namespace) -> None:
print("🎬 FastVideo Ray Serve App")
print("=" * 50)
print(f"T2V Models: {args.t2v_model_paths}")
print(f"T2V Model Replicas: {args.t2v_model_replicas}")
print(f"Output: {args.output_path}")
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
print("=" * 50)
def parse_arguments() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="FastVideo Ray Serve App")
parser.add_argument("--t2v_model_paths",
type=str,
default=DEFAULT_T2V_MODELS,
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default=DEFAULT_T2V_REPLICAS,
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--output_path",
type=str,
default=DEFAULT_OUTPUT_PATH,
help="Path to save generated videos")
parser.add_argument("--backend_host",
type=str,
default=DEFAULT_BACKEND_HOST,
help="Backend host to bind to")
parser.add_argument("--backend_port",
type=int,
default=DEFAULT_BACKEND_PORT,
help="Backend port to bind to")
parser.add_argument("--frontend_host",
type=str,
default=DEFAULT_FRONTEND_HOST,
help="Frontend host to bind to")
parser.add_argument("--frontend_port",
type=int,
default=DEFAULT_FRONTEND_PORT,
help="Frontend port to bind to")
parser.add_argument("--skip_backend_check",
action="store_true",
help="Skip backend health check")
return parser.parse_args()
def main() -> None:
args = parse_arguments()
os.makedirs(args.output_path, exist_ok=True)
print_startup_info(args)
service_manager = ServiceManager(args)
setup_signal_handlers(service_manager)
try:
service_manager.start_backend()
if not args.skip_backend_check:
if not service_manager.check_backend_health():
print("❌ Backend failed to start. Terminating...")
service_manager.shutdown_services()
sys.exit(1)
service_manager.start_frontend()
print("\n🎉 Both services are starting up!")
print(f"📺 Frontend will be available at: http://{args.frontend_host}:{args.frontend_port}")
print(f"🔧 Backend API will be available at: http://{args.backend_host}:{args.backend_port}")
print("\nPress Ctrl+C to stop both services...")
service_manager.monitor_processes()
except Exception as e:
print(f"❌ Unexpected error: {e}")
service_manager.shutdown_services()
sys.exit(1)
if __name__ == "__main__":
main()
@@ -14,9 +14,10 @@ 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-1000/transformer",
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-160/transformer",
lora_nickname="crush_smol"
)
generator.unmerge_lora_weights()
kwargs = {
"height": 480,
"width": 832,
@@ -7,9 +7,7 @@ These are e2e example scripts for finetuning Wan2.1 T2V with VSA to accelerate i
## Make sure you have installed VSA
```bash
cd csrc/attn
git submodule update --init --recursive
python setup_vsa.py install
pip install vsa
```
### Download the synthetic dataset:
@@ -0,0 +1,24 @@
#!/bin/bash
GPU_NUM=2 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
torchrun --nproc_per_node=$GPU_NUM \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type i2v \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 77 \
--preprocess.train_fps 16 \
--preprocess.samples_per_file 8 \
--preprocess.flush_frequency 8 \
--preprocess.video_length_tolerance_range 5
@@ -7,7 +7,7 @@ export WANDB_MODE=online
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=2
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
@@ -76,6 +76,7 @@ miscellaneous_args=(
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--resume_from_checkpoint "checkpoints/wan_t2v_finetune_lora/checkpoint-160"
)
torchrun \
@@ -1,19 +1,20 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
GPU_NUM=2 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATASET_PATH="data/crush-smol-test/"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocessing_new.py \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--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 4 \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
+79
View File
@@ -1,9 +1,59 @@
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:
@@ -12,6 +62,7 @@ class PreprocessConfig:
# Model and dataset configuration
model_path: str = ""
dataset_path: str = ""
dataset_type: DatasetType = DatasetType.HF
dataset_output_dir: str = "./output"
# Dataloader configuration
@@ -23,6 +74,7 @@ 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
@@ -35,6 +87,9 @@ class PreprocessConfig:
# Model configuration
training_cfg_rate: float = 0.0
# framework configuration
seed: int = 42
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "preprocess") -> FlexibleArgumentParser:
@@ -52,6 +107,12 @@ class PreprocessConfig:
type=str,
default=PreprocessConfig.dataset_path,
help="Path to the dataset directory for preprocessing")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-type",
type=str,
choices=DatasetType.choices(),
default=PreprocessConfig.dataset_type.value,
help="Type of the dataset")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-output-dir",
type=str,
@@ -83,6 +144,12 @@ 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,
@@ -123,6 +190,10 @@ class PreprocessConfig:
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
type=int,
default=PreprocessConfig.seed,
help="Seed for random number generator")
return parser
@@ -130,6 +201,14 @@ 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):
@@ -92,6 +92,12 @@ 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
+3
View File
@@ -85,6 +85,9 @@ 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
+12 -7
View File
@@ -7,10 +7,14 @@ 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
from fastvideo.configs.pipelines.wan import (FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
# isort: off
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
WanT2V720PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -31,9 +35,10 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": FastWan2_2_TI2V_5B_Config,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
"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,
# Add other specific weight variants
}
+10
View File
@@ -138,3 +138,13 @@ 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
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
+20 -3
View File
@@ -10,6 +10,7 @@ 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,
@@ -17,6 +18,7 @@ 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
@@ -28,16 +30,31 @@ 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,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-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,6 +107,21 @@ 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 =============
# =============================================
@@ -141,3 +156,11 @@ 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
+9 -3
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
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.validation_dataset import ValidationDataset
@@ -39,7 +39,13 @@ def getdataset(args) -> VideoCaptionMergedDataset:
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"
]
"VideoCaptionMergedDataset", "TextDataset"
]
+73
View File
@@ -78,3 +78,76 @@ pyarrow_schema_t2v = pa.schema([
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_ode_trajectory = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", 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()),
# I2V
pa.field("image_condition_latents_bytes", pa.binary()),
pa.field("image_condition_latents_shape", pa.list_(pa.int64())),
pa.field("image_condition_latents_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()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
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
])
pyarrow_schema_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()),
])
+131
View File
@@ -628,3 +628,134 @@ 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"]
+19 -2
View File
@@ -9,6 +9,7 @@ diffusion models.
import math
import os
import time
from copy import deepcopy
from typing import Any
import imageio
@@ -202,6 +203,8 @@ class VideoGenerator:
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
else:
sampling_param = deepcopy(sampling_param)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
@@ -275,6 +278,7 @@ class VideoGenerator:
width: {target_width}
video_length: {sampling_param.num_frames}
prompt: {prompt}
image_path: {sampling_param.image_path}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
infer_steps: {sampling_param.num_inference_steps}
@@ -304,7 +308,8 @@ class VideoGenerator:
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
@@ -334,9 +339,11 @@ class VideoGenerator:
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time
"generation_time": gen_time,
"logging_info": logging_info,
}
def set_lora_adapter(self,
@@ -344,6 +351,16 @@ class VideoGenerator:
lora_path: str | None = None) -> None:
self.executor.set_lora_adapter(lora_nickname, lora_path)
def unmerge_lora_weights(self) -> None:
"""
Use unmerged weights for inference to produce videos that align with
validation videos generated during training.
"""
self.executor.unmerge_lora_weights()
def merge_lora_weights(self) -> None:
self.executor.merge_lora_weights()
def shutdown(self):
"""
Shutdown the video generator.
+1 -1
View File
@@ -76,7 +76,7 @@ class BaseLayerWithLoRA(nn.Module):
lora_B = self.lora_B.to_local()
lora_A = self.lora_A.to_local()
if (self.training_mode or not self.merged) and not self.disable_lora:
if not self.merged and not self.disable_lora:
delta = x @ (
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
+9
View File
@@ -29,6 +29,9 @@ 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:
@@ -267,6 +270,7 @@ 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.
@@ -292,6 +296,9 @@ 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
@@ -370,6 +377,7 @@ 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.
@@ -413,6 +421,7 @@ 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
+5 -1
View File
@@ -86,12 +86,16 @@ class TimestepEmbedder(nn.Module):
dtype=dtype)
self.freq_dtype = freq_dtype
def forward(self, t: torch.Tensor) -> torch.Tensor:
def forward(self,
t: torch.Tensor,
timestep_seq_len: int | None = None) -> 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
+2 -2
View File
@@ -120,14 +120,14 @@ def _info(logger: Logger,
if not _warned_local_main_process and local_main_process_only:
logger.warning(
'%s is_local_main_process is set to True, logging only from the local main process.%s',
'%s By default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s',
GREEN,
RESET,
)
_warned_local_main_process = True
if not _warned_main_process and main_process_only:
logger.warning(
'%s is_main_process_only is set to True, logging only from the main process.%s',
'%s is_main_process_only is set to True, logging only from the main (RANK==0) process.%s',
GREEN,
RESET,
)
+648
View File
@@ -0,0 +1,648 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import numpy as np
import torch
import torch.nn as nn
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
from torch.nn.attention.flex_attention import BlockMask
# wan 1.3B model has a weird channel / head configurations and require max-autotune to work with flexattention
# see https://github.com/pytorch/pytorch/issues/133254
# change to default for other models
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
import torch.distributed as dist
import fastvideo.envs as envs
from fastvideo.attention import (DistributedAttention,
LocalAttention)
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
def __init__(self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm=True,
eps=1e-6,
parallel_attention=False) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.parallel_attention = parallel_attention
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Scaled dot product attention
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
seq_lens(Tensor): Shape [B]
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
if cache_start is None:
cache_start = current_start
cos, sin = freqs_cis
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
if kv_cache is None:
# Padding for flex attention
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_roped_query = torch.cat(
[roped_query,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_roped_key = torch.cat(
[roped_key, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
x = flex_attention(
query=padded_roped_query.transpose(2, 1),
key=padded_roped_key.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask
)[:, :, :-padded_length].transpose(2, 1)
else:
frame_seqlen = q.shape[1]
current_end = current_start + roped_query.shape[1]
sink_tokens = self.sink_size * frame_seqlen
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
else:
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
roped_query,
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
return x
class CausalWanTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = CausalWanSelfAttention(
dim,
num_heads,
local_attn_size=local_attn_size,
sink_size=sink_size,
qk_norm=qk_norm,
eps=eps)
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.local_attn_size = local_attn_size
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
print("QK Norm type not supported")
raise Exception
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
attn_output = self.attn1(query, key, value, freqs_cis, block_mask, kv_cache, current_start, cache_start)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None,
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class CausalWanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
param_names_mapping = WanVideoConfig().param_names_mapping
reverse_param_names_mapping = WanVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = WanVideoConfig().lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
self.local_attn_size = config.local_attn_size
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
CausalWanTransformerBlock(inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.local_attn_size,
config.sink_size,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.independent_first_frame = False
self.__post_init__()
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
) -> BlockMask:
"""
we will divide the token sequence into the following format
[1 latent frame] [1 latent frame] ... [1 latent frame]
We use flexattention to construct the attention mask
"""
total_length = num_frames * frame_seqlen
# we do right padding to get to a multiple of 128
padded_length = math.ceil(total_length / 128) * 128 - total_length
ends = torch.zeros(total_length + padded_length,
device=device, dtype=torch.long)
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
frame_indices = torch.arange(
start=0,
end=total_length,
step=frame_seqlen * num_frame_per_block,
device=device
)
for tmp in frame_indices:
ends[tmp:tmp + frame_seqlen * num_frame_per_block] = tmp + \
frame_seqlen * num_frame_per_block
def attention_mask(b, h, q_idx, kv_idx):
if local_attn_size == -1:
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)
else:
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx)
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
KV_LEN=total_length + padded_length, _compile=False, device=device)
if not dist.is_initialized() or dist.get_rank() == 0:
print(
f" cache a block wise causal mask with block size of {num_frame_per_block} frames")
print(block_mask)
# import imageio
# import numpy as np
# from torch.nn.attention.flex_attention import create_mask
# mask = create_mask(attention_mask, B=None, H=None, Q_LEN=total_length +
# padded_length, KV_LEN=total_length + padded_length, device=device)
# import cv2
# mask = cv2.resize(mask[0, 0].cpu().float().numpy(), (1024, 1024))
# imageio.imwrite("mask_%d.jpg" % (0), np.uint8(255. * mask))
return block_mask
def _forward_inference(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
kv_cache: dict = None,
crossattn_cache: dict = None,
current_start: int = 0,
cache_start: int = 0,
start_frame: int = 0,
**kwargs) -> torch.Tensor:
r"""
Run the diffusion model with kv caching.
See Algorithm 2 of CausVid paper https://arxiv.org/abs/2412.07772 for details.
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame # Assume that start_frame is 0 when kv_cache is None
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
encoder_hidden_states = encoder_hidden_states.to(
orig_dtype) if current_platform.is_mps(
) else encoder_hidden_states # cast to orig_dtype for MPS
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
for block_index, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
causal_kwargs = {
"kv_cache": kv_cache[block_index],
"current_start": current_start,
"cache_start": cache_start,
"block_mask": self.block_mask
}
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
**causal_kwargs)
else:
causal_kwargs = {
"kv_cache": kv_cache[block_index],
"crossattn_cache": crossattn_cache[block_index],
"current_start": current_start,
"cache_start": cache_start,
"block_mask": self.block_mask
}
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
**causal_kwargs)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
def _forward_train(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
start_frame: int = 0,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size
)
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
encoder_hidden_states = encoder_hidden_states.to(
orig_dtype) if current_platform.is_mps(
) else encoder_hidden_states # cast to orig_dtype for MPS
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
def forward(
self,
*args,
**kwargs
):
if kwargs.get('kv_cache', None) is not None:
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
+59 -11
View File
@@ -81,8 +81,9 @@ class WanTimeTextImageEmbedding(nn.Module):
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep)
temb = self.time_embedder(timestep, timestep_seq_len)
timestep_proj = self.time_modulation(temb)
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
@@ -145,7 +146,7 @@ class WanSelfAttention(nn.Module):
class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens):
def forward(self, x, context, context_lens, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -156,8 +157,20 @@ class WanT2VCrossAttention(WanSelfAttention):
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
@@ -307,9 +320,24 @@ class WanTransformerBlock(nn.Module):
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
@@ -637,9 +665,21 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten() # batch_size * seq_len
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
# batch_size, 6, inner_dim
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -676,8 +716,15 @@ class WanTransformer3DModel(CachableDiT):
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
if temb.dim() == 3:
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2)
shift = shift.squeeze(2)
scale = scale.squeeze(2)
else:
# batch_size, inner_dim
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -781,3 +828,4 @@ class WanTransformer3DModel(CachableDiT):
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
+2
View File
@@ -25,12 +25,14 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel")
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
}
_TEXT_ENCODER_MODELS = {
+21
View File
@@ -137,3 +137,24 @@ def modulate(x: torch.Tensor,
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(
1) # type: ignore[union-attr]
def pred_noise_to_pred_video(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
"""
timestep = timestep.expand(noise_input_latent.shape[0])
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
noise_input_latent = noise_input_latent.float().to(device)
sigmas = scheduler.sigmas.float().to(device)
timesteps = scheduler.timesteps.float().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
@@ -0,0 +1,69 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan causal DMD pipeline implementation.
This module wires the causal DMD denoising stage into the modular pipeline.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
CausalDMDDenosingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
# isort: on
logger = init_logger(__name__)
class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanCausalDMDPipeline
@@ -12,12 +12,10 @@ from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ImageEncodingStage, ConditioningStage,
DecodingStage, DmdDenoisingStage,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
@@ -67,7 +65,7 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
@@ -63,6 +63,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
self.add_stage(stage_name="decoding_stage",
@@ -282,7 +282,11 @@ class ComposedPipelineBase(ABC):
for module_name, (transformers_or_diffusers,
architecture) in model_index.items():
if transformers_or_diffusers is None:
self.required_config_modules.remove(module_name)
logger.warning(
"Module in model_index.json has null value, removing from required_config_modules"
)
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
+8
View File
@@ -217,3 +217,11 @@ class LoRAPipeline(ComposedPipelineBase):
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
def merge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.merge_lora_weights()
def unmerge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.unmerge_lora_weights()
+43 -2
View File
@@ -17,10 +17,47 @@ import torch
if TYPE_CHECKING:
from torchcodec.decoders import VideoDecoder
import time
from collections import OrderedDict
from fastvideo.attention import AttentionMetadata
from fastvideo.configs.sample.teacache import TeaCacheParams, WanTeaCacheParams
class PipelineLoggingInfo:
"""Simple approach using OrderedDict to track stage metrics."""
def __init__(self):
# OrderedDict preserves insertion order and allows easy access
self.stages: OrderedDict[str, dict[str, Any]] = OrderedDict()
def add_stage_execution_time(self, stage_name: str, execution_time: float):
"""Add execution time for a stage."""
if stage_name not in self.stages:
self.stages[stage_name] = {}
self.stages[stage_name]['execution_time'] = execution_time
self.stages[stage_name]['timestamp'] = time.time()
def add_stage_metric(self, stage_name: str, metric_name: str, value: Any):
"""Add any metric for a stage."""
if stage_name not in self.stages:
self.stages[stage_name] = {}
self.stages[stage_name][metric_name] = value
def get_stage_info(self, stage_name: str) -> dict[str, Any]:
"""Get all info for a specific stage."""
return self.stages.get(stage_name, {})
def get_execution_order(self) -> list[str]:
"""Get stages in execution order."""
return list(self.stages.keys())
def get_total_execution_time(self) -> float:
"""Get total pipeline execution time."""
return sum(
stage.get('execution_time', 0) for stage in self.stages.values())
@dataclass
class ForwardBatch:
"""
@@ -40,7 +77,7 @@ class ForwardBatch:
# Image inputs
image_path: str | None = None
image_embeds: list[torch.Tensor] = field(default_factory=list)
pil_image: PIL.Image.Image | None = None
pil_image: torch.Tensor | PIL.Image.Image | None = None
preprocessed_image: torch.Tensor | None = None
# Text inputs
@@ -132,6 +169,10 @@ class ForwardBatch:
# VSA parameters
VSA_sparsity: float = 0.0
# Logging info
logging_info: PipelineLoggingInfo = field(
default_factory=PipelineLoggingInfo)
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
@@ -200,5 +241,5 @@ class TrainingBatch:
@dataclass
class PreprocessBatch(ForwardBatch):
video_loader: list["VideoDecoder"] = field(default_factory=list)
video_loader: list["VideoDecoder"] | list[str] = field(default_factory=list)
video_file_name: list[str] = field(default_factory=list)
+8 -6
View File
@@ -21,10 +21,16 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
WorkloadType.I2V: "PreprocessPipelineI2V",
WorkloadType.T2V: "PreprocessPipelineT2V",
}
class PipelineType(str, Enum):
"""
@@ -68,12 +74,8 @@ class _PipelineRegistry:
def _load_preprocess_pipeline_cls(
self, workload_type: WorkloadType,
arch: str) -> type[ComposedPipelineBase] | None:
if workload_type == WorkloadType.I2V:
pipeline_name = "I2VPreprocessPipeline"
elif workload_type == WorkloadType.T2V:
pipeline_name = "T2VPreprocessPipeline"
else:
raise ValueError(f"Invalid workload type: {workload_type.value}")
pipeline_name = _PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME[
workload_type]
return self.pipelines[
PipelineType.PREPROCESS.value][arch][pipeline_name]
@@ -0,0 +1,762 @@
# SPDX-License-Identifier: Apache-2.0
"""
ODE Trajectory Data Preprocessing pipeline implementation.
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
using the modular pipeline architecture.
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
"""
import os
from collections.abc import Iterator
from typing import Any
import numpy as np
import pyarrow as pa
import torch
from PIL import Image
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import getdataset, gettextdataset
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory, pyarrow_schema_ode_trajectory_text_only
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.utils import shallow_asdict, save_decoded_latents_as_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DenoisingStage, ImageVAEEncodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
DecodingStage)
logger = init_logger(__name__)
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
def get_schema_fields(self):
"""Get the schema fields for ODE Trajectory pipeline."""
# Check if we're using text dataset by checking if the dataset is TextDataset
if hasattr(self, 'preprocess_dataloader') and hasattr(self.preprocess_dataloader.dataset, '_process_text_data'):
return [f.name for f in pyarrow_schema_ode_trajectory_text_only]
return [f.name for f in pyarrow_schema_ode_trajectory]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="vae_encoding_stage",
stage=ImageVAEEncodingStage(
vae=self.get_module("vae"), ))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def preprocess_video_and_text_and_trajectory(self,
fastvideo_args: FastVideoArgs,
args):
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
get_local_torch_device())).mean
# Get extra features if needed
extra_features = self.get_extra_features(
valid_data, fastvideo_args)
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
# # Get sequence lengths from attention masks (number of 1s)
# seq_lens = prompt_attention_mask.sum(dim=1)
# non_padded_embeds = []
# non_padded_masks = []
# # Process each item in the batch
# for i in range(prompt_embeds.size(0)):
# seq_len = seq_lens[i].item()
# # Slice the embeddings and masks to keep only non-padding parts
# non_padded_embeds.append(prompt_embeds[i, :seq_len])
# non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
# prompt_embeds = non_padded_embeds
# prompt_attention_masks = non_padded_masks
# prompt_embeds = prompt_embeds
# logger.info(f"===== prompt_embeds: {prompt_embeds[0].shape}")
# logger.info(f"===== prompt_attention_masks: {prompt_attention_masks[0].shape}")
sampling_params = SamplingParam.from_pretrained(
args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(zip(prompt_embeds, prompt_attention_masks)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
logger.info(f"what")
logger.info(f"===== prompt_embed: {prompt_embed.shape}")
logger.info(f"===== prompt_attention_mask: {prompt_attention_mask.shape}")
# Collect the trajectory data
batch = ForwardBatch(
**shallow_asdict(sampling_params),
# data_type="video",
# seed=args.seed,
# prompt=batch_captions[i],
# prompt_embeds=[prompt_embed],
# prompt_attention_mask=[prompt_attention_mask],
# height=args.max_height,
# width=args.max_width,
# num_frames=81,
# fps=args.train_fps,
# return_trajectory_latents=True,
# guidance_scale=3.0,
# do_classifier_free_guidance=True,
)
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [negative_prompt_attention_mask]
batch.return_trajectory_latents = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
# batch.num_frames = 81
batch.fps = args.train_fps
batch.guidance_scale = 3.0
batch.do_classifier_free_guidance = True
# fastvideo_args.pipeline_config.ti2v_task = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
# result_batch = self.prompt_encoding_stage(result_batch, fastvideo_args)
# result_batch = self.vae_encoding_stage(result_batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch, fastvideo_args)
# trajectory_latents = result_batch.trajectory_latents
trajectory_latents.append(result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
extra_features["trajectory_latents"] = trajectory_latents
extra_features["trajectory_timesteps"] = trajectory_timesteps
logger.info(f"===== trajectory_latents: {trajectory_latents[0].shape}")
logger.info(f"===== trajectory_latents len: {len(trajectory_latents)}")
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
logger.info(f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
if batch.return_trajectory_decoded:
logger.info(f"===== SAVING TRAJECTORY DECODED")
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
logger.info(f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}")
save_decoded_latents_as_video(decoded_frame, f"decoded_videos/trajectory_decoded_{i}_{j}.mp4", args.train_fps)
# assert False
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
logger.info(f"===== key: {key}")
if isinstance(value, torch.Tensor):
logger.info(f"===== value: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().numpy(
)
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
logger.info(f"===== value in list: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().float().numpy(
)
else:
logger.info(f"===== value in list: not tensor")
sample_extra_features[key] = value[idx]
# logger.info(f"===== value: not tensor")
# sample_extra_features[key] = value[idx]
# Create record for Parquet dataset
record = self.create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def preprocess_text_and_trajectory(self,
fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data and generate trajectory information."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
# Add fps and duration if available in data
if "fps" in data:
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
if "duration" in data:
valid_data["duration"] = [data["duration"][i] for i in valid_indices]
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(
args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(zip(prompt_embeds, prompt_attention_masks)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(
**shallow_asdict(sampling_params),
)
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [negative_prompt_attention_mask]
batch.return_trajectory_latents = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 3.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch, fastvideo_args)
trajectory_latents.append(result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
extra_features = {
"trajectory_latents": trajectory_latents,
"trajectory_timesteps": trajectory_timesteps
}
logger.info(f"===== trajectory_latents: {trajectory_latents[0].shape}")
logger.info(f"===== trajectory_latents len: {len(trajectory_latents)}")
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
logger.info(f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
if batch.return_trajectory_decoded:
logger.info(f"===== SAVING TRAJECTORY DECODED")
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
logger.info(f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}")
save_decoded_latents_as_video(decoded_frame, f"decoded_videos/trajectory_decoded_{i}_{j}.mp4", args.train_fps)
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
logger.info(f"===== key: {key}")
if isinstance(value, torch.Tensor):
logger.info(f"===== value: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().numpy()
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
logger.info(f"===== value in list: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().float().numpy()
else:
logger.info(f"===== value in list: not tensor")
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (without VAE latents for text-only)
record = self.create_text_only_record(
args,
video_name=video_name,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
# Final flush for any remaining samples
if hasattr(self, 'all_tables') and self.all_tables and self.num_processed_samples > 0:
logger.info(f"Final flush with {self.num_processed_samples} remaining samples")
self._flush_tables(self.num_processed_samples, args, self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def create_text_only_record(
self,
args,
video_name: str,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{video_name}_{idx}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "text",
}
# Add trajectory data if available
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features["trajectory_latents"][idx] if isinstance(extra_features["trajectory_latents"], list) else extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes": trajectory_latents.tobytes(),
"trajectory_latents_shape": list(trajectory_latents.shape),
"trajectory_latents_dtype": str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features["trajectory_timesteps"][idx] if isinstance(extra_features["trajectory_timesteps"], list) else extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
return record
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("vae").to(get_local_torch_device())
# generator = torch.Generator(device=get_local_torch_device(), seed=42)
generator = torch.Generator("cpu").manual_seed(42)
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
unprocessed_images = []
pil_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
pil_images.append(frame_pil)
# processed_img = self.get_module("image_processor")(
# images=frame_pil, return_tensors="pt")
unprocessed_images.append(frame_pil)
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in unprocessed_images:
latent = self.vae_encoding_stage.encode_image(
frame, height, width, fastvideo_args, generator)
video_conditions.append(latent)
features["image_condition_latents"] = video_conditions
features["pil_images"] = pil_images
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "image_condition_latents" in extra_features:
image_condition_latents = extra_features["image_condition_latents"]
record.update({
"image_condition_latents_bytes":
image_condition_latents.tobytes(),
"image_condition_latents_shape":
list(image_condition_latents.shape),
"image_condition_latents_dtype":
str(image_condition_latents.dtype),
})
else:
record.update({
"image_condition_latents_bytes": b"",
"image_condition_latents_shape": [],
"image_condition_latents_dtype": "",
})
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes": trajectory_latents.tobytes(),
"trajectory_latents_shape": list(trajectory_latents.shape),
"trajectory_latents_dtype": str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
#train_dataset = getdataset(args)
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for video preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing videos",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.video_data: dict[str, Any] = {} # Store video metadata and paths
self.latent_data: dict[str, Any] = {} # Store latent tensors
#self.preprocess_video_and_text_and_trajectory(fastvideo_args, args)
self.preprocess_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_ODE_Trajectory
@@ -1,14 +1,17 @@
import random
from collections.abc import Callable
from typing import cast
import numpy as np
import torch
import torchvision
from einops import rearrange
from torchvision import transforms
from fastvideo.configs.configs import VideoLoaderType
from fastvideo.dataset.transform import (CenterCropResizeVideo,
TemporalRandomCrop)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
from fastvideo.pipelines.pipeline_batch_info import (ForwardBatch,
PreprocessBatch)
from fastvideo.pipelines.stages.base import PipelineStage
@@ -60,7 +63,16 @@ class VideoTransformStage(PipelineStage):
else:
frame_indices = frame_indices[:self.num_frames]
video = batch.video_loader[i].get_frames_at(frame_indices).data
if fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
video = batch.video_loader[i].get_frames_at(frame_indices).data
elif fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHVISION:
video, _, _ = torchvision.io.read_video(batch.video_loader[i],
output_format="TCHW")
video = video[frame_indices]
else:
raise ValueError(
f"Invalid video loader type: {fastvideo_args.preprocess_config.video_loader_type}"
)
video = self.video_transform(video)
video_pixel_batch.append(video)
@@ -68,6 +80,39 @@ class VideoTransformStage(PipelineStage):
video_pixel_values = rearrange(video_pixel_values,
"b t c h w -> b c t h w")
video_pixel_values = video_pixel_values.to(torch.uint8)
if fastvideo_args.workload_type == WorkloadType.I2V:
batch.pil_image = video_pixel_values[:, :, 0, :, :]
video_pixel_values = video_pixel_values.float() / 255.0
batch.latents = video_pixel_values
batch.num_frames = [video_pixel_values.shape[2]] * len(
batch.video_loader)
batch.height = [video_pixel_values.shape[3]] * len(batch.video_loader)
batch.width = [video_pixel_values.shape[4]] * len(batch.video_loader)
return cast(ForwardBatch, batch)
class TextTransformStage(PipelineStage):
"""
Process text data according to the cfg rate.
"""
def __init__(self, cfg_uncondition_drop_rate: float, seed: int) -> None:
self.cfg_rate = cfg_uncondition_drop_rate
self.rng = random.Random(seed)
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
prompts = []
for prompt in batch.prompt:
if not isinstance(prompt, list):
prompt = [prompt]
prompt = self.rng.choice(prompt)
prompt = prompt if self.rng.random() > self.cfg_rate else ""
prompts.append(prompt)
batch.prompt = prompts
return cast(ForwardBatch, batch)
@@ -9,8 +9,11 @@ from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess_text import PreprocessPipeline_Text
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -21,12 +24,22 @@ def main(args) -> None:
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config.update_config_from_dict(kwargs)
if args.preprocess_task == "text_only":
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"text_encoder_cpu_offload": False,
}
pipeline_config.update_config_from_dict(kwargs)
else:
# Full config for video/image processing
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
@@ -35,7 +48,25 @@ def main(args) -> None:
text_encoder_cpu_offload=False,
pipeline_config=pipeline_config,
)
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
if args.preprocess_task == "t2v":
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "ode_trajectory":
print("Preprocess pipeline...")
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
elif args.preprocess_task == "text_only":
print("Text-only preprocessing pipeline...")
PreprocessPipeline = PreprocessPipeline_Text
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
logger.info(
f"Preprocess task: {args.preprocess_task} using {PreprocessPipeline.__name__}"
)
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -74,7 +105,11 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "ode_trajectory", "text_only"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
@@ -98,4 +133,4 @@ if __name__ == "__main__":
)
args = parser.parse_args()
main(args)
main(args)
@@ -1,21 +1,17 @@
import os
from fastvideo.distributed import (
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser
from fastvideo.workflow.preprocess.preprocess_workflow_t2v import (
PreprocessWorkflowT2V)
from fastvideo.workflow.workflow_base import WorkflowBase
logger = init_logger(__name__)
def main(fastvideo_args: FastVideoArgs) -> None:
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
preprocess_workflow = PreprocessWorkflowT2V(fastvideo_args)
preprocess_workflow_cls = WorkflowBase.get_workflow_cls(fastvideo_args)
preprocess_workflow = preprocess_workflow_cls(fastvideo_args)
preprocess_workflow.run()
@@ -1,36 +1,66 @@
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.preprocess.preprocess_stages import VideoTransformStage
from fastvideo.pipelines.preprocess.preprocess_stages import (
TextTransformStage, VideoTransformStage)
from fastvideo.pipelines.stages import (EncodingStage, ImageEncodingStage,
TextEncodingStage)
from fastvideo.pipelines.stages.image_encoding import ImageVAEEncodingStage
class I2VPreprocessPipeline(ComposedPipelineBase):
class PreprocessPipelineI2V(ComposedPipelineBase):
_required_config_modules = [
"image_encoder", "image_processor", "text_encoder", "tokenizer", "vae"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
assert fastvideo_args.preprocess_config is not None
self.add_stage(stage_name="text_transform_stage",
stage=TextTransformStage(
cfg_uncondition_drop_rate=fastvideo_args.
preprocess_config.training_cfg_rate,
seed=fastvideo_args.preprocess_config.seed,
))
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(
stage_name="video_transform_stage",
stage=VideoTransformStage(
train_fps=fastvideo_args.preprocess_config.train_fps,
num_frames=fastvideo_args.preprocess_config.num_frames,
max_height=fastvideo_args.preprocess_config.max_height,
max_width=fastvideo_args.preprocess_config.max_width,
do_temporal_sample=fastvideo_args.preprocess_config.
do_temporal_sample,
))
if (self.get_module("image_encoder") is not None
and self.get_module("image_processor") is not None):
self.add_stage(
stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="image_vae_encoding_stage",
stage=ImageVAEEncodingStage(
vae=self.get_module("vae"), ))
self.add_stage(stage_name="video_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae"), ))
class T2VPreprocessPipeline(ComposedPipelineBase):
class PreprocessPipelineT2V(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
assert fastvideo_args.preprocess_config is not None
self.add_stage(stage_name="text_transform_stage",
stage=TextTransformStage(
cfg_uncondition_drop_rate=fastvideo_args.
preprocess_config.training_cfg_rate,
seed=fastvideo_args.preprocess_config.seed,
))
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
@@ -50,4 +80,4 @@ class T2VPreprocessPipeline(ComposedPipelineBase):
stage=EncodingStage(vae=self.get_module("vae"), ))
EntryClass = [I2VPreprocessPipeline, T2VPreprocessPipeline]
EntryClass = [PreprocessPipelineI2V, PreprocessPipelineT2V]
+216
View File
@@ -0,0 +1,216 @@
# SPDX-License-Identifier: Apache-2.0
"""
Text-only Data Preprocessing pipeline implementation.
This module contains an implementation of the Text-only Data Preprocessing pipeline
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
"""
import os
from collections.abc import Iterator
from typing import Any
import numpy as np
import pyarrow as pa
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (TextEncodingStage)
logger = init_logger(__name__)
class PreprocessPipeline_Text(BasePreprocessPipeline):
"""Text-only preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
def get_schema_fields(self):
"""Get the schema fields for text-only pipeline."""
return [f.name for f in pyarrow_schema_text_only]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
def preprocess_text_only(self,
fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
logger.info(f"===== prompt_embeds: {prompt_embeds.shape}")
logger.info(f"===== prompt_attention_masks: {prompt_attention_masks.shape}")
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, text_path in save_pbar:
text_name = os.path.basename(text_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Create record for Parquet dataset (text-only)
record = self.create_text_only_record(
text_name=text_name,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
# Final flush for any remaining samples
if hasattr(self, 'all_tables') and self.all_tables and self.num_processed_samples > 0:
logger.info(f"Final flush with {self.num_processed_samples} remaining samples")
self._flush_tables(self.num_processed_samples, args, self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def create_text_only_record(
self,
text_name: str,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int) -> dict[str, Any]:
"""Create a record for text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{text_name}_{idx}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
}
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading text dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for text preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing text",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.text_data: dict[str, Any] = {} # Store text metadata and paths
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text

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