Compare commits
26
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
44840ac49d | ||
|
|
c99b1d4d97 | ||
|
|
edfe4dd1bf | ||
|
|
93ebd15a0d | ||
|
|
1110474065 | ||
|
|
80baffd540 | ||
|
|
918180048e | ||
|
|
b7dbd7cb9e | ||
|
|
71159b6416 | ||
|
|
ac11127397 | ||
|
|
e028dcc7c0 | ||
|
|
076f45c1ee | ||
|
|
85eb7265db | ||
|
|
d3ceb67e66 | ||
|
|
7ac153a5ca | ||
|
|
d1e7aa0abd | ||
|
|
2d846c55a1 | ||
|
|
b318063c0a | ||
|
|
4aa307be55 | ||
|
|
055e52e5ea | ||
|
|
7d2069596b | ||
|
|
c45009c9a4 | ||
|
|
b91020b407 | ||
|
|
2dcc5ea4f6 | ||
|
|
359151d9a0 | ||
|
|
ce67cd3729 |
+41
-18
@@ -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:
|
||||
@@ -176,3 +176,26 @@ steps:
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VMoBA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/vmoba/**"
|
||||
- "fastvideo/attention/backends/vmoba.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests VMoBA"
|
||||
env:
|
||||
- TEST_TYPE=inference_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -109,6 +109,15 @@ case "$TEST_TYPE" in
|
||||
log "Running distillation DMD tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
|
||||
;;
|
||||
# run_inference_tests_vmoba
|
||||
"inference_vmoba")
|
||||
log "Running V-MoBA inference tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
|
||||
;;
|
||||
"precision_vmoba")
|
||||
log "Running V-MoBA precision tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -64,3 +64,4 @@ docs/source/distillation/examples/
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
dmd_t2v_output/
|
||||
+6
-2
@@ -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
|
||||
|
||||
@@ -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/rG0QpZdw" 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">
|
||||
|
||||
+2
-2
@@ -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
|
||||
```
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
Submodule
+1
Submodule csrc/attn/sliding_tile_attn/tk added at 6c27e28c81
-1
Submodule csrc/attn/tk deleted from 1719fb7264
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
@@ -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/')))
|
||||
Submodule
+1
Submodule csrc/attn/video_sparse_attn/tk added at 6c27e28c81
@@ -0,0 +1,32 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
|
||||
|
||||
### Installation
|
||||
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
|
||||
|
||||
### Usage
|
||||
|
||||
You can use `moba_attn_varlen` in the following ways:
|
||||
|
||||
**Install from source:**
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
**Import after installation:**
|
||||
```python
|
||||
from vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
**Or import directly from the project root:**
|
||||
```python
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
python csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
```
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from setuptools import find_packages, setup
|
||||
|
||||
PACKAGE_NAME = "vmoba"
|
||||
VERSION = "0.0.0"
|
||||
AUTHOR = "JianzongWu"
|
||||
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
|
||||
URL = "https://github.com/KwaiVGI/VMoBA"
|
||||
|
||||
setup(
|
||||
name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[]
|
||||
)
|
||||
@@ -0,0 +1,97 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
|
||||
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
|
||||
"""
|
||||
Generates random data for testing the variable-length attention function.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
random.seed(42)
|
||||
torch.cuda.manual_seed_all(42)
|
||||
|
||||
# Generate sequence lengths for each item in the batch
|
||||
if batch_size > 1:
|
||||
# Ensure sequence lengths are reasonably distributed
|
||||
avg_seqlen = total_seqlen // batch_size
|
||||
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
|
||||
remaining_len = total_seqlen - sum(seqlens)
|
||||
if remaining_len > 0:
|
||||
seqlens.append(remaining_len)
|
||||
else: # Adjust if sum exceeds total_seqlen
|
||||
seqlens.append(avg_seqlen)
|
||||
current_sum = sum(seqlens)
|
||||
seqlens[-1] -= (current_sum - total_seqlen)
|
||||
# Ensure all lengths are positive
|
||||
seqlens = [max(1, s) for s in seqlens]
|
||||
# Final adjustment to match total_seqlen
|
||||
seqlens[-1] += total_seqlen - sum(seqlens)
|
||||
|
||||
else:
|
||||
seqlens = [total_seqlen]
|
||||
|
||||
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
|
||||
max_seqlen = max(seqlens) if seqlens else 0
|
||||
|
||||
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [1, 2])
|
||||
@pytest.mark.parametrize("total_seqlen", [512, 1024])
|
||||
@pytest.mark.parametrize("num_heads", [8])
|
||||
@pytest.mark.parametrize("head_dim", [64])
|
||||
@pytest.mark.parametrize("moba_chunk_size", [64])
|
||||
@pytest.mark.parametrize("moba_topk", [2, 4])
|
||||
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
|
||||
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
|
||||
def test_moba_attn_varlen_forward(
|
||||
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
|
||||
):
|
||||
"""
|
||||
Tests the forward pass of moba_attn_varlen for basic correctness.
|
||||
It checks output shape, dtype, and for the presence of NaNs/Infs.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
pytest.skip("float32 is not supported in flash attention")
|
||||
|
||||
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
|
||||
batch_size, total_seqlen, num_heads, head_dim, dtype
|
||||
)
|
||||
|
||||
# Ensure chunk size is not larger than the smallest sequence length
|
||||
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
|
||||
if moba_chunk_size > min_seqlen:
|
||||
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
|
||||
|
||||
try:
|
||||
output = moba_attn_varlen(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
moba_chunk_size=moba_chunk_size,
|
||||
moba_topk=moba_topk,
|
||||
select_mode=select_mode,
|
||||
threshold_type=threshold_type,
|
||||
simsum_threshold=0.5, # A reasonable default for threshold mode
|
||||
)
|
||||
except Exception as e:
|
||||
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
|
||||
|
||||
# 1. Check output shape
|
||||
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
|
||||
|
||||
# 2. Check output dtype
|
||||
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
|
||||
|
||||
# 3. Check for NaNs or Infs in the output
|
||||
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
|
||||
@@ -0,0 +1,860 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
|
||||
|
||||
import random
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def calc_chunks(cu_seqlen, moba_chunk_size):
|
||||
"""
|
||||
Calculate chunk boundaries.
|
||||
|
||||
For vision tasks we include all chunks (even the last one which might be shorter)
|
||||
so that every chunk can be selected.
|
||||
"""
|
||||
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
|
||||
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
|
||||
cu_num_chunk = torch.ones(
|
||||
batch_num_chunk.numel() + 1,
|
||||
device=cu_seqlen.device,
|
||||
dtype=batch_num_chunk.dtype,
|
||||
)
|
||||
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
|
||||
num_chunk = cu_num_chunk[-1]
|
||||
chunk_sizes = torch.full(
|
||||
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_sizes[0] = 0
|
||||
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
|
||||
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
|
||||
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
|
||||
chunk_to_batch = torch.zeros(
|
||||
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_to_batch[cu_num_chunk[1:-1]] = 1
|
||||
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
|
||||
|
||||
# Do not filter out any chunk
|
||||
filtered_chunk_indices = torch.arange(
|
||||
num_chunk, device=cu_seqlen.device, dtype=torch.int32
|
||||
)
|
||||
num_filtered_chunk = num_chunk
|
||||
|
||||
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
|
||||
|
||||
|
||||
# --- Threshold Selection Helper Functions ---
|
||||
|
||||
def _select_threshold_query_head(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects chunks for each <query, head> pair based on threshold.
|
||||
Normalization and sorting happen along the chunk dimension (dim=0).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per <head, query> (across chunks)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
row_min = gate_min_val.amin(dim=0) # (H, S)
|
||||
row_max = gate_masked.amax(dim=0) # (H, S)
|
||||
denom = row_max - row_min
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
|
||||
|
||||
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
|
||||
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
|
||||
|
||||
# 2) compute how much more normalized weight we need beyond self
|
||||
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
|
||||
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
|
||||
|
||||
# 3) zero out the self‐chunk in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0
|
||||
|
||||
# 4) sort the other chunks by descending norm, per <head,seq>
|
||||
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per <head,seq>
|
||||
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
|
||||
|
||||
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
|
||||
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
|
||||
any_cond = cond.any(dim=0) # (H, S)
|
||||
# Find the index of the first True value along dim 0. If none, use C-1.
|
||||
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff
|
||||
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
|
||||
|
||||
# 8) scatter it back to original chunk order
|
||||
others_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
others_mask.scatter_(0, sorted_idx, sorted_mask)
|
||||
|
||||
# 9) finally, include every self‐chunk plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_block(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <query, head> pairs for each block based on threshold.
|
||||
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
HS = H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per block (across heads and queries)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_denom = block_max - block_min
|
||||
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
|
||||
|
||||
gate_norm = (gate - block_min) / block_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights *per block*
|
||||
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
|
||||
|
||||
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
|
||||
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
|
||||
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort the other <head, seq> pairs by descending norm, per block
|
||||
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
|
||||
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per block
|
||||
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
|
||||
|
||||
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
|
||||
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
|
||||
any_cond = cond_flat.any(dim=1) # (C,)
|
||||
# Find the index of the first True value along dim 1. If none, use HS-1.
|
||||
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff per block
|
||||
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
|
||||
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
|
||||
|
||||
# 8) scatter it back to original <head, seq> order per block
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
|
||||
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_overall(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query, head> triplets globally based on threshold.
|
||||
Normalization and sorting happen across all valid entries.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
CHS = C * H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization globally across all valid entries
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
overall_max = gate_masked.max() # scalar
|
||||
overall_min = gate_min_val.min() # scalar
|
||||
overall_denom = overall_max - overall_min
|
||||
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
|
||||
|
||||
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights globally
|
||||
self_norm_sum_overall = self_norm_entries.sum() # scalar
|
||||
|
||||
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
|
||||
total_norm_sum_overall = gate_norm.sum() # scalar
|
||||
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort all other entries by descending norm, globally
|
||||
others_flat = others_norm.flatten() # (C*H*S,)
|
||||
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
|
||||
|
||||
# Only sort the valid 'other' entries
|
||||
valid_others_indices = torch.where(valid_others_mask_flat)[0]
|
||||
valid_others_values = others_flat[valid_others_indices]
|
||||
|
||||
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
|
||||
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
|
||||
|
||||
# 5) cumulative‑sum the sorted valid 'other' norms globally
|
||||
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
|
||||
|
||||
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
|
||||
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
|
||||
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
|
||||
any_cond = cond_values.any() # scalar
|
||||
|
||||
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
|
||||
cutoff_idx_in_sorted = torch.where(
|
||||
any_cond,
|
||||
cond_values.float().argmax(dim=0),
|
||||
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
|
||||
)
|
||||
|
||||
# 7) build a mask selecting the top-k others based on the cutoff
|
||||
# Select the original indices corresponding to the top entries in the sorted list
|
||||
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
|
||||
|
||||
# 8) create the mask in the original flat shape
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
|
||||
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
|
||||
others_mask_flat[selected_other_indices] = True
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_head_global(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query> globally for each head based on threshold.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# 1) LSE‐style normalization per head (across chunks and sequence dims)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
|
||||
|
||||
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
denom = max_per_head - min_per_head
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
|
||||
|
||||
gate_norm = (gate - min_per_head) / denom
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 2) sum normalized self‐chunk contributions per head
|
||||
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 3) total normalized sum per head
|
||||
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 4) how much more normalized weight needed per head
|
||||
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0)
|
||||
|
||||
# 5) zero out self‐chunk entries to focus on "others"
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
|
||||
|
||||
# 6) flatten chunk and sequence dims, per head
|
||||
CS = C * S
|
||||
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
|
||||
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
|
||||
# 7) vectorized selection of “others” per head
|
||||
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
|
||||
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
|
||||
|
||||
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
|
||||
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
|
||||
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
|
||||
|
||||
has_cutoff = cond.any(dim=1) # (H,)
|
||||
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
|
||||
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
|
||||
|
||||
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
|
||||
|
||||
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
|
||||
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
|
||||
|
||||
# 8) reshape selection mask back to (C, H, S)
|
||||
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
|
||||
|
||||
# 9) include self‐chunks plus selected others, and obey valid mask
|
||||
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
class MixedAttention(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
):
|
||||
ctx.max_seqlen = max_seqlen
|
||||
ctx.moba_chunk_size = moba_chunk_size
|
||||
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
|
||||
|
||||
# Non-causal self-attention branch
|
||||
# return out, softmax_lse, S_dmask, rng_state
|
||||
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
# MOBA attention branch (non-causal)
|
||||
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
|
||||
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
|
||||
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
|
||||
|
||||
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
output_2d = output.view(-1, q.shape[2])
|
||||
|
||||
max_lse_1d = self_attn_lse_sh.view(-1)
|
||||
max_lse_1d = max_lse_1d.index_reduce(
|
||||
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
|
||||
)
|
||||
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
|
||||
moba_attn_lse = (
|
||||
moba_attn_lse.view(-1)
|
||||
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
|
||||
.reshape_as(moba_attn_lse)
|
||||
)
|
||||
|
||||
mixed_attn_se_sh = self_attn_lse_sh.exp()
|
||||
moba_attn_se = moba_attn_lse.exp()
|
||||
|
||||
mixed_attn_se_sh.view(-1).index_add_(
|
||||
0, moba_q_sh_indices, moba_attn_se.view(-1)
|
||||
)
|
||||
mixed_attn_lse_sh = mixed_attn_se_sh.log()
|
||||
|
||||
# Combine self-attention output
|
||||
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
|
||||
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
|
||||
output_2d += self_attn_out_sh.reshape_as(output_2d)
|
||||
|
||||
# Combine MOBA attention output
|
||||
mixed_attn_lse = (
|
||||
mixed_attn_lse_sh.view(-1)
|
||||
.index_select(0, moba_q_sh_indices)
|
||||
.view_as(moba_attn_lse)
|
||||
)
|
||||
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
|
||||
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
|
||||
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
|
||||
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
|
||||
output = output.to(q.dtype)
|
||||
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
|
||||
ctx.save_for_backward(
|
||||
output,
|
||||
mixed_attn_lse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, d_output):
|
||||
|
||||
max_seqlen = ctx.max_seqlen
|
||||
moba_chunk_size = ctx.moba_chunk_size
|
||||
softmax_scale = ctx.softmax_scale
|
||||
|
||||
(
|
||||
output,
|
||||
mixed_attn_vlse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
) = ctx.saved_tensors
|
||||
|
||||
d_output = d_output.contiguous()
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_output,
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
out=output,
|
||||
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
|
||||
dq=dq,
|
||||
dk=dk,
|
||||
dv=dv,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
headdim = q.shape[-1]
|
||||
d_moba_output = (
|
||||
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
moba_output = (
|
||||
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
|
||||
mixed_attn_vlse = (
|
||||
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
|
||||
)
|
||||
|
||||
dmq = torch.empty_like(moba_q)
|
||||
dmkv = torch.empty_like(moba_kv)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_moba_output,
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
out=moba_output,
|
||||
softmax_lse=mixed_attn_vlse,
|
||||
dq=dmq,
|
||||
dk=dmkv[:,0],
|
||||
dv=dmkv[:,1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
|
||||
|
||||
|
||||
def moba_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int,
|
||||
moba_chunk_size: int,
|
||||
moba_topk: int,
|
||||
select_mode: str = 'threshold', # "topk" or "threshold"
|
||||
simsum_threshold: float = 0.25,
|
||||
threshold_type: str = 'query_head',
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Accelerated MOBA attention for vision tasks with proper LSE normalization.
|
||||
|
||||
This version:
|
||||
- Splits KV into chunks.
|
||||
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
|
||||
by amplifying the diagonal (self-chunk) logits.
|
||||
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
|
||||
reduction so that attending to each query over the selected chunks is equivalent
|
||||
to the original algorithm.
|
||||
"""
|
||||
# Stack keys and values.
|
||||
kv = torch.stack((k, v), dim=1)
|
||||
seqlen, num_head, head_dim = q.shape
|
||||
|
||||
# Compute chunk boundaries.
|
||||
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
|
||||
cu_seqlens, moba_chunk_size
|
||||
)
|
||||
|
||||
self_attn_cu_seqlen = cu_chunk
|
||||
|
||||
# Update top-k selection to include the self chunk.
|
||||
moba_topk = min(moba_topk, num_filtered_chunk)
|
||||
|
||||
# --- Build filtered KV from chunks ---
|
||||
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
|
||||
max_chunk_len = int(chunk_lengths.max().item())
|
||||
|
||||
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
|
||||
indices = chunk_starts.unsqueeze(1) + range_tensor
|
||||
indices = torch.clamp(indices, max=kv.shape[0] - 1)
|
||||
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
|
||||
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
|
||||
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
|
||||
|
||||
# Compute key_gate_weight over valid tokens.
|
||||
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
|
||||
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
|
||||
key_sum = (key_values * valid_mask_exp).sum(dim=1)
|
||||
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
|
||||
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
|
||||
|
||||
# Compute gate logits between key_gate_weight and queries.
|
||||
q_float = q.float()
|
||||
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
|
||||
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
|
||||
|
||||
# Amplify the diagonal (self chunk) contributions.
|
||||
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
|
||||
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
|
||||
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
|
||||
amplification_factor = 1e9 # Example factor; adjust as needed.
|
||||
origin_gate = gate.clone()
|
||||
gate = gate.clone()
|
||||
if select_mode == "topk":
|
||||
gate[gate_self_chunk_mask] += amplification_factor
|
||||
|
||||
# Exclude positions that are outside the valid batch boundaries.
|
||||
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
|
||||
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
|
||||
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
|
||||
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
|
||||
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
|
||||
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
|
||||
|
||||
if select_mode == 'topk':
|
||||
# We amplify self‐chunk in gate already, so self entries will rank highest.
|
||||
valid_gate_mask = gate != -float("inf")
|
||||
if threshold_type == 'query_head':
|
||||
# === per‐<head,seq> top-k across chunks (original behavior) ===
|
||||
# gate: (C, H, S)
|
||||
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
|
||||
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
gate_idx_mask.scatter_(0, gate_topk_idx, True)
|
||||
gate_mask = valid_gate_mask & gate_idx_mask
|
||||
elif threshold_type == 'overall':
|
||||
# === global top-k across all (chunk, head, seq) entries ===
|
||||
C, H, S = gate.shape
|
||||
flat_gate = gate.flatten()
|
||||
flat_mask = valid_gate_mask.flatten()
|
||||
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
|
||||
# pick topk global entries
|
||||
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
|
||||
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
|
||||
others_mask_flat[idx] = True
|
||||
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
|
||||
elif threshold_type == 'head_global':
|
||||
# per-head top-k across all chunks and sequence positions
|
||||
C, H, S = gate.shape
|
||||
CS = C * S
|
||||
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
|
||||
# pick top-k indices per head
|
||||
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
|
||||
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
|
||||
gate_idx_flat.scatter_(1, topk_idx, True)
|
||||
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid threshold_type for topk: {threshold_type}. "
|
||||
"Choose 'query_head', 'block', or 'overall'."
|
||||
)
|
||||
elif select_mode == 'threshold':
|
||||
# Delegate to the specific thresholding function
|
||||
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
|
||||
if threshold_type == 'query_head':
|
||||
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'block':
|
||||
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'overall':
|
||||
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'head_global':
|
||||
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
else:
|
||||
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
|
||||
else:
|
||||
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
|
||||
|
||||
# eliminate self_chunk in MoBA branch
|
||||
gate_mask = gate_mask & ~gate_self_chunk_mask
|
||||
# if gate_mask is all false, perform flash_attn instead
|
||||
if gate_mask.sum() == 0:
|
||||
return flash_attn_varlen_func(
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
|
||||
)
|
||||
|
||||
# Determine which query positions are selected.
|
||||
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
|
||||
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
|
||||
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
|
||||
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
|
||||
|
||||
# Build cumulative sequence lengths for the selected queries.
|
||||
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
|
||||
q_zero_mask = moba_seqlen_q == 0
|
||||
valid_expert_mask = ~q_zero_mask
|
||||
if q_zero_mask.sum() > 0:
|
||||
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
|
||||
moba_cu_seqlen_q = torch.cat(
|
||||
(
|
||||
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
|
||||
moba_seqlen_q.cumsum(dim=0),
|
||||
),
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
# Rearrange gathered KV for the MOBA branch.
|
||||
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
|
||||
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
|
||||
if q_zero_mask.sum() > 0:
|
||||
experts_tensor = experts_tensor[valid_expert_mask]
|
||||
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
|
||||
|
||||
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
|
||||
mask = seq_range < valid_expert_lengths.unsqueeze(1)
|
||||
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
|
||||
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
|
||||
|
||||
moba_cu_seqlen_kv = torch.cat(
|
||||
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
|
||||
valid_expert_lengths.cumsum(dim=0)],
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
assert (
|
||||
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
|
||||
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
|
||||
|
||||
return MixedAttention.apply(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
|
||||
def process_moba_input(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
"""
|
||||
Process inputs for the attention function.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
|
||||
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
|
||||
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed input tensor.
|
||||
"""
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
|
||||
else:
|
||||
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
|
||||
if len(chunk_size) == 2:
|
||||
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
|
||||
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
|
||||
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
|
||||
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
|
||||
elif len(chunk_size) == 3:
|
||||
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
|
||||
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
|
||||
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
|
||||
else:
|
||||
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
|
||||
|
||||
return x, moba_chunk_size
|
||||
|
||||
|
||||
def process_moba_output(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
pass
|
||||
elif len(chunk_size) == 2:
|
||||
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
|
||||
elif len(chunk_size) == 3:
|
||||
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# TEST
|
||||
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
|
||||
random.seed(0)
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
|
||||
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
|
||||
max_seqlen = q.shape[1]
|
||||
q = rearrange(q, "b s ... -> (b s) ...")
|
||||
k = rearrange(k, "b s ... -> (b s) ...")
|
||||
v = rearrange(v, "b s ... -> (b s) ...")
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
|
||||
"""Speed test comparing flash_attn vs moba_attention"""
|
||||
# Get data
|
||||
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
|
||||
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
|
||||
vo_grad = torch.randn_like(q)
|
||||
|
||||
# Warmup
|
||||
warmup_iters = 3
|
||||
perf_test_iters = 10
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_flash = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_moba = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
|
||||
|
||||
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
|
||||
print(f"Speedup: {time_flash / time_moba:.2f}x")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
CUDA_VISIBLE_DEVICES=1 \
|
||||
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
"""
|
||||
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 3 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -0,0 +1,151 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 16
|
||||
--num_height 448 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--num_frames 61 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 4 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 16
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 4
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -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:
|
||||
|
||||
@@ -98,6 +98,7 @@ dmd_args=(
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port $MASTER_PORT \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export MASTER_PORT=29501
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir="checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps=4000
|
||||
--train_batch_size=1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 31
|
||||
--num_height 704
|
||||
--num_width 1280
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--training_state_checkpointing_steps=500
|
||||
--weight_only_checkpointing_steps=500
|
||||
--lora_rank 32
|
||||
--lora_training True
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 1
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-4
|
||||
--mixed_precision="bf16"
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3.5
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port $MASTER_PORT \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -9,9 +9,9 @@ def main():
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
num_gpus=4,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
@@ -25,9 +25,7 @@ def main():
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
@@ -35,11 +33,7 @@ def main():
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
"The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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:
|
||||
|
||||
@@ -10,6 +10,7 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--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 \
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from flash_attn.bert_padding import pad_input
|
||||
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class VMOBAAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "VMOBA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["VMOBAAttentionImpl"]:
|
||||
return VMOBAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["VideoMobaAttentionMetadata"]:
|
||||
return VideoMobaAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["VideoMobaAttentionMetadataBuilder"]:
|
||||
return VideoMobaAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoMobaAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
|
||||
temporal_chunk_size: int
|
||||
temporal_topk: int
|
||||
spatial_chunk_size: tuple[int, int]
|
||||
spatial_topk: int
|
||||
st_chunk_size: tuple[int, int, int]
|
||||
st_topk: int
|
||||
|
||||
moba_select_mode: str
|
||||
moba_threshold: float
|
||||
moba_threshold_type: str
|
||||
patch_resolution: list[int]
|
||||
|
||||
first_full_step: int = 12
|
||||
first_full_layer: int = 0
|
||||
# temporal_layer -> spatial_layer -> st_layer
|
||||
temporal_layer: int = 1
|
||||
spatial_layer: int = 1
|
||||
st_layer: int = 1
|
||||
|
||||
|
||||
class VideoMobaAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
temporal_chunk_size: int,
|
||||
temporal_topk: int,
|
||||
spatial_chunk_size: tuple[int, int],
|
||||
spatial_topk: int,
|
||||
st_chunk_size: tuple[int, int, int],
|
||||
st_topk: int,
|
||||
moba_select_mode: str = 'threshold',
|
||||
moba_threshold: float = 0.25,
|
||||
moba_threshold_type: str = 'query_head',
|
||||
device: torch.device = None,
|
||||
first_full_layer: int = 0,
|
||||
first_full_step: int = 12,
|
||||
temporal_layer: int = 1,
|
||||
spatial_layer: int = 1,
|
||||
st_layer: int = 1,
|
||||
**kwargs,
|
||||
) -> VideoMobaAttentionMetadata:
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
assert raw_latent_shape[0] % patch_size[0] == 0 and raw_latent_shape[
|
||||
1] % patch_size[1] == 0 and raw_latent_shape[2] % patch_size[
|
||||
2] == 0, f"spatial patch_resolution {raw_latent_shape} should be divisible by patch_size {patch_size}"
|
||||
patch_resolution = [
|
||||
t // pt for t, pt in zip(raw_latent_shape, patch_size, strict=False)
|
||||
]
|
||||
|
||||
return VideoMobaAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
temporal_chunk_size=temporal_chunk_size,
|
||||
temporal_topk=temporal_topk,
|
||||
spatial_chunk_size=spatial_chunk_size,
|
||||
spatial_topk=spatial_topk,
|
||||
st_chunk_size=st_chunk_size,
|
||||
st_topk=st_topk,
|
||||
moba_select_mode=moba_select_mode,
|
||||
moba_threshold=moba_threshold,
|
||||
moba_threshold_type=moba_threshold_type,
|
||||
patch_resolution=patch_resolution,
|
||||
first_full_layer=first_full_layer,
|
||||
first_full_step=first_full_step,
|
||||
temporal_layer=temporal_layer,
|
||||
spatial_layer=spatial_layer,
|
||||
st_layer=st_layer,
|
||||
)
|
||||
|
||||
|
||||
class VMOBAAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(self,
|
||||
num_heads,
|
||||
head_size,
|
||||
softmax_scale,
|
||||
causal=False,
|
||||
num_kv_heads=None,
|
||||
prefix="",
|
||||
**extra_impl_args) -> None:
|
||||
self.prefix = prefix
|
||||
self.layer_idx = self._get_layer_idx(prefix)
|
||||
|
||||
def _get_layer_idx(self, prefix: str) -> int | None:
|
||||
match = re.search(r"blocks\.(\d+)", prefix)
|
||||
if not match:
|
||||
raise ValueError(f"Invalid prefix: {prefix}")
|
||||
return int(match.group(1))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
query: [B, L, H, D]
|
||||
key: [B, L, H, D]
|
||||
value: [B, L, H, D]
|
||||
attn_metadata: AttentionMetadata
|
||||
"""
|
||||
batch_size, sequence_length, num_heads, head_dim = query.shape
|
||||
|
||||
# select chunk type according to layer idx:
|
||||
loop_layer_num = attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer
|
||||
moba_layer = self.layer_idx - attn_metadata.first_full_layer
|
||||
if moba_layer % loop_layer_num < attn_metadata.temporal_layer:
|
||||
moba_chunk_size = attn_metadata.temporal_chunk_size
|
||||
moba_topk = attn_metadata.temporal_topk
|
||||
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer:
|
||||
moba_chunk_size = attn_metadata.spatial_chunk_size
|
||||
moba_topk = attn_metadata.spatial_topk
|
||||
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer:
|
||||
moba_chunk_size = attn_metadata.st_chunk_size
|
||||
moba_topk = attn_metadata.st_topk
|
||||
|
||||
# torch.distributed.breakpoint()
|
||||
query, chunk_size = process_moba_input(query,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
key, chunk_size = process_moba_input(key,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
value, chunk_size = process_moba_input(value,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
max_seqlen = query.shape[1]
|
||||
indices_q = torch.arange(0,
|
||||
query.shape[0] * query.shape[1],
|
||||
device=query.device)
|
||||
cu_seqlens = torch.arange(0,
|
||||
query.shape[0] * query.shape[1] + 1,
|
||||
query.shape[1],
|
||||
dtype=torch.int32,
|
||||
device=query.device)
|
||||
query = rearrange(query, "b s ... -> (b s) ...")
|
||||
key = rearrange(key, "b s ... -> (b s) ...")
|
||||
value = rearrange(value, "b s ... -> (b s) ...")
|
||||
|
||||
# current_timestep=attn_metadata.current_timestep
|
||||
hidden_states = moba_attn_varlen(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
moba_chunk_size=chunk_size,
|
||||
moba_topk=moba_topk,
|
||||
select_mode=attn_metadata.moba_select_mode,
|
||||
simsum_threshold=attn_metadata.moba_threshold,
|
||||
threshold_type=attn_metadata.moba_threshold_type,
|
||||
)
|
||||
hidden_states = pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = process_moba_output(hidden_states,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"temporal_chunk_size": 2,
|
||||
"temporal_topk": 2,
|
||||
"spatial_chunk_size": [4, 13],
|
||||
"spatial_topk": 6,
|
||||
"st_chunk_size": [4, 4, 13],
|
||||
"st_topk": 18,
|
||||
"moba_select_mode": "topk",
|
||||
"moba_threshold": 0.25,
|
||||
"moba_threshold_type": "query_head",
|
||||
"first_full_layer": 0,
|
||||
"first_full_step": 12,
|
||||
"temporal_layer": 1,
|
||||
"spatial_layer": 1,
|
||||
"st_layer": 1
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"temporal_chunk_size": 2,
|
||||
"temporal_topk": 3,
|
||||
"spatial_chunk_size": [3, 4],
|
||||
"spatial_topk": 20,
|
||||
"st_chunk_size": [4, 6, 4],
|
||||
"st_topk": 15,
|
||||
"moba_select_mode": "threshold",
|
||||
"moba_threshold": 0.25,
|
||||
"moba_threshold_type": "query_head",
|
||||
"first_full_layer": 0,
|
||||
"first_full_step": 12,
|
||||
"temporal_layer": 1,
|
||||
"spatial_layer": 1,
|
||||
"st_layer": 1
|
||||
}
|
||||
@@ -32,6 +32,29 @@ class DatasetType(str, Enum):
|
||||
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:
|
||||
"""Configuration for preprocessing operations."""
|
||||
@@ -51,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
|
||||
@@ -120,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,
|
||||
@@ -174,6 +204,10 @@ class PreprocessConfig:
|
||||
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(
|
||||
|
||||
@@ -15,9 +15,13 @@ class DiTArchConfig(ArchConfig):
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -12,13 +12,13 @@ from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
mask: torch.tensor = outputs.attention_mask
|
||||
hidden_state: torch.tensor = outputs.last_hidden_state
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
assert torch.isnan(hidden_state).sum() == 0
|
||||
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.tensor = torch.stack([
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
@@ -39,12 +39,12 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
vae_sp: bool = False
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 3
|
||||
flow_shift: float | None = 3.0
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (T5Config(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(t5_postprocess_text, ))
|
||||
|
||||
@@ -68,7 +68,7 @@ class WanT2V720PConfig(WanT2V480PConfig):
|
||||
# WanConfig-specific parameters with defaults
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
flow_shift: float | None = 5.0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -94,7 +94,7 @@ class WanI2V720PConfig(WanI2V480PConfig):
|
||||
# WanConfig-specific parameters with defaults
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
flow_shift: float | None = 5.0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -104,7 +104,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
|
||||
# WanConfig-specific parameters with defaults
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 8
|
||||
flow_shift: float | None = 8.0
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 757, 522])
|
||||
|
||||
@@ -115,7 +115,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
flow_shift: int = 5
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
@@ -125,7 +125,7 @@ class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
|
||||
@dataclass
|
||||
class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
|
||||
flow_shift: int = 5
|
||||
flow_shift: float | None = 5.0
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 757, 522])
|
||||
|
||||
@@ -138,3 +138,15 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Causal Self-Forcing =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
is_causal: bool = True
|
||||
flow_shift: float | None = 5.0
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
@@ -191,6 +191,13 @@ class SamplingParam:
|
||||
default=SamplingParam.image_path,
|
||||
help="Path to input image for image-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--moba-config-path",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -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,17 +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-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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
+122
-2
@@ -1,9 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py
|
||||
"""The arguments of FastVideo Inference."""
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
import json
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import field
|
||||
from enum import Enum
|
||||
@@ -139,6 +139,10 @@ class FastVideoArgs:
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
|
||||
# V-MoBA parameters
|
||||
moba_config_path: str | None = None
|
||||
moba_config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Master port for distributed training/inference
|
||||
master_port: int | None = None
|
||||
|
||||
@@ -166,6 +170,16 @@ class FastVideoArgs:
|
||||
return not self.inference_mode
|
||||
|
||||
def __post_init__(self):
|
||||
if self.moba_config_path:
|
||||
try:
|
||||
with open(self.moba_config_path) as f:
|
||||
self.moba_config = json.load(f)
|
||||
logger.info("Loaded V-MoBA config from %s",
|
||||
self.moba_config_path)
|
||||
except (FileNotFoundError, json.JSONDecodeError) as e:
|
||||
logger.error("Failed to load V-MoBA config from %s: %s",
|
||||
self.moba_config_path, e)
|
||||
raise
|
||||
self.check_fastvideo_args()
|
||||
|
||||
@staticmethod
|
||||
@@ -591,6 +605,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
generator_model_path: str = "" # path for generator (student) model
|
||||
real_score_model_path: str = "" # path for real score (teacher) model
|
||||
fake_score_model_path: str = "" # path for fake score (critic) model
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
@@ -613,6 +632,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
@@ -644,6 +664,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
linear_quadratic_threshold: float = 0.0
|
||||
linear_range: float = 0.0
|
||||
weight_decay: float = 0.0
|
||||
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
|
||||
use_ema: bool = False
|
||||
multi_phased_distill_schedule: str = ""
|
||||
pred_decay_weight: float = 0.0
|
||||
@@ -664,16 +685,29 @@ class TrainingArgs(FastVideoArgs):
|
||||
|
||||
# distillation args
|
||||
generator_update_interval: int = 5
|
||||
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
|
||||
min_timestep_ratio: float = 0.2
|
||||
max_timestep_ratio: float = 0.98
|
||||
real_score_guidance_scale: float = 3.5
|
||||
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
|
||||
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
|
||||
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
|
||||
training_state_checkpointing_steps: int = 0 # for resuming training
|
||||
weight_only_checkpointing_steps: int = 0 # for inference
|
||||
log_visualization: bool = False
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
warp_denoising_step: bool = False
|
||||
|
||||
# Self-forcing specific arguments
|
||||
num_frame_per_block: int = 3
|
||||
independent_first_frame: bool = False
|
||||
enable_gradient_masking: bool = True
|
||||
gradient_mask_last_n_frames: int = 21
|
||||
validate_cache_structure: bool = False # Debug flag for cache validation
|
||||
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -775,6 +809,20 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
help="Directory to cache models")
|
||||
|
||||
# DMD model paths - separate paths for each network
|
||||
parser.add_argument(
|
||||
"--generator-model-path",
|
||||
type=str,
|
||||
help="Path to generator (student) model for DMD distillation")
|
||||
parser.add_argument(
|
||||
"--real-score-model-path",
|
||||
type=str,
|
||||
help="Path to real score (teacher) model for DMD distillation")
|
||||
parser.add_argument(
|
||||
"--fake-score-model-path",
|
||||
type=str,
|
||||
help="Path to fake score (critic) model for DMD distillation")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
@@ -845,6 +893,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--resume-from-checkpoint",
|
||||
type=str,
|
||||
help="Path to checkpoint to resume from")
|
||||
parser.add_argument(
|
||||
"--init-weights-from-safetensors",
|
||||
type=str,
|
||||
help="Path to safetensors file for initial weight loading")
|
||||
parser.add_argument("--logging-dir",
|
||||
type=str,
|
||||
help="Directory for logging")
|
||||
@@ -949,6 +1001,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Linear quadratic threshold")
|
||||
parser.add_argument("--linear-range", type=float, help="Linear range")
|
||||
parser.add_argument("--weight-decay", type=float, help="Weight decay")
|
||||
parser.add_argument("--betas",
|
||||
type=str,
|
||||
default=TrainingArgs.betas,
|
||||
help="Betas for optimizer (format: 'beta1,beta2')")
|
||||
parser.add_argument("--use-ema",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA")
|
||||
@@ -985,11 +1041,27 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
|
||||
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
|
||||
|
||||
# V-MoBA parameters
|
||||
parser.add_argument(
|
||||
"--moba-config-path",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
|
||||
# Distillation arguments
|
||||
parser.add_argument("--generator-update-interval",
|
||||
type=int,
|
||||
default=TrainingArgs.generator_update_interval,
|
||||
help="Ratio of student updates to critic updates.")
|
||||
parser.add_argument(
|
||||
"--dfake-gen-update-ratio",
|
||||
type=int,
|
||||
default=TrainingArgs.dfake_gen_update_ratio,
|
||||
help=
|
||||
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
|
||||
)
|
||||
parser.add_argument("--min-timestep-ratio",
|
||||
type=float,
|
||||
default=TrainingArgs.min_timestep_ratio,
|
||||
@@ -1006,6 +1078,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=float,
|
||||
default=TrainingArgs.fake_score_learning_rate,
|
||||
help="Learning rate for fake score transformer")
|
||||
parser.add_argument(
|
||||
"--fake-score-betas",
|
||||
type=str,
|
||||
default=TrainingArgs.fake_score_betas,
|
||||
help="Betas for fake score optimizer (format: 'beta1,beta2')")
|
||||
parser.add_argument(
|
||||
"--fake-score-lr-scheduler",
|
||||
type=str,
|
||||
@@ -1018,6 +1095,49 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--simulate-generator-forward",
|
||||
action=StoreBoolean,
|
||||
help="Whether to simulate generator forward to match inference")
|
||||
parser.add_argument(
|
||||
"--warp-denoising-step",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Whether to warp denoising step according to the scheduler time shift"
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
parser.add_argument(
|
||||
"--num-frame-per-block",
|
||||
type=int,
|
||||
default=TrainingArgs.num_frame_per_block,
|
||||
help="Number of frames per block for causal generation")
|
||||
parser.add_argument(
|
||||
"--independent-first-frame",
|
||||
action=StoreBoolean,
|
||||
help="Whether the first frame is independent in causal generation")
|
||||
parser.add_argument(
|
||||
"--enable-gradient-masking",
|
||||
action=StoreBoolean,
|
||||
help="Whether to enable frame-level gradient masking")
|
||||
parser.add_argument(
|
||||
"--gradient-mask-last-n-frames",
|
||||
type=int,
|
||||
default=TrainingArgs.gradient_mask_last_n_frames,
|
||||
help="Number of last frames to enable gradients for")
|
||||
parser.add_argument(
|
||||
"--validate-cache-structure",
|
||||
action=StoreBoolean,
|
||||
help="Whether to validate KV cache structure (debug flag)")
|
||||
parser.add_argument(
|
||||
"--same-step-across-blocks",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use the same exit timestep for all blocks")
|
||||
parser.add_argument(
|
||||
"--last-step-only",
|
||||
action=StoreBoolean,
|
||||
help="Whether to only use the last timestep for training")
|
||||
parser.add_argument(
|
||||
"--context-noise",
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -1025,4 +1145,4 @@ class TrainingArgs(FastVideoArgs):
|
||||
def parse_int_list(value: str) -> list[int]:
|
||||
if not value:
|
||||
return []
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
@@ -100,7 +100,16 @@ class ScaleResidual(nn.Module):
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
gate: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply gated residual connection."""
|
||||
return residual + x * gate
|
||||
# x.shape: [batch_size, seq_len, inner_dim]
|
||||
if gate.dim() == 4:
|
||||
# gate.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
num_frames = gate.shape[1]
|
||||
frame_seqlen = x.shape[1] // num_frames
|
||||
return residual + (x.unflatten(
|
||||
dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
|
||||
else:
|
||||
# gate.shape: [batch_size, 1, inner_dim]
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
|
||||
@@ -159,7 +168,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
gate: torch.Tensor, shift: torch.Tensor,
|
||||
gate: torch.Tensor | int, shift: torch.Tensor,
|
||||
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply gated residual connection, followed by layernorm and
|
||||
@@ -171,12 +180,41 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
- residual value (value after residual connection
|
||||
but before normalization)
|
||||
"""
|
||||
# x.shape: [batch_size, seq_len, inner_dim]
|
||||
# Apply residual connection with gating
|
||||
residual_output = residual + x * gate
|
||||
if isinstance(gate, int):
|
||||
# used by cross-attention, should be 1
|
||||
assert gate == 1
|
||||
residual_output = residual + x
|
||||
elif isinstance(gate, torch.Tensor):
|
||||
if gate.dim() == 4:
|
||||
# gate.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
num_frames = gate.shape[1]
|
||||
frame_seqlen = x.shape[1] // num_frames
|
||||
residual_output = residual + (
|
||||
x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
gate).flatten(1, 2)
|
||||
else:
|
||||
# used by bidirectional self attention
|
||||
# gate.shape: [batch_size, 1, inner_dim]
|
||||
residual_output = residual + x * gate
|
||||
else:
|
||||
raise ValueError(f"Gate type {type(gate)} not supported")
|
||||
# residual_output.shape: [batch_size, seq_len, inner_dim]
|
||||
|
||||
# Apply normalization
|
||||
normalized = self.norm(residual_output)
|
||||
# Apply scale and shift
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
|
||||
# scale.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
# shift.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
num_frames = scale.shape[1]
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
modulated = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
modulated = normalized * (1 + scale) + shift
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
@@ -218,8 +256,24 @@ class LayerNormScaleShift(nn.Module):
|
||||
def forward(self, x: torch.Tensor, shift: torch.Tensor,
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
# x.shape: [batch_size, seq_len, inner_dim]
|
||||
normalized = self.norm(x)
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
normalized = normalized.float()
|
||||
|
||||
if scale.dim() == 4:
|
||||
# scale.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
num_frames = scale.shape[1]
|
||||
frame_seqlen = normalized.shape[1] // num_frames
|
||||
output = (
|
||||
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale) + shift).flatten(1, 2)
|
||||
else:
|
||||
return normalized * (1.0 + scale) + shift
|
||||
# scale.shape: [batch_size, 1, inner_dim]
|
||||
# shift.shape: [batch_size, 1, inner_dim]
|
||||
output = normalized * (1 + scale) + shift
|
||||
|
||||
if self.compute_dtype == torch.float32:
|
||||
output = output.to(x.dtype)
|
||||
|
||||
return output
|
||||
@@ -63,7 +63,7 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
device=self.base_layer.weight.device,
|
||||
dtype=self.base_layer.weight.dtype))
|
||||
torch.nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
|
||||
torch.nn.init.kaiming_uniform_(self.lora_B, a=math.sqrt(5))
|
||||
torch.nn.init.zeros_(self.lora_B)
|
||||
else:
|
||||
self.lora_A = None
|
||||
self.lora_B = None
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,682 @@
|
||||
# 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"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
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 = nn.LayerNorm(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)
|
||||
|
||||
# 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)
|
||||
|
||||
# 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:
|
||||
# hidden_states.shape: [batch_size, seq_length, inner_dim]
|
||||
# temb.shape: [batch_size, num_frames, 6, inner_dim]
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
num_frames = temb.shape[1]
|
||||
frame_seqlen = hidden_states.shape[1] // num_frames
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb
|
||||
# e.shape: [batch_size, num_frames, 6, inner_dim]
|
||||
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
|
||||
# assert shift_msa.dtype == torch.float32
|
||||
|
||||
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
|
||||
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
|
||||
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
|
||||
(1 + scale_msa) + shift_msa).flatten(1, 2)
|
||||
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
|
||||
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.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(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)
|
||||
|
||||
# 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)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
|
||||
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)
|
||||
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 = 3
|
||||
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,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
|
||||
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
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
|
||||
dim=2)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return torch.stack(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,
|
||||
freqs_sin) 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)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
|
||||
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
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
|
||||
dim=2)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return torch.stack(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)
|
||||
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
@@ -1,3 +1,5 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
@@ -37,16 +39,14 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
dtype = encoder_hidden_states_image.dtype
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states).to(dtype)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
@@ -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]
|
||||
@@ -155,9 +156,21 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# 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)
|
||||
q = self.norm_q.forward_native(self.to_q(x)[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.forward_native(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.forward_native(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)
|
||||
@@ -200,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# 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)
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
|
||||
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
|
||||
b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
@@ -234,7 +247,7 @@ class WanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = nn.LayerNorm(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)
|
||||
@@ -265,29 +278,29 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
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)
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -306,23 +319,35 @@ class WanTransformerBlock(nn.Module):
|
||||
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
|
||||
|
||||
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
|
||||
).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
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
|
||||
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)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -342,26 +367,20 @@ class WanTransformerBlock(nn.Module):
|
||||
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,
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
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 WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -378,7 +397,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = nn.LayerNorm(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)
|
||||
@@ -410,8 +429,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -431,8 +449,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -452,23 +469,22 @@ class WanTransformerBlock_VSA(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()
|
||||
e = self.scale_shift_table + temb
|
||||
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)
|
||||
norm_hidden_states = (self.norm1(hidden_states) *
|
||||
(1 + scale_msa) + shift_msa)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -493,8 +509,6 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
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,
|
||||
@@ -502,17 +516,15 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
context_lens=None)
|
||||
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 WanTransformer3DModel(CachableDiT):
|
||||
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = WanVideoConfig()._compile_conditions
|
||||
@@ -570,8 +582,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -631,15 +642,31 @@ class WanTransformer3DModel(CachableDiT):
|
||||
rope_theta=10000)
|
||||
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
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
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
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
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,19 +703,47 @@ 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)
|
||||
|
||||
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)
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output
|
||||
return torch.stack(output)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
@@ -780,4 +835,4 @@ class WanTransformer3DModel(CachableDiT):
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
return hidden_states + self.previous_residual_odd
|
||||
@@ -430,6 +430,16 @@ class TransformerLoader(ComponentLoader):
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
# Check if we should use custom initialization weights
|
||||
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
|
||||
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
|
||||
fastvideo_args.training_mode and
|
||||
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
|
||||
|
||||
if use_custom_weights:
|
||||
logger.info("Using custom initialization weights from: %s", custom_weights_path)
|
||||
safetensors_list = [custom_weights_path]
|
||||
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
@@ -59,6 +61,9 @@ _SCHEDULERS = {
|
||||
"FlowMatchEulerDiscreteScheduler"),
|
||||
"UniPCMultistepScheduler":
|
||||
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
|
||||
"SelfForcingFlowMatchScheduler":
|
||||
("schedulers", "scheduling_self_forcing_flow_match",
|
||||
"SelfForcingFlowMatchScheduler"),
|
||||
}
|
||||
|
||||
_FAST_VIDEO_MODELS = {
|
||||
|
||||
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
|
||||
"""
|
||||
Args:
|
||||
clean_latent: the clean latent with shape [B, C, H, W],
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise: the noise with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
|
||||
Returns:
|
||||
the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.base import BaseScheduler
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
"""
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
|
||||
config_name = "scheduler_config.json"
|
||||
order = 1
|
||||
@register_to_config
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.inverse_timesteps = inverse_timesteps
|
||||
self.extra_one_step = extra_one_step
|
||||
self.reverse_sigmas = reverse_sigmas
|
||||
self.set_timesteps(num_inference_steps, training=training)
|
||||
|
||||
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, return_dict=False, **kwargs):
|
||||
sigma_start = self.sigma_min + \
|
||||
(self.sigma_max - self.sigma_min) * denoising_strength
|
||||
if self.extra_one_step:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
|
||||
else:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps)
|
||||
if self.inverse_timesteps:
|
||||
self.sigmas = torch.flip(self.sigmas, dims=[0])
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
if self.reverse_sigmas:
|
||||
self.sigmas = 1 - self.sigmas
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
if training:
|
||||
x = self.timesteps
|
||||
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
|
||||
num_inference_steps) ** 2)
|
||||
y_shifted = y - y.min()
|
||||
bsmntw_weighing = y_shifted * \
|
||||
(num_inference_steps / y_shifted.sum())
|
||||
self.linear_timesteps_weights = bsmntw_weighing
|
||||
|
||||
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
|
||||
sigma_ = 1 if (
|
||||
self.inverse_timesteps or self.reverse_sigmas) else 0
|
||||
else:
|
||||
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
prev_sample = sample + model_output * (sigma_ - sigma)
|
||||
if isinstance(prev_sample, torch.Tensor | float) and not return_dict:
|
||||
return (prev_sample, )
|
||||
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B*T, C, H, W]
|
||||
- noise: the noise with shape [B*T, C, H, W]
|
||||
- timestep: the timestep with shape [B*T]
|
||||
Output: the corrupted latent with shape [B*T, C, H, W]
|
||||
"""
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def training_target(self, sample, noise, timestep):
|
||||
target = noise - sample
|
||||
return target
|
||||
|
||||
def training_weight(self, timestep):
|
||||
"""
|
||||
Input:
|
||||
- timestep: the timestep with shape [B*T]
|
||||
Output: the corresponding weighting [B*T]
|
||||
"""
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
self.linear_timesteps_weights = self.linear_timesteps_weights.to(timestep.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(1) - timestep.unsqueeze(0)).abs(), dim=0)
|
||||
weights = self.linear_timesteps_weights[timestep_id]
|
||||
return weights
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, timestep: int | None = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def set_shift(self, shift: float) -> None:
|
||||
self.shift = shift
|
||||
|
||||
@@ -137,3 +137,46 @@ 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.
|
||||
|
||||
Args:
|
||||
pred_noise: the predicted noise with shape [B, C, H, W]
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise_input_latent: the noisy latent with shape [B, C, H, W],
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
scheduler: the scheduler
|
||||
|
||||
Returns:
|
||||
the predicted video with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == noise_input_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(noise_input_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == noise_input_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
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,60 @@
|
||||
# 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)
|
||||
# isort: on
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
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="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
|
||||
@@ -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",
|
||||
|
||||
@@ -79,7 +79,7 @@ class ComposedPipelineBase(ABC):
|
||||
for name, module in self.modules.items():
|
||||
if not isinstance(module, torch.nn.Module):
|
||||
continue
|
||||
if name == "transformer":
|
||||
if "transformer" in name:
|
||||
module.requires_grad_(True)
|
||||
else:
|
||||
module.requires_grad_(False)
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file
|
||||
from torch.distributed.device_mesh import init_device_mesh
|
||||
from torch.distributed.device_mesh import DeviceMesh, init_device_mesh
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
@@ -32,6 +32,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
cur_adapter_name: str = ""
|
||||
cur_adapter_path: str = ""
|
||||
lora_layers: dict[str, BaseLayerWithLoRA] = {}
|
||||
lora_layers_critic: dict[str, BaseLayerWithLoRA] = {}
|
||||
fastvideo_args: FastVideoArgs | TrainingArgs
|
||||
exclude_lora_layers: list[str] = []
|
||||
device: torch.device = get_local_torch_device()
|
||||
@@ -81,6 +82,17 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
|
||||
def set_trainable(self) -> None:
|
||||
|
||||
def set_lora_grads(lora_layers: dict[str, BaseLayerWithLoRA],
|
||||
device_mesh: DeviceMesh):
|
||||
for name, layer in lora_layers.items():
|
||||
layer.lora_A.requires_grad_(True)
|
||||
layer.lora_B.requires_grad_(True)
|
||||
layer.base_layer.requires_grad_(False)
|
||||
layer.lora_A = nn.Parameter(
|
||||
DTensor.from_local(layer.lora_A, device_mesh=device_mesh))
|
||||
layer.lora_B = nn.Parameter(
|
||||
DTensor.from_local(layer.lora_B, device_mesh=device_mesh))
|
||||
|
||||
is_lora_training = self.training_mode and getattr(
|
||||
self.fastvideo_args, "lora_training", False)
|
||||
if not is_lora_training:
|
||||
@@ -88,18 +100,12 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
return
|
||||
|
||||
self.modules["transformer"].requires_grad_(False)
|
||||
if "fake_score_transformer" in self.modules:
|
||||
self.modules["fake_score_transformer"].requires_grad_(False)
|
||||
device_mesh = init_device_mesh("cuda", (dist.get_world_size(), 1),
|
||||
mesh_dim_names=["fake", "replicate"])
|
||||
for name, layer in self.lora_layers.items():
|
||||
# Enable grads for lora weights only
|
||||
# Must convert to DTensor for compatibility with other FSDP modules in grad calculation
|
||||
layer.lora_A.requires_grad_(True)
|
||||
layer.lora_B.requires_grad_(True)
|
||||
layer.base_layer.requires_grad_(False)
|
||||
layer.lora_A = nn.Parameter(
|
||||
DTensor.from_local(layer.lora_A, device_mesh=device_mesh))
|
||||
layer.lora_B = nn.Parameter(
|
||||
DTensor.from_local(layer.lora_B, device_mesh=device_mesh))
|
||||
set_lora_grads(self.lora_layers, device_mesh)
|
||||
set_lora_grads(self.lora_layers_critic, device_mesh)
|
||||
|
||||
def convert_to_lora_layers(self) -> None:
|
||||
"""
|
||||
@@ -131,6 +137,24 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
converted_count += 1
|
||||
logger.info("Converted %d layers to LoRA layers", converted_count)
|
||||
|
||||
if "fake_score_transformer" in self.modules:
|
||||
for name, layer in self.modules[
|
||||
"fake_score_transformer"].named_modules():
|
||||
if not self.is_target_layer(name):
|
||||
continue
|
||||
layer = get_lora_layer(layer,
|
||||
lora_rank=self.lora_rank,
|
||||
lora_alpha=self.lora_alpha,
|
||||
training_mode=self.training_mode)
|
||||
if layer is not None:
|
||||
self.lora_layers_critic[name] = layer
|
||||
replace_submodule(self.modules["fake_score_transformer"],
|
||||
name, layer)
|
||||
converted_count += 1
|
||||
logger.info(
|
||||
"Converted %d layers to LoRA layers in the critic model",
|
||||
converted_count)
|
||||
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None): # type: ignore
|
||||
@@ -224,4 +248,4 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
|
||||
def unmerge_lora_weights(self) -> None:
|
||||
for name, layer in self.lora_layers.items():
|
||||
layer.unmerge_lora_weights()
|
||||
layer.unmerge_lora_weights()
|
||||
|
||||
@@ -241,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)
|
||||
|
||||
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanPipeline": "wan",
|
||||
"WanDMDPipeline": "wan",
|
||||
"WanImageToVideoPipeline": "wan",
|
||||
"WanCausalDMDPipeline": "wan",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
}
|
||||
|
||||
@@ -4,9 +4,11 @@ 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, WorkloadType
|
||||
@@ -61,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)
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
|
||||
from fastvideo.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.pipelines.stages.denoising import (DenoisingStage,
|
||||
@@ -30,6 +31,7 @@ __all__ = [
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
"CausalDMDDenosingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
|
||||
@@ -0,0 +1,445 @@
|
||||
import torch # type: ignore
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
SlidingTileAttentionBackend = None # type: ignore
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
VideoSparseAttentionBackend = None # type: ignore
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CausalDMDDenosingStage(DenoisingStage):
|
||||
"""
|
||||
Denoising stage for causal diffusion.
|
||||
"""
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__(transformer, scheduler)
|
||||
# KV and cross-attention cache state (initialized on first forward)
|
||||
self.kv_cache1: list | None = None
|
||||
self.crossattn_cache: list | None = None
|
||||
# Model-dependent constants (aligned with causal_inference.py assumptions)
|
||||
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
|
||||
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
|
||||
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
|
||||
|
||||
try:
|
||||
self.local_attn_size = getattr(self.transformer.model,
|
||||
"local_attn_size",
|
||||
-1) # type: ignore
|
||||
except Exception:
|
||||
self.local_attn_size = -1
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
|
||||
patch_ratio = self.transformer.config.arch_config.patch_size[
|
||||
-1] * self.transformer.config.arch_config.patch_size[-2]
|
||||
self.frame_seq_length = latent_seq_length // patch_ratio
|
||||
# TODO(will): make this a parameter once we add i2v support
|
||||
independent_first_frame = self.transformer.independent_first_frame
|
||||
|
||||
# Timesteps for DMD
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long).cpu()
|
||||
|
||||
if fastvideo_args.pipeline_config.warp_denoising_step:
|
||||
logger.info("Warping timesteps...")
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
logger.info("Using timesteps: %s", timesteps)
|
||||
|
||||
# Image kwargs (kept empty unless caller provides compatible args)
|
||||
image_kwargs: dict = {}
|
||||
|
||||
pos_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
# "encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
# STA
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
# Latents and prompts
|
||||
assert batch.latents is not None, "latents must be provided"
|
||||
latents = batch.latents # [B, C, T, H, W]
|
||||
b, c, t, h, w = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
|
||||
# Initialize or reset caches
|
||||
if self.kv_cache1 is None:
|
||||
self._initialize_kv_cache(batch_size=latents.shape[0],
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
self._initialize_crossattn_cache(
|
||||
batch_size=latents.shape[0],
|
||||
max_text_len=fastvideo_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.text_len,
|
||||
dtype=target_dtype,
|
||||
device=latents.device)
|
||||
else:
|
||||
assert self.crossattn_cache is not None
|
||||
# reset cross-attention cache
|
||||
for block_index in range(self.num_transformer_blocks):
|
||||
self.crossattn_cache[block_index][
|
||||
"is_init"] = False # type: ignore
|
||||
# reset kv cache pointers
|
||||
for block_index in range(len(self.kv_cache1)):
|
||||
self.kv_cache1[block_index][
|
||||
"global_end_index"] = torch.tensor( # type: ignore
|
||||
[0],
|
||||
dtype=torch.long,
|
||||
device=latents.device)
|
||||
self.kv_cache1[block_index][
|
||||
"local_end_index"] = torch.tensor( # type: ignore
|
||||
[0],
|
||||
dtype=torch.long,
|
||||
device=latents.device)
|
||||
|
||||
# Optional: cache context features from provided image latents prior to generation
|
||||
current_start_frame = 0
|
||||
if getattr(batch, "image_latent", None) is not None:
|
||||
image_latent = batch.image_latent
|
||||
assert image_latent is not None
|
||||
input_frames = image_latent.shape[2]
|
||||
# timestep zero (or configured context noise) for cache warm-up
|
||||
t_zero = torch.zeros([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.long)
|
||||
if independent_first_frame and input_frames >= 1:
|
||||
# warm-up with the very first frame independently
|
||||
image_first_btchw = image_latent[:, :, :1, :, :].to(
|
||||
target_dtype).permute(0, 2, 1, 3, 4)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
_ = self.transformer(
|
||||
image_first_btchw,
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame *
|
||||
self.frame_seq_length,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
current_start_frame += 1
|
||||
remaining_frames = input_frames - 1
|
||||
else:
|
||||
remaining_frames = input_frames
|
||||
|
||||
# process remaining input frames in blocks of num_frame_per_block
|
||||
while remaining_frames > 0:
|
||||
block = min(self.num_frames_per_block, remaining_frames)
|
||||
ref_btchw = image_latent[:, :, current_start_frame:
|
||||
current_start_frame +
|
||||
block, :, :].to(target_dtype).permute(
|
||||
0, 2, 1, 3, 4)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
_ = self.transformer(
|
||||
ref_btchw,
|
||||
prompt_embeds,
|
||||
t_zero,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=current_start_frame *
|
||||
self.frame_seq_length,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
current_start_frame += block
|
||||
remaining_frames -= block
|
||||
|
||||
# Base position offset from any cache warm-up
|
||||
pos_start_base = current_start_frame
|
||||
|
||||
# Determine block sizes
|
||||
if not independent_first_frame or (independent_first_frame
|
||||
and batch.image_latent is not None):
|
||||
if t % self.num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
|
||||
)
|
||||
num_blocks = t // self.num_frames_per_block
|
||||
block_sizes = [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
else:
|
||||
if (t - 1) % self.num_frames_per_block != 0:
|
||||
raise ValueError(
|
||||
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
|
||||
)
|
||||
num_blocks = (t - 1) // self.num_frames_per_block
|
||||
block_sizes = [1] + [self.num_frames_per_block] * num_blocks
|
||||
start_index = 0
|
||||
|
||||
# DMD loop in causal blocks
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
for current_num_frames in block_sizes:
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
# use BTCHW for DMD conversion routines
|
||||
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
|
||||
video_raw_latent_shape = noise_latents_btchw.shape
|
||||
|
||||
for i, t_cur in enumerate(timesteps):
|
||||
# Copy for pred conversion
|
||||
noise_latents = noise_latents_btchw.clone()
|
||||
latent_model_input = current_latents.to(target_dtype)
|
||||
|
||||
if batch.image_latent is not None and independent_first_frame and start_index == 0:
|
||||
latent_model_input = torch.cat([
|
||||
latent_model_input,
|
||||
batch.image_latent.to(target_dtype)
|
||||
],
|
||||
dim=2)
|
||||
|
||||
# Prepare inputs
|
||||
t_expand = t_cur.repeat(latent_model_input.shape[0])
|
||||
|
||||
# Attention metadata if needed
|
||||
if (vsa_available and self.attn_backend
|
||||
== VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=i, # type: ignore
|
||||
raw_latent_shape=(current_num_frames, h,
|
||||
w), # type: ignore
|
||||
patch_size=fastvideo_args.pipeline_config.
|
||||
dit_config.patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.
|
||||
VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(), # type: ignore
|
||||
) # type: ignore
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
# Run transformer; follow DMD stage pattern
|
||||
t_expanded_noise = t_cur * torch.ones(
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
pred_noise_btchw = self.transformer(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expanded_noise,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Convert pred noise to pred video with FM Euler scheduler utilities
|
||||
pred_video_btchw = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise_btchw.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
timestep=t_expand,
|
||||
scheduler=self.scheduler).unflatten(
|
||||
0, pred_noise_btchw.shape[:2])
|
||||
|
||||
if i < len(timesteps) - 1:
|
||||
next_timestep = timesteps[i + 1] * torch.ones(
|
||||
[1],
|
||||
dtype=torch.long,
|
||||
device=pred_video_btchw.device)
|
||||
noise = torch.randn(
|
||||
video_raw_latent_shape,
|
||||
dtype=pred_video_btchw.dtype,
|
||||
generator=(batch.generator[0] if isinstance(
|
||||
batch.generator, list) else
|
||||
batch.generator)).to(self.device)
|
||||
noise_btchw = noise
|
||||
noise_latents_btchw = self.scheduler.add_noise(
|
||||
pred_video_btchw.flatten(0, 1),
|
||||
noise_btchw.flatten(0, 1),
|
||||
next_timestep).unflatten(0,
|
||||
pred_video_btchw.shape[:2])
|
||||
current_latents = noise_latents_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
else:
|
||||
current_latents = pred_video_btchw.permute(
|
||||
0, 2, 1, 3, 4)
|
||||
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
|
||||
# Write back and advance
|
||||
latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :] = current_latents
|
||||
|
||||
# Re-run with context timestep to update KV cache using clean context
|
||||
context_noise = getattr(fastvideo_args.pipeline_config,
|
||||
"context_noise", 0)
|
||||
t_context = torch.ones([latents.shape[0]],
|
||||
device=latents.device,
|
||||
dtype=torch.long) * int(context_noise)
|
||||
context_bcthw = current_latents.to(target_dtype)
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled), \
|
||||
set_forward_context(current_timestep=0,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
_ = self.transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
t_expanded_context,
|
||||
kv_cache=self.kv_cache1,
|
||||
crossattn_cache=self.crossattn_cache,
|
||||
current_start=(pos_start_base + start_index) *
|
||||
self.frame_seq_length,
|
||||
start_frame=start_index,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
def _initialize_kv_cache(self, batch_size, dtype, device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
kv_cache1 = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
if self.local_attn_size != -1:
|
||||
kv_cache_size = self.local_attn_size * self.frame_seq_length
|
||||
else:
|
||||
kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames
|
||||
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
kv_cache1.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, kv_cache_size, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
|
||||
self.kv_cache1 = kv_cache1
|
||||
|
||||
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
|
||||
device) -> None:
|
||||
"""
|
||||
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
|
||||
"""
|
||||
crossattn_cache = []
|
||||
num_attention_heads = self.transformer.num_attention_heads
|
||||
attention_head_dim = self.transformer.attention_head_dim
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
crossattn_cache.append({
|
||||
"k":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"v":
|
||||
torch.zeros([
|
||||
batch_size, max_text_len, num_attention_heads,
|
||||
attention_head_dim
|
||||
],
|
||||
dtype=dtype,
|
||||
device=device),
|
||||
"is_init":
|
||||
False,
|
||||
})
|
||||
self.crossattn_cache = crossattn_cache
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
V.none_or_tensor_with_dims(5))
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("eta", batch.eta, V.non_negative_float)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
@@ -4,6 +4,7 @@ Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import math
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
@@ -24,12 +25,13 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
@@ -38,6 +40,13 @@ try:
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
|
||||
from fastvideo.utils import is_vmoba_available
|
||||
vmoba_attn_available = is_vmoba_available()
|
||||
except ImportError:
|
||||
vmoba_attn_available = False
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
@@ -60,11 +69,13 @@ class DenoisingStage(PipelineStage):
|
||||
transformer,
|
||||
scheduler,
|
||||
pipeline=None,
|
||||
transformer_2=None) -> None:
|
||||
transformer_2=None,
|
||||
vae=None) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.transformer_2 = transformer_2
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
|
||||
self.attn_backend = get_attn_backend(
|
||||
@@ -73,6 +84,7 @@ class DenoisingStage(PipelineStage):
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
|
||||
) # hack
|
||||
)
|
||||
@@ -145,7 +157,8 @@ class DenoisingStage(PipelineStage):
|
||||
# Prepare image latents and embeddings for I2V generation
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
assert not torch.isnan(
|
||||
image_embeds[0]).any(), "image_embeds contains nan"
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
@@ -182,17 +195,57 @@ class DenoisingStage(PipelineStage):
|
||||
# Get latents and embeddings
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
assert not torch.isnan(
|
||||
prompt_embeds[0]).any(), "prompt_embeds contains nan"
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds
|
||||
assert neg_prompt_embeds is not None
|
||||
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
|
||||
assert not torch.isnan(
|
||||
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||
if fastvideo_args.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
# TI2V directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
assert self.vae is not None, "VAE is not provided for TI2V task"
|
||||
z = self.vae.encode(batch.pil_image).mean.float()
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
z -= self.vae.shift_factor.to(z.device, z.dtype)
|
||||
else:
|
||||
z -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
z = z * self.vae.scaling_factor.to(z.device, z.dtype)
|
||||
else:
|
||||
z = z * self.vae.scaling_factor
|
||||
|
||||
latent_model_input = latent_model_input.squeeze(0)
|
||||
_, mask2 = masks_like([latent_model_input], zero=True)
|
||||
|
||||
latent_model_input = (1. -
|
||||
mask2[0]) * z + mask2[0] * latent_model_input
|
||||
# latent_model_input = latent_model_input.unsqueeze(0)
|
||||
latent_model_input = latent_model_input.to(get_local_torch_device())
|
||||
latents = latent_model_input
|
||||
F = batch.num_frames
|
||||
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
|
||||
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
seq_len = ((F - 1) // temporal_scale +
|
||||
1) * (batch.height // spatial_scale) * (
|
||||
batch.width // spatial_scale) // (patch_size[1] *
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -217,19 +270,34 @@ class DenoisingStage(PipelineStage):
|
||||
self.transformer.to('cpu')
|
||||
current_model = self.transformer_2
|
||||
current_guidance_scale = batch.guidance_scale_2
|
||||
assert current_model is not None, "current_model is None"
|
||||
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
|
||||
assert not torch.isnan(
|
||||
latent_model_input).any(), "latent_model_input contains nan"
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
timestep = torch.stack([t]).to(get_local_torch_device())
|
||||
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
|
||||
temp_ts = torch.cat([
|
||||
temp_ts,
|
||||
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
|
||||
])
|
||||
timestep = temp_ts.unsqueeze(0)
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
@@ -270,6 +338,31 @@ class DenoisingStage(PipelineStage):
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
elif (vmoba_attn_available
|
||||
and self.attn_backend == VMOBAAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
# Prepare V-MoBA parameters from config
|
||||
moba_params = fastvideo_args.moba_config.copy()
|
||||
moba_params.update({
|
||||
"current_timestep":
|
||||
i,
|
||||
"raw_latent_shape":
|
||||
batch.raw_latent_shape[2:5],
|
||||
"patch_size":
|
||||
fastvideo_args.pipeline_config.dit_config.
|
||||
patch_size,
|
||||
"device":
|
||||
get_local_torch_device(),
|
||||
})
|
||||
attn_metadata = self.attn_metadata_builder.build(
|
||||
**moba_params)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
# TODO(will): finalize the interface. vLLM uses this to
|
||||
@@ -329,6 +422,11 @@ class DenoisingStage(PipelineStage):
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
@@ -677,7 +775,8 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
|
||||
video_raw_latent_shape = latents.shape
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
assert not torch.isnan(
|
||||
prompt_embeds[0]).any(), "prompt_embeds contains nan"
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.pipeline_config.dmd_denoising_steps,
|
||||
dtype=torch.long,
|
||||
@@ -716,7 +815,8 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
batch.image_latent.permute(0, 2, 1, 3, 4)
|
||||
],
|
||||
dim=2).to(target_dtype)
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
assert not torch.isnan(
|
||||
latent_model_input).any(), "latent_model_input contains nan"
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
@@ -778,8 +878,6 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
**pos_cond_kwargs,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
from fastvideo.training.training_utils import (
|
||||
pred_noise_to_pred_video)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noise_latents.flatten(0, 1),
|
||||
|
||||
@@ -92,24 +92,6 @@ class EncodingStage(PipelineStage):
|
||||
latents = latents.to(vae_dtype)
|
||||
latents = self.vae.encode(latents).mean
|
||||
|
||||
# Apply shifting if needed (reverse of decoding)
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents -= self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents -= self.vae.shift_factor
|
||||
|
||||
# Apply scaling factor
|
||||
if (hasattr(self.vae, "scaling_factor")
|
||||
and self.vae.scaling_factor is not None):
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents * self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents * self.vae.scaling_factor
|
||||
|
||||
# Update batch with encoded latents
|
||||
batch.latents = latents
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@ Input validation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -12,6 +14,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import (StageValidators,
|
||||
VerificationResult)
|
||||
from fastvideo.utils import best_output_size
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -94,6 +97,7 @@ class InputValidationStage(PipelineStage):
|
||||
)
|
||||
|
||||
# for i2v, get image from image_path
|
||||
# @TODO(Wei) hard-coded for wan2.2 5b ti2v for now. Should put this in image_encoding stage
|
||||
if batch.image_path is not None:
|
||||
if batch.image_path.endswith(".mp4"):
|
||||
image = load_video(batch.image_path)[0]
|
||||
@@ -101,6 +105,36 @@ class InputValidationStage(PipelineStage):
|
||||
image = load_image(batch.image_path)
|
||||
batch.pil_image = image
|
||||
|
||||
# further processing for ti2v task
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
img = batch.pil_image
|
||||
ih, iw = img.height, img.width
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
|
||||
max_area = 704 * 1280
|
||||
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||
|
||||
scale = max(ow / iw, oh / ih)
|
||||
img = img.resize((round(iw * scale), round(ih * scale)),
|
||||
Image.LANCZOS)
|
||||
logger.info("resized img height: %s, img width: %s", img.height,
|
||||
img.width)
|
||||
|
||||
# center-crop
|
||||
x1 = (img.width - ow) // 2
|
||||
y1 = (img.height - oh) // 2
|
||||
img = img.crop((x1, y1, x1 + ow, y1 + oh))
|
||||
assert img.width == ow and img.height == oh
|
||||
|
||||
# to tensor
|
||||
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
|
||||
self.device).unsqueeze(1)
|
||||
img = img.unsqueeze(0)
|
||||
batch.height = oh
|
||||
batch.width = ow
|
||||
batch.pil_image = img
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
|
||||
@@ -59,58 +59,37 @@ class TextEncodingStage(PipelineStage):
|
||||
assert len(self.text_encoders) == len(
|
||||
fastvideo_args.pipeline_config.text_encoder_configs)
|
||||
|
||||
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
|
||||
self.tokenizers,
|
||||
self.text_encoders,
|
||||
fastvideo_args.pipeline_config.text_encoder_configs,
|
||||
fastvideo_args.pipeline_config.preprocess_text_funcs,
|
||||
fastvideo_args.pipeline_config.postprocess_text_funcs,
|
||||
strict=True):
|
||||
# Encode positive prompt with all available encoders
|
||||
assert batch.prompt is not None
|
||||
prompt_text: str | list[str] = batch.prompt
|
||||
all_indices: list[int] = list(range(len(self.text_encoders)))
|
||||
prompt_embeds_list, prompt_masks_list = self.encode_text(
|
||||
prompt_text,
|
||||
fastvideo_args,
|
||||
encoder_index=all_indices,
|
||||
return_attention_mask=True,
|
||||
)
|
||||
for pe in prompt_embeds_list:
|
||||
batch.prompt_embeds.append(pe)
|
||||
if batch.prompt_attention_mask is not None:
|
||||
for am in prompt_masks_list:
|
||||
batch.prompt_attention_mask.append(am)
|
||||
|
||||
assert isinstance(batch.prompt, str | list)
|
||||
if isinstance(batch.prompt, str):
|
||||
batch.prompt = [batch.prompt]
|
||||
texts = []
|
||||
for prompt_str in batch.prompt:
|
||||
texts.append(preprocess_func(prompt_str))
|
||||
text_inputs = tokenizer(texts,
|
||||
**encoder_config.tokenizer_kwargs).to(
|
||||
get_local_torch_device())
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = text_encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
prompt_embeds = postprocess_func(outputs)
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
if batch.prompt_attention_mask is not None:
|
||||
batch.prompt_attention_mask.append(attention_mask)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
assert isinstance(batch.negative_prompt, str)
|
||||
negative_text = preprocess_func(batch.negative_prompt)
|
||||
negative_text_inputs = tokenizer(
|
||||
negative_text, **encoder_config.tokenizer_kwargs).to(
|
||||
get_local_torch_device())
|
||||
negative_input_ids = negative_text_inputs["input_ids"]
|
||||
negative_attention_mask = negative_text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0,
|
||||
attn_metadata=None):
|
||||
negative_outputs = text_encoder(
|
||||
input_ids=negative_input_ids,
|
||||
attention_mask=negative_attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
negative_prompt_embeds = postprocess_func(negative_outputs)
|
||||
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_prompt_embeds)
|
||||
if batch.negative_attention_mask is not None:
|
||||
batch.negative_attention_mask.append(
|
||||
negative_attention_mask)
|
||||
# Encode negative prompt if CFG is enabled
|
||||
if batch.do_classifier_free_guidance:
|
||||
assert isinstance(batch.negative_prompt, str)
|
||||
neg_embeds_list, neg_masks_list = self.encode_text(
|
||||
batch.negative_prompt,
|
||||
fastvideo_args,
|
||||
encoder_index=all_indices,
|
||||
return_attention_mask=True,
|
||||
)
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
for ne in neg_embeds_list:
|
||||
batch.negative_prompt_embeds.append(ne)
|
||||
if batch.negative_attention_mask is not None:
|
||||
for nm in neg_masks_list:
|
||||
batch.negative_attention_mask.append(nm)
|
||||
|
||||
return batch
|
||||
|
||||
@@ -129,6 +108,171 @@ class TextEncodingStage(PipelineStage):
|
||||
V.none_or_list)
|
||||
return result
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_text(
|
||||
self,
|
||||
text: str | list[str],
|
||||
fastvideo_args: FastVideoArgs,
|
||||
encoder_index: int | list[int] | None = None,
|
||||
return_attention_mask: bool = False,
|
||||
return_type: str = "list", # one of: "list", "dict", "stack"
|
||||
device: torch.device | str | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
max_length: int | None = None,
|
||||
truncation: bool | None = None,
|
||||
padding: bool | str | None = None,
|
||||
):
|
||||
"""
|
||||
Encode plain text using selected text encoder(s) and return embeddings.
|
||||
|
||||
Args:
|
||||
text: A single string or a list of strings to encode.
|
||||
fastvideo_args: The inference arguments providing pipeline config,
|
||||
including tokenizer and encoder settings, preprocess and postprocess
|
||||
functions.
|
||||
encoder_index: Encoder selector by index. Accepts an int or list of ints.
|
||||
return_attention_mask: If True, also return attention masks for each
|
||||
selected encoder.
|
||||
return_type: "list" (default) returns a list aligned with selection;
|
||||
"dict" returns a dict keyed by encoder index as a string; "stack" stacks along a
|
||||
new first dimension (requires matching shapes).
|
||||
device: Optional device override for inputs; defaults to local torch device.
|
||||
dtype: Optional dtype to cast returned embeddings to.
|
||||
max_length: Optional per-call tokenizer override.
|
||||
truncation: Optional per-call tokenizer override.
|
||||
padding: Optional per-call tokenizer override.
|
||||
|
||||
Returns:
|
||||
Depending on return_type and return_attention_mask:
|
||||
- list: List[Tensor] or (List[Tensor], List[Tensor])
|
||||
- dict: Dict[str, Tensor] or (Dict[str, Tensor], Dict[str, Tensor])
|
||||
- stack: Tensor of shape [num_encoders, ...] or a tuple with stacked
|
||||
attention masks
|
||||
"""
|
||||
|
||||
assert len(self.tokenizers) == len(self.text_encoders)
|
||||
assert len(self.text_encoders) == len(
|
||||
fastvideo_args.pipeline_config.text_encoder_configs)
|
||||
|
||||
# Resolve selection into indices
|
||||
encoder_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
|
||||
if encoder_index is None:
|
||||
indices: list[int] = [0]
|
||||
elif isinstance(encoder_index, int):
|
||||
indices = [encoder_index]
|
||||
else:
|
||||
indices = list(encoder_index)
|
||||
# validate range
|
||||
num_encoders = len(self.text_encoders)
|
||||
for idx in indices:
|
||||
if idx < 0 or idx >= num_encoders:
|
||||
raise IndexError(
|
||||
f"encoder index {idx} out of range [0, {num_encoders-1}]")
|
||||
|
||||
# Validate indices are within range
|
||||
num_encoders = len(self.text_encoders)
|
||||
|
||||
# Normalize input to list[str]
|
||||
assert isinstance(text, str | list)
|
||||
if isinstance(text, str):
|
||||
texts: list[str] = [text]
|
||||
else:
|
||||
texts = text
|
||||
|
||||
embeds_list: list[torch.Tensor] = []
|
||||
attn_masks_list: list[torch.Tensor] = []
|
||||
|
||||
preprocess_funcs = fastvideo_args.pipeline_config.preprocess_text_funcs
|
||||
postprocess_funcs = fastvideo_args.pipeline_config.postprocess_text_funcs
|
||||
encoder_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
|
||||
|
||||
if return_type not in ("list", "dict", "stack"):
|
||||
raise ValueError(
|
||||
f"Invalid return_type '{return_type}'. Expected one of: 'list', 'dict', 'stack'"
|
||||
)
|
||||
|
||||
target_device = device if device is not None else get_local_torch_device(
|
||||
)
|
||||
|
||||
for i in indices:
|
||||
tokenizer = self.tokenizers[i]
|
||||
text_encoder = self.text_encoders[i]
|
||||
encoder_config = encoder_cfgs[i]
|
||||
preprocess_func = preprocess_funcs[i]
|
||||
postprocess_func = postprocess_funcs[i]
|
||||
|
||||
processed_texts: list[str] = []
|
||||
for prompt_str in texts:
|
||||
processed_texts.append(preprocess_func(prompt_str))
|
||||
|
||||
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
|
||||
if max_length is not None:
|
||||
tok_kwargs["max_length"] = max_length
|
||||
if truncation is not None:
|
||||
tok_kwargs["truncation"] = truncation
|
||||
if padding is not None:
|
||||
tok_kwargs["padding"] = padding
|
||||
|
||||
text_inputs = tokenizer(processed_texts,
|
||||
**tok_kwargs).to(target_device)
|
||||
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = text_encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
prompt_embeds = postprocess_func(outputs)
|
||||
if dtype is not None:
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype)
|
||||
embeds_list.append(prompt_embeds)
|
||||
if return_attention_mask:
|
||||
attn_masks_list.append(attention_mask)
|
||||
|
||||
# Shape results according to return_type
|
||||
if return_type == "list":
|
||||
if return_attention_mask:
|
||||
return embeds_list, attn_masks_list
|
||||
return embeds_list
|
||||
|
||||
if return_type == "dict":
|
||||
key_strs = [str(i) for i in indices]
|
||||
embeds_dict = {
|
||||
k: v
|
||||
for k, v in zip(key_strs, embeds_list, strict=False)
|
||||
}
|
||||
if return_attention_mask:
|
||||
attn_dict = {
|
||||
k: v
|
||||
for k, v in zip(key_strs, attn_masks_list, strict=False)
|
||||
}
|
||||
return embeds_dict, attn_dict
|
||||
return embeds_dict
|
||||
|
||||
# return_type == "stack"
|
||||
# Validate shapes are compatible
|
||||
base_shape = list(embeds_list[0].shape)
|
||||
for t in embeds_list[1:]:
|
||||
if list(t.shape) != base_shape:
|
||||
raise ValueError(
|
||||
f"Cannot stack embeddings with differing shapes: {[list(t.shape) for t in embeds_list]}"
|
||||
)
|
||||
stacked_embeds = torch.stack(embeds_list, dim=0)
|
||||
if return_attention_mask:
|
||||
base_mask_shape = list(attn_masks_list[0].shape)
|
||||
for m in attn_masks_list[1:]:
|
||||
if list(m.shape) != base_mask_shape:
|
||||
raise ValueError(
|
||||
f"Cannot stack attention masks with differing shapes: {[list(m.shape) for m in attn_masks_list]}"
|
||||
)
|
||||
stacked_masks = torch.stack(attn_masks_list, dim=0)
|
||||
return stacked_embeds, stacked_masks
|
||||
return stacked_embeds
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify text encoding stage outputs."""
|
||||
|
||||
@@ -159,6 +159,20 @@ class CudaPlatformBase(Platform):
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Video Sparse Attention backend is not installed. ") from e
|
||||
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
||||
try:
|
||||
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
|
||||
moba_attn_varlen)
|
||||
from fastvideo.attention.backends.vmoba import ( # noqa: F401
|
||||
VMOBAAttentionBackend)
|
||||
logger.info("Using Video MOBA Attention backend.")
|
||||
|
||||
return "fastvideo.attention.backends.vmoba.VMOBAAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.error(
|
||||
"Failed to import Video MoBA Attention backend: %s", str(e))
|
||||
raise ImportError(
|
||||
"Video MoBA Attention backend is not installed. ") from e
|
||||
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
@@ -19,6 +19,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
def test_inference_vmoba():
|
||||
"""Test FastVideo VMOBA_ATTN inference pipeline"""
|
||||
|
||||
num_gpus = "1"
|
||||
model_base = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
output_dir = Path("outputs_video/vmoba_1.3B/")
|
||||
moba_config = "fastvideo/configs/backend/vmoba/wan_1.3B_77_480_832.json"
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VMOBA_ATTN"
|
||||
|
||||
cmd = [
|
||||
"fastvideo", "generate",
|
||||
"--model-path", model_base,
|
||||
"--sp-size", num_gpus,
|
||||
"--tp-size", "1",
|
||||
"--num-gpus", num_gpus,
|
||||
"--dit-cpu-offload", "False",
|
||||
"--vae-cpu-offload", "False",
|
||||
"--text-encoder-cpu-offload", "True",
|
||||
"--pin-cpu-memory", "False",
|
||||
"--height", "480",
|
||||
"--width", "832",
|
||||
"--num-frames", "77",
|
||||
"--num-inference-steps", "50",
|
||||
"--moba-config-path", moba_config,
|
||||
"--fps", "16",
|
||||
"--guidance-scale", "6.0",
|
||||
"--flow-shift", "8.0",
|
||||
"--prompt", "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.",
|
||||
"--negative-prompt", (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, "
|
||||
"works, paintings, images, static, overall gray, worst quality, low quality, "
|
||||
"JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, "
|
||||
"poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, "
|
||||
"still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
),
|
||||
"--seed", "1024",
|
||||
"--output-path", str(output_dir),
|
||||
]
|
||||
|
||||
subprocess.run(cmd, check=True)
|
||||
|
||||
assert output_dir.exists(), f"Output directory {output_dir} does not exist"
|
||||
|
||||
video_files = list(output_dir.glob("*.mp4"))
|
||||
assert len(video_files) > 0, "No video files were generated"
|
||||
|
||||
for video_file in video_files:
|
||||
assert video_file.stat().st_size > 0, f"Video file {video_file} is empty"
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_inference_vmoba()
|
||||
@@ -102,10 +102,18 @@ def run_precision_tests_STA():
|
||||
def run_precision_tests_VSA():
|
||||
run_test("python csrc/attn/tests/test_vsa.py")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_precision_tests_vmoba():
|
||||
run_test("pytest csrc/attn/vmoba_attn/tests/test_vmoba_attn.py")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_inference_tests_vmoba():
|
||||
run_test('python fastvideo/tests/inference/vmoba/test_vmoba_inference.py')
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=3600)
|
||||
def run_inference_lora_tests():
|
||||
run_test("pytest ./fastvideo/tests/inference/lora/test_lora_inference_similarity.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900)
|
||||
def run_distill_dmd_tests():
|
||||
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
|
||||
BIN
Binary file not shown.
@@ -4,7 +4,8 @@ The reference videos in the `*_reference_videos` directory are used as part of a
|
||||
|
||||
run `bash update_reference_videos.sh` from inside the `fastvideo/tests/ssim/` directory after running `test_inference_similarity.py` to update reference videos. Note: make sure to update the path to the corresponding device.
|
||||
|
||||
all reference videos are were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
|
||||
reference videos were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
|
||||
causal videos were generated on commit b318063c0a4618f1d5d99ea82ca67a06aad0d19d
|
||||
|
||||
## Generation Details
|
||||
|
||||
@@ -76,4 +77,4 @@ Wan2.1-I2V-14B-480P-Diffusers: {
|
||||
### Image-to-Video Prompts
|
||||
|
||||
1. "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
Image path: "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
Image path: "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
from fastvideo.worker.multiproc_executor import MultiprocExecutor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
device_name = torch.cuda.get_device_name()
|
||||
device_reference_folder_suffix = '_reference_videos'
|
||||
|
||||
if "A40" in device_name:
|
||||
device_reference_folder = "A40" + device_reference_folder_suffix
|
||||
elif "L40S" in device_name:
|
||||
device_reference_folder = "L40S" + device_reference_folder_suffix
|
||||
|
||||
# Base parameters from the shell script
|
||||
|
||||
SF_WAN_T2V_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"num_inference_steps": 4,
|
||||
"seed": 1024,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
}
|
||||
|
||||
|
||||
MODEL_TO_PARAMS = {
|
||||
"SFWan2.1-T2V-1.3B-Diffusers": SF_WAN_T2V_PARAMS,
|
||||
}
|
||||
|
||||
I2V_MODEL_TO_PARAMS = {
|
||||
}
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
|
||||
# "A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
|
||||
]
|
||||
|
||||
I2V_TEST_PROMPTS = [
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot.",
|
||||
]
|
||||
|
||||
I2V_IMAGE_PATHS = [
|
||||
"https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg",
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
|
||||
def test_causal_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"""
|
||||
Test that runs inference with different parameters and compares the output
|
||||
to reference videos using SSIM.
|
||||
"""
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
|
||||
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
|
||||
output_video_name = f"{prompt[:100]}.mp4"
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
|
||||
num_inference_steps = BASE_PARAMS["num_inference_steps"]
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": BASE_PARAMS["num_gpus"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"dit_cpu_offload": True,
|
||||
}
|
||||
if BASE_PARAMS.get("vae_sp"):
|
||||
init_kwargs["vae_sp"] = True
|
||||
init_kwargs["vae_tiling"] = True
|
||||
#if "text-encoder-precision" in BASE_PARAMS:
|
||||
# init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
|
||||
|
||||
generation_kwargs = {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"output_path": output_dir,
|
||||
"height": BASE_PARAMS["height"],
|
||||
"width": BASE_PARAMS["width"],
|
||||
"num_frames": BASE_PARAMS["num_frames"],
|
||||
"seed": BASE_PARAMS["seed"],
|
||||
}
|
||||
if "neg_prompt" in BASE_PARAMS:
|
||||
generation_kwargs["neg_prompt"] = BASE_PARAMS["neg_prompt"]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_path=BASE_PARAMS["model_path"], **init_kwargs)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
if isinstance(generator.executor, MultiprocExecutor):
|
||||
generator.executor.shutdown()
|
||||
|
||||
assert os.path.exists(
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}")
|
||||
|
||||
# Find the matching reference video based on the prompt
|
||||
reference_video_name = None
|
||||
|
||||
for filename in os.listdir(reference_folder):
|
||||
if filename.endswith('.mp4') and prompt[:100] in filename:
|
||||
reference_video_name = filename
|
||||
break
|
||||
|
||||
if not reference_video_name:
|
||||
logger.error(f"Reference video not found for prompt: {prompt} with backend: {ATTENTION_BACKEND}")
|
||||
raise FileNotFoundError(f"Reference video missing")
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
|
||||
logger.info(
|
||||
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
|
||||
)
|
||||
ssim_values = compute_video_ssim_torchvision(reference_video_path,
|
||||
generated_video_path,
|
||||
use_ms_ssim=True)
|
||||
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info(f"SSIM mean value: {mean_ssim}")
|
||||
logger.info(f"Writing SSIM results to directory: {output_dir}")
|
||||
|
||||
success = write_ssim_results(output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path, num_inference_steps,
|
||||
prompt)
|
||||
|
||||
if not success:
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
min_acceptable_ssim = 0.98
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
@@ -101,7 +101,7 @@ I2V_IMAGE_PATHS = [
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", I2V_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN", "TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
@pytest.mark.parametrize("model_id", list(I2V_MODEL_TO_PARAMS.keys()))
|
||||
def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"""
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user