Compare commits

..
Author SHA1 Message Date
SolitaryThinker 104a539a22 update docker 2025-12-24 09:32:42 +00:00
219 changed files with 7617 additions and 14289 deletions
+60 -21
View File
@@ -22,7 +22,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- TEST_TYPE=encoder
@@ -35,7 +35,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- TEST_TYPE=vae
@@ -61,7 +61,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
command: "timeout 45m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- TEST_TYPE=ssim
@@ -76,7 +76,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
@@ -129,7 +129,11 @@ steps:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "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:
@@ -141,7 +145,10 @@ steps:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "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:
@@ -152,16 +159,48 @@ steps:
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "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:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Kernel Tests"
label: "Precision Tests STA"
env:
- TEST_TYPE=kernel_tests
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "csrc/attn/video_sparse_attn/**"
- "csrc/attn/video_sparse_attn/tk/**"
- "csrc/attn/tests/test_vsa.py"
- "csrc/attn/video_sparse_attn/setup.py"
- "csrc/attn/video_sparse_attn/config_vsa.py"
- "csrc/attn/video_sparse_attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- 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"
@@ -183,14 +222,14 @@ steps:
- TEST_TYPE=unit_test
agents:
queue: "default"
# - path:
# - "scripts/lora_extraction/**"
# - "pyproject.toml"
# - "docker/Dockerfile.python3.12"
# config:
# command: "timeout 90m .buildkite/scripts/pr_test.sh"
# label: "LoRA Extraction Tests"
# env:
# - TEST_TYPE=lora_extraction
# agents:
# queue: "default"
- path:
- "scripts/lora_extraction/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
+11 -3
View File
@@ -93,9 +93,13 @@ case "$TEST_TYPE" in
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"kernel_tests")
log "Running kernel tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_kernel_tests"
"precision_sta")
log "Running precision STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
;;
"precision_vsa")
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
;;
"inference_lora")
log "Running LoRA tests..."
@@ -114,6 +118,10 @@ case "$TEST_TYPE" in
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"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
+76 -47
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "fastvideo-kernel/pyproject.toml"
- "csrc/fastvideo_kernel/pyproject.toml"
workflow_dispatch:
jobs:
@@ -23,15 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd fastvideo-kernel
cd csrc/fastvideo_kernel
# Get current commit's version from pyproject.toml
# Use ^ to match start of line to avoid matching minimum-version
NEW_VERSION=$(grep -oP '^version\s*=\s*"\K[^"]+' pyproject.toml)
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
echo "New version: $NEW_VERSION"
# Get previous version from git history
# Note: git show expects path relative to repo root
OLD_VERSION=$(git show HEAD~1:fastvideo-kernel/pyproject.toml | grep -oP '^version\s*=\s*"\K[^"]+' || echo "0.0.0")
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -53,18 +51,15 @@ jobs:
fail-fast: false
matrix:
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12']
python-version: ['3.10', '3.11', '3.12', '3.13']
torch-cuda:
# - torch-version: '2.5.1'
# cuda-version: '12.4.1'
# torch-cuda-short: 'cu124'
# - torch-version: '2.6.0'
# cuda-version: '12.6.3'
# torch-cuda-short: 'cu126'
# - torch-version: '2.7.1'
# cuda-version: '12.8.0'
# torch-cuda-short: 'cu128'
- torch-version: '2.9.1'
- torch-version: '2.5.1'
cuda-version: '12.4.1'
torch-cuda-short: 'cu124'
- torch-version: '2.6.0'
cuda-version: '12.6.3'
torch-cuda-short: 'cu126'
- torch-version: '2.7.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
@@ -143,32 +138,30 @@ jobs:
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
pip install setuptools ninja packaging wheel triton
cd fastvideo-kernel
cd csrc/fastvideo_kernel
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
# Release builds are produced on GPU-less runners, so force-enable TK and target Hopper.
export TORCH_CUDA_ARCH_LIST="9.0a"
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
# Build standard wheel (no local version suffix) for PyPI
python -m build --wheel --outdir dist
# Fix the wheel to be manylinux compliant
pip install auditwheel
# Target manylinux_2_35 (Ubuntu 22.04 native)
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist
# Move fixed wheels back to dist for upload consistency
rm dist/*.whl
mv fixed_dist/*.whl dist/
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/fastvideo_kernel
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
# Get the correct version format
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
# Rename with version information
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
- name: Upload wheel artifact
# Only upload if it's the "main" CUDA version we want on PyPI
# We upload all to artifacts for inspection/GH releases, but give them distinct artifact names
uses: actions/upload-artifact@v4
with:
name: fastvideo_kernel-py${{ matrix.python-version }}-${{ matrix.torch-cuda.torch-cuda-short }}-torch${{ matrix.torch-cuda.torch-version }}
path: fastvideo-kernel/dist/*.whl
name: ${{ env.wheel_name }}-py${{ matrix.python-version }}
path: csrc/fastvideo_kernel/dist/*.whl
retention-days: 90
publish_package:
@@ -186,22 +179,58 @@ jobs:
with:
python-version: '3.10'
- name: Download PyPI wheels
uses: actions/download-artifact@v4
- name: Install CUDA 12.4.1
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
path: fastvideo-kernel/dist/
pattern: 'fastvideo_kernel-py*'
merge-multiple: true
cuda: 12.4.1
linux-local-args: '["--toolkit"]'
method: 'network'
sub-packages: '["nvcc"]'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-12.4.1
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch 2.5.1+cu12.4.1
run: |
pip install --upgrade pip
pip install typing-extensions==4.12.2
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build source distribution
run: |
pip install build scikit-build-core cmake ninja
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
cd fastvideo-kernel
# We don't need full CUDA/Torch to just package the source (sdist)
python -m build --sdist --outdir dist
pip install setuptools ninja packaging wheel triton
cd csrc/fastvideo_kernel
git submodule update --init --recursive
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: fastvideo-kernel/dist/
packages-dir: csrc/fastvideo_kernel/dist/
+6 -5
View File
@@ -1,6 +1,7 @@
[submodule "fastvideo-kernel/include/tk"]
path = fastvideo-kernel/include/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
[submodule "fastvideo-kernel/include/cutlass"]
path = fastvideo-kernel/include/cutlass
url = https://github.com/NVIDIA/cutlass.git
+1 -1
View File
@@ -4,7 +4,7 @@ default_stages:
exclude: |
(?x)(
fastvideo/third_party/.*|
fastvideo-kernel/.*|
csrc/.*|
assets/.*|
tests/.*|
demo/.*|
+15 -21
View File
@@ -1,47 +1,41 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
</div>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/sv3MMKyv" 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/"><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-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/c7g1qdD" target="_blank"> <b> WeChat </b> </a> |
</p>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
<div align="center">
<img src=assets/fastwan.png width="90%"/>
</div>
## NEWS
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
<details>
<summary>More</summary>
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
</details>
## Key Features
FastVideo has the following features:
- End-to-end post-training support for bidirectional and autoregressive models:
- End-to-end post-training support:
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
- Data preprocessing pipeline for video data
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
- Data preprocessing pipeline for video, image, and text data
- Distribution Matching Distillation (DMD2) stepwise distillation.
- Sparse attention with [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achineve >50x denoising speedup
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
- Causal distillation through Self-Forcing
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
- State-of-the-art performance optimizations for inference
- Sequence Parallelism for distributed inference
- Multiple state-of-the-art attention backends
- User-friendly CLI and Python API
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/optimizations/) for full list of supported optimizations.
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- Diverse hardware and OS support
- Support H100, A100, 4090
- Support Linux, Windows, MacOS
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/hardware_support/) for full list of supported hardware and OS.
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 490 KiB

Binary file not shown.
+3 -6
View File
@@ -38,8 +38,7 @@ python -m benchmarks.fvd.cli \
--num-frames 32 \
--clip-strategy random \
--batch-size 32 \
--seed 42 \
--extractor clip
--seed 42
```
**Standard protocols:**
@@ -52,8 +51,6 @@ python -m benchmarks.fvd.cli \
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
```
This would use i3d model by default as the feature extractor
**Feature caching** (speed up repeated evaluations):
```bash
@@ -61,7 +58,7 @@ python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f \
--cache-real-features fvd-cache/extractor_name # Directory path (will save/load fvd-cache/extractor_name/extractor-name_real_features.pkl)
--cache-real-features cache/real # Directory path (will save/load cache/real/real_features.pkl)
```
Run `python -m benchmarks.fvd.cli --help` for all options.
@@ -86,7 +83,6 @@ batch_size=32, # GPU batch size
device='cuda', # cuda|cpu
cache_real_features=None, # Cache path for speed
seed=42, # Reproducibility
extractor='i3d', # i3d|clip|videomae
```
## Programmatic Usage
@@ -101,6 +97,7 @@ print(f"FVD: {results['fvd']:.2f}")
## Notes
- I3D model auto-downloads from Hugging Face on first run
- Requires minimum 10 frames per clip
- Supports both video files (.mp4, .avi, etc.) and frame directories
- `--cache-real-features` expects a **directory path** (e.g., `cache/real`), it will automatically create/load `real_features.pkl` inside that directory
+1 -4
View File
@@ -13,8 +13,7 @@ from .fvd import (
compute_statistics,
FVDConfig,
)
from .feature_extractors import (BaseFeatureExtractor, I3DFeatureExtractor,
load_extractor)
from .i3d_model import I3DFeatureExtractor
from .video_utils import (
load_video_auto,
sample_clips_from_video,
@@ -28,9 +27,7 @@ __all__ = [
'compute_frechet_distance',
'compute_statistics',
'FVDConfig',
'BaseFeatureExtractor',
'I3DFeatureExtractor',
'load_extractor',
'load_video_auto',
'sample_clips_from_video',
'load_video_clips_streaming',
+119 -41
View File
@@ -1,64 +1,122 @@
import argparse
import json
import sys
import traceback
from pathlib import Path
from .fvd import compute_fvd_with_config, FVDConfig
def main() -> int:
parser = argparse.ArgumentParser(
description='Compute Fréchet Video Distance (FVD)')
description='Compute Fréchet Video Distance (FVD)',
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Standard FVD2048_16f protocol
python -m fastvideo.benchmarks.fvd.cli \\
--real-path data/real/ \\
--gen-path outputs/gen/ \\
--protocol fvd2048_16f
# Custom configuration
python -m fastvideo.benchmarks.fvd.cli \\
--real-path data/real/ \\
--gen-path outputs/gen/ \\
--num-videos 1024 \\
--num-frames 32 \\
--clip-strategy random \\
--frame-stride 2
""")
# Required arguments
parser.add_argument('--real-path',
type=str,
required=True,
help='Path to real videos')
help='Path to real videos directory')
parser.add_argument('--gen-path',
type=str,
required=True,
help='Path to generated videos')
help='Path to generated videos directory')
# Extractor selection
parser.add_argument('--extractor',
type=str,
default='i3d',
choices=['i3d', 'clip', 'videomae'],
help='Feature extractor model to use (default: i3d)')
# Reproducibility
parser.add_argument(
'--seed',
type=int,
default=None,
help='Random seed for reproducibility (np.random, random, torch)')
# Standard args
parser.add_argument('--seed',
type=int,
default=None,
help='Random seed for reproducibility')
# Protocol presets
parser.add_argument('--protocol',
type=str,
default=None,
choices=['fvd2048_16f', 'fvd2048_128f', 'quick_test'],
choices=[
'fvd2048_16f', 'fvd2048_128f',
'fvd2048_128f_subsample8', 'quick_test'
],
help='Use standard protocol (overrides other settings)')
# Video selection
parser.add_argument('--num-videos',
type=int,
default=2048,
help='Number of videos to use')
help='Number of videos to use (default: 2048)')
# Clip sampling
parser.add_argument('--num-frames',
type=int,
default=16,
help='Number of frames per clip')
parser.add_argument('--clip-strategy',
type=str,
default='beginning',
help='Clip sampling strategy')
help='Number of frames per clip (default: 16)')
parser.add_argument('--num-clips',
type=int,
default=1,
help='Number of clips per video (default: 1)')
parser.add_argument(
'--clip-strategy',
type=str,
default='beginning',
choices=['beginning', 'random', 'uniform', 'middle', 'sliding', 'all'],
help='Clip sampling strategy (default: beginning)')
parser.add_argument(
'--frame-stride',
type=int,
default=1,
help='Frame stride for FPS subsampling (default: 1, no subsampling)')
parser.add_argument('--temporal-stride',
type=int,
default=1,
help='Temporal stride for sliding window (default: 1)')
# Data processing
parser.add_argument('--no-frame-dirs',
action='store_true',
help='Disable frame directory support')
# Computation
parser.add_argument('--batch-size',
type=int,
default=32,
help='Batch size for feature extraction')
help='Batch size for feature extraction (default: 32)')
parser.add_argument('--device',
type=str,
default='cuda',
help='Device to use (cuda or cpu)')
choices=['cuda', 'cpu'],
help='Device to use (default: cuda)')
# Caching
parser.add_argument('--cache-real-features',
type=str,
default=None,
help='Path to cache real video features')
parser.add_argument('--i3d-model-path',
type=str,
default=None,
help='Custom cache path for I3D model')
# Output
parser.add_argument('--output',
type=str,
default='fvd_results.json',
help='Output JSON file (default: fvd_results.json)')
parser.add_argument('--quiet',
action='store_true',
help='Suppress progress output')
@@ -70,36 +128,56 @@ def main() -> int:
protocol_map = {
'fvd2048_16f': FVDConfig.fvd2048_16f,
'fvd2048_128f': FVDConfig.fvd2048_128f,
'fvd2048_128f_subsample8': FVDConfig.fvd2048_128f_subsample8,
'quick_test': FVDConfig.quick_test,
}
config = protocol_map[args.protocol]()
# Apply overrides
# Override device and caching from args
config.device = args.device
config.cache_real_features = args.cache_real_features
config.extractor_model = args.extractor # Apply extractor arg
config.i3d_model_path = args.i3d_model_path
config.batch_size = args.batch_size
config.seed = args.seed
else:
config = FVDConfig(
num_videos=args.num_videos,
num_frames_per_clip=args.num_frames,
extractor_model=args.extractor, # Apply extractor arg
clip_strategy=args.clip_strategy,
batch_size=args.batch_size,
device=args.device,
cache_real_features=args.cache_real_features,
seed=args.seed)
# Custom config from args
config = FVDConfig(num_videos=args.num_videos,
num_frames_per_clip=args.num_frames,
num_clips_per_video=args.num_clips,
clip_strategy=args.clip_strategy,
frame_stride=args.frame_stride,
temporal_stride=args.temporal_stride,
support_frame_dirs=not args.no_frame_dirs,
batch_size=args.batch_size,
device=args.device,
cache_real_features=args.cache_real_features,
i3d_model_path=args.i3d_model_path,
seed=args.seed)
# Compute FVD
try:
_ = compute_fvd_with_config(
args.real_path, # noqa: F841
args.gen_path,
config,
verbose=not args.quiet)
results = compute_fvd_with_config(real_videos=args.real_path,
gen_videos=args.gen_path,
config=config,
verbose=not args.quiet)
# Save results
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w') as f:
json.dump(results, f, indent=2)
print(f"\nResults saved to {output_path}")
print(f"FVD: {results['fvd']:.2f}")
print(f"Protocol: {results['protocol']}")
return 0
except Exception as e:
print(f"Error: {e}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
import traceback
traceback.print_exc()
return 1
-264
View File
@@ -1,264 +0,0 @@
"""
Pluggable Feature Extractors for FVD Computation.
Supports I3D (standard), CLIP, and VideoMAE via a common interface.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from abc import ABC, abstractmethod
from huggingface_hub import hf_hub_download
from tqdm import tqdm
try:
from transformers import CLIPModel, CLIPProcessor, VideoMAEModel
TRANSFORMERS_AVAILABLE = True
except ImportError:
TRANSFORMERS_AVAILABLE = False
class BaseFeatureExtractor(ABC, nn.Module):
"""Abstract base class for all video feature extractors."""
def __init__(self, device: str = 'cuda'):
super().__init__()
self.device = torch.device(
device if torch.cuda.is_available() else 'cpu')
@property
@abstractmethod
def feature_dim(self) -> int:
"""Dimension of the output feature vector."""
pass
@abstractmethod
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Args:
videos: [B, T, C, H, W] in [0, 255] range.
Returns:
Preprocessed tensor ready for the model.
"""
pass
@abstractmethod
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
"""
Extract features for a single batch.
Args:
videos: [B, T, C, H, W] (raw input)
Returns:
Features: [B, feature_dim]
"""
pass
@torch.no_grad()
def extract_features(self,
videos: torch.Tensor,
batch_size: int = 32,
verbose: bool = True) -> torch.Tensor:
"""
Extract features for a large tensor of videos by batching.
"""
N = len(videos)
all_features = []
iterator = range(0, N, batch_size)
if verbose:
iterator = tqdm(
iterator,
desc=f"Extracting features ({self.__class__.__name__})")
for i in iterator:
batch = videos[i:i + batch_size].to(self.device)
features = self.extract_features_batch(batch)
all_features.append(features.cpu())
return torch.cat(all_features, dim=0)
# 1. I3D Extractor (The Standard FVD Metric)
class I3DFeatureExtractor(BaseFeatureExtractor):
REPO_ID = 'flateon/FVD-I3D-torchscript'
MODEL_FILENAME = 'i3d_torchscript.pt'
def __init__(self, device: str = 'cuda', cache_dir: str | None = None):
super().__init__(device)
self.cache_dir = cache_dir
self.model = self._load_model()
self.model.eval()
self.model.to(self.device)
@property
def feature_dim(self) -> int:
return 400
def _load_model(self) -> torch.nn.Module:
try:
model_path = hf_hub_download(repo_id=self.REPO_ID,
filename=self.MODEL_FILENAME,
cache_dir=self.cache_dir)
return torch.jit.load(model_path, map_location=self.device)
except Exception as e:
raise RuntimeError(f"Failed to load I3D model: {e}") from e
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""Standard I3D preprocessing: Resize to 224, Norm to [-1, 1]."""
B, T, C, H, W = videos.shape
if T < 10:
raise ValueError(f"I3D requires at least 10 frames, got {T}")
# Normalize to [0, 1]
if videos.max() > 1.0:
videos = videos / 255.0
# Scale to [-1, 1]
videos = videos * 2.0 - 1.0
# Resize to 224x224
if H != 224 or W != 224:
videos = videos.reshape(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.reshape(B, T, C, 224, 224)
# [B, T, C, H, W] -> [B, C, T, H, W]
return videos.permute(0, 2, 1, 3, 4).contiguous()
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
batch = self.preprocess(videos)
# TorchScript I3D returns raw logits when return_features=True
return self.model(batch,
rescale=False,
resize=False,
return_features=True)
# 2. CLIP Extractor (Semantic/Content Quality)
class CLIPFeatureExtractor(BaseFeatureExtractor):
def __init__(self,
device: str = 'cuda',
model_name: str = "openai/clip-vit-base-patch32"):
if not TRANSFORMERS_AVAILABLE:
raise ImportError(
"Please install transformers: pip install transformers")
super().__init__(device)
self.processor = CLIPProcessor.from_pretrained(model_name)
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
self.model.eval()
self._feature_dim = self.model.config.projection_dim
@property
def feature_dim(self) -> int:
return self._feature_dim
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
# Ensure values are [0, 255]
if videos.max() <= 1.0:
videos = videos * 255.0
return videos.to(torch.uint8)
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
# Input: [B, T, C, H, W]
B, T, C, H, W = videos.shape
videos = self.preprocess(videos)
# Flatten B*T to treat frames as images
images = videos.view(B * T, C, H, W)
# HF Processor
inputs = self.processor(images=images,
return_tensors="pt",
padding=True)
inputs = {k: v.to(self.device) for k, v in inputs.items()}
# Extract features [B*T, Dim]
outputs = self.model.get_image_features(**inputs)
# Reshape [B, T, Dim] and Average Pooling over time
outputs = outputs.view(B, T, -1)
return outputs.mean(dim=1)
# 3. VideoMAE Extractor (Structure/Motion Quality)
class VideoMAEFeatureExtractor(BaseFeatureExtractor):
def __init__(self,
device: str = 'cuda',
model_name: str = "MCG-NJU/videomae-base"):
if not TRANSFORMERS_AVAILABLE:
raise ImportError(
"Please install transformers: pip install transformers")
super().__init__(device)
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
self.model.eval()
self.register_buffer(
'mean',
torch.tensor([0.485, 0.456, 0.406],
device=self.device).view(1, 1, 3, 1, 1))
self.register_buffer(
'std',
torch.tensor([0.229, 0.224, 0.225],
device=self.device).view(1, 1, 3, 1, 1))
@property
def feature_dim(self) -> int:
return self.model.config.hidden_size
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Efficient GPU-based preprocessing.
Input: [B, T, C, H, W] in range [0, 255]
"""
B, T, C, H, W = videos.shape
# 1. Resize to 224x224
if H != 224 or W != 224:
videos = videos.view(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.view(B, T, C, 224, 224)
# 2. Normalize to [0, 1]
if videos.dtype != torch.float32:
videos = videos.float()
if videos.max() > 1.0:
videos = videos / 255.0
# 3. Apply ImageNet Mean/Std
return (videos - self.mean) / self.std
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
# Input: [B, T, C, H, W]
# Fast GPU Preprocessing
pixel_values = self.preprocess(videos)
# Forward pass
outputs = self.model(pixel_values)
# Global Average Pooling of last hidden state [B, T_patches, 768] -> [B, 768]
return outputs.last_hidden_state.mean(dim=1)
# Factory
def load_extractor(name: str, device: str = 'cuda') -> BaseFeatureExtractor:
name = name.lower()
if name == 'i3d':
return I3DFeatureExtractor(device)
elif name == 'clip':
return CLIPFeatureExtractor(device)
elif name == 'videomae':
return VideoMAEFeatureExtractor(device)
else:
raise ValueError(
f"Unknown extractor: {name}. Options: i3d, clip, videomae")
+150 -108
View File
@@ -5,7 +5,8 @@ from pathlib import Path
from collections.abc import Iterator
import pickle
from dataclasses import dataclass, field
from .feature_extractors import BaseFeatureExtractor, load_extractor
from .i3d_model import I3DFeatureExtractor
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
@@ -54,9 +55,6 @@ class FVDConfig:
# Video selection
num_videos: int = 2048
# Feature Extractor Selection
extractor_model: str = 'i3d' # Options: 'i3d', 'clip', 'videomae'
# Clip sampling
num_frames_per_clip: int = 16
num_clips_per_video: int = 1
@@ -87,7 +85,11 @@ class FVDConfig:
@classmethod
def fvd2048_16f(cls) -> 'FVDConfig':
"""Standard FVD protocol: 2048 videos, 16 frames, beginning clip."""
"""
Standard FVD protocol: 2048 videos, 16 frames, beginning clip.
most common FVD configuration used in papers
"""
return cls(num_videos=2048,
num_frames_per_clip=16,
clip_strategy='beginning',
@@ -101,6 +103,18 @@ class FVDConfig:
clip_strategy='beginning',
use_streaming=True)
@classmethod
def fvd2048_128f_subsample8(cls) -> 'FVDConfig':
"""
Long video with FPS subsampling: 2048 videos, 128 frames (every 8th).
Used for very long videos - samples every 8th frame
"""
return cls(num_videos=2048,
num_frames_per_clip=16,
frame_stride=8,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def quick_test(cls) -> 'FVDConfig':
"""Quick test config: 100 videos, 16 frames."""
@@ -110,13 +124,22 @@ class FVDConfig:
def to_dict(self) -> dict:
"""Export config to dict for logging"""
d = self.__dict__.copy()
d['clip_strategy'] = str(self.clip_strategy)
return d
return {
'num_videos': self.num_videos,
'num_frames_per_clip': self.num_frames_per_clip,
'num_clips_per_video': self.num_clips_per_video,
'clip_strategy': str(self.clip_strategy),
'frame_stride': self.frame_stride,
'temporal_stride': self.temporal_stride,
'batch_size': self.batch_size,
'device': self.device,
'seed': self.seed,
'use_streaming': self.use_streaming,
}
def __str__(self) -> str:
"""Human-readable protocol name"""
desc = f"FVD_{self.extractor_model.upper()}_{self.num_videos}_{self.num_frames_per_clip}f"
desc = f"FVD{self.num_videos}_{self.num_frames_per_clip}f"
if self.frame_stride > 1:
desc += f"_subsample{self.frame_stride}"
if self.num_clips_per_video > 1:
@@ -127,42 +150,57 @@ class FVDConfig:
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
extractor: BaseFeatureExtractor,
extractor: I3DFeatureExtractor,
batch_size: int = 32,
max_clips: int | None = None,
verbose: bool = True) -> np.ndarray:
"""
Extract features from a video clip generator using streaming.
Args:
video_generator: Iterator yielding clips [T, C, H, W]
extractor: I3D feature extractor
batch_size: Batch size for processing
max_clips: Maximum clips to process (for validation)
verbose: Show progress
Returns:
features: [N, 400] numpy array
"""
all_features = []
batch = []
clip_count = 0
if verbose:
print(f"Extracting features with batch_size={batch_size}...")
with torch.no_grad():
for clip_count, clip in enumerate(video_generator):
batch.append(clip)
for clip_count, clip in enumerate(video_generator):
batch.append(clip)
# Process batch when full
if len(batch) == batch_size:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features_batch(batch_tensor)
all_features.append(features.detach().cpu().numpy())
batch = []
if verbose and clip_count % (batch_size * 10) == 0:
print(f"Processed {clip_count} clips...")
if max_clips is not None and clip_count >= max_clips:
break
# Process remaining clips
if len(batch) > 0:
# Process batch when full
if len(batch) == batch_size:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features_batch(batch_tensor)
all_features.append(features.detach().cpu().numpy())
features = extractor.extract_features(batch_tensor,
batch_size=batch_size,
verbose=False)
all_features.append(features.cpu().numpy())
batch = [] # Clear batch
if verbose and clip_count % (batch_size * 10) == 0:
print(f"Processed {clip_count} clips...")
# Stop if we've reached max_clips
if max_clips is not None and clip_count >= max_clips:
break
# Process remaining clips
if len(batch) > 0:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features(batch_tensor,
batch_size=len(batch),
verbose=False)
all_features.append(features.cpu().numpy())
if len(all_features) == 0:
raise RuntimeError("No features extracted - check video loading")
@@ -176,17 +214,14 @@ def extract_features_streaming(video_generator: Iterator[torch.Tensor],
def load_or_compute_features(videos: str | Path | torch.Tensor,
extractor: BaseFeatureExtractor,
extractor: I3DFeatureExtractor,
config: FVDConfig,
cache_path: str | None = None,
cache_name: str = "real_features") -> np.ndarray:
"""Load features from cache or compute (with streaming support)"""
if cache_path is not None:
script_dir = Path(__file__).parent
cache_dir = script_dir / cache_path
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
cache_file = Path(cache_path) / f"{cache_name}.pkl"
if cache_file.exists():
print(f"Loading cached features from {cache_file}")
with open(cache_file, 'rb') as f:
@@ -194,25 +229,19 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
# Validate and limit based on config
max_features = config.num_videos * config.num_clips_per_video
if len(features) < max_features:
print(
f"WARNING: Cache has {len(features)} features but need {max_features}"
)
print("Cached features insufficient - will recompute...")
print("Recomputing features...")
elif len(features) > max_features:
print(
f"Using {max_features} features from cache (truncated from {len(features)})"
)
features = features[:max_features]
return features
else:
print(f"Using all {len(features)} cached features")
return features
print("Computing features from scratch...")
if isinstance(videos, (str | Path)):
# Compute features
if isinstance(videos, str | Path):
target_size = (224, 224) if config.resize_before_extraction else None
video_generator = load_video_clips_streaming(
@@ -233,7 +262,9 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
batch_size=config.batch_size,
max_clips=max_clips,
verbose=True)
else:
# Already a tensor
print(f"Extracting features from {len(videos)} video tensors...")
features = extractor.extract_features(videos,
batch_size=config.batch_size,
@@ -252,10 +283,9 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
# Cache features if requested
if cache_path is not None:
script_dir = Path(__file__).parent
cache_dir = script_dir / cache_path
cache_dir = Path(cache_path)
cache_dir.mkdir(parents=True, exist_ok=True)
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
cache_file = cache_dir / f"{cache_name}.pkl"
print(f"Caching features to {cache_file}")
with open(cache_file, 'wb') as f:
pickle.dump(features, f)
@@ -263,34 +293,79 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
return features
def compute_fvd(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
num_frames: int = 16,
batch_size: int = 32,
device: str = 'cuda',
num_videos: int | None = 2048,
cache_real_features: str | None = None,
i3d_model_path: str | None = None,
seed: int | None = None,
verbose: bool = True) -> float:
"""
Compute Fréchet Video Distance (FVD)
For advanced control, use compute_fvd_with_config() instead.
Args:
real_videos: Path to real videos or tensor [N, T, C, H, W]
gen_videos: Path to generated videos or tensor [N, T, C, H, W]
num_frames: Frames per video (default: 16)
batch_size: Batch size (default: 32)
device: 'cuda' or 'cpu' (default: 'cuda')
num_videos: Max videos (default: 2048)
cache_real_features: Cache path for real features
i3d_model_path: Custom I3D model cache path
seed: Random seed for reproducibility
verbose: Print progress
Returns:
FVD score (float). Lower is better.
"""
num_videos = num_videos if num_videos is not None else 2048
config = FVDConfig(
num_videos=num_videos,
num_frames_per_clip=num_frames,
batch_size=batch_size,
device=device,
cache_real_features=cache_real_features,
i3d_model_path=i3d_model_path,
seed=seed,
)
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
return result['fvd']
def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
config: FVDConfig,
verbose: bool = True) -> dict:
"""
Compute FVD using a standardized configuration.
This is the recommended way to compute FVD for reproducibility.
Args:
real_videos: Path or tensors
gen_videos: Path or tensors
config: FVDConfig specifying protocol
verbose: Print progress
Returns:
results: Dictionary with:
- 'fvd': FVD score (float)
- 'protocol': Protocol name (str)
- 'model': Feature extractor model name (str)
- 'config': Configuration dict
Example:
>>> config = FVDConfig.fvd2048_16f()
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
Compute FVD using a standardized configuration.
This is the recommended way to compute FVD for reproducibility.
Args:
real_videos: Path or tensors
gen_videos: Path or tensors
config: FVDConfig specifying protocol
verbose: Print progress
Returns:
results: Dictionary with:
- 'fvd': FVD score (float)
- 'protocol': Protocol name (str)
- 'config': Configuration dict
Example:
>>> config = FVDConfig.fvd2048_16f()
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
>>> print(f"Protocol: {results['protocol']}") # "FVD2048_16f"
"""
# Seed for reproducibility
if config.seed is not None:
import random as _rnd
@@ -303,20 +378,18 @@ def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
if verbose:
print("=" * 70)
print(f"Computing FVD with protocol: {config}")
print(f"Model: {config.extractor_model.upper()}")
print("=" * 70)
print("\nConfiguration:")
for key, value in config.to_dict().items():
print(f" {key}: {value}")
print()
# Initialize Extractor using Factory
# Initialize I3D
if verbose:
print(
f"\nInitializing {config.extractor_model.upper()} model on {config.device}..."
)
print(f"\nInitializing I3D model on {config.device}...")
extractor = load_extractor(config.extractor_model, device=config.device)
extractor = I3DFeatureExtractor(device=config.device,
cache_dir=config.i3d_model_path)
# Extract features
if verbose:
@@ -361,45 +434,14 @@ def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
if verbose:
print(f"\n{'='*70}")
print(f"FVD Score ({config.extractor_model.upper()}): {fvd:.4f}")
print(f"FVD Score: {fvd:.4f}")
print(f"Protocol: {config}")
print(f"{'='*70}\n")
results = {
'fvd': fvd,
'protocol': str(config),
'model': config.extractor_model,
'config': config.to_dict(),
}
return results
def compute_fvd(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
num_frames: int = 16,
batch_size: int = 32,
device: str = 'cuda',
num_videos: int | None = 2048,
cache_real_features: str | None = None,
i3d_model_path: str | None = None,
seed: int | None = None,
verbose: bool = True) -> float:
"""
Backward compatibility wrapper for computing FVD (defaults to I3D).
"""
num_videos = num_videos if num_videos is not None else 2048
config = FVDConfig(
num_videos=num_videos,
num_frames_per_clip=num_frames,
extractor_model='i3d', # Default to I3D
batch_size=batch_size,
device=device,
cache_real_features=cache_real_features,
i3d_model_path=i3d_model_path,
seed=seed,
)
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
return result['fvd']
+17 -37
View File
@@ -1,53 +1,33 @@
import sys
from pathlib import Path
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config
root_dir = Path(__file__).parent.parent.parent
sys.path.insert(0, str(root_dir))
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config # noqa: E402
def main() -> None:
# Get script directory
script_dir = Path(__file__).parent.resolve()
# Define directories
clip_strategy = 'beginning' # Options: 'uniform', 'random', 'beginning', 'end', 'all'
cfg = FVDConfig(
num_videos=650,
num_frames_per_clip=16,
num_clips_per_video=1,
clip_strategy=clip_strategy,
frame_stride=1,
batch_size=32,
device='cuda',
seed=42,
cache_real_features=str(script_dir / f'fvd-cache/{clip_strategy}'),
)
real_dir = "benchmarks/data/real_videos"
gen_dir = "benchmarks/data/generated_videos"
# Compare all 3 models
models_to_test = ['i3d', 'clip', 'videomae']
print(f"\n{'='*60}")
print("STARTING COMPARISON BENCHMARK")
print(f"{'='*60}")
for model_name in models_to_test:
print(f"\n>>> Running evaluation with {model_name.upper()}...")
try:
cfg = FVDConfig(
num_videos=650,
num_frames_per_clip=16,
extractor_model=model_name,
clip_strategy='beginning',
device='cuda',
seed=42,
# Use separate cache folders for each model to avoid conflicts
cache_real_features=str(script_dir / f'fvd-cache/{model_name}'),
)
results = compute_fvd_with_config(real_dir,
gen_dir,
cfg,
verbose=False)
print(f"FVD: {results['fvd']}\nModel: {results['model']}")
except Exception as e:
print(f"{model_name.upper()} Failed: {e}")
print(f"\n{'='*60}")
print("BENCHMARK COMPLETE")
print(f"{'='*60}")
results = compute_fvd_with_config(real_dir, gen_dir, cfg, verbose=True)
print(f"FVD = {results['fvd']:.2f}")
if __name__ == '__main__':
+1 -1
View File
@@ -1,7 +1,7 @@
#!/bin/bash
# 1. Install missing dependency
pip install -q opencv-python-headless transformers huggingface_hub
pip install -q opencv-python-headless
# 2. Run FVD script
python benchmarks/fvd/run_fvd.py
+113
View File
@@ -0,0 +1,113 @@
# Attention Kernel Used in FastVideo
## Video Sparse Attention (VSA)
### Installation
We support H100 (via TK) and any other GPU (via triton) for VSA.
```bash
git submodule update --init --recursive
python setup_vsa.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
```
## Sliding Tile Attention (STA)
We only support H100 for STA.
```bash
git submodule update --init --recursive
python setup_sta.py install
```
### 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.
+145
View File
@@ -0,0 +1,145 @@
import os
from collections import defaultdict
import matplotlib.pyplot as plt
import numpy as np
import torch
from st_attn import sliding_tile_attention
from triton.testing import do_bench
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
assert mode in ["fwd", "bwd", "fwd_bwd"]
f = 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
def compute_TFLOPS(flops, ms):
flops = flops / 1e12
ms = ms / 1e3
return flops / ms
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
# # Warmup for forward pass
# for _ in range(10):
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# # Time the forward pass
# for i in range(10):
# start_events_fwd[i].record()
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# end_events_fwd[i].record()
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
# time_us_fwd = np.mean(times_fwd) * 1000
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass (ms): {ms:.2f}")
print(f"Average TFLOPS: {tflops_fwd}")
print("-" * 60)
# torch.cuda.empty_cache()
# torch.cuda.synchronize()
# # Prepare for timing backward pass
# start_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# end_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# # Warmup for backward pass
# for _ in range(10):
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# # Time the backward pass
# for i in range(10):
# start_events_bwd[i].record()
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# end_events_bwd[i].record()
# torch.cuda.synchronize()
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
# time_us_bwd = np.mean(times_bwd) * 1000
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass(ms): {ms:.2f}")
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
return results
def plot_results(results):
os.makedirs('benchmark_results', exist_ok=True)
for mode in ['fwd', 'bwd']:
for (D, causal), values in results[mode].items():
seq_lens = [x[0] for x in values]
tflops = [x[1] for x in values]
plt.figure(figsize=(10, 6))
bars = plt.bar(range(len(seq_lens)), tflops, tick_label=seq_lens)
plt.xlabel('Sequence Length')
plt.ylabel('TFLOPS')
plt.title(f'{mode.upper()} Pass - Head Dim: {D}, Causal: {causal}')
plt.grid(True)
# Adding the numerical y value on top of each bar
for bar in bars:
yval = bar.get_height()
plt.text(bar.get_x() + bar.get_width() / 2, yval, round(yval, 2), ha='center', va='bottom')
filename = f'benchmark_results/{mode}_D{D}_causal{causal}.png'
plt.savefig(filename)
plt.close()
# Example list of configurations to test
configurations = [
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
# (16, 16, 768*8, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768, 128, True),
# (16, 16, 768*2, 128, True),
# (16, 16, 768*4, 128, True),
# (16, 16, 768*8, 128, True),
# (16, 16, 768*16, 128, True),
# (16, 32, 768, 64, False),
# (16, 32, 768*2, 64, False),
# (16, 32, 768*4, 64, False),
# (16, 32, 768*8, 64, False),
# (16, 32, 768*16, 64, False),
# (16, 32, 768, 64, True),
# (16, 32, 768*2, 64, True),
# (16, 32, 768*4, 64, True),
# (16, 32, 768*8, 64, True),
# (16, 32, 768*16, 64, True),
]
results = benchmark_attention(configurations)
# plot_results(results)
+224
View File
@@ -0,0 +1,224 @@
import torch
import argparse
from triton.testing import do_bench
from vsa import block_sparse_fwd, block_sparse_bwd
from vsa import BLOCK_M, BLOCK_N
import triton
import numpy as np
import random
def set_seed(seed: int = 42):
# Python random module
random.seed(seed)
# NumPy
np.random.seed(seed)
# PyTorch
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
return parser.parse_args()
def create_input_tensors(batch, head, seq_len, headdim):
"""Create random input tensors for attention."""
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
return q, k, v
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
Args:
bs: batch size
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key-value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to k).
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
Binary mask where 1 indicates attention connection.
"""
# Ensure k is not larger than num_kv_blocks
k = min(k, num_kv_blocks)
# Create random scores for sampling
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
# Get top-k indices for each q block
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
# sort q2k_block_sparse_index
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
# All q blocks attend to exactly k kv blocks
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
# Create the corresponding mask
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
# Fill in the mask based on the indices
for b in range(bs):
for head in range(h):
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx]
block_sparse_mask[b, head, q_idx, kv_indices] = True
# Create the reverse mapping (k2q)
# First, initialize lists to collect q indices for each kv block
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
# Populate the lists based on q2k mapping
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
for kv_idx in kv_indices:
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
# Find the maximum number of q blocks that attend to any kv block
max_q_per_kv = 0
for flat_idx in range(bs * h):
for kv_idx in range(num_kv_blocks):
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
# Create tensors for k2q mapping
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
dtype=torch.int32, device=device)
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
dtype=torch.int32, device=device)
# Fill the tensors
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for kv_idx in range(num_kv_blocks):
q_indices = k2q_indices_list[flat_idx][kv_idx]
num_q = len(q_indices)
k2q_block_sparse_num[b, head, kv_idx] = num_q
if num_q > 0:
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
q_indices, dtype=torch.int32, device=device)
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
"""Benchmark block sparse attention forward and backward passes."""
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
# Forward pass
# Warm-up run
variable_block_sizes = torch.ones(q2k_block_sparse_index.shape[2], device=q.device).int() * BLOCK_M
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
torch.cuda.synchronize()
# Benchmark forward
fwd_time = do_bench(
lambda: block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes),
warmup=5,
rep=20,
quantiles=None
)
sparse_tflops = flops / fwd_time * 1e-12 * 1e3
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
# Backward pass
grad_output = torch.randn_like(o)
# Warm-up runs
for _ in range(5):
block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
torch.cuda.synchronize()
# Benchmark backward
bwd_time = do_bench(
lambda: block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes),
warmup=5,
rep=20,
quantiles=None
)
bwd_flops = 2.5 * flops # Approximation
sparse_bwd_tflops = bwd_flops / bwd_time * 1e-12 * 1e3
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
return sparse_tflops, sparse_bwd_tflops
def main():
args = parse_arguments()
set_seed(42)
# Extract parameters
batch = args.batch_size
head = args.num_heads
headdim = args.head_dim
print(f"Block Sparse Attention Benchmark")
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
# Test with different sequence lengths
for seq_len in args.seq_lengths:
# Skip very long sequences if they might cause OOM
if seq_len > 16384 and batch > 1:
continue
print("="*100)
print(f"\nSequence length: {seq_len}")
# Calculate theoretical FLOPs for attention
flops = 4 * batch * head * headdim * seq_len * seq_len
# Create input tensors
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
# Setup block sparse parameters
num_q_blocks = seq_len // BLOCK_M
num_kv_blocks = seq_len // BLOCK_N
# Determine k value (number of kv blocks per q block)
topk = args.topk
if topk is None:
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
topk = max(1, topk)
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
# Generate block sparse pattern
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
# Benchmark block sparse attention
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
)
# Print results
print("\n=== PERFORMANCE RESULTS ===")
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
if __name__ == "__main__":
main()
+217
View File
@@ -0,0 +1,217 @@
import torch
import argparse
import triton.testing
from vsa import block_sparse_attn
from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
def set_seed(seed: int = 42):
# Python random module
random.seed(seed)
# NumPy
np.random.seed(seed)
# PyTorch
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
return parser.parse_args()
def create_input_tensors(batch, head, seq_len, headdim):
"""Create random input tensors for attention."""
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
return q, k, v
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
Args:
bs: batch size
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key-value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to k).
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
Binary mask where 1 indicates attention connection.
"""
# Ensure k is not larger than num_kv_blocks
k = min(k, num_kv_blocks)
# Create random scores for sampling
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
# Get top-k indices for each q block
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
# sort q2k_block_sparse_index
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
# All q blocks attend to exactly k kv blocks
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
# Create the corresponding mask
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
# Fill in the mask based on the indices
for b in range(bs):
for head in range(h):
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx]
block_sparse_mask[b, head, q_idx, kv_indices] = True
# Create the reverse mapping (k2q)
# First, initialize lists to collect q indices for each kv block
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
# Populate the lists based on q2k mapping
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
for kv_idx in kv_indices:
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
# Find the maximum number of q blocks that attend to any kv block
max_q_per_kv = 0
for flat_idx in range(bs * h):
for kv_idx in range(num_kv_blocks):
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
# Create tensors for k2q mapping
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
dtype=torch.int32, device=device)
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
dtype=torch.int32, device=device)
# Fill the tensors
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for kv_idx in range(num_kv_blocks):
q_indices = k2q_indices_list[flat_idx][kv_idx]
num_q = len(q_indices)
k2q_block_sparse_num[b, head, kv_idx] = num_q
if num_q > 0:
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
q_indices, dtype=torch.int32, device=device)
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
"""Benchmark block sparse attention forward+backward pass."""
print("\n=== BLOCK SPARSE ATTENTION FORWARD+BACKWARD BENCHMARK ===")
# Combined forward+backward pass
# Warm-up run
q_fwd = q.clone().requires_grad_(True)
k_fwd = k.clone().requires_grad_(True)
v_fwd = v.clone().requires_grad_(True)
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
grad_output = torch.randn_like(o)
o.backward(grad_output)
torch.cuda.synchronize()
# Benchmark forward+backward
def forward_backward_fn():
q_fwd = q.clone().requires_grad_(True)
k_fwd = k.clone().requires_grad_(True)
v_fwd = v.clone().requires_grad_(True)
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
grad_output = torch.randn_like(o)
o.backward(grad_output)
total_time = triton.testing.do_bench(
forward_backward_fn,
warmup=25,
rep=100,
return_mode='mean'
)
# Total flops for forward + backward (forward + 2.5x backward approximation)
total_flops = flops + 2.5 * flops # 3.5x the forward flops
sparse_tflops = total_flops / total_time * 1e-12 * 1e3
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_tflops:.2f}")
return sparse_tflops
def main():
args = parse_arguments()
set_seed(42)
# Extract parameters
batch = args.batch_size
head = args.num_heads
headdim = args.head_dim
print(f"Block Sparse Attention Benchmark")
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
# Test with different sequence lengths
for seq_len in args.seq_lengths:
# Skip very long sequences if they might cause OOM
if seq_len > 16384 and batch > 1:
continue
print("="*100)
print(f"\nSequence length: {seq_len}")
# Calculate theoretical FLOPs for attention
flops = 4 * batch * head * headdim * seq_len * seq_len
# Create input tensors
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
# Setup block sparse parameters
num_q_blocks = seq_len // BLOCK_M
num_kv_blocks = seq_len // BLOCK_N
# Determine k value (number of kv blocks per q block)
topk = args.topk
if topk is None:
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
topk = max(1, topk)
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
# Generate block sparse pattern
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
# Benchmark block sparse attention
sparse_fwd = benchmark_block_sparse_attention(
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
)
# Print results
print("\n=== PERFORMANCE RESULTS ===")
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_fwd:.2f}")
if __name__ == "__main__":
main()
+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config_sta.py
+103
View File
@@ -0,0 +1,103 @@
# Attention Kernel Used in FastVideo
## Sliding Tile Attention (STA)
We support H100 (via TK) and any other GPU (via triton) for STA.
### Installation
```bash
pip install st_attn
```
Install from source:
```bash
git submodule update --init --recursive
python setup.py install
```
If you want to skip the compilation of the TK kernel and only use the Triton version, try below:
```bash
SKIP_SM90_EXT=1 python setup.py install
or
SKIP_SM90_EXT=1 pip install --no-build-isolation .
```
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
```
### 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
## STA Configuration Logic
Here is a diagram of how the window is configured and passed through the FastVideo pipeline:
<div align="center">
<img src="../../../docs/assets/images/STA_configuration.png" width="80%"/>
</div>
## 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.
+15
View File
@@ -0,0 +1,15 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'st_attn': {
'source_files': {
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['st_attn']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
+83
View File
@@ -0,0 +1,83 @@
import os
import subprocess
from config_sta import kernels, sources, target
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
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"
# Set environment variables
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
python_include = subprocess.check_output(['python', '-c',
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
torch_include = subprocess.check_output([
'python', '-c',
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip()
print('st_attn root:', tk_root)
print('Python include:', python_include)
print('Torch include directories:', torch_include)
# CUDA flags
cuda_flags = [
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
] + torch_include.split()
cpp_flags = ['-std=c++20', '-O3']
if target == 'h100':
cuda_flags.append('-DKITTENS_HOPPER')
cuda_flags.append('-arch=sm_90a')
else:
raise ValueError(f'Target {target} not supported')
source_files = ['st_attn.cpp']
for k in kernels:
if target not in sources[k]['source_files']:
raise KeyError(f'Target {target} not found in source files for kernel {k}')
if isinstance(sources[k]['source_files'][target], list):
source_files.extend(sources[k]['source_files'][target])
else:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
ext_modules = []
if os.environ.get("SKIP_SM90_EXT", "0") != "1":
ext_modules.append(
CUDAExtension('st_attn_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
)
else:
print("ENV SKIP_SM90_EXT=1, skip st_attn_cuda compile")
setup(name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
ext_modules=ext_modules,
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.10',
install_requires=["torch>=2.5.0"])
+23
View File
@@ -0,0 +1,23 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ST_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
@@ -0,0 +1,63 @@
import math
import torch
from torch.utils.checkpoint import detach_variable
try:
from st_attn_cuda import sta_fwd
except ImportError:
sta_fwd = None
try:
from st_attn.st_attn_triton import sliding_tile_attention_triton
except ImportError:
sliding_tile_attention_triton = None
def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
seq_length = q_all.shape[2]
dit_seq_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
2] >= 115200 and q_all.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '30x48x80' for HunyuanVideo"
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
target_size = math.ceil(seq_length / 384) * 384
pad_size = target_size - seq_length
if pad_size > 0:
q_all = torch.cat([q_all, q_all[:, :, -pad_size:]], dim=2)
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
else:
if dit_seq_shape == '36x48x48': # Stepvideo 204x768x68
assert q_all.shape[2] == 82944
elif dit_seq_shape == '18x48x80': # Wan 69x768x1280
assert q_all.shape[2] == 69120
else:
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
kernel_aspect_ratio_flag = dit_seq_shape_mapping[dit_seq_shape]
hidden_states = torch.empty_like(q_all)
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
for batch in range(q_all.shape[0]):
q_head, k_head, v_head, o_head = (q_all[batch:batch + 1, head_index:head_index + 1],
k_all[batch:batch + 1,
head_index:head_index + 1], v_all[batch:batch + 1,
head_index:head_index + 1],
hidden_states[batch:batch + 1, head_index:head_index + 1])
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
return hidden_states[:, :, :seq_length]
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
major, minor = torch.cuda.get_device_capability(q_all.device)
if major == 9 and minor == 0 and sta_fwd is not None:
return sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
elif sliding_tile_attention_triton is not None:
return sliding_tile_attention_triton(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
else:
raise ImportError("No suitable sliding tile attention implementation found.")
@@ -0,0 +1,841 @@
// # Define TORCH_COMPILE macro
#include "kittens.cuh"
#include <cooperative_groups.h>
#include <iostream>
#include <stdio.h>
#include <c10/cuda/CUDAGuard.h>
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
return (value < min) ? min : ((value > max) ? max : value);
}
// #define ABS(x) ((x) < 0 ? -(x) : (x))
__device__ __forceinline__ int abs_int(int value) {
return (value < 0) ? -value : value;
}
constexpr int CONSUMER_WARPGROUPS = (3);
constexpr int PRODUCER_WARPGROUPS = (1);
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
using namespace kittens;
namespace cg = cooperative_groups;
template<int D> struct fwd_attend_ker_tile_dims {};
template<> struct fwd_attend_ker_tile_dims<64> {
constexpr static int tile_width = (64);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (4);
};
template<> struct fwd_attend_ker_tile_dims<128> {
constexpr static int tile_width = (128);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (2);
};
template<int D> struct fwd_globals {
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
q_gl q;
k_gl k;
v_gl v;
l_gl l;
o_gl o;
const int N;
const int text_L;
const int hr;
};
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
extern __shared__ int __shm[];
tma_swizzle_allocator al((int*)&__shm[0]);
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
using K = fwd_attend_ker_tile_dims<D>;
using q_tile = st_bf<K::qo_height, K::tile_width>;
using k_tile = st_bf<K::kv_height, K::tile_width>;
using v_tile = st_bf<K::kv_height, K::tile_width>;
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
using o_tile = st_bf<K::qo_height, K::tile_width>;
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
int img_kv_blocks;
int kv_blocks = g.N / (K::kv_height);
if constexpr (text_kv) {
img_kv_blocks = kv_blocks - 3;
} else {
img_kv_blocks = kv_blocks;
}
int kv_head_idx = blockIdx.y / g.hr;
int seq_idx;
if constexpr (text_q) {
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
} else {
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
}
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
if (threadIdx.x == 0) {
init_semaphore(qsmem_semaphore, 0, 1);
for(int j = 0; j < K::stages; j++) {
init_semaphore(k_smem_arrived[j], 0, 1);
init_semaphore(v_smem_arrived[j], 0, 1);
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
}
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
}
if constexpr (text_q){
for (int j = 0; j < K::stages - 1; j++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int count = 0;
int j = 0;
while (count < K::stages - 1) {
int kt = j / 3 / (CH * CW);
int kh = (j / 3) % (CH * CW) / CW;
int kw = (j / 3) % CW;
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
if (mask){
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
count += 1;
}
j += 1;
}
}
}
__syncthreads();
int pipe_idx = K::stages - 1;
if(warpgroupid == NUM_WARPGROUPS-1) {
warpgroup::decrease_registers<32>();
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
}
else { kv_iters = kv_blocks-2;}
if(warpid == NUM_WORKERS-4) {
if constexpr (text_q){
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int k_t_min = clamp_int(qt-DT, 0, CT-1);
int k_t_max = clamp_int(qt+DT, 0, CT-1);
int k_h_min = clamp_int(qh-DH, 0, CH-1);
int k_h_max = clamp_int(qh+DH, 0, CH-1);
int k_w_min = clamp_int(qw-DW, 0, CW-1);
int k_w_max = clamp_int(qw+DW, 0, CW-1);
int count = 0;
for (int kt = k_t_min; kt <= k_t_max; kt++) {
for (int kh = k_h_min; kh <= k_h_max; kh++) {
for (int kw = k_w_min; kw <= k_w_max; kw++) {
for (int j = 0; j <= 2; j++){
if (count >= K::stages - 1) {
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
} else {
count += 1;
}
}
}
}
}
// for text
for (int index = img_kv_blocks; index < kv_blocks; index++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
}
}
}
}
else {
warpgroup::increase_registers<160>();
rt_fl<16, K::kv_height> att_block;
rt_bf<16, K::kv_height> att_block_mma;
rt_fl<16, K::tile_width> o_reg;
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
neg_infty(max_vec);
zero(norm_vec);
zero(o_reg);
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
kv_iters = (kv_iters/8);
}
else if constexpr (text_q){
// the last three kv blocks are for text, we process them separately
kv_iters = img_kv_blocks - 1;
} else {
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
// the last three kv blocks are for text, we process them separately
if constexpr(text_kv) {
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
// apply non-pad mask
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
}
div_row(o_reg, o_reg, norm_vec);
warpgroup::store(o_smem[warpgroupid], o_reg);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
}
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
log(norm_vec, norm_vec);
add(norm_vec, norm_vec, max_vec_scaled);
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
else { mul(norm_vec, norm_vec, -11.313708499f); }
warpgroup::store(l_smem[warpgroupid], norm_vec);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
}
tma::store_async_wait();
}
}
#include "pyutils/torch_helpers.cuh"
#include <ATen/cuda/CUDAContext.h>
#include <iostream>
torch::Tensor
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
CHECK_INPUT(v);
auto batch = q.size(0);
auto seq_len = q.size(2);
auto head_dim = q.size(3);
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
auto hr = qo_heads / kv_heads;
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<128>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
dim3 grid_text(2, qo_heads, batch);
if (!process_text) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 1, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true,1, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==5 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==5 && kernel_h_size == 3 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 0, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true,2, 0, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
} else {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
}
} else {
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (kernel_aspect_ratio_flag == 2){
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else if (kernel_aspect_ratio_flag == 3) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else {
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
}
}
CHECK_CUDA_ERROR(cudaGetLastError());
// cudaStreamSynchronize(stream);
}
return o;
//cudadevicesynchronize();
}
@@ -11,20 +11,12 @@ def is_hip():
target = triton.runtime.driver.active.get_current_target()
return target.backend == 'hip'
def is_cdna3_cdna4():
target = triton.runtime.driver.active.get_current_target()
return (target.arch == 'gfx1201' or target.arch == 'gfx1101' or target.arch == 'gfx1100' or target.arch == 'gfx1030')
def get_common_autotune_config():
# cdna arch does not support a 4-stage software pipeline, see https://github.com/ROCm/triton/issues/916
supported_num_staged = [1, 2] if is_cdna3_cdna4() else [1, 2, 3, 4]
configs = [
triton.Config({'BLOCK_Q': BLOCK_Q, 'BLOCK_KV': BLOCK_KV}, num_stages=s, num_warps=w) \
for BLOCK_Q in [32, 64, 128]\
for BLOCK_KV in [32, 64, 128]\
for s in supported_num_staged\
for s in [1, 2, 3, 4]\
for w in [4, 8]\
]
return configs
@@ -1,6 +1,6 @@
import torch
from .support_flex_sta import get_sliding_tile_attention_mask
from fastvideo_kernel import sliding_tile_attention
from flex_sta_ref import get_sliding_tile_attention_mask
from st_attn import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
@@ -73,23 +73,15 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
# Example usage
def test_sliding_tile_attention():
if not torch.cuda.is_available():
return
b, h, d = 2, 24, 128
n = 69120 # Sequence length
causal = False
mean = 1e-1
std = 10
# Run correctness check directly
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
if __name__ == "__main__":
test_sliding_tile_attention()
b, h, d = 2, 24, 128
n = 69120 # Sequence length
causal = False
mean = 1e-1
std = 10
# Run correctness check directly
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
@@ -1,13 +1,13 @@
import torch
import pytest
import sys
import os
import numpy as np
from tqdm import tqdm
from .utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
# Use installed package
# from fastvideo_kernel import video_sparse_attn as block_sparse_attn
# Add the parent directory to the path to import block_sparse_attn
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from tests.utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
from vsa import block_sparse_attn
BLOCK_M = 64
BLOCK_N = 64
@@ -42,36 +42,7 @@ def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, q
q_padded = vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
k_padded = vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
# Use raw kernel or triton
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
except ImportError:
raw_kernel = None
from fastvideo_kernel.triton_kernels.index import map_to_index
# Convert mask to indices
# block_sparse_mask is [H, M, N] bool
# We need to map it to index.
# block_sparse_mask needs to be expanded/reshaped?
# generate_block_sparse_mask_for_function returns [H, NumBlocksQ, NumBlocksKV]
# Ops.py logic:
# mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
# idx, num = map_to_index(mask)
idx, num = map_to_index(block_sparse_mask.unsqueeze(0)) # Add batch dim [1, H, M, N]
if raw_kernel:
out_s = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())
output = out_s[0]
else:
# Fallback to triton testing if C++ not available
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
output, _ = triton_block_sparse_attn_forward(q_padded, k_padded, v_padded, idx, num, variable_block_sizes)
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
output = output[:, :, q_non_pad_index, :]
output.backward(dO)
return output, Q.grad, K.grad, V.grad
@@ -264,19 +235,11 @@ def generate_error_graphs_qkdiff(h, d, error_mode='all'):
print("-" * 150)
@pytest.mark.skip()
def test_video_sparse_attention_backward():
if not torch.cuda.is_available():
return
if __name__ == "__main__":
h, d = 16, 128
print("Block Sparse Attention with Variable Block Sizes Analysis")
print("=" * 60)
for mode in ['backward']:
generate_error_graphs(h, d, error_mode=mode)
generate_error_graphs_qkdiff(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
if __name__ == "__main__":
test_video_sparse_attention_backward()
print("\nAnalysis completed for all modes.")
@@ -4,12 +4,21 @@ from typing import Tuple
import torch
from .utils import (
# Make sure we can import from the project root (`vsa`, `tests.utils`, etc.)
CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = os.path.dirname(CURRENT_DIR)
if PROJECT_ROOT not in sys.path:
sys.path.append(PROJECT_ROOT)
if CURRENT_DIR not in sys.path:
sys.path.append(CURRENT_DIR)
from tests.utils import (
generate_block_sparse_mask_for_function,
create_full_mask_from_block_mask,
)
from .test_vsa import BLOCK_M # Import from local test_vsa
from . import test_vsa as ref
from vsa import block_sparse_attn, BLOCK_M
import test_vsa as ref # reuse helper functions from backward test
def pytorch_forward(
Q: torch.Tensor,
@@ -47,7 +56,8 @@ def block_sparse_forward_test(
kv_num_blocks: int,
) -> torch.Tensor:
"""
Forward-only wrapper
Forward-only wrapper around `block_sparse_attn`, mirroring `block_sparse_kernel_test`
but without any backward / grad logic.
"""
Q = Q.detach()
K = K.detach()
@@ -57,24 +67,9 @@ def block_sparse_forward_test(
k_padded = ref.vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = ref.vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
# Use raw kernel or triton
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
raw_kernel = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
except ImportError:
raw_kernel = None
from fastvideo_kernel.triton_kernels.index import map_to_index
idx, num = map_to_index(block_sparse_mask)
if raw_kernel:
out_padded = raw_kernel(q_padded, k_padded, v_padded, idx, num, variable_block_sizes.int())[0]
else:
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
out_padded, _ = triton_block_sparse_attn_forward(
q_padded, k_padded, v_padded, idx, num, variable_block_sizes
)
out_padded, _ = block_sparse_attn(
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
)
# Remove padding on the query side
out = out_padded[:, :, q_non_pad_index, :]
return out
@@ -223,10 +218,7 @@ def run_forward_qk_diff(
return avg_abs_err, max_rel_diff
def test_video_sparse_attention_forward():
if not torch.cuda.is_available():
return
if __name__ == "__main__":
h, d = 16, 128
print("Forward Block Sparse Attention Check (QK Equal)")
print("=" * 80)
@@ -242,6 +234,3 @@ def test_video_sparse_attention_forward():
f"QK diff: avg |ΔO| = {avg_err_diff:.6e}, max rel ΔO = {max_rel_diff:.6e}"
)
if __name__ == "__main__":
test_video_sparse_attention_forward()
+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config_vsa.py
+61
View File
@@ -0,0 +1,61 @@
# Attention Kernel Used in FastVideo
## Video Sparse Attention (VSA)
### Installation
We support H100 (via TK) and any other GPU (via triton) for VSA.
```bash
pip install vsa
```
Install from source:
```bash
git submodule update --init --recursive
python setup.py install
```
If you encounter error during installation, try below:
Install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
(If you use CUDA12.8)
```bash
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
### Verify if you have successfully installed
```bash
# test numerical
python ../tests/test_vsa.py
# (For H100) test speed
python ../benchmarks/bench_vsa_hopper.py
```
bench_vsa_hopper.py should print something like this:
```bash
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
=== BLOCK SPARSE ATTENTION BENCHMARK ===
Block Sparse Forward - TFLOPS: 5622.26
Block Sparse Backward - TFLOPS: 3865.68
```
## Acknowledgement
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
+15
View File
@@ -0,0 +1,15 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'block_sparse': {
'source_files': {
'h100': 'vsa/block_sparse_h100.cu'
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['block_sparse']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
+81
View File
@@ -0,0 +1,81 @@
import os
import subprocess
from config_vsa import kernels, sources, target
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
target = target.lower()
# Package metadata
PACKAGE_NAME = "vsa"
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/video_sparse_attn"
# Set environment variables
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
python_include = subprocess.check_output(['python', '-c',
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
torch_include = subprocess.check_output([
'python', '-c',
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip()
print('vsa root:', tk_root)
print('Python include:', python_include)
print('Torch include directories:', torch_include)
# CUDA flags
cuda_flags = [
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
] + torch_include.split()
cpp_flags = ['-std=c++20', '-O3']
if target == 'h100':
cuda_flags.append('-DKITTENS_HOPPER')
cuda_flags.append('-arch=sm_90a')
else:
raise ValueError(f'Target {target} not supported')
source_files = ['vsa.cpp']
for k in kernels:
if target not in sources[k]['source_files']:
raise KeyError(f'Target {target} not found in source files for kernel {k}')
if isinstance(sources[k]['source_files'][target], list):
source_files.extend(sources[k]['source_files'][target])
else:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
ext_modules = [
CUDAExtension('vsa_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
]
setup(name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
ext_modules=ext_modules,
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.10',
install_requires=["torch>=2.5.0"])
+27
View File
@@ -0,0 +1,27 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_BLOCK_SPARSE
extern std::vector<torch::Tensor> block_sparse_attention_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
);
extern std::vector<torch::Tensor> block_sparse_attention_backward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_BLOCK_SPARSE
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
#endif
}
@@ -0,0 +1,80 @@
import torch
from typing import Tuple
block_sparse_attn=None
import torch
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
block_sparse_attn = block_sparse_attn_SM90
else:
from vsa.block_sparse_wrapper import block_sparse_attn_triton
block_sparse_fwd = None
block_sparse_bwd = None
block_sparse_attn = block_sparse_attn_triton
BLOCK_M = 64
BLOCK_N = 64
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight=None):
"""
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
NOTE: We assume q, k, v is zero padded!!
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements == 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = (q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
k_compress = (k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
v_compress = (v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
topK_indices = torch.topk(block_attn_score, topk, dim=-1).indices
block_mask = torch.zeros_like(block_attn_score, dtype=torch.bool).scatter_(-1, topK_indices, True)
output_select, _ = block_sparse_attn(q, k, v, block_mask, variable_block_sizes)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
@@ -0,0 +1,450 @@
"""
Fused Attention
===============
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
(https://tridao.me/publications/flash2/flash2.pdf)
Credits: OpenAI kernel team
"""
import pytest
import torch
import triton
import triton.language as tl
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
import math # small utility needed by the sparse wrapper
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
# the code below and commenting out the equivalent parameters is convenient for
# re-tuning.
configs = [
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BM in [64]\
for BN in [64]\
for s in [3, 4, 7]\
for w in [4, 8]\
]
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(Q, K, V, sm_scale, #
q2k_index, q2k_num, max_kv_blks, #
variable_block_sizes,
M, Out, #
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vk, stride_vn,
stride_oz, stride_oh, stride_om, stride_on,
Z, H, N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
STAGE: tl.constexpr):
"""
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
(32×64 and 64×32) – memory footprint unchanged.
"""
# ----- program-id mapping -----
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
qvk_off = (b.to(tl.int64) * stride_qz +
h.to(tl.int64) * stride_qh)
Q_ptr = tl.make_block_ptr(
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
K_base = tl.make_block_ptr(
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
V_base = tl.make_block_ptr(
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
O_ptr = tl.make_block_ptr(
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
# ----- accumulators -----
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
qk_scale = sm_scale * 1.44269504 # 1/ln2
q = tl.load(Q_ptr)
# ----- sparse loop over valid K/V tiles -----
for i in range(0, kv_blocks):
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
block_size = tl.load(variable_block_sizes + kv_idx)
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
k = tl.load(K_ptr)
qk = tl.dot(q, k)
# mask out invalid columns
mask = tl.arange(0, BLOCK_N) < block_size
qk = tl.where(mask[None, :], qk, -float("inf"))
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
l_ij = tl.sum(p, 1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
v = tl.load(V_ptr)
acc = tl.dot(p.to(tl.bfloat16), v, acc)
m_i = m_ij
# ----- epilogue -----
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
@triton.jit
def _attn_bwd_preprocess(O, DO, #
Delta, #
Z, H, N_CTX, #
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
):
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
off_hz = tl.program_id(1)
off_n = tl.arange(0, HEAD_DIM)
# load
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
delta = tl.sum(o * do, axis=1)
# write-back
tl.store(Delta + off_hz * N_CTX + off_m, delta)
# The main inner-loop logic for computing dK and dV.
@triton.jit
def _attn_bwd_dkdv(dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_tok, stride_d, #
H, N_CTX, BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
# Filled in by the wrapper.
start_n, start_m, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M1)
offs_n = start_n + tl.arange(0, BLOCK_N1)
offs_k = tl.arange(0, HEAD_DIM)
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
step_m = BLOCK_M1
kv_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_N1
meta_base = ((b * H + h) * q_tiles + kv_blk)
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
block_size = tl.load(variable_block_sizes + kv_blk)
for blk_idx in range(q_blocks*2):
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
qkT = tl.dot(k, qT)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
# Compute dV.
ppT = pT
ppT = ppT.to(tl.bfloat16)
dv += tl.dot(ppT, do)
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# Compute dP and dS.
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
dsT = pT * (dpT - Di[None, :])
dsT = dsT.to(tl.bfloat16)
dk += tl.dot(dsT, tl.trans(qT))
# Increment pointers.
return dk, dv
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq(dq, q, K, V, #
do, m, D,
# shared by Q/K/V/DO.
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
# Filled in by the wrapper.
start_m, start_n, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M2)
offs_n = start_n + tl.arange(0, BLOCK_N2)
offs_k = tl.arange(0, HEAD_DIM)
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
step_n = BLOCK_N2
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M2
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
for blk_idx in range(kv_blocks*2):
kv_idx = tl.load(kv_ptr + blk_idx//2).to(tl.int32)
block_size = tl.load(variable_block_sizes + kv_idx) - (blk_idx % 2) * step_n
block_sparse_offset = (kv_idx*2 + blk_idx%2) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
p = tl.where(mask[None, :], p , 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@triton.jit
def _attn_bwd(Q, K, V, sm_scale, #
DO, #
DQ, DK, DV, #
M, D,
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_z, stride_h, stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
off_chz = (bhid * N_CTX).to(tl.int64)
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
pid = tl.program_id(0)
# offset pointers for batch/head
Q += adj
K += adj
V += adj
DO += adj
DQ += adj
DK += adj
DV += adj
M += off_chz
D += off_chz
# load scales
offs_k = tl.arange(0, HEAD_DIM)
start_n = pid * BLOCK_N1
start_m = 0
offs_n = start_n + tl.arange(0, BLOCK_N1)
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
num_steps = N_CTX // BLOCK_M1
dk, dv = _attn_bwd_dkdv( #
dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
start_n, start_m, num_steps #
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv)
# Write back dK.
dk *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk)
# THIS BLOCK DOES DQ:
start_m = pid * BLOCK_M2
end_n = 0
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)
m = m[:, None]
num_steps = N_CTX // BLOCK_N2
dq = _attn_bwd_dq(dq, q, K, V, #
do, m, D, #
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
start_m, end_n, num_steps #
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= LN2
tl.store(dq_ptrs, dq)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
o = torch.empty_like(q)
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
_attn_fwd_sparse[grid](
q, k, v, sm_scale,
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
M, o,
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
B, H, T,
HEAD_DIM=D, STAGE=3
)
return o, M
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
assert do.is_contiguous()
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
BATCH, N_HEAD, N_CTX = q.shape[:3]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert N_CTX % PRE_BLOCK == 0
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o, do, #
delta, #
BATCH, N_HEAD, N_CTX, #
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
)
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
q, arg_k, v, sm_scale, do, dq, dk, dv, #
M, delta, #
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
N_HEAD, N_CTX, #
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
HEAD_DIM=D #
)
return dq, dk, dv
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,185 @@
import torch
try:
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
except ImportError:
block_sparse_fwd = None
block_sparse_bwd = None
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
from vsa.index import map_to_index
from typing import Tuple, Optional
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
block_map = block_map.int()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
return o, M
@torch.library.register_fake("vsa::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
return o, M
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
def block_sparse_attn_backward_triton(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
return dq, dk, dv
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
def _block_sparse_attn_backward_triton_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
dq = torch.empty_like(grad_output_padded)
dk = torch.empty_like(grad_output_padded)
dv = torch.empty_like(grad_output_padded)
return dq, dk, dv
def backward_triton(ctx, grad_output1, grad_output2):
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def setup_context_triton(ctx, inputs, output):
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
o_padded, M = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_SM90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor]:
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
variable_block_sizes = variable_block_sizes.int()
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
return o_padded, lse_padded
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
def _block_sparse_attn_SM90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
B, H, S, D = q_padded.shape
o_padded = torch.empty_like(q_padded)
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
return o_padded, lse_padded
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_backward_SM90(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
)
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
return grad_q_padded, grad_k_padded, grad_v_padded
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
def _block_sparse_attn_backward_SM90_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch._check(grad_output_padded.dtype == torch.bfloat16)
torch._check(lse_padded.dtype == torch.float32)
grad_output_padded = grad_output_padded.contiguous()
dq = torch.empty_like(grad_output_padded)
dk = torch.empty_like(grad_output_padded)
dv = torch.empty_like(grad_output_padded)
return dq, dk, dv
def backward_SM90(ctx, grad_output1, grad_output2):
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def setup_context_SM90(ctx, inputs, output):
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
o_padded, lse_padded = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
@@ -1,9 +1,9 @@
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
@@ -26,7 +26,6 @@ def topk_index_to_map_kernel(
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
@@ -60,7 +59,6 @@ def map_to_index_kernel(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
@@ -108,7 +106,6 @@ def topk_index_to_map(index: torch.Tensor,
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
+32
View File
@@ -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
```
+26
View File
@@ -0,0 +1,26 @@
# 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=[
"flash-attn >= 2.7.1",
]
)
@@ -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"
+2
View File
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
+7
View File
@@ -0,0 +1,7 @@
build/
dist/
*.egg-info/
__pycache__/
*.so
*.pyc
.ipynb_checkpoints/
+6
View File
@@ -0,0 +1,6 @@
include LICENSE
include README.md
include pyproject.toml
recursive-include src/fastvideo_kernel *.cu *.cuh *.cpp *.h
recursive-include csrc *.cu *.cuh *.cpp *.h
recursive-include tk *.cu *.cuh *.cpp *.h
+31
View File
@@ -0,0 +1,31 @@
# FastVideo Kernel
CUDA kernels for FastVideo video generation.
## Installation
```bash
git submodule update --init --recursive
cd csrc/fastvideo_kernel
pip install .
```
## Usage
```python
from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_attn_varlen
# Example: Sliding Tile Attention
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
# Example: Video Sparse Attention (with Triton fallback)
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
# Example: VMoBA
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
```
## Requirements
- H100 GPU (sm_90a) for CUDA kernels
- Triton for non-H100 fallback
+23
View File
@@ -0,0 +1,23 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ST_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
+27
View File
@@ -0,0 +1,27 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_BLOCK_SPARSE
extern std::vector<torch::Tensor> block_sparse_attention_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
);
extern std::vector<torch::Tensor> block_sparse_attention_backward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_BLOCK_SPARSE
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
#endif
}
+27
View File
@@ -0,0 +1,27 @@
[build-system]
requires = ["setuptools>=61.0", "torch>=2.5.0", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "fastvideo-kernel"
version = "0.1.0"
description = "CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
license = {text = "Apache-2.0"}
authors = [{name = "Hao AI Lab"}]
classifiers = [
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
]
dependencies = [
"torch>=2.5.0",
"triton>=2.0.0"
]
[project.urls]
Repository = "https://github.com/hao-ai-lab/FastVideo"
[tool.setuptools.packages.find]
where = ["src"]
+132
View File
@@ -0,0 +1,132 @@
import os
import subprocess
import sys
from pathlib import Path
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
ROOT = Path(__file__).parent.absolute()
CSRC_DIR = ROOT / "csrc"
# Path to ThunderKittens (TK)
def get_tk_dir():
tk_env = os.getenv("THUNDERKITTENS_ROOT")
if tk_env:
return tk_env
# Check common locations
possible_paths = [
ROOT / "tk",
ROOT / "csrc" / "tk",
ROOT.parent / "attn" / "sliding_tile_attn" / "tk",
ROOT.parent / "attn" / "video_sparse_attn" / "tk",
]
for p in possible_paths:
if (p / "include" / "kittens.cuh").exists():
return str(p)
# Default fallback
return str(ROOT.parent / "attn" / "sliding_tile_attn" / "tk")
TK_DIR = get_tk_dir()
def get_cuda_flags(tk_root: str) -> list:
python_include = subprocess.check_output(
["python", "-c", "import sysconfig; print(sysconfig.get_path('include'))"]
).decode().strip()
torch_includes = subprocess.check_output([
"python", "-c",
"import torch; from torch.utils.cpp_extension import include_paths; "
"print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip().split()
return [
"-DNDEBUG",
"-Xcompiler=-Wno-psabi",
"-Xcompiler=-fno-strict-aliasing",
"--expt-extended-lambda",
"--expt-relaxed-constexpr",
"-forward-unknown-to-host-compiler",
"--use_fast_math",
"-std=c++20",
"-O3",
"-Xnvlink=--verbose",
"-Xptxas=--verbose",
"-Xptxas=--warn-on-spills",
f"-I{tk_root}/include",
f"-I{tk_root}/prototype",
f"-I{python_include}",
"-DTORCH_COMPILE",
"-DKITTENS_HOPPER",
"-arch=sm_90a",
] + torch_includes
def get_extensions():
if not torch.cuda.is_available():
return []
extensions = []
cpp_flags = ["-std=c++20", "-O3"]
# Check if TK is available
if not os.path.exists(os.path.join(TK_DIR, "include", "kittens.cuh")):
print(f"Warning: ThunderKittens not found at {TK_DIR}. CUDA kernels will not be built.")
return []
cuda_flags = get_cuda_flags(TK_DIR)
# STA Extension
extensions.append(CUDAExtension(
"fastvideo_kernel._C.st_attn",
sources=[
"csrc/st_attn.cpp",
"csrc/st_attn_h100.cu",
],
extra_compile_args={
"cxx": cpp_flags + ["-DTK_COMPILE_ST_ATTN"],
"nvcc": cuda_flags + ["-DTK_COMPILE_ST_ATTN"]
},
libraries=["cuda"],
))
# VSA Extension
extensions.append(CUDAExtension(
"fastvideo_kernel._C.vsa",
sources=[
"csrc/vsa.cpp",
"csrc/block_sparse_h100.cu",
],
extra_compile_args={
"cxx": cpp_flags + ["-DTK_COMPILE_BLOCK_SPARSE"],
"nvcc": cuda_flags + ["-DTK_COMPILE_BLOCK_SPARSE"]
},
libraries=["cuda"],
))
return extensions
ext_modules = []
if not any(arg in sys.argv for arg in ["clean", "egg_info", "--version"]):
try:
import torch
ext_modules = get_extensions()
except Exception as e:
print(f"Warning: Failed to configure CUDA extensions: {e}")
setup(
name="fastvideo-kernel",
version="0.1.0",
description="Unified CUDA kernels for FastVideo",
long_description=open("README.md").read(),
long_description_content_type="text/markdown",
license="Apache-2.0",
author="Hao AI Lab",
url="https://github.com/hao-ai-lab/FastVideo",
package_dir={"": "src"},
packages=find_packages(where="src"),
ext_modules=ext_modules,
cmdclass={"build_ext": BuildExtension} if ext_modules else {},
python_requires=">=3.10",
install_requires=["torch>=2.5.0", "triton>=2.0.0"],
)
@@ -1,4 +1,4 @@
from .version import __version__
__version__ = "0.1.0"
from fastvideo_kernel.ops import (
sliding_tile_attention,
@@ -11,24 +11,11 @@ from fastvideo_kernel.vmoba import (
process_moba_output,
)
from fastvideo_kernel.turbodiffusion_ops import (
Int8Linear,
FastRMSNorm,
FastLayerNorm,
int8_linear,
int8_quant,
)
__all__ = [
"sliding_tile_attention",
"video_sparse_attn",
"moba_attn_varlen",
"process_moba_input",
"process_moba_output",
"Int8Linear",
"FastRMSNorm",
"FastLayerNorm",
"int8_linear",
"int8_quant",
"__version__",
]
@@ -1,20 +1,20 @@
import math
import torch
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
from .triton_kernels.index import map_to_index
# Try to load the C++ extension
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
sta_fwd = getattr(fastvideo_kernel_ops, "sta_fwd", None)
block_sparse_fwd = getattr(fastvideo_kernel_ops, "block_sparse_fwd", None)
block_sparse_bwd = getattr(fastvideo_kernel_ops, "block_sparse_bwd", None)
from fastvideo_kernel._C.st_attn import sta_fwd
except ImportError:
sta_fwd = None
try:
from fastvideo_kernel._C.vsa import block_sparse_fwd, block_sparse_bwd
except ImportError:
block_sparse_fwd = None
block_sparse_bwd = None
def sliding_tile_attention(
q: torch.Tensor,
k: torch.Tensor,
@@ -24,15 +24,12 @@ def sliding_tile_attention(
has_text: bool = True,
seq_shape: str = "30x48x80",
) -> torch.Tensor:
# Check if the specific op is available
if sta_fwd is None:
return sliding_tile_attention_triton(
q, k, v, window_size, text_length, has_text, seq_shape
)
raise RuntimeError("STA kernel not compiled. Requires H100 and ThunderKittens at build time.")
seq_length = q.shape[2]
shape_map = {"30x48x80": 1, "36x48x48": 2, "18x48x80": 3}
if has_text:
target_size = math.ceil(seq_length / 384) * 384
pad_size = target_size - seq_length
@@ -40,20 +37,22 @@ def sliding_tile_attention(
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
output = torch.empty_like(q)
flag = shape_map[seq_shape]
for head_idx, (t, h, w) in enumerate(window_size):
sta_fwd(
q[:, head_idx:head_idx + 1], k[:, head_idx:head_idx + 1],
v[:, head_idx:head_idx + 1], output[:, head_idx:head_idx + 1],
q[:, head_idx:head_idx+1],
k[:, head_idx:head_idx+1],
v[:, head_idx:head_idx+1],
output[:, head_idx:head_idx+1],
t, h, w, text_length, False, has_text, flag
)
if has_text:
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
return output[:, :, :seq_length]
@@ -68,45 +67,37 @@ def video_sparse_attn(
) -> torch.Tensor:
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
batch, heads, seq_len, dim = q.shape
# Compression branch
q_c = q.view(batch, heads, seq_len // block_elements, block_elements, dim)
k_c = k.view(batch, heads, seq_len // block_elements, block_elements, dim)
v_c = v.view(batch, heads, seq_len // block_elements, block_elements, dim)
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
q.dtype)
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
k.dtype)
v_c = (v_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(
v.dtype)
scores = torch.matmul(q_c, k_c.transpose(-2, -1)) / (dim**0.5)
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
v_c = (v_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
scores = torch.matmul(q_c, k_c.transpose(-2, -1)) / (dim ** 0.5)
attn = torch.softmax(scores, dim=-1)
out_c = torch.matmul(attn, v_c)
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, seq_len, dim)
out_c = out_c.repeat(1, 1, 1, block_elements, 1).view(batch, heads, seq_len, dim)
# Sparse branch
topk_idx = torch.topk(scores, topk, dim=-1).indices
mask = torch.zeros_like(scores,
dtype=torch.bool).scatter_(-1, topk_idx, True)
idx, num = map_to_index(mask)
mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
if block_sparse_fwd is not None:
out_s = block_sparse_fwd(
q, k, v, idx, num, variable_block_sizes.int()
)[0] # block_sparse_fwd returns vector<Tensor>
idx, num = map_to_index(mask)
out_s, _ = block_sparse_fwd(q, k, v, idx, num, variable_block_sizes.int())
else:
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num,
variable_block_sizes)
idx, num = map_to_index(mask)
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
return out_c + out_s
@@ -8,6 +8,7 @@ This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
Credits: OpenAI kernel team
"""
import torch
import triton
import triton.language as tl
@@ -16,6 +17,7 @@ import triton.language as tl
import math # small utility needed by the sparse wrapper
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
# the code below and commenting out the equivalent parameters is convenient for
# re-tuning.
@@ -27,92 +29,65 @@ configs = [
for w in [4, 8]\
]
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(
Q,
K,
V,
sm_scale, #
q2k_index,
q2k_num,
max_kv_blks, #
variable_block_sizes,
M,
Out, #
stride_qz,
stride_qh,
stride_qm,
stride_qk,
stride_kz,
stride_kh,
stride_kn,
stride_kk,
stride_vz,
stride_vh,
stride_vk,
stride_vn,
stride_oz,
stride_oh,
stride_om,
stride_on,
Z,
H,
N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
STAGE: tl.constexpr):
def _attn_fwd_sparse(Q, K, V, sm_scale, #
q2k_index, q2k_num, max_kv_blks, #
variable_block_sizes,
M, Out, #
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vk, stride_vn,
stride_oz, stride_oh, stride_om, stride_on,
Z, H, N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
STAGE: tl.constexpr):
"""
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
(32×64 and 64×32) – memory footprint unchanged.
"""
# ----- program-id mapping -----
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
qvk_off = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
qvk_off = (b.to(tl.int64) * stride_qz +
h.to(tl.int64) * stride_qh)
Q_ptr = tl.make_block_ptr(base=Q + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0))
Q_ptr = tl.make_block_ptr(
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
K_base = tl.make_block_ptr(base=K + qvk_off,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N),
order=(0, 1))
K_base = tl.make_block_ptr(
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1,
0)
V_base = tl.make_block_ptr(base=V + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=v_order)
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
V_base = tl.make_block_ptr(
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
O_ptr = tl.make_block_ptr(base=Out + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0))
O_ptr = tl.make_block_ptr(
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
# ----- accumulators -----
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
@@ -152,30 +127,23 @@ def _attn_fwd_sparse(
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
@triton.jit
def _attn_bwd_preprocess(
O,
DO, #
Delta, #
Z,
H,
N_CTX, #
BLOCK_M: tl.constexpr,
HEAD_DIM: tl.constexpr #
):
def _attn_bwd_preprocess(O, DO, #
Delta, #
Z, H, N_CTX, #
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
):
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
off_hz = tl.program_id(1)
off_n = tl.arange(0, HEAD_DIM)
# load
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM +
off_n[None, :])
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM +
off_n[None, :]).to(tl.float32)
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
delta = tl.sum(o * do, axis=1)
# write-back
tl.store(Delta + off_hz * N_CTX + off_m, delta)
@@ -183,32 +151,19 @@ def _attn_bwd_preprocess(
# The main inner-loop logic for computing dK and dV.
@triton.jit
def _attn_bwd_dkdv(
dk,
dv, #
Q,
k,
v,
sm_scale, #
DO, #
M,
D, #
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_tok,
stride_d, #
H,
N_CTX,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
# Filled in by the wrapper.
start_n,
start_m,
num_steps):
def _attn_bwd_dkdv(dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_tok, stride_d, #
H, N_CTX, BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
# Filled in by the wrapper.
start_n, start_m, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M1)
offs_n = start_n + tl.arange(0, BLOCK_N1)
offs_k = tl.arange(0, HEAD_DIM)
@@ -217,20 +172,21 @@ def _attn_bwd_dkdv(
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
step_m = BLOCK_M1
kv_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
kv_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_N1
meta_base = ((b * H + h) * q_tiles + kv_blk)
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
block_size = tl.load(variable_block_sizes + kv_blk)
for blk_idx in range(q_blocks * 2):
block_sparse_offset = (tl.load(q_ptr + blk_idx // 2).to(tl.int32) * 2 +
blk_idx % 2) * step_m
for blk_idx in range(q_blocks*2):
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
@@ -256,32 +212,21 @@ def _attn_bwd_dkdv(
return dk, dv
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq(
dq,
q,
K,
V, #
do,
m,
D,
# shared by Q/K/V/DO.
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
# Filled in by the wrapper.
start_m,
start_n,
num_steps):
def _attn_bwd_dq(dq, q, K, V, #
do, m, D,
# shared by Q/K/V/DO.
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
# Filled in by the wrapper.
start_m, start_n, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M2)
offs_n = start_n + tl.arange(0, BLOCK_N2)
offs_k = tl.arange(0, HEAD_DIM)
@@ -292,27 +237,27 @@ def _attn_bwd_dq(
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
step_n = BLOCK_N2
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M2
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
block_size = tl.load(variable_block_sizes + q_blk)
for blk_idx in range(kv_blocks * 2):
block_sparse_offset = (tl.load(kv_ptr + blk_idx // 2).to(tl.int32) * 2 +
blk_idx % 2) * step_n * stride_tok
for blk_idx in range(kv_blocks*2):
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
p = tl.where(mask[None, :], p, 0.0)
p = tl.where(mask[None, :], p , 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
@@ -324,37 +269,23 @@ def _attn_bwd_dq(
return dq
@triton.jit
def _attn_bwd(
Q,
K,
V,
sm_scale, #
DO, #
DQ,
DK,
DV, #
M,
D,
q2k_index,
q2k_num,
max_kv_blks,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_z,
stride_h,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
def _attn_bwd(Q, K, V, sm_scale, #
DO, #
DQ, DK, DV, #
M, D,
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_z, stride_h, stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
@@ -388,32 +319,20 @@ def _attn_bwd(
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
num_steps = N_CTX // BLOCK_M1
dk, dv = _attn_bwd_dkdv( #
dk,
dv, #
Q,
k,
v,
sm_scale, #
dk, dv, #
Q, k, v, sm_scale, #
DO, #
M,
D, #
k2q_index,
k2q_num,
max_q_blks,
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M1,
BLOCK_N1,
HEAD_DIM, #
start_n,
start_m,
num_steps #
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
start_n, start_m, num_steps #
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
@@ -438,88 +357,54 @@ def _attn_bwd(
m = m[:, None]
num_steps = N_CTX // BLOCK_N2
dq = _attn_bwd_dq(
dq,
q,
K,
V, #
do,
m,
D, #
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d, #
H,
N_CTX, #
BLOCK_M2,
BLOCK_N2,
HEAD_DIM, #
start_m,
end_n,
num_steps #
)
dq = _attn_bwd_dq(dq, q, K, V, #
do, m, D, #
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
start_m, end_n, num_steps #
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= LN2
tl.store(dq_ptrs, dq)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
variable_block_sizes):
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
assert q2k_num.shape[
-1] == T // 64, f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
o = torch.empty_like(q)
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
_attn_fwd_sparse[grid](q,
k,
v,
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
M,
o,
q.stride(0),
q.stride(1),
q.stride(2),
q.stride(3),
k.stride(0),
k.stride(1),
k.stride(2),
k.stride(3),
v.stride(0),
v.stride(1),
v.stride(2),
v.stride(3),
o.stride(0),
o.stride(1),
o.stride(2),
o.stride(3),
B,
H,
T,
HEAD_DIM=D,
STAGE=3)
_attn_fwd_sparse[grid](
q, k, v, sm_scale,
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
M, o,
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
B, H, T,
HEAD_DIM=D, STAGE=3
)
return o, M
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
k2q_index, k2q_num, variable_block_sizes):
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
assert do.is_contiguous()
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
@@ -535,49 +420,30 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o,
do, #
o, do, #
delta, #
BATCH,
N_HEAD,
N_CTX, #
BLOCK_M=PRE_BLOCK,
HEAD_DIM=D #
BATCH, N_HEAD, N_CTX, #
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
)
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
q,
arg_k,
v,
sm_scale,
do,
dq,
dk,
dv, #
M,
delta, #
q2k_index,
q2k_num,
max_kv_blks,
k2q_index,
k2q_num,
max_q_blks,
q, arg_k, v, sm_scale, do, dq, dk, dv, #
M, delta, #
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
q.stride(0),
q.stride(1),
q.stride(2),
q.stride(3), #
N_HEAD,
N_CTX, #
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2, #
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
N_HEAD, N_CTX, #
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
HEAD_DIM=D #
)
return dq, dk, dv
@@ -0,0 +1,152 @@
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
@@ -0,0 +1,868 @@
# 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
try:
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
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
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')
@@ -0,0 +1,71 @@
from typing import Tuple
import torch
from torch import BoolTensor, IntTensor
from torch.nn.attention.flex_attention import create_block_mask
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
torch._inductor.config.realize_opcount_threshold = 100
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
"""Generates a 3D NATTEN attention mask with a given kernel size.
Args:
canvas_t: The time dimension of the canvas.
canvas_h: The height of the canvas.
canvas_w: The width of the canvas.
kernel_t: The time dimension of the kernel.
kernel_h: The height of the kernel.
kernel_w: The width of the kernel.
"""
canvas_t, canvas_h, canvas_w = canvas_twh
kernel_t, kernel_h, kernel_w = kernel_twh
tile_t_size, tile_h_size, tile_w_size = tile_twh
total_tile_size = tile_t_size * tile_h_size * tile_w_size
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
img_seq_len = canvas_t * canvas_h * canvas_w
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
tile_id = idx // total_tile_size
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
tile_w = tile_id % canvas_tile_w
return tile_t, tile_h, tile_w
def sta_mask_3d(
b: IntTensor,
h: IntTensor,
q_idx: IntTensor,
kv_idx: IntTensor,
) -> BoolTensor:
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
# kernel nominally attempts to center itself on the query, but kernel center
# is clamped to a fixed distance (kernel half-length) from the canvas edge
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
return sta_mask_3d
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
img_seq_len = img_size[0] * img_size[1] * img_size[2]
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
mask = create_block_mask(image_mask,
B=None,
H=None,
Q_LEN=img_seq_len + text_max_len,
KV_LEN=img_seq_len + text_max_len,
device=device,
_compile=True)
return mask
@@ -0,0 +1,63 @@
import torch
import sys
import os
from tqdm import tqdm
# Local support import
from .support_flex_sta import get_sliding_tile_attention_mask
# USE OUR NEW PACKAGE!
from fastvideo_kernel import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
# Using the same parameters as the original test
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
for kernel_size in kernel_size_ls:
print(f"Testing kernel_size: {kernel_size}")
for xi in tqdm(range(num_iterations)):
torch.manual_seed(xi)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
pt_o = flex_test(Q, K, V, kernel_size)
diff = pt_o - tk_o
abs_diff = torch.abs(diff)
max_d = torch.max(abs_diff).item()
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
if max_d > 0.1:
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
if __name__ == "__main__":
b, h, d = 2, 24, 128
n = 69120
causal = False
mean = 1e-1
std = 10
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
@@ -3,7 +3,7 @@
import torch
import pytest
import random
from fastvideo_kernel import moba_attn_varlen
from fastvideo_kernel.vmoba import moba_attn_varlen
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
"""
+9 -3
View File
@@ -55,12 +55,18 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
./build.sh
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
EXPOSE 22
+3 -4
View File
@@ -55,12 +55,11 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
./build.sh
python setup.py install
EXPOSE 22
+3 -3
View File
@@ -55,11 +55,11 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
./build.sh
python setup.py install
EXPOSE 22
+3 -4
View File
@@ -55,12 +55,11 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
./build.sh
python setup.py install
EXPOSE 22
+4 -4
View File
@@ -36,7 +36,7 @@ RUN echo "# Placeholder" > README.md
RUN source $HOME/.local/bin/env && \
uv venv --python 3.10 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip
uv pip install --no-cache-dir --upgrade pip
COPY . .
@@ -48,11 +48,11 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
./build.sh --rocm
python setup.py install
EXPOSE 22
+3 -5
View File
@@ -2,11 +2,9 @@
writing-mode: sideways-lr;
white-space: nowrap;
max-width: 0;
}
/* Keep header cell paragraph content tight (avoid CSS nesting for compatibility) */
.vertical-table-header th.head:not(.stub) p {
margin: 0;
p {
margin: 0;
}
}
/* Image sizing classes */
-186
View File
@@ -1,186 +0,0 @@
# Adding a New Attention Backend
FastVideo allows integrating new attention mechanisms easily. This guide walks you through adding a new backend (e.g., `MyNewAttn`).
## 1. Implement the Backend (Python)
Create a new file in `fastvideo/attention/backends/` (e.g., `mynew_attn.py`).
Your implementation should inherit from `AttentionBackend` defined in `abstract.py`.
```python
# fastvideo/attention/backends/mynew_attn.py
import torch
from .abstract import AttentionBackend
# Import the context manager to access metadata (optional)
from fastvideo.forward_context import get_forward_context
# Import compiled kernel if applicable (see Section 2)
try:
# Import from the top-level package
from fastvideo_kernel import my_compiled_attn_func
except ImportError:
my_compiled_attn_func = None
class MyNewAttnBackend(AttentionBackend):
def process_inputs(self, q, k, v, **kwargs):
# Pre-process inputs if necessary
return q, k, v
def forward(self, q, k, v, **kwargs):
# Optional: Access extra metadata passed via ForwardContext
# Only needed if your backend requires global state (e.g. window_size)
try:
context = get_forward_context()
metadata = context.attn_metadata
# Example: window_size = metadata.window_size
except (AssertionError, AttributeError):
# Handle case where context is not set (e.g. standard inference)
pass
if my_compiled_attn_func is not None:
return my_compiled_attn_func(q, k, v)
else:
# Fallback implementation (e.g., Triton or pure PyTorch)
return self.fallback_impl(q, k, v)
```
## 2. Passing Extra Information via ForwardContext (Optional)
FastVideo uses a `ForwardContext` to pass global metadata (like current timestep, batch info, or custom attention configurations) to attention backends without changing the `forward` signature of every layer. **This is optional and only required if your backend needs dynamic per-step information.**
To use this:
1. **Set Context**: In your pipeline or generation loop, use the `set_forward_context` context manager.
2. **Access Context**: Inside your attention backend, use `get_forward_context()`.
See [`docs/attention/sta/index.md`](../sta/index.md) (Sliding Tile Attention) for an example of how complex configuration (window sizes) is passed this way.
## 3. Adding Compiled Kernels (C++/CUDA)
If your backend requires custom CUDA kernels, you need to add them to the `fastvideo-kernel` package.
### A. Add Source Files
Place your kernel implementation files in `fastvideo-kernel/csrc/attention/`.
* `mynew_attn.cu` (CUDA implementation)
* `mynew_attn.h` (Optional headers)
### B. Register in Extension
Update `fastvideo-kernel/csrc/common_extension.cpp` to expose your function to Python.
```cpp
// 1. Declare external function
#ifdef COMPILE_MYNEW_ATTN
extern torch::Tensor mynew_attn_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v);
#endif
// 2. Register in module
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
// ... other kernels ...
#ifdef COMPILE_MYNEW_ATTN
m.def("mynew_attn_fwd", torch::wrap_pybind_function(mynew_attn_forward), "My New Attention Forward");
#endif
}
```
### C. Update CMakeLists.txt
Update `fastvideo-kernel/CMakeLists.txt` to compile your new files.
**Case 1: General CUDA Kernel (Runs on all GPUs)**
Add your source file directly to `EXTENSION_SOURCES` and define the compilation flag.
```cmake
# Add to EXTENSION_SOURCES
list(APPEND EXTENSION_SOURCES csrc/attention/mynew_attn.cu)
# Add compilation definition for common_extension.cpp
list(APPEND COMPILE_DEFS COMPILE_MYNEW_ATTN)
```
**Case 2: ThunderKittens Kernel (Hopper H100 Only)**
If your kernel uses ThunderKittens (TK), it requires specific architecture flags (`sm_90a`). Add it inside the `ENABLE_TK_KERNELS` block.
```cmake
if(ENABLE_TK_KERNELS)
# Add source only if TK is enabled
list(APPEND EXTENSION_SOURCES csrc/attention/mynew_attn_tk.cu)
# Add definition to guard registration
list(APPEND COMPILE_DEFS TK_COMPILE_MYNEW_ATTN)
endif()
```
### D. Expose in Python Ops
Update `fastvideo-kernel/python/fastvideo_kernel/ops.py` to make the function importable and handle fallbacks gracefully.
```python
# fastvideo-kernel/python/fastvideo_kernel/ops.py
# Try to load C++ extension symbols
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
mynew_attn_fwd = getattr(fastvideo_kernel_ops, "mynew_attn_fwd", None)
except ImportError:
mynew_attn_fwd = None
def my_compiled_attn_func(q, k, v):
# Runtime check: use C++ kernel if available, else fallback
if mynew_attn_fwd is not None:
return mynew_attn_fwd(q, k, v)
else:
# Call Triton/Python fallback
return mynew_attn_triton(q, k, v)
```
### E. Expose in Package Init
Update `fastvideo-kernel/python/fastvideo_kernel/__init__.py` to export the function.
```python
from fastvideo_kernel.ops import (
my_compiled_attn_func,
# ...
)
__all__ = [
"my_compiled_attn_func",
# ...
]
```
## 4. Register the Backend
Update `fastvideo/attention/backends/__init__.py` to export your new class.
```python
from .mynew_attn import MyNewAttnBackend
```
## 5. Platform Integration
If your backend requires specific platform checks (e.g., checking for H100 support), handle that in `fastvideo/platforms/cuda.py` or within your backend's `__init__`.
## 6. Add Documentation
Create a new documentation page for your backend to explain its usage, installation (if custom kernels are needed), and features.
1. **Create Directory**: `docs/attention/mynew_attn/`
2. **Create Index**: `docs/attention/mynew_attn/index.md`
3. **Update Navigation**: Add an entry to `mkdocs.yml` under the "Attention" tab.
## Checklist
* [ ] Created `fastvideo/attention/backends/mynew_attn.py`.
* [ ] (Optional) Added CUDA kernels in `fastvideo-kernel/csrc/attention/`.
* [ ] (Optional) Updated `common_extension.cpp` and `CMakeLists.txt`.
* [ ] (Optional) Exposed kernel in `fastvideo-kernel/python/fastvideo_kernel/ops.py`.
* [ ] (Optional) Exported kernel in `fastvideo-kernel/python/fastvideo_kernel/__init__.py`.
* [ ] Implemented `forward` method respecting the standard signature.
* [ ] Added unit tests in `tests/`.
* [ ] Added documentation in `docs/attention/` and updated `mkdocs.yml`.
-53
View File
@@ -1,53 +0,0 @@
# FastVideo Attention Kernels
FastVideo provides highly optimized custom attention kernels to accelerate video generation.
## Supported Kernels
* **[Video Sparse Attention (VSA)](vsa/index.md)**: Sparse attention mechanism selecting top-k blocks.
* **[Sliding Tile Attention (STA)](sta/index.md)**: Optimized attention for window-based video generation.
## General Build Instructions
These instructions apply to building the `fastvideo-kernel` package from source, which includes both STA and VSA kernels.
### Prerequisites
* **PyTorch**: 2.5.0+
* **CUDA**: 12.4+ (12.8 recommended for best performance)
* **C++ Compiler**: GCC 11+ (C++20 support required for ThunderKittens)
Install system dependencies:
```bash
sudo apt update
sudo apt install -y gcc-11 g++-11 clang-11 ninja-build
# Set gcc-11 as default
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
```
Set up your CUDA environment variables (adjust version as needed):
```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
```
### Compile and Install
Clone the repository and build the kernel:
```bash
# Clone recursively to get ThunderKittens submodule
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo/fastvideo-kernel
# Build and install
./build.sh
```
The build script automatically detects your GPU architecture:
* **H100 (sm_90a)**: Compiles optimized C++ ThunderKittens kernels.
* **Other (A100, etc.)**: Skips C++ compilation; installs Python package with Triton kernels.
-36
View File
@@ -1,36 +0,0 @@
# Sliding Tile Attention (STA)
Optimized attention for window-based video generation (e.g., HunyuanVideo).
## Installation
STA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
## Usage
```python
from fastvideo_kernel import sliding_tile_attention
# q, k, v: [batch_size, num_heads, seq_length, head_dim]
# window_size: List of (t, h, w) tiles. Tile size is (6, 8, 8).
# text_length: Number of text tokens (0-256)
out = sliding_tile_attention(
q, k, v,
window_size=[(3, 3, 3)], # Example window
text_length=256
)
```
## Citation
If you use Sliding Tile Attention in your research, please cite:
```bibtex
@article{zhang2025fast,
title={Fast video generation with sliding tile attention},
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
journal={arXiv preprint arXiv:2502.04507},
year={2025}
}
```
-36
View File
@@ -1,36 +0,0 @@
# Video Sparse Attention (VSA)
Sparse attention mechanism selecting top-k blocks.
## Installation
VSA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
## Usage
```python
from fastvideo_kernel import video_sparse_attn
# q, k, v: [batch_size, num_heads, seq_len, head_dim]
# variable_block_sizes: Number of valid tokens per block
# topk: Number of blocks to attend
output = video_sparse_attn(
q, k, v,
variable_block_sizes=block_sizes,
topk=32
)
```
## Citation
If you use Video Sparse Attention in your research, please cite:
```bibtex
@article{zhang2025vsa,
title={Vsa: Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
```
+14 -35
View File
@@ -4,26 +4,25 @@ This document outlines FastVideo's architecture for developers interested in fra
## Table of Contents - Directory Structure and Files
- [`fastvideo/pipelines/`](#pipeline-system) - Core diffusion pipeline components
- [`fastvideo/models/`](#model-components) - Model implementations
- [`dits/`](#transformer-models) - Transformer-based diffusion models
- [`vaes/`](#vae-variational-auto-encoder) - Variational autoencoders
- [`encoders/`](#text-and-image-encoders) - Text and image encoders
- [`schedulers/`](#schedulers) - Diffusion schedulers
- [`fastvideo/attention/`](#optimized-attention) - Optimized attention implementations
- [`fastvideo/distributed/`](#distributed-processing) - Distributed computing utilities
- [`fastvideo/layers/`](#tensor-parallelism) - Custom neural network layers
- [`fastvideo/platforms/`](#platforms) - Hardware platform abstractions
- [`fastvideo/worker/`](#executor-and-worker-system) - Multi-GPU process management
- [`fastvideo/fastvideo_args.py`](#fastvideoargs) - Argument handling
- [`fastvideo/forward_context.py`](#forward-context-management) - Forward pass context management
- [`fastvideo/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
- [`fastvideo/models/`](#design-model-components) - Model implementations
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
- [`fastvideo/attention/`](#design-optimized-attention) - Optimized attention implementations
- [`fastvideo/distributed/`](#design-distributed-processing) - Distributed computing utilities
- [`fastvideo/layers/`](#design-tensor-parallelism) - Custom neural network layers
- [`fastvideo/platforms/`](#design-platforms) - Hardware platform abstractions
- [`fastvideo/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
- [`fastvideo/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
- [`fastvideo/forward_context.py`](#design-forwardcontext) - Forward pass context management
- `fastvideo/utils.py` - Utility functions
- [`fastvideo/logger.py`](#logger) - Logging infrastructure
- [`fastvideo/logger.py`](#design-logger) - Logging infrastructure
## Core Architecture
FastVideo separates model components from execution logic with these principles:
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
- **Modular Design**: Components can be independently replaced
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
@@ -35,14 +34,12 @@ FastVideo separates model components from execution logic with these principles:
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
Key features include:
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
- **Parameter Validation**: Ensures valid combinations of settings
Common configuration areas:
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
@@ -93,9 +90,7 @@ class MyCustomPipeline(ComposedPipelineBase):
```
### Pipeline Stages
Each stage handles a specific diffusion process component:
- **Input Validation**: Parameter verification
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
- **Image Encoding**: Image input processing
@@ -138,7 +133,6 @@ Transformer networks perform the actual denoising during diffusion:
- `HunyuanVideoTransformer3DModel`
Features include:
- Text/image conditioning
- Standardized interface for model-specific optimizations
@@ -167,7 +161,6 @@ VAEs handle conversion between pixel space and latent space:
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
FastVideo's VAE implementations include:
- Efficient video batch processing
- Memory optimization
- Optional tiling for large frames
@@ -186,7 +179,6 @@ Encoders process conditioning inputs into embeddings:
- `CLIPVisionModel`
FastVideo implements optimizations such as:
- Vocab parallelism for distributed processing
- Caching for common prompts
- Precision-tuned computation
@@ -201,7 +193,6 @@ Schedulers manage the diffusion sampling process:
- `FlowMatchEulerDiscreteScheduler`
These components control:
- Diffusion timestep sequences
- Noise prediction to latent update conversions
- Quality/speed trade-offs
@@ -228,9 +219,7 @@ This diagram shows how models are discovered, validated, and loaded across entry
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
### Attention Backends
Multiple implementations with automatic selection:
- **FLASH_ATTN**: Optimized for supporting hardware
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
- **SLIDING_TILE_ATTN**: For very long sequences
@@ -251,9 +240,7 @@ self.attn = LocalAttention(
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
- Chunking, progressive computation, optimized masking
@@ -309,7 +296,6 @@ self.attn = DistributedAttention(
```
### Communication Primitives
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
Efficient communication primitives minimize distributed overhead:
@@ -328,7 +314,6 @@ Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-sp
- **Profiling Data**: Potential hooks for performance metrics collection
This context-based approach enables:
- Dynamic optimization based on execution state (e.g., attention backend selection)
- Step-specific customizations within model components
@@ -354,14 +339,12 @@ FastVideo implements a flexible execution model for distributed processing:
- **GPU Workers**: Handle actual model execution on individual GPUs
The MultiProcExecutor implementation:
1. Spawns worker processes for each GPU
2. Establishes communication channels via pipes
3. Coordinates distributed operations across workers
4. Handles graceful startup and shutdown of the process group
Each GPU worker:
1. Initializes the distributed environment
2. Builds the pipeline for the specified model
3. Executes requested operations on its assigned GPU
@@ -376,13 +359,11 @@ The `fastvideo/platforms/` directory provides hardware platform abstractions tha
### Platform Abstraction
FastVideo's platform abstraction layer enables:
- **Hardware Detection**: Automatic detection of available hardware
- **Backend Selection**: Appropriate selection of compute kernels
- **Memory Management**: Efficient utilization of hardware-specific memory features
The primary components include:
- **Platform Interface**: Defines the common API for all platform implementations
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
- **Backend Enum**: Used throughout the codebase for feature selection
@@ -402,7 +383,6 @@ else:
The platform system is designed to be extensible for future hardware targets.
## Logger
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
@@ -417,7 +397,6 @@ If you're a new contributor, here are some common areas to explore:
4. **Hardware support**: Extend the `platforms` module for new hardware targets
When adding code, follow these practices:
- Use type hints for better code readability
- Add appropriate docstrings
- Maintain the separation between model components and execution logic
+1 -1
View File
@@ -13,7 +13,7 @@ We provide two distilled models:
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
## ⚙️ Inference
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
```bash
bash scripts/inference/v1_inference_wan_dmd.sh
+7 -11
View File
@@ -11,13 +11,15 @@ FastVideo supports the following hardware platforms:
### Using pip
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
pip install fastvideo
```
### Using conda
```bash
conda install -c conda-forge fastvideo
```
### From source
```bash
@@ -26,12 +28,6 @@ cd FastVideo
pip install -e .
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
## Hardware Requirements
- **NVIDIA GPUs**: CUDA 11.8+ with compute capability 7.0+
@@ -42,4 +38,4 @@ pip install flash-attn --no-build-isolation
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
- [Examples](../inference/examples/) - Explore example scripts and notebooks
+2 -2
View File
@@ -84,12 +84,12 @@ pip install flash-attn --no-build-isolation
## Set up using Docker
We also have prebuilt docker images with FastVideo dependencies pre-installed:
[Docker Images](../../contributing/developer_env/docker.md)
[Docker Images](#docker)
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](../../contributing/overview.md)
[Contributor Guide](#developer-overview)
## Hardware Requirements
+1 -1
View File
@@ -78,7 +78,7 @@ uv pip install -e .
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](../../contributing/overview.md)
[Contributor Guide](#developer-overview)
## Hardware Requirements
-6
View File
@@ -15,12 +15,6 @@ conda activate fastvideo
pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
## Basic Usage
### Text-to-Video Generation
-6
View File
@@ -45,7 +45,6 @@ FastVideo uses the Hugging Face Diffusers format for model organization:
### Implementing Modules
Place new modules in the appropriate directories:
- Encoders: `fastvideo/models/encoders/`
- VAEs: `fastvideo/models/vaes/`
- Transformer models: `fastvideo/models/dits/`
@@ -54,15 +53,12 @@ Place new modules in the appropriate directories:
### Adapting Model Layers
#### Layer Replacements
Replace standard PyTorch layers with FastVideo optimized versions:
- nn.LayerNorm → fastvideo.layers.layernorm.RMSNorm
- Embedding layers → fastvideo.layers.vocab_parallel_embedding modules
- Activation functions → versions from fastvideo.layers.activation
#### Distributed Linear Layers
Use appropriate parallel layers for distribution:
```python
@@ -95,7 +91,6 @@ self.out_proj = RowParallelLinear(
```
### Attention Layers
Replace standard attention with FastVideo's optimized attention:
```python
@@ -309,7 +304,6 @@ EntryClass = [MyCustomPipeline, MyOtherPipeline]
```
The registry will automatically:
1. Scan all packages under `fastvideo/pipelines/`
2. Look for `EntryClass` variables
3. Register pipelines using their class names as identifiers
+1 -1
View File
@@ -1,7 +1,7 @@
# FastVideo CLI Inference
The FastVideo CLI provides a quick way to access the FastVideo inference pipeline for video generation. For more advanced usage,
see the Python interface [here](examples/basic.md).
see the Python interface [here](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html).
## Basic Usage
+1 -1
View File
@@ -74,4 +74,4 @@ if __name__ == '__main__':
## Performance Optimization
For configuring optimizations, please see our [optimizations guide](optimizations.md)
For configuring optimizations, please see our [optimizations guide](#inference-optimizations)
+5 -16
View File
@@ -3,7 +3,6 @@
This page contains step-by-step instructions to get you quickly started with video generation using FastVideo.
## Requirements
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.8
@@ -22,10 +21,9 @@ conda activate fastvideo
pip install fastvideo
```
For advanced installation options, see the [Installation Guide](../getting_started/installation.md).
For advanced installation options, see the [Installation Guide](installation.md).
## Generating Your First Video
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
```python
@@ -62,10 +60,9 @@ python example.py
The generated video will be saved in the current directory under `my_videos/`
More inference example scripts can be found in `scripts/inference/`
## Available Models
Please see the [support matrix](support_matrix.md) for the list of supported models and their available optimizations.
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
## Image-to-Video Generation
@@ -99,28 +96,20 @@ if __name__ == '__main__':
Common issues and their solutions:
### Out of Memory Errors
If you encounter CUDA out of memory errors:
- Reduce `num_frames` or video resolution
- Enable memory optimization with `enable_model_cpu_offload`
- Try a smaller model or use distilled versions
- Use `num_gpus` > 1 if multiple GPUs are available
- Try enabling FSDP inference with `use_fsdp_inference=True` (may slow down generation)
- Try enabling DiT layerwise offload with `dit_layerwise_offload=True` (now only a few models support this, but may introduce less overhead than FSDP)
### Slow Generation
To speed up generation:
- Reduce `num_inference_steps` (20-30 is usually sufficient)
- Use half precision (`fp16`) for the VAE
- Use multiple GPUs if available
### Unexpected Results
If the generated video doesn't match your prompt:
- Try increasing `guidance_scale` (7.0-9.0 works well)
- Make your prompt more detailed and specific
- Experiment with different random seeds
@@ -128,8 +117,8 @@ If the generated video doesn't match your prompt:
## Next Steps
- Learn about [Advanced Inference Configurations](configuration.md)
- Learn about using [Optimizations](optimizations.md)
- See [Examples](examples/examples_inference_index.md) for more usage scenarios
- Learn about [Advanced Inference Configurations](#inference-configuration)
- Learn about using [Optimizations](#inference-optimizations)
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
+7 -7
View File
@@ -7,13 +7,13 @@ This page describes the various options for speeding up generation times in Fast
- Optimized Attention Backends
- [Flash Attention](#flash-attention)
- [Sliding Tile Attention](#sliding-tile-attention)
- [Sage Attention](#sage-attention)
- [Sage Attention 3](#sage-attention-3)
- [Flash Attention](#optimizations-flash)
- [Sliding Tile Attention](#optimizations-sta)
- [Sage Attention](#optimizations-sage)
- [Sage Attention 3](#optimizations-sage3)
- Caching Techniques
- [TeaCache](#teacache)
- [TeaCache](#optimizations-teacache)
## Attention Backends
@@ -74,7 +74,7 @@ python setup.py install
pip install st_attn==0.0.4
```
Please see [this page](../attention/sta/index.md) for more installation instructions.
Please see [this page](#sta-installation) for more installation instructions.
### Video Sparse Attention
@@ -85,7 +85,7 @@ git submodule update --init --recursive
python setup_vsa.py install
```
Please see [this page](../attention/vsa/index.md) for more installation instructions.
Please see [this page](#vsa-installation) for more installation instructions.
### Sage Attention
+14 -30
View File
@@ -40,26 +40,20 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
}
</style>
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA | BSA |
|------------|---------------------|-------------|----------|-------------------|-----------|-----|-----|
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| TurboWan2.1 T2V 1.3B | `loayrashid/TurboWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| TurboWan2.1 T2V 14B | `loayrashid/TurboWan2.1-T2V-14B-Diffusers` | 480P, 720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| LongCat T2V 13.6B | See note** | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
| Matrix Game 2.0 Base | `FastVideo/Matrix-Game-2.0-Base-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Matrix Game 2.0 GTA | `FastVideo/Matrix-Game-2.0-GTA-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Matrix Game 2.0 TempleRun | `FastVideo/Matrix-Game-2.0-TempleRun-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA |
|------------|---------------------|-------------|----------|-------------------|-----------|-----|
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ |
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ |
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ |
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ |
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
@@ -70,13 +64,3 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
### Sliding Tile Attention
- Currently only Hopper GPUs (H100s) are supported.
### TurboWan2.1 (TurboDiffusion)
- Uses TurboDiffusionPipeline with RCM scheduler for 1-4 step generation
- Requires SLA attention backend: `export FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
- Uses `guidance_scale=1.0` (no classifier-free guidance)
### Matrix Game 2.0
- Image-to-video game world models with keyboard/mouse control input
- Three variants available: Base (universal), GTA, and TempleRun
- Each variant has different keyboard dimensions for control inputs
+16
View File
@@ -0,0 +1,16 @@
# 🔍 Demo
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
<div style="text-align: center;">
<video controls width="800">
<source src="https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747" type="video/mp4">
Your browser does not support the video tag.
</video>
</div>
You can run STA using the following command:
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
@@ -0,0 +1,65 @@
# 🔧 Installation
You can install the Sliding Tile Attention package using
```
pip install st_attn
```
# Building from Source
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently, we only have an implementation for H100s.
First, 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
```
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.py install
```
# 🧪 Test
```bash
python csrc/attn/tests/test_sta.py
```
# 📋 Usage
```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)
```
# 🚀Inference
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
+21 -106
View File
@@ -1,130 +1,45 @@
# 🧱 Data Preprocessing
To save GPU memory during training, FastVideo precomputes text embeddings and VAE latents. This eliminates the need to load the text encoder and VAE during training.
# 🧱 Data Preprocess
## Quick Start
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
Download the sample dataset and run preprocessing:
We provide a sample dataset to help you get started. Download the source media using the following command:
```bash
# Download the crush-smol dataset
python scripts/huggingface/download_hf.py \
--repo_id "wlsaidhi/crush-smol-merged" \
--local_dir "data/crush-smol" \
--repo_type "dataset"
# Run preprocessing
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v_new.sh
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=data/mini_i2v_dataset --repo_type=dataset
```
## Preprocessing Pipeline
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
The new preprocessing pipeline supports multiple dataset formats and video loaders:
```bash
GPU_NUM=2
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.video_loader_type torchvision \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 77 \
--preprocess.train_fps 16 \
--preprocess.samples_per_file 8 \
--preprocess.flush_frequency 8 \
--preprocess.video_length_tolerance_range 5
```
### Key Parameters
| Parameter | Description |
|-----------|-------------|
| `--workload_type` | Task type: `t2v` (text-to-video) or `i2v` (image-to-video) |
| `--preprocess.dataset_type` | Input format: `hf` (HuggingFace) or `merged` (local folder) |
| `--preprocess.dataset_path` | Path to dataset (HF repo ID or local folder) |
| `--preprocess.dataset_output_dir` | Output directory for Parquet files |
| `--preprocess.video_loader_type` | Video decoder: `torchcodec` or `torchvision` |
| `--preprocess.max_height` / `max_width` | Target resolution for videos |
| `--preprocess.num_frames` | Number of frames to extract per video |
| `--preprocess.train_fps` | Target FPS for frame extraction |
## Dataset Formats
### Merged Dataset (Local Folder)
Structure your dataset as follows:
To preprocess the dataset for fine-tuning or distillation, run:
```
your_dataset/
├── videos/
│ ├── video_001.mp4
│ ├── video_002.mp4
│ └── ...
└── videos2caption.json
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
```
The `videos2caption.json` maps video filenames to captions:
## Process your own dataset
```json
[
{"path": "video_001.mp4", "cap": "A cat playing with yarn..."},
{"path": "video_002.mp4", "cap": "Ocean waves at sunset..."}
]
```
### HuggingFace Dataset
Use `--preprocess.dataset_type hf` and point `--preprocess.dataset_path` to a HuggingFace dataset with `video` and `caption` columns.
## Creating Your Own Dataset
If you have raw videos and captions in separate files, generate the `videos2caption.json`:
```bash
python scripts/dataset_preparation/prepare_json_file.py \
--data_folder path/to/your_raw_data/ \
--output path/to/output_folder
```
Your raw data folder should contain:
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
```
your_raw_data/
path_to_your_dataset_folder/
├── videos/
│ ├── 0.mp4
│ ├── 1.mp4
│ └── ...
├── videos.txt # list of video filenames
└── prompt.txt # corresponding captions (one per line)
├── videos.txt
└── prompt.txt
```
## Output Format
To generate the `videos2caption.json` and `merge.txt`, run
Preprocessing outputs Parquet files in the `combined_parquet_dataset/` subdirectory containing:
``` python
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
```
- `vae_latent_bytes` — VAE-encoded video latent
- `text_embedding_bytes` — text encoder output
- `clip_feature_bytes` — CLIP image features (I2V only)
- `first_frame_latent_bytes` — first frame latent (I2V only)
- Metadata: shapes, dtypes, and sample identifiers
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
## Examples
```
bash scripts/preprocess/v1_preprocess_****.sh
```
See ready-to-run preprocessing scripts in the training examples:
- **T2V**: `examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v_new.sh`
- **I2V**: `examples/training/finetune/wan_i2v_14B_480p/crush_smol/preprocess_wan_data_i2v_new.sh`
**→ [Browse all training examples](examples/examples_training_index.md)**
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
+55 -153
View File
@@ -1,176 +1,78 @@
# 🧠 Finetuning
This guide covers finetuning video diffusion models with FastVideo, including full finetuning and LoRA.
## Training Arguments
FastVideo training scripts use several argument groups:
### Training Arguments
| Argument | Description |
|----------|-------------|
| `--max_train_steps` | Total training steps |
| `--train_batch_size` | Batch size per GPU |
| `--gradient_accumulation_steps` | Steps to accumulate before optimizer update |
| `--num_latent_t` | Temporal latent dimension (reduce to save memory) |
| `--num_height` / `--num_width` | Video resolution |
| `--num_frames` | Number of frames per video |
| `--output_dir` | Directory for checkpoints |
### Parallelism Arguments
| Argument | Description |
|----------|-------------|
| `--num_gpus` | Total number of GPUs |
| `--sp_size` | Sequence parallel size (increase to reduce memory per GPU) |
| `--tp_size` | Tensor parallel size |
| `--hsdp_replicate_dim` | HSDP replication dimension |
| `--hsdp_shard_dim` | HSDP sharding dimension |
### Optimizer Arguments
| Argument | Description |
|----------|-------------|
| `--learning_rate` | Base learning rate |
| `--mixed_precision` | Precision mode (`bf16` recommended) |
| `--weight_decay` | Weight decay for regularization |
| `--max_grad_norm` | Gradient clipping threshold |
### Validation Arguments
| Argument | Description |
|----------|-------------|
| `--log_validation` | Enable validation logging |
| `--validation_dataset_file` | JSON file with validation prompts |
| `--validation_steps` | Run validation every N steps |
| `--validation_sampling_steps` | Inference steps for validation |
| `--validation_guidance_scale` | CFG scale for validation |
## Full Finetuning
Full finetuning updates all model weights. This provides the best quality but requires more GPU memory.
# 🧠 Finetune
## ⚡ Full Finetune
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
```bash
# Example: Wan2.1 T2V 1.3B full finetune (4 GPUs)
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
```
**Typical settings:**
Download the original model weights as specified in the [Distillation Section](../distillation/dmd.md):
- Learning rate: `1e-5` to `5e-5`
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
Then you can run the finetune with:
## LoRA Finetuning
```
bash scripts/finetune/finetune_mochi.sh # for mochi
```
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
### LoRA-Specific Arguments
| Argument | Description |
|----------|-------------|
| `--lora_training True` | Enable LoRA mode |
| `--lora_rank` | Rank of LoRA adapters (16, 32, 64, 128) |
### Learning Rate for LoRA
**Important:** LoRA typically requires a **10–20× higher learning rate** than full finetuning because only the low-rank adapters are being trained while the base model is frozen.
| Training Mode | Recommended Learning Rate |
|---------------|---------------------------|
| Full finetune | `1e-5` to `5e-5` |
| LoRA | `1e-4` to `2e-4` |
### Example LoRA Training
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
## ⚡ Finetune with VSA
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
```bash
# Example: Wan2.1 T2V 1.3B LoRA finetune (1 GPU)
bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v_lora.sh
bash scripts/finetune/finetune_v1_VSA.sh
```
Key differences from full finetune:
## ⚡ Lora Finetune
- Add `--lora_training True --lora_rank 32`
- Use higher learning rate (10–20× full finetune)
- Can run on fewer GPUs (even single GPU)
- Outputs adapter weights instead of full model
## LoRA Extraction and Merging
FastVideo provides tools to extract LoRA adapters from finetuned models and merge them back.
### Extract LoRA Adapter
Extract a LoRA adapter by comparing a finetuned model to its base:
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
```bash
python scripts/lora_extraction/extract_lora.py \
--base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--ft path/to/your/finetuned_model \
--out adapter_r32.safetensors \
--rank 32
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
```
| Argument | Description |
|----------|-------------|
| `--base` | Base model (HuggingFace ID or local path) |
| `--ft` | Finetuned model path |
| `--out` | Output adapter file (.safetensors) |
| `--rank` | LoRA rank (16, 32, 64, 128) |
| `--full-rank` | Extract full-rank adapter (optional) |
### Minimum Hardware Requirement
- 40 GB GPU memory each for 2 GPUs with lora.
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
### Merge LoRA Adapter
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
Merge an adapter back into a base model:
### Dataset Preparation
We provide scripts to better help you get started to train on your own characters!
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
```
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
```
Also, we provide script to resize your videos:
```
python scripts/data_preprocess/resize_videos.py
```
### Finetuning
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
```
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
```
### Inference
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
```
bash scripts/inference/inference_hunyuan_hf.sh
```
**We also provide scripts for Mochi in the same directory.**
### Finetune with Both Image and Video
Our codebase support finetuning with both image and video.
```bash
python scripts/lora_extraction/merge_lora.py \
--base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--adapter adapter_r32.safetensors \
--ft path/to/your/finetuned_model \
--output merged_model
bash scripts/finetune/finetune_hunyuan.sh
bash scripts/finetune/finetune_mochi_lora_mix.sh
```
| Argument | Description |
|----------|-------------|
| `--base` | Base model path |
| `--adapter` | LoRA adapter file |
| `--ft` | Finetuned model (for config reference) |
| `--output` | Output directory for merged model |
### Validate Merged Model
Compare the merged model against the original finetuned model:
```bash
python scripts/lora_extraction/lora_inference_comparison.py \
--base merged_model \
--ft path/to/your/finetuned_model \
--adapter NONE \
--output-dir results \
--prompt "A cat sitting on a windowsill" \
--compute-ssim \
--compute-lpips
```
## Training Examples
Ready-to-run training scripts are available for multiple models:
**→ [Browse all training examples](examples/examples_training_index.md)**
| Model | Type | Example |
|-------|------|---------|
| Wan2.1 T2V 1.3B | T2V | `examples/training/finetune/wan_t2v_1.3B/crush_smol/` |
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
Each example includes:
- `download_dataset.sh` — download sample data
- `preprocess_*.sh` — run preprocessing
- `finetune_*.sh` — full finetune launcher
- `finetune_*_lora.sh` — LoRA finetune launcher
- `validation.json` — validation prompts
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
-67
View File
@@ -1,67 +0,0 @@
# Training Overview
FastVideo supports finetuning video diffusion models on custom datasets. This page explains what data you need and how to get started.
## Data Requirements
To save GPU memory during training, FastVideo precomputes embeddings and latents ahead of time. This eliminates the need to load the text encoder and VAE during training, significantly reducing memory usage.
### Text-to-Video (T2V) Finetuning
For T2V models, you need:
| Component | Description |
|-----------|-------------|
| **Text embeddings** | Precomputed embeddings from the model's text encoder (e.g., T5 or LLaMA). Stored as numpy arrays in Parquet files. |
| **Video latents** | VAE-encoded representations of your training videos. Each video is encoded into a compressed latent tensor. |
### Image-to-Video (I2V) Finetuning
For I2V models, you need everything from T2V plus additional image conditioning. Note that not all I2V architectures require encoded images—this depends on how the model conditions on the input frame. Wan2.1 and Wan2.2 A14B I2V models do require these additional components:
| Component | Description |
|-----------|-------------|
| **Text embeddings** | Same as T2V—precomputed from the text encoder. |
| **Video latents** | Same as T2V—VAE-encoded video representations. |
| **First frame latent** | VAE-encoded representation of the first frame, used as the conditioning image. |
| **CLIP features** | Image embeddings from a CLIP vision encoder for the conditioning frame. |
## Preprocessing
Before training, you need to preprocess your raw videos and captions into Parquet files containing precomputed latents and embeddings.
FastVideo supports two input formats:
- **HuggingFace datasets** — load directly from HF Hub or local HF datasets
- **Merged datasets** — local folder with videos and a `videos2caption.json` metadata file
**→ See [Data Preprocessing](data_preprocess.md) for full details and examples.**
## Training Examples
Ready-to-run examples with preprocessing scripts, training launchers, and validation configs are available for multiple models and datasets:
**→ [Browse all training examples](examples/examples_training_index.md)**
Each example includes:
- `download_dataset.sh` — download sample data
- `preprocess_*.sh` — run preprocessing
- `finetune_*.sh` — launch training (full finetune or LoRA)
- `validation.json` — validation prompts for checkpoints
## Training Methods
FastVideo supports several training approaches:
| Method | Use Case |
|--------|----------|
| **Full finetune** | Adapt entire model to a new domain or style |
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
## Next Steps
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
-71
View File
@@ -1,71 +0,0 @@
# LoRA Extraction and Merging
Tools for extracting and merging LoRA adapters for FastVideo models.
## Extract LoRA Adapter
```bash
python scripts/lora_extraction/extract_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--out adapter_r32.safetensors \
--rank 32
```
**Options:**
- `--base`: Base model (HuggingFace ID or local path)
- `--ft`: Fine-tuned model (HuggingFace ID or local path)
- `--out`: Output adapter file
- `--rank`: LoRA rank (16, 32, 64, 128)
- `--full-rank`: Extract full-rank adapter (optional)
## Merge Adapter
```bash
python scripts/lora_extraction/merge_lora.py \
--base Wan-AI/Wan2.2-TI2V-5B-Diffusers \
--adapter adapter_r32.safetensors \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--output merged_model
```
**Options:**
- `--base`: Base model (HuggingFace ID or local path)
- `--adapter`: LoRA adapter file (.safetensors)
- `--ft`: Fine-tuned model (for configuration)
- `--output`: Output directory
## Validate Quality (Optional)
```bash
python scripts/lora_extraction/lora_inference_comparison.py \
--base merged_model \
--ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \
--adapter NONE \
--output-dir results \
--prompt "A cat sitting on a windowsill" \
--seed 42 \
--height 480 \
--width 480 \
--num-frames 49 \
--num-inference-steps 32 \
--compute-ssim \
--compute-lpips
```
**Options:**
- `--base`: Merged model or base model path
- `--ft`: Fine-tuned model (reference)
- `--adapter`: Path to adapter or NONE
- `--output-dir`: Output directory
- `--prompt`: Text prompt (default: "A cat sitting on a windowsill")
- `--seed`: Random seed (default: 42)
- `--height`: Video height (default: 480)
- `--width`: Video width (default: 832)
- `--num-frames`: Number of frames (default: 49)
- `--num-inference-steps`: Inference steps (default: 32)
- `--compute-ssim`: Compute SSIM metric
- `--compute-lpips`: Compute LPIPS metric
@@ -0,0 +1,65 @@
# 🔧 Installation
You can install the Video Sparse Attention package using
```bash
pip install vsa
```
# Building from Source
We support H100s (via ThunderKittens) and any other GPU (via Triton) for VSA.
First, install C++20 for ThunderKittens (if using an H100):
```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
```
Set up CUDA environment (if using CUDA 12.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
```
Install VSA:
```bash
cd csrc/attn/video_sparse_attn/
git submodule update --init --recursive
python setup.py install
```
# 🧪 Test
```bash
python csrc/attn/tests/test_vsa.py
```
# 📋 Usage
```python
from vsa import video_sparse_attn
# q, k, v: [batch_size, num_heads, seq_len, head_dim]
# variable_block_sizes: [num_blocks] - number of valid tokens in each block
# topk: int - number of top-k blocks to attend to
# block_size: int or tuple of 3 ints - size of each block (default: 64 tokens)
# compress_attn_weight: optional weight for compressed attention branch
output = video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight)
```
# 🚀Inference
```bash
bash scripts/inference/v1_inference_wan_VSA.sh
```
-19
View File
@@ -1,19 +0,0 @@
# Self-Forcing Distillation for SFWan2.1 T2V 1.3B
These scripts demonstrate self-forcing distillation (SFwan) for the causal Wan2.1 T2V 1.3B model. The workflow mirrors DMD2 while injecting self-forcing blocks so the student can autoregressively refine later frames.
## Run the recipe
1. Download the preprocessed text-video dataset:
```bash
bash examples/distill/SFWan2.1-T2V/download_dataset.sh
```
2. (Optional) Regenerate parquet shards locally:
```bash
bash examples/distill/SFWan2.1-T2V/preprocess_data.sh
```
3. Launch self-forcing distillation with your cluster settings:
```bash
sbatch examples/distill/SFWan2.1-T2V/distill_dmd_t2v_1.3B.sh
```
Update the dataset paths and wandb credentials inside the script before running on your environment.
+1 -1
View File
@@ -12,7 +12,7 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
+1 -1
View File
@@ -14,7 +14,7 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=True,
# Adjust these offload parameters if you have < 32GB of VRAM
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"

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