Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
19c1d164c3 |
+58
-19
@@ -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
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -1,207 +0,0 @@
|
||||
name: Publish FastVideo Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "fastvideo-kernel/pyproject.toml"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd 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)
|
||||
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")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12']
|
||||
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'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
cd 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/
|
||||
|
||||
- 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
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Download PyPI wheels
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
path: fastvideo-kernel/dist/
|
||||
pattern: 'fastvideo_kernel-py*'
|
||||
merge-multiple: true
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
pip install build scikit-build-core cmake ninja
|
||||
|
||||
cd fastvideo-kernel
|
||||
# We don't need full CUDA/Torch to just package the source (sdist)
|
||||
python -m build --sdist --outdir dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: fastvideo-kernel/dist/
|
||||
@@ -30,8 +30,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
+6
-5
@@ -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
|
||||
|
||||
@@ -4,7 +4,7 @@ default_stages:
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/third_party/.*|
|
||||
fastvideo-kernel/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
|
||||
@@ -15,7 +15,7 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
|
||||
|
||||
## 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/).
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
@@ -125,7 +125,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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,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
|
||||
|
||||
@@ -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.
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_sta.py
|
||||
@@ -0,0 +1,96 @@
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python ../benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
|
||||
## 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.
|
||||
@@ -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'
|
||||
@@ -0,0 +1,76 @@
|
||||
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()}')
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
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"])
|
||||
@@ -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,49 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
try:
|
||||
from st_attn_cuda import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
def sliding_tile_attention(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]
|
||||
+347
-79
@@ -451,10 +451,6 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
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)};
|
||||
|
||||
// Shared memory size for the kernel.
|
||||
// We use the maximum available shared memory (kittens::MAX_SHARED_MEMORY)
|
||||
// which is approximately 227KB on H100, necessary for the high-performance
|
||||
// TMA-based attention tiles with multiple stages.
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
@@ -462,31 +458,104 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
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) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(2, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 3, 0); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(2, 2, 3); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 0, 5); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
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;
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
} else {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
|
||||
@@ -499,67 +568,266 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
if (kernel_aspect_ratio_flag == 2){
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
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);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(3, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 3, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 3, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 0, 0); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 0, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 3, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 0, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(0, 3, 0); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
} 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;
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
else if (kernel_aspect_ratio_flag == 3) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 3, 0); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(1, 2, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(1, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 2, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 2, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(0, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 0, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(0, 3, 2); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
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;
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
|
||||
else {
|
||||
TORCH_CHECK(false, "Unsupported kernel_aspect_ratio_flag: ", kernel_aspect_ratio_flag);
|
||||
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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']}")
|
||||
@@ -0,0 +1,156 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
# 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
|
||||
|
||||
def pytorch_test(Q, K, V, block_sparse_mask, dO):
|
||||
q_ = Q.clone().float().requires_grad_()
|
||||
k_ = K.clone().float().requires_grad_()
|
||||
v_ = V.clone().float().requires_grad_()
|
||||
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
QK = QK.masked_fill(~block_sparse_mask.unsqueeze(0), float('-inf'))
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
dO_ = dO
|
||||
output.backward(dO_)
|
||||
return (
|
||||
output.to(torch.bfloat16),
|
||||
q_.grad.to(torch.bfloat16),
|
||||
k_.grad.to(torch.bfloat16),
|
||||
v_.grad.to(torch.bfloat16),
|
||||
)
|
||||
|
||||
|
||||
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
|
||||
Q = Q.detach().requires_grad_()
|
||||
K = K.detach().requires_grad_()
|
||||
V = V.detach().requires_grad_()
|
||||
|
||||
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
|
||||
output = output[:, :, non_pad_index, :]
|
||||
output.backward(dO)
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def get_non_pad_index(
|
||||
vid_len: torch.LongTensor,
|
||||
n_win: int,
|
||||
win_size: int,
|
||||
):
|
||||
device = vid_len.device
|
||||
starts_pad = torch.arange(n_win, device=device) * win_size
|
||||
index_pad = starts_pad[:, None] + torch.arange(win_size, device=device)[None, :]
|
||||
index_mask = torch.arange(win_size, device=device)[None, :] < vid_len[:, None]
|
||||
|
||||
return index_pad[index_mask]
|
||||
|
||||
def generate_tensor(shape, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
return tensor
|
||||
|
||||
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
|
||||
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
|
||||
|
||||
|
||||
def vsa_pad(x, non_pad_index, num_blocks, block_size):
|
||||
padded_x = torch.zeros((1, x.shape[1], num_blocks * BLOCK_M, x.shape[3]), device=x.device, dtype=x.dtype)
|
||||
padded_x[:, :, non_pad_index, :] = x
|
||||
return padded_x
|
||||
|
||||
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
|
||||
results = {
|
||||
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
}
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
variable_block_sizes = generate_variable_block_sizes(num_blocks, device=device)
|
||||
S = int(variable_block_sizes.sum().item())
|
||||
padded_S = num_blocks * BLOCK_M
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
|
||||
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
|
||||
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
|
||||
for _ in range(num_iterations):
|
||||
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
|
||||
# dO_padded = torch.zeros_like(dO_padded)
|
||||
# dO_padded[:, :, non_pad_index, :] = dO
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
|
||||
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes,non_pad_index, dO)
|
||||
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
|
||||
if bs is not None:
|
||||
diff = pt - bs
|
||||
abs_diff = torch.abs(diff)
|
||||
results[name]['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
|
||||
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
total_elements = h * S * d * num_iterations
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_graphs(h, d, error_mode='all'):
|
||||
test_configs = [
|
||||
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
|
||||
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
|
||||
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
|
||||
]
|
||||
|
||||
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
|
||||
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
|
||||
f"{'gK Avg':<12} {'Rel gK Max':<12} "
|
||||
f"{'gV Avg':<12} {'Rel gV Max':<12} "
|
||||
f"{'gO Avg':<12} {'Rel gO Max':<12}")
|
||||
print("-" * 150)
|
||||
|
||||
for config in test_configs:
|
||||
num_blocks = config["num_blocks"]
|
||||
k = config["k"]
|
||||
description = config["description"]
|
||||
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
|
||||
print(f"{description:<20} {num_blocks:<8} {k:<4} "
|
||||
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
|
||||
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
|
||||
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
|
||||
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
|
||||
|
||||
print("-" * 150)
|
||||
|
||||
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)
|
||||
print("\nAnalysis completed for all modes.")
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch
|
||||
|
||||
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate block sparse mask of shape [h, num_blocks, num_blocks].
|
||||
|
||||
Args:
|
||||
h: number of heads
|
||||
num_blocks: number of blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
"""
|
||||
k = min(k, num_blocks)
|
||||
scores = torch.rand(h, num_blocks, num_blocks, device=device)
|
||||
_, indices = torch.topk(scores, k, dim=-1)
|
||||
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
|
||||
return block_sparse_mask
|
||||
|
||||
|
||||
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
|
||||
"""
|
||||
Convert block-level sparse mask to full attention mask.
|
||||
|
||||
Args:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
variable_block_sizes: [num_blocks] tensor
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
full_mask: [h, S, S] bool tensor where S = total sequence length
|
||||
"""
|
||||
h, num_blocks, _ = block_sparse_mask.shape
|
||||
total_seq_len = variable_block_sizes.sum().item()
|
||||
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
|
||||
|
||||
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
|
||||
|
||||
for head in range(h):
|
||||
for q_block in range(num_blocks):
|
||||
q_start = cumsum[q_block]
|
||||
q_end = q_start + variable_block_sizes[q_block]
|
||||
|
||||
for kv_block in range(num_blocks):
|
||||
if block_sparse_mask[head, q_block, kv_block]:
|
||||
kv_start = cumsum[kv_block]
|
||||
kv_end = kv_start + variable_block_sizes[kv_block]
|
||||
full_mask[head, q_start:q_end, kv_start:kv_end] = True
|
||||
|
||||
return full_mask
|
||||
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
@@ -0,0 +1,61 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python ../tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python ../benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -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'
|
||||
@@ -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"])
|
||||
Submodule
+1
Submodule csrc/attn/video_sparse_attn/tk added at 6c27e28c81
@@ -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
|
||||
|
||||
+172
-305
@@ -8,6 +8,7 @@ This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
import pytest
|
||||
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,28 @@ 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
|
||||
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
|
||||
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)
|
||||
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 +270,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 +320,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 +358,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 +421,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
|
||||
|
||||
|
||||
+99
-158
@@ -639,7 +639,8 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
|
||||
// store kq and vq
|
||||
|
||||
// ensuring all writes are finished
|
||||
// ! the following two line seems unnecessary.
|
||||
// tma::store_async_wait(); // ensure qg is finished
|
||||
__syncthreads();
|
||||
|
||||
warpgroup::store(kg_smem[0], kg_reg);
|
||||
@@ -660,145 +661,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
tma::store_async_wait();
|
||||
}
|
||||
|
||||
|
||||
template<int D>
|
||||
void block_sparse_attention_forward_impl(
|
||||
bf16* d_q, bf16* d_k, bf16* d_v, float* d_l, bf16* d_o,
|
||||
int batch, int qo_heads, int kv_heads, int seq_len, int hr,
|
||||
int max_kv_blocks_per_q,
|
||||
int32_t* q2k_block_sparse_index_ptr,
|
||||
int32_t* q2k_block_sparse_num_ptr,
|
||||
int32_t* block_size_ptr,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
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>;
|
||||
|
||||
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<D>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
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), static_cast<uint32_t>(D)};
|
||||
|
||||
globals g{
|
||||
qg_arg, kg_arg, vg_arg, lg_arg, og_arg,
|
||||
static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q),
|
||||
q2k_block_sparse_index_ptr, q2k_block_sparse_num_ptr, block_size_ptr
|
||||
};
|
||||
|
||||
// Shared memory size for the kernel
|
||||
// 54000 bytes is calibrated for H100 shared memory constraints for these tile sizes
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<D>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
|
||||
fwd_attend_ker<D><<<grid, (128), mem_size, stream>>>(g);
|
||||
}
|
||||
|
||||
template<int D>
|
||||
void block_sparse_attention_backward_impl(
|
||||
bf16* d_q, bf16* d_k, bf16* d_v, bf16* d_o, bf16* d_og, float* d_l, float* d_d, float* d_qg, float* d_kg, float* d_vg,
|
||||
int batch, int qo_heads, int kv_heads, int seq_len, int hr, int max_q_blocks_per_kv,
|
||||
int32_t* k2q_block_sparse_index_ptr,
|
||||
int32_t* k2q_block_sparse_num_ptr,
|
||||
int32_t* block_size_ptr,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
using G = bwd_attend_ker_tile_dims<D>;
|
||||
using og_tile = st_bf<4*16, D>;
|
||||
using o_tile = st_bf<4*16, D>;
|
||||
using d_tile = col_vec<st_fl<4*16, D>>;
|
||||
|
||||
using og_global = gl<bf16, -1, -1, -1, -1, og_tile>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
using d_global = gl<float, -1, -1, -1, -1, d_tile>;
|
||||
|
||||
using prep_globals = bwd_prep_globals<D>;
|
||||
|
||||
constexpr int mem_size_prep = kittens::MAX_SHARED_MEMORY;
|
||||
int threads_prep = PREP_NUM_WARPS * kittens::WARP_THREADS;
|
||||
dim3 grid_bwd_prep(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
bwd_attend_prep_ker<D>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size_prep
|
||||
);
|
||||
bwd_attend_prep_ker<D><<<grid_bwd_prep, threads_prep, mem_size_prep, stream>>>(bwd_g);
|
||||
|
||||
using bwd_q_tile = st_bf<G::tile_h_qo, G::tile_width>;
|
||||
using bwd_k_tile = st_bf<G::tile_h, G::tile_width>;
|
||||
using bwd_v_tile = st_bf<G::tile_h, G::tile_width>;
|
||||
using bwd_og_tile = st_bf<G::tile_h_qo, G::tile_width>;
|
||||
using bwd_qg_tile = st_fl<G::tile_h_qo, G::tile_width>;
|
||||
using bwd_kg_tile = st_fl<G::tile_h, G::tile_width>;
|
||||
using bwd_vg_tile = st_fl<G::tile_h, G::tile_width>;
|
||||
using bwd_l_tile = row_vec<st_fl<G::tile_h_qo, G::tile_h>>;
|
||||
using bwd_d_tile = row_vec<st_fl<G::tile_h_qo, G::tile_h>>;
|
||||
|
||||
using bwd_q_global = gl<bf16, -1, -1, -1, -1, bwd_q_tile>;
|
||||
using bwd_k_global = gl<bf16, -1, -1, -1, -1, bwd_k_tile>;
|
||||
using bwd_v_global = gl<bf16, -1, -1, -1, -1, bwd_v_tile>;
|
||||
using bwd_og_global = gl<bf16, -1, -1, -1, -1, bwd_og_tile>;
|
||||
using bwd_qg_global = gl<float, -1, -1, -1, -1, bwd_qg_tile>;
|
||||
using bwd_kg_global = gl<float, -1, -1, -1, -1, bwd_kg_tile>;
|
||||
using bwd_vg_global = gl<float, -1, -1, -1, -1, bwd_vg_tile>;
|
||||
using bwd_l_global = gl<float, -1, -1, -1, -1, bwd_l_tile>;
|
||||
using bwd_d_global = gl<float, -1, -1, -1, -1, bwd_d_tile>;
|
||||
|
||||
using bwd_global_args = bwd_globals<D>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), static_cast<uint32_t>(D)};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
|
||||
bwd_global_args bwd_global{bwd_q_arg, bwd_k_arg, bwd_v_arg, bwd_og_arg, bwd_qg_arg, bwd_kg_arg, bwd_vg_arg, bwd_l_arg, bwd_d_arg,
|
||||
static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_q_blocks_per_kv),
|
||||
k2q_block_sparse_index_ptr, k2q_block_sparse_num_ptr, block_size_ptr};
|
||||
|
||||
dim3 grid_bwd_main(seq_len/64, qo_heads, batch);
|
||||
int threads_main = 128;
|
||||
// Calibrated shared memory sizes for different head dimensions
|
||||
int bwd_mem_size = (D == 64) ? 72000 : 113000;
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
bwd_attend_ker<D>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
bwd_mem_size
|
||||
);
|
||||
bwd_attend_ker<D><<<grid_bwd_main, threads_main, bwd_mem_size, stream>>>(bwd_global);
|
||||
}
|
||||
|
||||
#include "pyutils/torch_helpers.cuh"
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <iostream>
|
||||
@@ -840,6 +702,7 @@ block_sparse_attention_forward(
|
||||
TORCH_CHECK(q2k_block_sparse_index.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_index idx 2 - must match seq_len / BLOCK_M");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_num idx 2 - must match seq_len / BLOCK_M");
|
||||
|
||||
|
||||
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");
|
||||
@@ -880,32 +743,110 @@ block_sparse_attention_forward(
|
||||
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();
|
||||
|
||||
// Temporated implementation to avoid code duplication between head_dim=64 and 128
|
||||
if (head_dim == 64) {
|
||||
block_sparse_attention_forward_impl<64>(
|
||||
d_q, d_k, d_v, d_l, d_o,
|
||||
batch, qo_heads, kv_heads, seq_len, hr,
|
||||
max_kv_blocks_per_q,
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<64>::kv_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<64>::kv_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::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<64>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
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), 64U};
|
||||
|
||||
globals g{
|
||||
qg_arg,
|
||||
kg_arg,
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr()),
|
||||
stream
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<64>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
} else if (head_dim == 128) {
|
||||
block_sparse_attention_forward_impl<128>(
|
||||
d_q, d_k, d_v, d_l, d_o,
|
||||
batch, qo_heads, kv_heads, seq_len, hr,
|
||||
max_kv_blocks_per_q,
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
|
||||
fwd_attend_ker<64><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(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>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr()),
|
||||
stream
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
} else {
|
||||
TORCH_CHECK(false, "Unsupported head_dim: ", head_dim, ". Only 64 and 128 are supported.");
|
||||
|
||||
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return {o, l_vec};
|
||||
@@ -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
-4
@@ -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.
|
||||
@@ -0,0 +1,32 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
|
||||
|
||||
### Installation
|
||||
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
|
||||
|
||||
### Usage
|
||||
|
||||
You can use `moba_attn_varlen` in the following ways:
|
||||
|
||||
**Install from source:**
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
**Import after installation:**
|
||||
```python
|
||||
from vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
**Or import directly from the project root:**
|
||||
```python
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
python csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
```
|
||||
|
||||
@@ -0,0 +1,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",
|
||||
]
|
||||
)
|
||||
+2
-2
@@ -3,7 +3,7 @@
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from fastvideo_kernel import moba_attn_varlen
|
||||
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"):
|
||||
"""
|
||||
@@ -51,7 +51,7 @@ def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, dev
|
||||
@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.float16, torch.bfloat16])
|
||||
@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
|
||||
):
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
|
||||
+248
-415
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,4 @@
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp310-cp310-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp311-cp311-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp312-cp312-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -55,11 +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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
FROM rocm/pytorch:rocm7.1_ubuntu22.04_py3.10_pytorch_release_2.9.1
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject_other.toml ./pyproject.toml
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[rocm] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Unified Kernel
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd fastvideo-kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
./build.sh --rocm
|
||||
|
||||
EXPOSE 22
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r requirements-mkdocs.txt
|
||||
pip install -r docs/requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -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 */
|
||||
|
||||
@@ -1,179 +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` (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`.
|
||||
@@ -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 --recursive 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.
|
||||
@@ -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}
|
||||
}
|
||||
```
|
||||
@@ -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}
|
||||
}
|
||||
```
|
||||
@@ -6,7 +6,7 @@ Thank you for your interest in contributing to FastVideo. We want to make the pr
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Profiling FastVideo
|
||||
|
||||
!!! warning
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down inference.
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
|
||||
|
||||
## Profiling with PyTorch
|
||||
|
||||
@@ -49,5 +49,5 @@ Traces can be visualized using <https://ui.perfetto.dev/>.
|
||||
### Best Practices
|
||||
|
||||
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
|
||||
- After profiling, clean up trace directories to avoid filling disk storage.
|
||||
- After profiling, clean up trace directories to avoid filling disks.
|
||||
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
|
||||
|
||||
@@ -74,11 +74,9 @@ To add a new SSIM test, follow these steps:
|
||||
generator.generate_video(prompt, ...)
|
||||
|
||||
# Compare with Reference
|
||||
ssim_values = compute_video_ssim_torchvision(
|
||||
reference_path, generated_path, use_ms_ssim=True
|
||||
)
|
||||
assert ssim_values[0] >= 0.98 # Threshold
|
||||
```
|
||||
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
|
||||
assert ssim_values[0] >= 0.98 # Threshold
|
||||
```
|
||||
|
||||
4. **Reference Videos**:
|
||||
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# 🎯 Distillation
|
||||
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computations, enabling much faster video generation.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
@@ -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](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
|
||||
@@ -7,11 +7,6 @@ Get up and running with FastVideo in minutes!
|
||||
First, install FastVideo:
|
||||
|
||||
```bash
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
@@ -20,53 +15,36 @@ pip install fastvideo
|
||||
### Text-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo import FastVideoPipeline
|
||||
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-t2v-1.3B")
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
# Generate a video
|
||||
prompt = "A cat playing with a ball of yarn"
|
||||
video = pipe(prompt, num_frames=16, height=512, width=512)
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
from fastvideo import FastVideoPipeline
|
||||
from PIL import Image
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
# Load an image
|
||||
image = Image.open("input.jpg")
|
||||
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-i2v-14B-480p")
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
# Generate a video from the image
|
||||
video = pipe(image, num_frames=16, height=480, width=480)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
```
|
||||
|
||||
## Next Steps
|
||||
@@ -75,4 +53,3 @@ if __name__ == '__main__':
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/) - Explore more examples
|
||||
- [Optimizations](../inference/optimizations.md) - Performance optimization tips
|
||||
- [Low VRAM Inference](../inference/low_vram_inference.md) - Memory-saving settings (CPU offload, sharded loading, etc.)
|
||||
|
||||
@@ -98,7 +98,7 @@ 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 CPU-offload and sharded loading flags (see [Low VRAM Inference](low_vram_inference.md))
|
||||
- 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
|
||||
|
||||
|
||||
@@ -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 implementation on H100.
|
||||
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
|
||||
```
|
||||
@@ -30,7 +30,7 @@ path_to_your_dataset_folder/
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
To generate the `videos2caption.json` and `merge.txt`, run
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
|
||||
# 🔧 Installation
|
||||
You can install the Video Sparse Attention package using
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We support H100 (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
|
||||
First, install C++20 for ThunderKittens (if using 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
|
||||
```
|
||||
@@ -6,7 +6,7 @@ The `VideoGenerator` class provides the primary Python interface for doing offli
|
||||
- Python 3.10-3.12
|
||||
|
||||
## Installation
|
||||
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation) first.
|
||||
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) first.
|
||||
|
||||
## Usage
|
||||
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
|
||||
|
||||
@@ -1,42 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,75 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.wan import MatrixGameI2V480PConfig
|
||||
from fastvideo.models.dits.matrix_game.utils import create_action_presets
|
||||
|
||||
import torch
|
||||
|
||||
# Available variants: "base_distilled_model", "gta_distilled_model", "templerun_distilled_model"
|
||||
# Each variant has different keyboard_dim:
|
||||
# - base_distilled_model: keyboard_dim=4
|
||||
# - gta_distilled_model: keyboard_dim=2
|
||||
# - templerun_distilled_model: keyboard_dim=7 (keyboard only, no mouse)
|
||||
MODEL_VARIANT = "base_distilled_model"
|
||||
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
config = VARIANT_CONFIG[MODEL_VARIANT]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
num_frames = 597
|
||||
actions = create_action_presets(num_frames, keyboard_dim=config["keyboard_dim"])
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
|
||||
generator.generate_video(
|
||||
prompt="",
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,107 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=4
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
mkdir -p ../profiler_traces/wan_t2v_finetune/
|
||||
|
||||
# Torch Profiler Configuration
|
||||
export FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_training_train_one_step"
|
||||
export FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_STACK=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WAIT_STEPS=2
|
||||
export FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS=1
|
||||
export FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS=1
|
||||
export FASTVIDEO_TORCH_PROFILER_DIR="../profiler_traces/wan_t2v_finetune/"
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_finetune"
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 8
|
||||
--num_latent_t 20
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port 29502 \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,159 +0,0 @@
|
||||
cmake_minimum_required(VERSION 3.26 FATAL_ERROR)
|
||||
project(fastvideo-kernel LANGUAGES CXX)
|
||||
|
||||
# Prefer environment variable (used by CI or pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
else()
|
||||
enable_language(CUDA)
|
||||
endif()
|
||||
|
||||
# Import common utils if needed, but we keep it simple for now
|
||||
|
||||
# Find Python and Torch
|
||||
find_package(Python COMPONENTS Interpreter Development.Module REQUIRED)
|
||||
|
||||
# Robustly find Torch include paths using Python
|
||||
execute_process(
|
||||
COMMAND "${Python_EXECUTABLE}" -c "import torch; from torch.utils.cpp_extension import include_paths; print(';'.join(include_paths()))"
|
||||
OUTPUT_VARIABLE TORCH_INCLUDE_PATHS
|
||||
OUTPUT_STRIP_TRAILING_WHITESPACE
|
||||
)
|
||||
list(APPEND TORCH_INCLUDE_DIRS ${TORCH_INCLUDE_PATHS})
|
||||
|
||||
# Find Torch package (still useful for libraries)
|
||||
find_package(Torch REQUIRED)
|
||||
|
||||
# Include directories
|
||||
include_directories(
|
||||
${CMAKE_SOURCE_DIR}/include
|
||||
${CMAKE_SOURCE_DIR}/include/cutlass/include
|
||||
${CMAKE_SOURCE_DIR}/include/tk/include
|
||||
${CMAKE_SOURCE_DIR}/include/tk/prototype
|
||||
${CMAKE_SOURCE_DIR}/csrc
|
||||
${CMAKE_SOURCE_DIR}/csrc/turbodiffusion
|
||||
${TORCH_INCLUDE_DIRS}
|
||||
)
|
||||
|
||||
# ---------------------------
|
||||
# ThunderKittens (TK) toggles
|
||||
# ---------------------------
|
||||
# AUTO: enable TK only when we can confidently target Hopper (sm_90a).
|
||||
# ON: force-enable TK kernels (intended for release wheels/images; does NOT require a GPU).
|
||||
# OFF: never build TK kernels.
|
||||
set(FASTVIDEO_KERNEL_BUILD_TK "AUTO" CACHE STRING "Build ThunderKittens kernels: AUTO/ON/OFF")
|
||||
set_property(CACHE FASTVIDEO_KERNEL_BUILD_TK PROPERTY STRINGS AUTO ON OFF)
|
||||
|
||||
# Prefer environment variable (used by CI) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED TORCH_CUDA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
|
||||
set(TORCH_CUDA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
|
||||
endif()
|
||||
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST (cmake/env): ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "FASTVIDEO_KERNEL_BUILD_TK: ${FASTVIDEO_KERNEL_BUILD_TK}")
|
||||
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
if(FASTVIDEO_KERNEL_BUILD_TK STREQUAL "ON")
|
||||
set(ENABLE_TK_KERNELS ON)
|
||||
elseif(FASTVIDEO_KERNEL_BUILD_TK STREQUAL "OFF")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
else()
|
||||
# AUTO: detect Hopper if possible.
|
||||
if(TORCH_CUDA_ARCH_LIST)
|
||||
# Accept common spellings: 9.0a, 90a, sm_90a.
|
||||
string(REGEX MATCH "(^|[; ,])((9\\.0a)|(90a)|(sm_90a))([; ,]|$)" _HAS_90A "${TORCH_CUDA_ARCH_LIST}")
|
||||
if(_HAS_90A)
|
||||
set(ENABLE_TK_KERNELS ON)
|
||||
endif()
|
||||
else()
|
||||
# Best-effort local detection (works when a CUDA device is visible).
|
||||
execute_process(
|
||||
COMMAND "${Python_EXECUTABLE}" -c "import torch; import sys; \nprint('1' if (torch.cuda.is_available() and torch.version.cuda and torch.cuda.get_device_capability()[0] >= 9) else '0')"
|
||||
OUTPUT_VARIABLE _LOCAL_HAS_HOPPER
|
||||
OUTPUT_STRIP_TRAILING_WHITESPACE
|
||||
ERROR_QUIET
|
||||
)
|
||||
if(_LOCAL_HAS_HOPPER STREQUAL "1")
|
||||
set(ENABLE_TK_KERNELS ON)
|
||||
endif()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
message(STATUS "ThunderKittens kernels: DISABLED (will use Triton fallbacks at runtime)")
|
||||
endif()
|
||||
|
||||
# Always try to build the extension if CUDA is available, but conditionally add sources/flags
|
||||
set(BUILD_CXX_KERNELS ON)
|
||||
|
||||
# Compiler flags
|
||||
set(CUDA_FLAGS
|
||||
"-DNDEBUG"
|
||||
"-O3"
|
||||
"-std=c++20"
|
||||
"--use_fast_math"
|
||||
"--expt-extended-lambda"
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
"-Xptxas=--warn-on-spills"
|
||||
)
|
||||
|
||||
# If TK is enabled, ensure we target Hopper. This is required even on GPU-less builders (CI).
|
||||
if(ENABLE_TK_KERNELS)
|
||||
if(NOT DEFINED CMAKE_CUDA_ARCHITECTURES OR CMAKE_CUDA_ARCHITECTURES STREQUAL "")
|
||||
set(CMAKE_CUDA_ARCHITECTURES "90a" CACHE STRING "CUDA architectures" FORCE)
|
||||
endif()
|
||||
list(APPEND CUDA_FLAGS "-DKITTENS_HOPPER")
|
||||
message(STATUS "CMAKE_CUDA_ARCHITECTURES: ${CMAKE_CUDA_ARCHITECTURES}")
|
||||
endif()
|
||||
|
||||
if(BUILD_CXX_KERNELS)
|
||||
# Source files
|
||||
set(EXTENSION_SOURCES
|
||||
csrc/common_extension.cpp
|
||||
csrc/turbodiffusion/gemm/gemm.cu
|
||||
csrc/turbodiffusion/norm/rmsnorm.cu
|
||||
csrc/turbodiffusion/norm/layernorm.cu
|
||||
csrc/turbodiffusion/quant/quant.cu
|
||||
)
|
||||
|
||||
# Conditionally add TK kernels
|
||||
if(ENABLE_TK_KERNELS)
|
||||
list(APPEND EXTENSION_SOURCES
|
||||
csrc/attention/st_attn_h100.cu
|
||||
csrc/attention/block_sparse_h100.cu
|
||||
)
|
||||
endif()
|
||||
|
||||
# Combined FastVideo Extension
|
||||
# Using name 'fastvideo_kernel_ops' to distinguish from the python package namespace
|
||||
Python_add_library(fastvideo_kernel_ops MODULE USE_SABI ${SKBUILD_SABI_VERSION} WITH_SOABI
|
||||
${EXTENSION_SOURCES}
|
||||
)
|
||||
|
||||
# Build compile definitions list
|
||||
set(COMPILE_DEFS TORCH_EXTENSION_NAME=fastvideo_kernel_ops)
|
||||
if(ENABLE_TK_KERNELS)
|
||||
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
|
||||
endif()
|
||||
|
||||
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
|
||||
|
||||
target_compile_options(fastvideo_kernel_ops PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${CUDA_FLAGS}>
|
||||
)
|
||||
|
||||
# We install it to fastvideo_kernel/_C so we can load it to register the ops
|
||||
install(TARGETS fastvideo_kernel_ops LIBRARY DESTINATION fastvideo_kernel/_C)
|
||||
endif()
|
||||
|
||||
@@ -1,187 +0,0 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
@@ -1,6 +0,0 @@
|
||||
include LICENSE
|
||||
include README.md
|
||||
include pyproject.toml
|
||||
recursive-include python/fastvideo_kernel *.py
|
||||
recursive-include csrc *.cu *.cuh *.cpp *.h
|
||||
recursive-include include/tk *.cu *.cuh *.cpp *.h *.src
|
||||
@@ -1,69 +0,0 @@
|
||||
# FastVideo Kernel
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
This will automatically detect your GPU architecture. If an NVIDIA Hopper (H100/sm_90a) GPU is detected, ThunderKittens kernels will be enabled. Otherwise, they will be skipped, and the package will use Triton fallbacks at runtime.
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
cd fastvideo-kernel
|
||||
./build.sh
|
||||
```
|
||||
|
||||
### Rocm Build
|
||||
If you are in a rocm environment without the compilation toolchaine of CUDA.
|
||||
|
||||
```bash
|
||||
cd fastvideo-kernel
|
||||
./build.sh --rocm
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
### Sliding Tile Attention (STA) & Video Sparse Attention (VSA)
|
||||
|
||||
For detailed usage, please check the [Attention Documentation](../docs/attention/index.md).
|
||||
|
||||
```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, ...)
|
||||
```
|
||||
|
||||
### TurboDiffusion Kernels
|
||||
|
||||
This package also includes kernels from [TurboDiffusion](https://github.com/thu-ml/TurboDiffusion), including INT8 GEMM, Quantization, RMSNorm and LayerNorm.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **Runtime**:
|
||||
- NVIDIA H100 (sm_90a) for C++ optimized kernels.
|
||||
- Any CUDA GPU for Triton-based fallbacks.
|
||||
- **Build**:
|
||||
- CUDA Toolkit 12.3+
|
||||
- C++20 compatible compiler (GCC 10+, Clang 11+)
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
This package structure and build system are based on [sgl-kernel](https://github.com/sgl-project/sglang/tree/main/sgl-kernel) from the SGLang project.
|
||||
|
||||
The implementation of `turbodiffusion` kernels is adapted from [TurboDiffusion](https://github.com/thu-ml/TurboDiffusion). If you use these kernels, please cite:
|
||||
|
||||
```bibtex
|
||||
@article{zhang2025turbodiffusion,
|
||||
title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
journal={arXiv preprint arXiv:2512.16093},
|
||||
year={2025}
|
||||
}
|
||||
```
|
||||
@@ -1,38 +0,0 @@
|
||||
#!/bin/bash
|
||||
set -ex
|
||||
|
||||
# Simple build script wrapping uv/pip
|
||||
# Usage:
|
||||
# ./build.sh # local dev build (auto-detect / skip TK kernels when not available)
|
||||
# ./build.sh --release # force-enable Hopper/TK kernels for release builds (no GPU required)
|
||||
|
||||
echo "Building fastvideo-kernel..."
|
||||
|
||||
# Ensure submodules are initialized if needed (tk)
|
||||
git submodule update --init --recursive
|
||||
|
||||
# Install build dependencies
|
||||
pip install scikit-build-core cmake ninja
|
||||
|
||||
RELEASE=0
|
||||
GPU_BACKEND=CUDA
|
||||
for arg in "$@"; do
|
||||
case "$arg" in
|
||||
--rocm)
|
||||
GPU_BACKEND=ROCM
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
# Force-enable ThunderKittens kernels and compile for Hopper.
|
||||
export TORCH_CUDA_ARCH_LIST="9.0a"
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
|
||||
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DGPU_BACKEND=${GPU_BACKEND}"
|
||||
|
||||
echo "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST:-<unset>}"
|
||||
echo "CMAKE_ARGS: ${CMAKE_ARGS:-<unset>}"
|
||||
echo "GPU_BACKEND: ${GPU_BACKEND:-<unset>}"
|
||||
# Build and install
|
||||
# Use -v for verbose output
|
||||
pip install . -v --no-build-isolation
|
||||
@@ -1,48 +0,0 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
#include <vector>
|
||||
|
||||
// Forward declarations
|
||||
#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
|
||||
|
||||
#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
|
||||
|
||||
// TurboDiffusion kernels
|
||||
void register_quant(pybind11::module_ &);
|
||||
void register_rms_norm(pybind11::module_ &);
|
||||
void register_layer_norm(pybind11::module_ &);
|
||||
void register_gemm(pybind11::module_ &);
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "FastVideo CUDA Kernels";
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention (Hopper)");
|
||||
#endif
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention forward (Hopper)");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward (Hopper)");
|
||||
#endif
|
||||
|
||||
// TurboDiffusion
|
||||
register_quant(m);
|
||||
register_rms_norm(m);
|
||||
register_layer_norm(m);
|
||||
register_gemm(m);
|
||||
}
|
||||
@@ -1,85 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
#include <torch/torch.h>
|
||||
#include <torch/extension.h>
|
||||
|
||||
CUTLASS_HOST_DEVICE int64_t cdiv(int64_t const& a, int64_t const &b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <class T>
|
||||
CUTLASS_HOST_DEVICE T max(T a, T b) { return a > b ? a : b; }
|
||||
|
||||
template <class T>
|
||||
CUTLASS_HOST_DEVICE T min(T a, T b) { return a > b ? b : a; }
|
||||
|
||||
#define MIN(a, b) ((a) > (b) ? (b) : (a))
|
||||
#define MAX(a, b) ((a) > (b) ? (a) : (b))
|
||||
|
||||
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
|
||||
[&] { \
|
||||
if (COND) { \
|
||||
static constexpr bool CONST_NAME = true; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} else { \
|
||||
static constexpr bool CONST_NAME = false; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define CUDA_CHECK(call) \
|
||||
{ \
|
||||
cudaError_t err = call; \
|
||||
if (err != cudaSuccess) { \
|
||||
fprintf(stderr, "CUDA Error at %s:%d: %s\n", __FILE__, __LINE__, cudaGetErrorString(err)); \
|
||||
exit(err); \
|
||||
} \
|
||||
}
|
||||
|
||||
#define CONFIG_SWITCH(N, ...) \
|
||||
[&] { \
|
||||
if (N <= 1024) { \
|
||||
constexpr int NUM_THR_PER_CTA = 128; \
|
||||
constexpr int MAX_HIDDEN_SIZE = 1024; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} else if (N <= 2048) { \
|
||||
constexpr int NUM_THR_PER_CTA = 128; \
|
||||
constexpr int MAX_HIDDEN_SIZE = 2048; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} else if (N <= 4096) { \
|
||||
constexpr int NUM_THR_PER_CTA = 128; \
|
||||
constexpr int MAX_HIDDEN_SIZE = 4096; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} else if (N <= 8192) { \
|
||||
constexpr int NUM_THR_PER_CTA = 256; \
|
||||
constexpr int MAX_HIDDEN_SIZE = 8192; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} else { \
|
||||
constexpr int NUM_THR_PER_CTA = 256; \
|
||||
constexpr int MAX_HIDDEN_SIZE = 16384; \
|
||||
return (__VA_ARGS__)(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
|
||||
template <int BlockSize>
|
||||
void create_tensor(
|
||||
torch::Device const &device,
|
||||
std::optional<at::Tensor> &output,
|
||||
std::optional<at::Tensor> &scale,
|
||||
int m, int n
|
||||
) {
|
||||
int num_block_m = cdiv(m, BlockSize);
|
||||
int num_block_n = cdiv(n, BlockSize);
|
||||
if (!output.has_value()) {
|
||||
output.emplace(torch::empty(
|
||||
{m, n},
|
||||
torch::TensorOptions().device(device).dtype(torch::kInt8)
|
||||
));
|
||||
scale.emplace(torch::empty(
|
||||
{num_block_m, num_block_n},
|
||||
torch::TensorOptions().device(device).dtype(torch::kFloat32)
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,44 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
|
||||
template <class Kernel>
|
||||
__global__ void device_kernel(
|
||||
__grid_constant__ typename Kernel::Params const params
|
||||
) {
|
||||
extern __shared__ char smem[];
|
||||
Kernel op;
|
||||
op(params, smem);
|
||||
}
|
||||
|
||||
template <class Kernel>
|
||||
__global__ __launch_bounds__(Kernel::MaxThreadsPerBlock, Kernel::MinBlocksPerMultiprocessor)
|
||||
void device_kernel_with_launch_bounds(
|
||||
__grid_constant__ typename Kernel::Params const params
|
||||
) {
|
||||
extern __shared__ char smem[];
|
||||
Kernel op;
|
||||
op(params, smem);
|
||||
}
|
||||
|
||||
template <class Kernel>
|
||||
void launch_kernel(
|
||||
typename Kernel::Params const ¶ms,
|
||||
dim3 grid_shape,
|
||||
dim3 cta_shape,
|
||||
size_t ShmSize,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
auto func = device_kernel<Kernel>;
|
||||
if (ShmSize >= 48 * 1024) {
|
||||
CUDA_CHECK(cudaFuncSetAttribute(
|
||||
func,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
ShmSize
|
||||
));
|
||||
}
|
||||
func<<<grid_shape, cta_shape, ShmSize, stream>>>(params);
|
||||
CUDA_CHECK(cudaGetLastError());
|
||||
}
|
||||
@@ -1,75 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
template <
|
||||
class InputDtype_,
|
||||
int TileM_,
|
||||
int TileN_,
|
||||
int NumThrPerCta_,
|
||||
bool IsEvenM,
|
||||
bool IsEvenN
|
||||
>
|
||||
class Loader {
|
||||
public:
|
||||
using InputDtype = InputDtype_;
|
||||
static constexpr int TileM = TileM_;
|
||||
static constexpr int TileN = TileN_;
|
||||
static constexpr int NumThrPerCta = NumThrPerCta_;
|
||||
static constexpr int NumElementPerThread = TileM * TileN / NumThrPerCta;
|
||||
static constexpr int NumThrPerRow = TileN / NumElementPerThread;
|
||||
|
||||
static_assert(NumThrPerCta % TileM == 0);
|
||||
static_assert(TileM * TileN % NumThrPerCta == 0);
|
||||
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
load(void const *input_ptr, void *thr_output_reg, int64_t m, int64_t n, int blk_m, int blk_n, int tid) {
|
||||
int n_alignment = (n & 31) * sizeof(InputDtype);
|
||||
int thr_m_offset = tid / NumThrPerRow;
|
||||
int thr_n_offset = (tid % NumThrPerRow) * NumElementPerThread;
|
||||
void const *cta_input_ptr = (void*)((InputDtype*)input_ptr + blk_m * TileM * n + blk_n * TileN);
|
||||
void const *thr_input_ptr = (void*)((InputDtype*)cta_input_ptr + thr_m_offset * n + thr_n_offset);
|
||||
InputDtype tmp_reg[NumElementPerThread];
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
tmp_reg[i] = InputDtype(0.f);
|
||||
bool pred = IsEvenM ? true : thr_m_offset + blk_m * TileM < m;
|
||||
int limit = IsEvenN ? NumElementPerThread : MIN(NumElementPerThread, n - (blk_n * TileN + thr_n_offset));
|
||||
if (n_alignment % 128 == 0)
|
||||
_load<int4, IsEvenN>(thr_input_ptr, (void*)tmp_reg, limit, pred);
|
||||
else if (n_alignment % 64 == 0)
|
||||
_load<int2, IsEvenN>(thr_input_ptr, (void*)tmp_reg, limit, pred);
|
||||
else
|
||||
_load<InputDtype, IsEvenN>(thr_input_ptr, (void*)tmp_reg, limit, pred);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
*((float*)thr_output_reg + i) = static_cast<float>(reinterpret_cast<InputDtype const&>(tmp_reg[i]));
|
||||
}
|
||||
|
||||
private:
|
||||
template <class LoadDataType, bool IsEven>
|
||||
CUTLASS_DEVICE void
|
||||
_load(void const *thr_input_ptr, void *thr_output_reg, int limit, bool pred) {
|
||||
static constexpr int NumElementPerLoad = sizeof(LoadDataType) / sizeof(InputDtype);
|
||||
if (pred) {
|
||||
if constexpr (IsEven) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; i += NumElementPerLoad) {
|
||||
*(LoadDataType*)((InputDtype*)thr_output_reg + i) = *(LoadDataType*)((InputDtype*)thr_input_ptr + i);
|
||||
}
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < limit; i += NumElementPerLoad) {
|
||||
if (limit - i > NumElementPerLoad)
|
||||
*(LoadDataType*)((InputDtype*)thr_output_reg + i) = *(LoadDataType*)((InputDtype*)thr_input_ptr + i);
|
||||
else {
|
||||
for (int j = 0; j < NumElementPerLoad; ++j) {
|
||||
if (i + j < limit)
|
||||
*((InputDtype*)thr_output_reg + i + j) = *((InputDtype*)thr_input_ptr + i + j);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -1,77 +0,0 @@
|
||||
#pragma once
|
||||
#include "common/common.hpp"
|
||||
|
||||
|
||||
template <
|
||||
class OutputDtype_,
|
||||
int TileM_,
|
||||
int TileN_,
|
||||
int NumThrPerCta_,
|
||||
bool IsEvenM,
|
||||
bool IsEvenN,
|
||||
bool Round = true,
|
||||
bool SaveScale = true
|
||||
>
|
||||
class Saver {
|
||||
public:
|
||||
using OutputDtype = OutputDtype_;
|
||||
|
||||
static constexpr int TileM = TileM_;
|
||||
static constexpr int TileN = TileN_;
|
||||
static constexpr int NumThrPerCta = NumThrPerCta_;
|
||||
static constexpr int NumElementPerThread = TileM * TileN / NumThrPerCta;
|
||||
static constexpr int NumThrPerRow = TileN / NumElementPerThread;
|
||||
|
||||
static_assert(TileM * TileN % NumThrPerCta == 0);
|
||||
static_assert(NumThrPerCta % TileM == 0);
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
store(void *Optr, void *OSptr, void *reg, float scale_inv, int64_t m, int64_t n, int blk_m, int blk_n, int tid) {
|
||||
int n_alignment = (n & 31) * sizeof(OutputDtype);
|
||||
int thr_m_offset = tid / NumThrPerRow;
|
||||
int thr_n_offset = (tid % NumThrPerRow) * NumElementPerThread;
|
||||
void *cta_output_ptr = (void*)((OutputDtype*)Optr + blk_m * TileM * (Round ? cdiv(n, TileN) * TileN : n) + blk_n * TileN);
|
||||
void *thr_output_ptr = (void*)((OutputDtype*)cta_output_ptr + thr_m_offset * (Round ? cdiv(n, TileN) * TileN : n) + thr_n_offset);
|
||||
bool pred = IsEvenM ? true : thr_m_offset + blk_m * TileM < m;
|
||||
int limit = IsEvenN ? NumElementPerThread : MIN(NumElementPerThread, n - (blk_n * TileN + thr_n_offset));
|
||||
if (n_alignment % 128 == 0)
|
||||
_store<int4, IsEvenN>(thr_output_ptr, reg, limit, pred);
|
||||
else if (n_alignment % 64 == 0)
|
||||
_store<int2, IsEvenN>(thr_output_ptr, reg, limit, pred);
|
||||
else
|
||||
_store<OutputDtype, IsEvenN>(thr_output_ptr, reg, limit, pred);
|
||||
|
||||
if constexpr (SaveScale) {
|
||||
if (tid == 0) {
|
||||
*((float*)OSptr + blk_m * cdiv(n, TileN)+ blk_n) = scale_inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
template <class StoreDataType, bool IsEven>
|
||||
CUTLASS_DEVICE void
|
||||
_store(void *thr_output_ptr, void *reg, int limit, bool pred) {
|
||||
static constexpr int NumElementPerStore = sizeof(StoreDataType) / sizeof(OutputDtype);
|
||||
if (pred) {
|
||||
if constexpr (IsEven) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; i += NumElementPerStore) {
|
||||
*(StoreDataType*)((OutputDtype*)thr_output_ptr + i) = *(StoreDataType*)((OutputDtype*)reg + i);
|
||||
}
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < limit; i += NumElementPerStore) {
|
||||
if (limit - i > NumElementPerStore)
|
||||
*(StoreDataType*)((OutputDtype*)thr_output_ptr + i) = *(StoreDataType*)((OutputDtype*)reg + i);
|
||||
else {
|
||||
for (int j = 0; j < limit - i; ++j) {
|
||||
*((OutputDtype*)thr_output_ptr + i + j) = *((OutputDtype*)reg + i + j);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
@@ -1,77 +0,0 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by TurboDiffusion team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
*
|
||||
* Citation (please cite if you use this code):
|
||||
*
|
||||
* @article{zhang2025turbodiffusion,
|
||||
* title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
* author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
* journal={arXiv preprint arXiv:2512.16093},
|
||||
* year={2025}
|
||||
* }
|
||||
*/
|
||||
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
#include "common/common.hpp"
|
||||
#include "gemm/launch.hpp"
|
||||
|
||||
void int8_gemm(
|
||||
at::Tensor const& A, at::Tensor const& A_S,
|
||||
at::Tensor const& B, at::Tensor const& B_S,
|
||||
torch::Tensor& C
|
||||
) {
|
||||
|
||||
|
||||
static constexpr int swizzle_dir = 1;
|
||||
static constexpr int swizzle_size_log = 5;
|
||||
|
||||
int k = B.size(1);
|
||||
int m = A.size(0);
|
||||
int n = B.size(0);
|
||||
|
||||
switch (C.scalar_type()) {
|
||||
case torch::kHalf:{
|
||||
int8_gemm_<cutlass::half_t> (
|
||||
(int8_t*)A.data_ptr(), A_S.data_ptr<float>(),
|
||||
(int8_t*)B.data_ptr(), B_S.data_ptr<float>(),
|
||||
(cutlass::half_t*)C.data_ptr(),
|
||||
m, n, k, swizzle_dir, swizzle_size_log, at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
case torch::kBFloat16:{
|
||||
int8_gemm_<cutlass::bfloat16_t> (
|
||||
(int8_t*)A.data_ptr(), A_S.data_ptr<float>(),
|
||||
(int8_t*)B.data_ptr(), B_S.data_ptr<float>(),
|
||||
(cutlass::bfloat16_t*)C.data_ptr(),
|
||||
m, n, k, swizzle_dir, swizzle_size_log, at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
default: {
|
||||
std::cerr << "Observing: " << C.scalar_type() << " for the output datatype which is invalid";
|
||||
throw std::runtime_error("Unsupported output data type for int8 gemm.");
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
void register_gemm(pybind11::module_ &m) {
|
||||
m.def("gemm_cuda", &int8_gemm);
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,522 +0,0 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by TurboDiffusion team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
*
|
||||
* Citation (please cite if you use this code):
|
||||
*
|
||||
* @article{zhang2025turbodiffusion,
|
||||
* title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
* author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
* journal={arXiv preprint arXiv:2512.16093},
|
||||
* year={2025}
|
||||
* }
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "common/common.hpp"
|
||||
#include "gemm/utils.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <
|
||||
class OutputDtype_,
|
||||
bool IsEvenM,
|
||||
bool IsEvenN
|
||||
>
|
||||
struct GemmKernel {
|
||||
using ElementA = int8_t;
|
||||
using ElementB = int8_t;
|
||||
using OutputDtype = OutputDtype_;
|
||||
using AccumulatorDtype = int32_t;
|
||||
static constexpr int BlockSize = 128;
|
||||
static constexpr int TileM = 128;
|
||||
static constexpr int TileN = 128;
|
||||
static constexpr int TileK = 128;
|
||||
static constexpr int Stage = 3;
|
||||
static constexpr int EpiStage = 2;
|
||||
|
||||
static_assert(
|
||||
BlockSize % TileM == 0
|
||||
&& BlockSize % TileN == 0
|
||||
&& BlockSize % TileK == 0
|
||||
);
|
||||
|
||||
static constexpr int NumTilePerBlock = BlockSize / TileK;
|
||||
|
||||
using SmemLayoutAtom = decltype(
|
||||
composition(
|
||||
Swizzle<3, 4, 3>{},
|
||||
make_layout(
|
||||
make_shape(Int<8>{}, Int<TileK>{}),
|
||||
make_stride(Int<TileK>{}, Int<1>{})
|
||||
)
|
||||
)
|
||||
);
|
||||
|
||||
using SmemLayoutA = decltype(
|
||||
tile_to_shape(
|
||||
SmemLayoutAtom{},
|
||||
make_shape(Int<TileM>{}, Int<TileK>{}, Int<Stage>{})
|
||||
)
|
||||
);
|
||||
|
||||
using SmemLayoutB = decltype(
|
||||
tile_to_shape(
|
||||
SmemLayoutAtom{},
|
||||
make_shape(Int<TileN>{}, Int<TileK>{}, Int<Stage>{})
|
||||
)
|
||||
);
|
||||
|
||||
using MmaOP = cute::SM80_16x8x32_S32S8S8S32_TN;
|
||||
using TiledMma = decltype(
|
||||
make_tiled_mma(
|
||||
MMA_Atom<MMA_Traits<MmaOP>>{},
|
||||
make_layout(make_shape(
|
||||
_4{}, _2{}, _1{}
|
||||
)),
|
||||
make_tile(Int<64>{}, Int<32>{}, Int<32>{})
|
||||
)
|
||||
);
|
||||
|
||||
using G2SCopyAtomA = Copy_Atom<Copy_Traits<SM80_CP_ASYNC_CACHEGLOBAL<cute::uint128_t>>, ElementA>;
|
||||
using G2SCopyAtomB = Copy_Atom<Copy_Traits<SM80_CP_ASYNC_CACHEGLOBAL<cute::uint128_t>>, ElementB>;
|
||||
using G2STiledCopyA = decltype(
|
||||
make_tiled_copy(
|
||||
G2SCopyAtomA{},
|
||||
make_layout(
|
||||
make_shape(Int<64>{}, Int<4>{}),
|
||||
make_stride(Int<4>{}, Int<1>{})
|
||||
),
|
||||
make_layout(make_shape(Int<1>{}, Int<16>{}))
|
||||
)
|
||||
);
|
||||
using G2STiledCopyB = decltype(
|
||||
make_tiled_copy(
|
||||
G2SCopyAtomB{},
|
||||
make_layout(
|
||||
make_shape(Int<64>{}, Int<4>{}),
|
||||
make_stride(Int<4>{}, Int<1>{})
|
||||
),
|
||||
make_layout(make_shape(Int<1>{}, Int<16>{}))
|
||||
)
|
||||
);
|
||||
|
||||
using S2RCopyAtomA = Copy_Atom<Copy_Traits<SM75_U32x4_LDSM_N>, ElementA>;
|
||||
using S2RCopyAtomB = Copy_Atom<Copy_Traits<SM75_U32x4_LDSM_N>, ElementB>;
|
||||
using S2RTiledCopyA = decltype(make_tiled_copy_A(S2RCopyAtomA{}, TiledMma{}));
|
||||
using S2RTiledCopyB = decltype(make_tiled_copy_B(S2RCopyAtomB{}, TiledMma{}));
|
||||
|
||||
// epilogue
|
||||
using SmemLayoutAtomD = decltype(
|
||||
composition(
|
||||
Swizzle<2, 3, 3>{},
|
||||
make_layout(
|
||||
make_shape(Int<32>{}, Int<32>{}),
|
||||
LayoutRight{}
|
||||
)
|
||||
)
|
||||
);
|
||||
|
||||
using SmemLayoutD = decltype(
|
||||
tile_to_shape(
|
||||
SmemLayoutAtomD{},
|
||||
make_shape(Int<64>{}, Int<32>{}, Int<EpiStage>{})
|
||||
)
|
||||
);
|
||||
|
||||
using R2SCopyAtomD = Copy_Atom<UniversalCopy<std::conditional_t<sizeof(OutputDtype) == 4, int32_t, int16_t>>, OutputDtype>;
|
||||
using R2STiledCopyD = decltype(make_tiled_copy_C(R2SCopyAtomD{}, TiledMma{}));
|
||||
|
||||
using S2GCopyAtomD = Copy_Atom<UniversalCopy<uint128_t>, OutputDtype>;
|
||||
using S2GCopyD = decltype(make_tiled_copy(
|
||||
S2GCopyAtomD{},
|
||||
make_layout(Shape<_64, _4>{}),
|
||||
make_layout(Shape<_1, _8>{})
|
||||
));
|
||||
|
||||
using TileShape = decltype(make_shape(Int<TileM>{}, Int<TileN>{}, Int<TileK>{}));
|
||||
|
||||
struct SharedStorageAB: cute::aligned_struct<128> {
|
||||
array_aligned<typename TiledMma::ValTypeA, cosize_v<SmemLayoutA>, 128> smem_A;
|
||||
array_aligned<typename TiledMma::ValTypeB, cosize_v<SmemLayoutB>, 128> smem_B;
|
||||
array_aligned<float, 1> smem_AS;
|
||||
array_aligned<float, 1> smem_BS;
|
||||
array_aligned<int32_t, 1> smem_AF;
|
||||
};
|
||||
|
||||
struct SharedStorageD: cute::aligned_struct<128> {
|
||||
array_aligned<OutputDtype, cosize_v<SmemLayoutD>> smem_D;
|
||||
};
|
||||
|
||||
union SharedStorage {
|
||||
SharedStorageAB storage_AB;
|
||||
SharedStorageD storage_D;
|
||||
};
|
||||
|
||||
|
||||
struct Params {
|
||||
void const* Aptr;
|
||||
void const* ASptr;
|
||||
void const* Bptr;
|
||||
void const* BSptr;
|
||||
void* Dptr;
|
||||
int64_t const m;
|
||||
int64_t const n;
|
||||
int64_t const k;
|
||||
int const swizzle_dir;
|
||||
int const swizzle_size;
|
||||
};
|
||||
|
||||
using Arguments = Params;
|
||||
|
||||
static constexpr int ThreadNum = size(TiledMma{});
|
||||
static constexpr int ShmSize = sizeof(SharedStorage);
|
||||
static constexpr bool FastInt2Float = false;
|
||||
|
||||
static bool can_implement(int64_t m, int64_t n, int64_t k) {
|
||||
if (k % BlockSize != 0) return false;
|
||||
if ((n * sizeof(OutputDtype)) % 16 != 0)
|
||||
return false;
|
||||
return true;
|
||||
}
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args) {
|
||||
return args;
|
||||
}
|
||||
|
||||
static dim3 get_grid_size(int64_t m, int64_t n) {
|
||||
return dim3(cdiv(m, TileM) * cdiv(n, TileN));
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static auto get_block_coord(
|
||||
int64_t m_blocks,
|
||||
int64_t n_blocks,
|
||||
int const swizzle_dir,
|
||||
int64_t const swizzle_size_log
|
||||
) {
|
||||
int64_t blk_m;
|
||||
int64_t blk_n;
|
||||
|
||||
if (swizzle_dir == 1)
|
||||
std::swap(m_blocks, n_blocks);
|
||||
|
||||
if (swizzle_size_log == 0) {
|
||||
blk_m = blockIdx.x % m_blocks;
|
||||
blk_n = blockIdx.x / m_blocks;
|
||||
} else {
|
||||
int64_t group_size = n_blocks << swizzle_size_log;
|
||||
int64_t num_groups = m_blocks >> swizzle_size_log;
|
||||
int64_t group_idx = blockIdx.x / group_size;
|
||||
int64_t local_idx = blockIdx.x % group_size;
|
||||
if (group_idx == num_groups) {
|
||||
blk_m = (num_groups << swizzle_size_log) + local_idx % (m_blocks - (num_groups << swizzle_size_log));
|
||||
blk_n = local_idx / (m_blocks - (num_groups << swizzle_size_log));
|
||||
} else {
|
||||
blk_m = (local_idx & ((1LL << swizzle_size_log) - 1)) + (group_idx << swizzle_size_log);
|
||||
blk_n = local_idx >> swizzle_size_log;
|
||||
}
|
||||
}
|
||||
|
||||
if (swizzle_dir == 1)
|
||||
std::swap(blk_m, blk_n);
|
||||
|
||||
return make_coord(blk_m, blk_n);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
Params const& params, char* smem_data
|
||||
) {
|
||||
|
||||
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_data);
|
||||
|
||||
auto t_idx = threadIdx.x;
|
||||
|
||||
int64_t const m = params.m;
|
||||
int64_t const n = params.n;
|
||||
int64_t const k = params.k;
|
||||
int const swizzle_dir = params.swizzle_dir;
|
||||
int const swizzle_size = params.swizzle_size;
|
||||
|
||||
Tensor A = make_tensor(
|
||||
make_gmem_ptr<ElementA>(params.Aptr),
|
||||
make_shape(m, k),
|
||||
make_stride(k, _1{})
|
||||
);
|
||||
Tensor B = make_tensor(
|
||||
make_gmem_ptr<ElementB>(params.Bptr),
|
||||
make_shape(m, k),
|
||||
make_stride(k, _1{})
|
||||
);
|
||||
Tensor AS = make_tensor(
|
||||
make_gmem_ptr<float>(params.ASptr),
|
||||
make_shape(cdiv(m, BlockSize), cdiv(k, BlockSize)),
|
||||
make_stride(cdiv(k, BlockSize), _1{})
|
||||
);
|
||||
Tensor BS = make_tensor(
|
||||
make_gmem_ptr<float>(params.BSptr),
|
||||
make_shape(cdiv(n, BlockSize), cdiv(k, BlockSize)),
|
||||
make_stride(cdiv(k, BlockSize), _1{})
|
||||
);
|
||||
Tensor D = make_tensor(
|
||||
make_gmem_ptr<OutputDtype>(params.Dptr),
|
||||
make_shape(m, n),
|
||||
LayoutRight{}
|
||||
);
|
||||
|
||||
auto [m_coord, n_coord] = get_block_coord(
|
||||
cdiv(m, size<0>(TileShape{})),
|
||||
cdiv(n, size<1>(TileShape{})),
|
||||
swizzle_dir, swizzle_size
|
||||
);
|
||||
|
||||
int32_t blk_m_coord = m_coord / (BlockSize / TileM);
|
||||
int32_t blk_n_coord = n_coord / (BlockSize / TileN);
|
||||
|
||||
// local tile
|
||||
auto gA = local_tile(A, TileShape{}, make_coord(m_coord, n_coord, _), Step<_1, X, _1>{});
|
||||
auto gB = local_tile(B, TileShape{}, make_coord(m_coord, n_coord, _), Step<X, _1, _1>{});
|
||||
auto gD = local_tile(D, TileShape{}, make_coord(m_coord, n_coord, _), Step<_1, _1, X>{});
|
||||
|
||||
// shared memory
|
||||
Tensor sA = make_tensor(
|
||||
make_smem_ptr<ElementA>(shared_storage.storage_AB.smem_A.data()),
|
||||
SmemLayoutA{}
|
||||
);
|
||||
|
||||
Tensor sB = make_tensor(
|
||||
make_smem_ptr<ElementB>(shared_storage.storage_AB.smem_B.data()),
|
||||
SmemLayoutB{}
|
||||
);
|
||||
|
||||
// register
|
||||
TiledMma tiled_mma;
|
||||
auto thr_mma = tiled_mma.get_slice(t_idx);
|
||||
auto tCrA = thr_mma.partition_fragment_A(gA(_, _, 0));
|
||||
auto tCrB = thr_mma.partition_fragment_B(gB(_, _, 0));
|
||||
auto tDrC = thr_mma.partition_fragment_C(gD); // mma accumulator
|
||||
auto tDrD = make_tensor_like<float>(tDrC); // float accumulator
|
||||
|
||||
if constexpr (FastInt2Float) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(tDrC); ++i)
|
||||
tDrC(i) = 0x4B400000;
|
||||
} else {
|
||||
clear(tDrC);
|
||||
}
|
||||
|
||||
clear(tDrD);
|
||||
|
||||
|
||||
// global to shared copy
|
||||
G2STiledCopyA g2s_tiled_copy_a;
|
||||
auto g2s_thr_copy_a = g2s_tiled_copy_a.get_slice(t_idx);
|
||||
auto tAgA = g2s_thr_copy_a.partition_S(gA);
|
||||
auto tAsA = g2s_thr_copy_a.partition_D(sA);
|
||||
auto cA = make_identity_tensor(make_shape(size<0>(sA), size<1>(sA)));
|
||||
auto tAcA = g2s_thr_copy_a.partition_S(cA);
|
||||
int const m_limit = m - TileM * m_coord;
|
||||
int const n_limit = n - TileN * n_coord;
|
||||
|
||||
G2STiledCopyB g2s_tiled_copy_b;
|
||||
auto g2s_thr_copy_b = g2s_tiled_copy_b.get_slice(t_idx);
|
||||
auto tBgB = g2s_thr_copy_b.partition_S(gB);
|
||||
auto tBsB = g2s_thr_copy_b.partition_D(sB);
|
||||
auto cB = make_identity_tensor(make_shape(size<0>(sB), size<1>(sB)));
|
||||
auto tBcB = g2s_thr_copy_a.partition_S(cB);
|
||||
|
||||
|
||||
// shared to register copy
|
||||
S2RTiledCopyA s2r_tiled_copy_a;
|
||||
auto s2r_thr_copy_a = s2r_tiled_copy_a.get_slice(t_idx);
|
||||
auto tCsA = s2r_thr_copy_a.partition_S(sA);
|
||||
auto tCrA_view = s2r_thr_copy_a.retile_D(tCrA);
|
||||
|
||||
S2RTiledCopyB s2r_tiled_copy_b;
|
||||
auto s2r_thr_copy_b = s2r_tiled_copy_b.get_slice(t_idx);
|
||||
auto tCsB = s2r_thr_copy_b.partition_S(sB);
|
||||
auto tCrB_view = s2r_thr_copy_b.retile_D(tCrB);
|
||||
|
||||
// pipeline status
|
||||
int64_t g2s_a_tile = 0;
|
||||
int64_t g2s_b_tile = 0;
|
||||
int g2s_a_smem = 0;
|
||||
int g2s_b_smem = 0;
|
||||
|
||||
int g2s_tile_in_block = 0;
|
||||
int g2s_block = 0; // b block idx
|
||||
|
||||
int s2r_a_smem = 0;
|
||||
int s2r_b_smem = 0;
|
||||
int s2r_tile_in_block = 0;
|
||||
|
||||
int mma_block_a = 0;
|
||||
int mma_block_b = 0;
|
||||
|
||||
int ntile = k / TileK;
|
||||
// load scale and fallback
|
||||
// we assume all ptrs are 128bit aligned
|
||||
// auto smem_fallback_A = raw_pointer_cast(make_smem_ptr<int32_t>(shared_storage.storage_AB.smem_AF.data()));
|
||||
// auto smem_scale_A = raw_pointer_cast(make_smem_ptr<float>(shared_storage.storage_AB.smem_AS.data()));
|
||||
// auto smem_scale_B = raw_pointer_cast(make_smem_ptr<float>(shared_storage.storage_AB.smem_BS.data()));
|
||||
__syncthreads();
|
||||
|
||||
|
||||
int32_t fallbackA_load = 0;
|
||||
int32_t fallbackA_mma = 0;
|
||||
|
||||
// copy first Stage - 1 tile
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0, _i = min(Stage - 1, ntile); i < _i; ++i) {
|
||||
if (g2s_b_tile < ntile) {
|
||||
g2s_tile_in_block = (g2s_tile_in_block + 1) % NumTilePerBlock;
|
||||
copy_AB<IsEvenM>(g2s_tiled_copy_a, tAgA, tAsA, tAcA, g2s_a_tile, g2s_a_smem, m_limit);
|
||||
copy_AB<IsEvenN>(g2s_tiled_copy_b, tBgB, tBsB, tBcB, g2s_b_tile, g2s_b_smem, n_limit);
|
||||
++g2s_b_tile;
|
||||
++g2s_b_smem;
|
||||
++g2s_block;
|
||||
g2s_a_tile = g2s_block * NumTilePerBlock;
|
||||
++g2s_a_smem;
|
||||
}
|
||||
cp_async_fence();
|
||||
}
|
||||
|
||||
constexpr int nk = size<2>(tCrA);
|
||||
float scale_a = AS(blk_m_coord, 0);
|
||||
float scale_b = BS(blk_n_coord, 0);
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (int64_t mma_b_tile = 0; mma_b_tile < ntile; ++mma_b_tile) {
|
||||
s2r_tile_in_block = (s2r_tile_in_block + 1) % NumTilePerBlock;
|
||||
cp_async_wait<Stage - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// do mma first
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ik = 0; ik < nk; ++ik) {
|
||||
cute::copy(s2r_tiled_copy_a, tCsA(_, _, ik, s2r_a_smem),
|
||||
tCrA_view(_, _, ik));
|
||||
cute::copy(s2r_tiled_copy_b, tCsB(_, _, ik, s2r_b_smem),
|
||||
tCrB_view(_, _, ik));
|
||||
cute::gemm(tiled_mma, tDrC, tCrA(_, _, ik), tCrB(_, _, ik), tDrC);
|
||||
}
|
||||
|
||||
// a s2r increase anyway
|
||||
s2r_a_smem = (s2r_a_smem + 1) % Stage;
|
||||
|
||||
// get next s2r b tile int64_t
|
||||
// end of a block
|
||||
|
||||
// dequant first
|
||||
dequant<AccumulatorDtype, TileM * TileN / ThreadNum, FastInt2Float>(
|
||||
tDrC.data(), tDrD.data(), scale_a * scale_b
|
||||
);
|
||||
|
||||
s2r_b_smem = (s2r_b_smem + 1) % Stage;
|
||||
// b advance
|
||||
++mma_block_b;
|
||||
if (mma_block_b < size<1>(BS)) scale_b = BS(blk_n_coord, mma_block_b);
|
||||
mma_block_a = mma_block_b;
|
||||
if (mma_block_a < size<1>(AS)) scale_a = AS(blk_m_coord, mma_block_a);
|
||||
|
||||
// load next stage
|
||||
if (g2s_b_tile < ntile) {
|
||||
g2s_tile_in_block = (g2s_tile_in_block + 1) % NumTilePerBlock;
|
||||
copy_AB<IsEvenM>(g2s_tiled_copy_a, tAgA, tAsA, tAcA, g2s_a_tile, g2s_a_smem, m_limit);
|
||||
copy_AB<IsEvenN>(g2s_tiled_copy_b, tBgB, tBsB, tBcB, g2s_b_tile, g2s_b_smem, n_limit);
|
||||
++g2s_b_tile;
|
||||
g2s_b_smem = (g2s_b_smem + 1) % Stage;
|
||||
++g2s_block;
|
||||
g2s_a_tile = g2s_block * NumTilePerBlock;
|
||||
g2s_a_smem = (g2s_a_smem + 1) % Stage;
|
||||
}
|
||||
cp_async_fence();
|
||||
}
|
||||
|
||||
|
||||
// epilogue
|
||||
|
||||
Tensor sD = make_tensor(
|
||||
make_smem_ptr<OutputDtype>(shared_storage.storage_D.smem_D.data()),
|
||||
SmemLayoutD{}
|
||||
);
|
||||
|
||||
R2STiledCopyD r2s_tiled_copy_d;
|
||||
auto r2s_thr_copy_d = r2s_tiled_copy_d.get_slice(t_idx);
|
||||
auto tDrD_r2s = r2s_thr_copy_d.retile_S(tDrD);
|
||||
auto tDsD_r2s = r2s_thr_copy_d.partition_D(sD);
|
||||
|
||||
S2GCopyD s2g_tiled_copy_d;
|
||||
auto s2g_thr_copy_d = s2g_tiled_copy_d.get_slice(t_idx);
|
||||
auto tDsD_s2g = s2g_thr_copy_d.partition_S(sD);
|
||||
auto tDgD_s2g = s2g_thr_copy_d.partition_D(gD);
|
||||
Tensor cD = make_identity_tensor(make_shape(Int<TileM>{}, Int<TileN>{}));
|
||||
auto tDcD_s2g = s2g_thr_copy_d.partition_D(cD);
|
||||
|
||||
auto tDgD_s2gx = group_modes<1, 3>(tDgD_s2g); // (CPY_, CPY_MN)
|
||||
auto tDrD_r2sx = group_modes<1, 3>(tDrD_r2s); // (CPY_, CPY_MN)
|
||||
auto tDcD_s2gx = group_modes<1, 3>(tDcD_s2g);
|
||||
|
||||
int32_t step = size<3>(tDsD_r2s); // pipe
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t i = 0; i < size<1>(tDrD_r2sx); i += step) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t j = 0; j < step; ++j) {
|
||||
if constexpr (std::is_same<OutputDtype, float>::value) {
|
||||
cute::copy(r2s_tiled_copy_d, tDrD_r2sx(_, i + j), tDsD_r2s(_, 0, 0, j));
|
||||
} else {
|
||||
auto t = make_tensor_like<OutputDtype>(tDrD_r2sx(_, i + j));
|
||||
cute::copy(tDrD_r2sx(_, i + j), t);
|
||||
cute::copy(r2s_tiled_copy_d, t, tDsD_r2s(_, 0, 0, j));
|
||||
}
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
|
||||
// shm -> global
|
||||
if constexpr (IsEvenM && IsEvenN) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t j = 0; j < step; ++j)
|
||||
cute::copy(s2g_tiled_copy_d, tDsD_s2g(_, 0, 0, j), tDgD_s2gx(_, i + j));
|
||||
} else if constexpr (IsEvenN) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t j = 0; j < step; ++j) {
|
||||
if (get<0>(tDcD_s2gx(0, i + j)) < m_limit)
|
||||
cute::copy(s2g_tiled_copy_d, tDsD_s2g(_, 0, 0, j), tDgD_s2gx(_, i + j));
|
||||
}
|
||||
} else if constexpr (IsEvenM) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t j = 0; j < step; ++j)
|
||||
if (get<1>(tDcD_s2gx(size<0>(tDsD_s2g) - 1, i + j)) < n_limit) {
|
||||
cute::copy(s2g_tiled_copy_d, tDsD_s2g(_, 0, 0, j), tDgD_s2gx(_, i + j));
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k = 0; k < size<0>(tDsD_s2g); ++k)
|
||||
if (get<1>(tDcD_s2gx(k, i + j)) < n_limit)
|
||||
tDgD_s2gx(k, i + j) = tDsD_s2g(k, 0, 0, j);
|
||||
}
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t j = 0; j < step; ++j)
|
||||
if (get<0>(tDcD_s2gx(0, i + j)) < m_limit) {
|
||||
if (get<1>(tDcD_s2gx(size<0>(tDsD_s2g) - 1, i + j)) < n_limit) {
|
||||
cute::copy(s2g_tiled_copy_d, tDsD_s2g(_, 0, 0, j), tDgD_s2gx(_, i + j));
|
||||
} else {
|
||||
for (int32_t k = 0; k < size<0>(tDsD_s2g); ++k)
|
||||
if (get<1>(tDcD_s2gx(k, i + j)) < n_limit)
|
||||
tDgD_s2gx(k, i + j) = tDsD_s2g(k, 0, 0, j);
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
@@ -1,64 +0,0 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by TurboDiffusion team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
*
|
||||
* Citation (please cite if you use this code):
|
||||
*
|
||||
* @article{zhang2025turbodiffusion,
|
||||
* title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
* author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
* journal={arXiv preprint arXiv:2512.16093},
|
||||
* year={2025}
|
||||
* }
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "common/common.hpp"
|
||||
#include "common/launch.hpp"
|
||||
#include "gemm/kernel.hpp"
|
||||
|
||||
|
||||
template <class OutputDtype>
|
||||
bool int8_gemm_(
|
||||
int8_t const *Aptr, float const *ASptr,
|
||||
int8_t const *Bptr, float const *BSptr,
|
||||
OutputDtype* Dptr, int64_t m, int64_t n, int64_t k,
|
||||
int swizzle_dir = 1, int swizzle_size_log = 0,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
BOOL_SWITCH(m % 128 == 0, IsEvenM, [&] {
|
||||
BOOL_SWITCH(n % 128 == 0, IsEvenN, [&] {
|
||||
using Kernel = GemmKernel<OutputDtype, IsEvenM, IsEvenN>;
|
||||
if (!Kernel::can_implement(m, n, k))
|
||||
return false;
|
||||
using Args = typename Kernel::Arguments;
|
||||
Args args {
|
||||
(void*)Aptr, (void*)ASptr,
|
||||
(void*)Bptr, (void*)BSptr, (void*)Dptr,
|
||||
m, n, k, swizzle_dir,
|
||||
swizzle_size_log
|
||||
};
|
||||
|
||||
auto params = Kernel::to_underlying_arguments(args);
|
||||
|
||||
static constexpr size_t ShmSize = Kernel::ShmSize;
|
||||
dim3 grid_shape = Kernel::get_grid_size(m, n);
|
||||
dim3 block_shape = dim3(Kernel::ThreadNum);
|
||||
auto func = device_kernel<Kernel>;
|
||||
if (ShmSize >= 48 * 1024) {
|
||||
cudaFuncSetAttribute(
|
||||
func,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
ShmSize
|
||||
);
|
||||
}
|
||||
func<<<grid_shape, block_shape, ShmSize, stream>>>(
|
||||
params
|
||||
);
|
||||
return true;
|
||||
});
|
||||
});
|
||||
return true;
|
||||
}
|
||||
@@ -1,129 +0,0 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by TurboDiffusion team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
*
|
||||
* Citation (please cite if you use this code):
|
||||
*
|
||||
* @article{zhang2025turbodiffusion,
|
||||
* title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
* author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
* journal={arXiv preprint arXiv:2512.16093},
|
||||
* year={2025}
|
||||
* }
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
template <
|
||||
bool IsEven,
|
||||
class TiledCopy,
|
||||
class SrcTensor,
|
||||
class DstTensor,
|
||||
class PrdTensor
|
||||
>
|
||||
CUTLASS_DEVICE void
|
||||
copy_AB(
|
||||
TiledCopy const& _copy,
|
||||
SrcTensor const &S,
|
||||
DstTensor &D,
|
||||
PrdTensor const &ID,
|
||||
const int64_t &i_read,
|
||||
const int64_t &i_write,
|
||||
const int64_t &limit
|
||||
) {
|
||||
using namespace cute;
|
||||
if constexpr (IsEven)
|
||||
cute::copy(_copy, S(_, _, _, i_read), D(_, _, _, i_write));
|
||||
else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size<1>(ID); ++i)
|
||||
if (get<0>(ID(0, i, 0)) < limit)
|
||||
cute::copy(_copy, S(_, i, _, i_read), D(_, i, _, i_write));
|
||||
}
|
||||
}
|
||||
|
||||
template <int N>
|
||||
CUTLASS_DEVICE void copy_async(
|
||||
void const* gmem_src,
|
||||
void* smem_dst
|
||||
) {
|
||||
uint32_t smem_int_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(smem_dst));;
|
||||
asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], %2;\n"
|
||||
:: "r"(smem_int_ptr),
|
||||
"l"(gmem_src),
|
||||
"n"(N));
|
||||
}
|
||||
|
||||
template<class LoadType, class T, int NumThreads>
|
||||
CUTLASS_DEVICE void copy_aligned(const void* src, void* dst, size_t N, int64_t thread_idx) {
|
||||
static constexpr int NumElementPerLoad = sizeof(LoadType) / sizeof(T);
|
||||
for (int64_t i = thread_idx * NumElementPerLoad; i < N; i += NumElementPerLoad * NumThreads) {
|
||||
if (i + NumElementPerLoad <= N) {
|
||||
copy_async<sizeof(LoadType)>(
|
||||
(void*)((T*)src + i),
|
||||
(void*)((T*)dst + i)
|
||||
);
|
||||
} else {
|
||||
for (int64_t j = 0; j < N - i; ++j)
|
||||
copy_async<sizeof(T)>(
|
||||
(void*)((T*)src + i + j),
|
||||
(void*)((T*)dst + i + j)
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<class T, int NumThreads, bool Wait = true, bool Commit = true>
|
||||
CUTLASS_DEVICE void g2s_vector_copy(const void* src, void* dst, size_t N, int64_t thread_idx) {
|
||||
|
||||
uintptr_t src_addr = reinterpret_cast<uintptr_t>(src);
|
||||
|
||||
if (src_addr % 16 == 0) {
|
||||
copy_aligned<int4, T, NumThreads>(src, dst, N, thread_idx);
|
||||
} else if (src_addr % 8 == 0) {
|
||||
copy_aligned<int2, T, NumThreads>(src, dst, N, thread_idx);
|
||||
} else if (src_addr % 4 == 0) {
|
||||
copy_aligned<int, T, NumThreads>(src, dst, N, thread_idx);
|
||||
} else {
|
||||
assert(0);
|
||||
}
|
||||
if constexpr (Commit) {
|
||||
asm volatile("cp.async.commit_group;\n" ::);
|
||||
}
|
||||
if constexpr (Wait) {
|
||||
asm volatile("cp.async.wait_all;\n" ::);
|
||||
}
|
||||
}
|
||||
|
||||
template<class T, int N, bool FastInt2Float>
|
||||
CUTLASS_DEVICE
|
||||
static void dequant(
|
||||
T* mma_accum_ptr,
|
||||
float* float_accum_ptr,
|
||||
float scale
|
||||
) {
|
||||
static int const ic = 0x4B400000;
|
||||
if constexpr (FastInt2Float && std::is_same_v<T, int32_t>) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
*(float_accum_ptr + i) += (__int_as_float(*(mma_accum_ptr + i)) - __int_as_float(ic)) * scale;
|
||||
*(mma_accum_ptr + i) = ic;
|
||||
}
|
||||
} else if constexpr (std::is_same_v<T, int32_t>) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
*(float_accum_ptr + i) += __int2float_rn(*(mma_accum_ptr + i)) * scale;
|
||||
*(mma_accum_ptr + i) = 0;
|
||||
}
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
*(float_accum_ptr + i) += (*(mma_accum_ptr + i)) * scale;
|
||||
*(mma_accum_ptr + i) = 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,63 +0,0 @@
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include "common/common.hpp"
|
||||
#include "norm/layernorm.hpp"
|
||||
|
||||
auto layer_norm(
|
||||
at::Tensor const Input,
|
||||
float eps,
|
||||
std::optional<at::Tensor const> W,
|
||||
std::optional<at::Tensor const> const B,
|
||||
std::optional<at::Tensor> Output
|
||||
) {
|
||||
using ElementIn = float;
|
||||
using ElementOut = float;
|
||||
using ElementWeight = float;
|
||||
|
||||
int64_t const m = Input.size(0);
|
||||
int64_t const n = Input.size(1);
|
||||
torch::Device const input_device = Input.device();
|
||||
|
||||
if (!Output.has_value()) {
|
||||
Output.emplace(
|
||||
torch::empty(
|
||||
{m, n},
|
||||
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
void *Iptr = Input.data_ptr();
|
||||
void *Wptr = W.has_value() ? W.value().data_ptr() : nullptr;
|
||||
void *Bptr = B.has_value() ? B.value().data_ptr() : nullptr;
|
||||
void *Optr = Output.value().data_ptr();
|
||||
|
||||
BOOL_SWITCH(B.has_value(), BIAS, [&]{
|
||||
BOOL_SWITCH(W.has_value(), AFFINE, [&]{
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
layernorm<
|
||||
ElementIn, ElementOut, ElementWeight,
|
||||
AFFINE, BIAS,
|
||||
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA> (
|
||||
Iptr, Wptr, Bptr,
|
||||
Optr, eps, m, n,
|
||||
at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
|
||||
|
||||
return Output;
|
||||
}
|
||||
|
||||
void register_layer_norm(pybind11::module_ &m) {
|
||||
m.def("layer_norm_cuda", &layer_norm);
|
||||
}
|
||||
|
||||
@@ -1,202 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
#include "common/load.hpp"
|
||||
#include "common/store.hpp"
|
||||
#include "common/launch.hpp"
|
||||
|
||||
|
||||
template <
|
||||
class InputDtype_,
|
||||
class OutputDtype_,
|
||||
class WeightDtype_,
|
||||
bool Affine_,
|
||||
bool Bias_,
|
||||
int MaxHiddenSize_,
|
||||
int NumThrPerCta_,
|
||||
bool IsEven
|
||||
>
|
||||
class LayerNorm {
|
||||
public:
|
||||
using InputDtype = InputDtype_;
|
||||
using OutputDtype = OutputDtype_;
|
||||
using WeightDtype = WeightDtype_;
|
||||
static constexpr int NumThrPerCta = NumThrPerCta_;
|
||||
static constexpr int MaxHiddenSize = MaxHiddenSize_;
|
||||
static constexpr bool Affine = Affine_;
|
||||
static constexpr bool Bias = Bias_;
|
||||
|
||||
static constexpr size_t ShmSize = 32;
|
||||
static constexpr int NumElementPerThread = MaxHiddenSize / NumThrPerCta;
|
||||
|
||||
static_assert(MaxHiddenSize % NumThrPerCta == 0);
|
||||
|
||||
struct Params {
|
||||
void const *Iptr;
|
||||
void const *Wptr;
|
||||
void const *Bptr;
|
||||
void *Optr;
|
||||
float eps;
|
||||
int64_t m;
|
||||
int64_t n;
|
||||
};
|
||||
|
||||
using Arguments = Params;
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args) {
|
||||
return args;
|
||||
}
|
||||
|
||||
static dim3 get_grid_size(int64_t m, int64_t n) {
|
||||
return dim3(m);
|
||||
}
|
||||
|
||||
static dim3 get_cta_size(int64_t m, int64_t n) {
|
||||
return dim3(NumThrPerCta, 1, 1);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const& params, char *shared_data) {
|
||||
int const blk_m = blockIdx.x;
|
||||
int const blk_n = 1;
|
||||
int tidx = threadIdx.x;
|
||||
float x[NumElementPerThread];
|
||||
|
||||
// load
|
||||
Loader<InputDtype, 1, MaxHiddenSize, NumThrPerCta, true, IsEven> loader;
|
||||
loader.load(params.Iptr, x, params.m, params.n, blk_m, 0, tidx);
|
||||
|
||||
// mean reduction
|
||||
float u = _reduce_sum(x, shared_data) / params.n;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
x[i] -= u;
|
||||
|
||||
__syncthreads();
|
||||
// var reduction
|
||||
float v = sqrtf(_reduce_square(x, shared_data) / params.n + params.eps);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
x[i] /= v;
|
||||
|
||||
if constexpr (Affine) {
|
||||
// load weight
|
||||
Loader<WeightDtype, 1, MaxHiddenSize, NumThrPerCta, true, IsEven> weight_loader;
|
||||
float w[NumElementPerThread];
|
||||
weight_loader.load(params.Wptr, w, 1, params.n, 0, 0, tidx);
|
||||
if constexpr (Bias) {
|
||||
float b[NumElementPerThread];
|
||||
weight_loader.load(params.Bptr, b, 1, params.n, 0, 0, tidx);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
x[i] = x[i] * w[i] + b[i];
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
x[i] = x[i] * w[i];
|
||||
}
|
||||
}
|
||||
|
||||
// save y
|
||||
{
|
||||
Saver<OutputDtype, 1, MaxHiddenSize, NumThrPerCta, true, IsEven, false, false> saver;
|
||||
if constexpr (std::is_same_v<OutputDtype, float>) {
|
||||
saver.store(params.Optr, nullptr, x, 0, params.m, params.n, blk_m, 0, tidx);
|
||||
} else {
|
||||
OutputDtype tmp[NumElementPerThread];
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
tmp[i] = OutputDtype(x[i]);
|
||||
saver.store(params.Optr, nullptr, tmp, 0, params.m, params.n, blk_m, 0, tidx);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
private:
|
||||
CUTLASS_DEVICE
|
||||
float _reduce_square(float *reg, char *shared_data) {
|
||||
// thread
|
||||
float sum_square = 0;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
sum_square += reg[i] * reg[i];
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 16; i >= 1; i >>= 1) {
|
||||
sum_square += __shfl_down_sync(0xFFFFFFFF, sum_square, i);
|
||||
}
|
||||
if (threadIdx.x == 0) {
|
||||
*(float*)shared_data = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x % 32 == 0) {
|
||||
atomicAdd((float*)shared_data, sum_square);
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
sum_square = *(float*)shared_data;
|
||||
return sum_square;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
float _reduce_sum(float *reg, char *shared_data) {
|
||||
// thread
|
||||
float sum = 0;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
sum += reg[i];
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 16; i >= 1; i >>= 1) {
|
||||
sum += __shfl_down_sync(0xFFFFFFFF, sum, i);
|
||||
}
|
||||
if (threadIdx.x == 0) {
|
||||
*(float*)shared_data = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x % 32 == 0) {
|
||||
atomicAdd((float*)shared_data, sum);
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
sum = *(float*)shared_data;
|
||||
return sum;
|
||||
}
|
||||
};
|
||||
|
||||
template <
|
||||
class InputDtype,
|
||||
class OutputDtype,
|
||||
class WeightDtype,
|
||||
bool Affine,
|
||||
bool Bias,
|
||||
int MaxHiddenSize,
|
||||
int NumThrPerCta
|
||||
>
|
||||
bool layernorm(
|
||||
void const *Iptr, void const *Wptr, void const *Bptr,
|
||||
void *Optr, float eps, int64_t m, int64_t n,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
BOOL_SWITCH(n % MaxHiddenSize == 0, IsEven, [&] {
|
||||
using Kernel = LayerNorm<
|
||||
InputDtype, OutputDtype, WeightDtype,
|
||||
Affine, Bias,
|
||||
MaxHiddenSize, NumThrPerCta,
|
||||
IsEven>;
|
||||
using Arguments = typename Kernel::Arguments;
|
||||
Arguments args = {
|
||||
Iptr, Wptr, Bptr, Optr,
|
||||
eps, m, n
|
||||
};
|
||||
auto params = Kernel::to_underlying_arguments(args);
|
||||
auto grid_shape = Kernel::get_grid_size(m, n);
|
||||
auto cta_shape = Kernel::get_cta_size(m, n);
|
||||
static constexpr size_t ShmSize = Kernel::ShmSize;
|
||||
launch_kernel<Kernel>(params, grid_shape, cta_shape, ShmSize, stream);
|
||||
});
|
||||
return true;
|
||||
}
|
||||
@@ -1,59 +0,0 @@
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <pybind11/pybind11.h>
|
||||
|
||||
#include "common/common.hpp"
|
||||
#include "norm/rmsnorm.hpp"
|
||||
|
||||
auto rms_norm(
|
||||
at::Tensor const& Input,
|
||||
float eps,
|
||||
const std::optional<at::Tensor>& Weight,
|
||||
std::optional<at::Tensor>& Output
|
||||
) {
|
||||
|
||||
using ElementIn = float;
|
||||
using ElementOut = float;
|
||||
using ElementWeight = float;
|
||||
|
||||
int64_t const m = Input.size(0);
|
||||
int64_t const n = Input.size(1);
|
||||
torch::Device const input_device = Input.device();
|
||||
|
||||
if (!Output.has_value()) {
|
||||
Output.emplace(
|
||||
torch::empty(
|
||||
{m, n},
|
||||
torch::TensorOptions().device(input_device).dtype(torch::kFloat32)
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
void *Iptr = Input.data_ptr();
|
||||
void *Wptr = Weight.has_value() ? Weight.value().data_ptr() : nullptr;
|
||||
void *Optr = Output.value().data_ptr();
|
||||
|
||||
|
||||
CONFIG_SWITCH(n, [&]{
|
||||
rmsnorm<
|
||||
ElementIn, ElementOut, ElementWeight,
|
||||
MAX_HIDDEN_SIZE, NUM_THR_PER_CTA
|
||||
> (
|
||||
Iptr, Wptr,
|
||||
Optr,
|
||||
eps, m, n,
|
||||
at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
});
|
||||
|
||||
|
||||
return Output;
|
||||
}
|
||||
|
||||
void register_rms_norm(pybind11::module_ &m) {
|
||||
m.def("rms_norm_cuda", &rms_norm);
|
||||
}
|
||||
@@ -1,147 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
#include "common/load.hpp"
|
||||
#include "common/store.hpp"
|
||||
#include "common/launch.hpp"
|
||||
|
||||
|
||||
template <
|
||||
class InputDtype_,
|
||||
class OutputDtype_,
|
||||
class WeightDtype_,
|
||||
int MaxHiddenSize_,
|
||||
int NumThrPerCta_,
|
||||
bool IsEven
|
||||
>
|
||||
class RMSNorm {
|
||||
public:
|
||||
using InputDtype = InputDtype_;
|
||||
using OutputDtype = OutputDtype_;
|
||||
using WeightDtype = WeightDtype_;
|
||||
static constexpr int NumThrPerCta = NumThrPerCta_;
|
||||
static constexpr int MaxHiddenSize = MaxHiddenSize_;
|
||||
|
||||
static constexpr size_t ShmSize = 32;
|
||||
static constexpr int NumElementPerThread = MaxHiddenSize / NumThrPerCta;
|
||||
|
||||
static_assert(MaxHiddenSize % NumThrPerCta == 0);
|
||||
|
||||
struct Params {
|
||||
void const *Iptr;
|
||||
void const *Wptr;
|
||||
void *Optr;
|
||||
float eps;
|
||||
int64_t m;
|
||||
int64_t n;
|
||||
};
|
||||
|
||||
using Arguments = Params;
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args) {
|
||||
return args;
|
||||
}
|
||||
|
||||
static dim3 get_grid_size(int64_t m, int64_t n) {
|
||||
return dim3(m);
|
||||
}
|
||||
|
||||
static dim3 get_cta_size(int64_t m, int64_t n) {
|
||||
return dim3(NumThrPerCta, 1, 1);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const& params, char *shared_data) {
|
||||
int const blk_m = blockIdx.x;
|
||||
int const blk_n = 1;
|
||||
int tidx = threadIdx.x;
|
||||
float x[NumElementPerThread];
|
||||
|
||||
// load
|
||||
Loader<InputDtype, 1, MaxHiddenSize, NumThrPerCta, true, IsEven> loader;
|
||||
loader.load(params.Iptr, x, params.m, params.n, blk_m, 0, tidx);
|
||||
|
||||
// rms reduction
|
||||
float rms = sqrtf(_reduce_square(x, shared_data) / params.n + params.eps);
|
||||
|
||||
// load weight
|
||||
Loader<WeightDtype, 1, MaxHiddenSize, NumThrPerCta, true, IsEven> weight_loader;
|
||||
float w[NumElementPerThread];
|
||||
loader.load(params.Wptr, w, 1, params.n, 0, 0, tidx);
|
||||
|
||||
// norm
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
x[i] = w[i] * x[i] / rms ;
|
||||
|
||||
// save y
|
||||
OutputDtype *output_reg = (OutputDtype*)x;
|
||||
if constexpr (!std::is_same_v<OutputDtype, float>) {
|
||||
output_reg = (OutputDtype*)w;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
output_reg[i] = OutputDtype(x[i]);
|
||||
}
|
||||
Saver<OutputDtype, 1, MaxHiddenSize, NumThrPerCta, true, IsEven, false, false> saver;
|
||||
saver.store(params.Optr, nullptr, output_reg, 0, params.m, params.n, blk_m, 0, tidx);
|
||||
|
||||
}
|
||||
|
||||
private:
|
||||
CUTLASS_DEVICE
|
||||
float _reduce_square(float *reg, char *shared_data) {
|
||||
// thread
|
||||
float sum_square = 0;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
sum_square += reg[i] * reg[i];
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 16; i >= 1; i >>= 1) {
|
||||
sum_square += __shfl_down_sync(0xFFFFFFFF, sum_square, i);
|
||||
}
|
||||
if (threadIdx.x == 0) {
|
||||
*(float*)shared_data = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x % 32 == 0) {
|
||||
atomicAdd((float*)shared_data, sum_square);
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
sum_square = *(float*)shared_data;
|
||||
return sum_square;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
template <
|
||||
class InputDtype,
|
||||
class OutputDtype,
|
||||
class WeightDtype,
|
||||
int MaxHiddenSize,
|
||||
int NumThrPerCta
|
||||
>
|
||||
bool rmsnorm(
|
||||
void const *Iptr, void const *Wptr,
|
||||
void *Optr, float eps,
|
||||
int64_t m, int64_t n,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
BOOL_SWITCH(n % MaxHiddenSize == 0, IsEven, [&] {
|
||||
using Kernel = RMSNorm<
|
||||
InputDtype, OutputDtype, WeightDtype,
|
||||
MaxHiddenSize, NumThrPerCta,
|
||||
IsEven>;
|
||||
using Arguments = typename Kernel::Arguments;
|
||||
Arguments args = {
|
||||
Iptr, Wptr, Optr, eps, m, n
|
||||
};
|
||||
auto params = Kernel::to_underlying_arguments(args);
|
||||
auto grid_shape = Kernel::get_grid_size(m, n);
|
||||
auto cta_shape = Kernel::get_cta_size(m, n);
|
||||
static constexpr size_t ShmSize = Kernel::ShmSize;
|
||||
launch_kernel<Kernel>(params, grid_shape, cta_shape, ShmSize, stream);
|
||||
});
|
||||
return true;
|
||||
}
|
||||
@@ -1,75 +0,0 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by TurboDiffusion team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
*
|
||||
* Citation (please cite if you use this code):
|
||||
*
|
||||
* @article{zhang2025turbodiffusion,
|
||||
* title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
* author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
* journal={arXiv preprint arXiv:2512.16093},
|
||||
* year={2025}
|
||||
* }
|
||||
*/
|
||||
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include <pybind11/pybind11.h>
|
||||
|
||||
#include "common/common.hpp"
|
||||
#include "quant/quant.hpp"
|
||||
|
||||
auto quant(
|
||||
torch::Tensor const& Input,
|
||||
std::optional<torch::Tensor>& Output,
|
||||
std::optional<torch::Tensor>& Output_S
|
||||
) {
|
||||
|
||||
using ElementOut = int8_t;
|
||||
static constexpr int BlockSize = 128;
|
||||
static constexpr int NumThrPerCta = 256;
|
||||
|
||||
int64_t m = Input.size(0);
|
||||
int64_t n = Input.size(1);
|
||||
torch::Device const input_device = Input.device();
|
||||
|
||||
create_tensor<BlockSize>(input_device, Output, Output_S, m, n);
|
||||
|
||||
ElementOut *Optr = (ElementOut*)Output.value().data_ptr();
|
||||
float *OSptr = Output_S.value().data_ptr<float>();
|
||||
|
||||
switch (Input.scalar_type()) {
|
||||
case torch::kHalf:{
|
||||
cutlass::half_t *Iptr = (cutlass::half_t*)Input.data_ptr();
|
||||
quantization<cutlass::half_t, BlockSize, NumThrPerCta> (
|
||||
Iptr, Optr, OSptr, m, n, at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
case torch::kBFloat16:{
|
||||
cutlass::bfloat16_t *Iptr = (cutlass::bfloat16_t*)Input.data_ptr();
|
||||
quantization<cutlass::bfloat16_t, BlockSize, NumThrPerCta> (
|
||||
Iptr, Optr, OSptr, m, n, at::cuda::getCurrentCUDAStream().stream()
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
default: {
|
||||
std::cerr << "Observing: " << Input.scalar_type() << " for the input datatype which is invalid";
|
||||
throw std::runtime_error("Unsupported input data type for quantize_to_fp4.");
|
||||
}
|
||||
}
|
||||
|
||||
return std::make_tuple(Output, Output_S);
|
||||
}
|
||||
|
||||
void register_quant(pybind11::module_ &m) {
|
||||
m.def("quant_cuda", &quant);
|
||||
}
|
||||
@@ -1,195 +0,0 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by TurboDiffusion team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
*
|
||||
* Citation (please cite if you use this code):
|
||||
*
|
||||
* @article{zhang2025turbodiffusion,
|
||||
* title={TurboDiffusion: Accelerating Video Diffusion Models by 100-200 Times},
|
||||
* author={Zhang, Jintao and Zheng, Kaiwen and Jiang, Kai and Wang, Haoxu and Stoica, Ion and Gonzalez, Joseph E and Chen, Jianfei and Zhu, Jun},
|
||||
* journal={arXiv preprint arXiv:2512.16093},
|
||||
* year={2025}
|
||||
* }
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
|
||||
#include "common/load.hpp"
|
||||
#include "common/store.hpp"
|
||||
#include "common/launch.hpp"
|
||||
|
||||
template <
|
||||
class InputDtype_,
|
||||
int NumThrPerCta_,
|
||||
bool IsEvenM,
|
||||
bool IsEvenN
|
||||
>
|
||||
class Quantization {
|
||||
public:
|
||||
using InputDtype = InputDtype_;
|
||||
using OutputDtype = int8_t;
|
||||
using FPConverter = cutlass::NumericConverter<int8_t, float, cutlass::FloatRoundStyle::round_to_nearest>;
|
||||
|
||||
static constexpr int BlockSize = 128;
|
||||
static constexpr int NumThrPerCta = NumThrPerCta_;
|
||||
static constexpr int NumElementPerThread = BlockSize * BlockSize / NumThrPerCta;
|
||||
static constexpr int NumThrPerRow = BlockSize / NumElementPerThread;
|
||||
|
||||
static_assert(BlockSize * BlockSize % NumThrPerCta == 0);
|
||||
static_assert(NumThrPerCta % BlockSize == 0);
|
||||
|
||||
static constexpr size_t ShmSize = 32;
|
||||
|
||||
static constexpr float int8_max = 128.f;
|
||||
|
||||
struct Params {
|
||||
void const *Iptr;
|
||||
void *Optr;
|
||||
void *OSptr;
|
||||
int64_t const m;
|
||||
int64_t const n;
|
||||
};
|
||||
|
||||
using Arguments = Params;
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args) {
|
||||
return args;
|
||||
}
|
||||
|
||||
static dim3 get_grid_size(int64_t m, int64_t n) {
|
||||
return dim3(
|
||||
cdiv(n, BlockSize),
|
||||
cdiv(m, BlockSize)
|
||||
);
|
||||
}
|
||||
|
||||
static dim3 get_cta_size(int64_t m, int64_t n) {
|
||||
return dim3(
|
||||
NumThrPerCta, 1, 1
|
||||
);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void quantization(
|
||||
float *float_reg,
|
||||
void *Optr, void *OSptr,
|
||||
int64_t const m, int64_t const n,
|
||||
int blk_m, int blk_n, int tidx,
|
||||
char *shared_data
|
||||
) {
|
||||
|
||||
OutputDtype output_reg[NumElementPerThread];
|
||||
|
||||
|
||||
Saver<OutputDtype, BlockSize, BlockSize, NumThrPerCta, IsEvenM, IsEvenN> saver;
|
||||
|
||||
float amax = _reduce_amax(float_reg, (float*)shared_data);
|
||||
|
||||
|
||||
_quantization(float_reg, output_reg, int8_max / amax);
|
||||
|
||||
float scale_inv = amax / int8_max;
|
||||
|
||||
saver.store(Optr, OSptr, output_reg, scale_inv, m, n, blk_m, blk_n, tidx);
|
||||
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const& params, char *shared_data) {
|
||||
int blk_m = blockIdx.y;
|
||||
int blk_n = blockIdx.x;
|
||||
int tidx = threadIdx.x;
|
||||
|
||||
float float_reg[NumElementPerThread];
|
||||
|
||||
// load float32 data
|
||||
Loader<InputDtype, BlockSize, BlockSize, NumThrPerCta, IsEvenM, IsEvenN> loader;
|
||||
loader.load(params.Iptr, float_reg, params.m, params.n, blk_m, blk_n, tidx);
|
||||
quantization(
|
||||
float_reg, params.Optr, params.OSptr, params.m, params.n, blk_m, blk_n, tidx, shared_data
|
||||
);
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
CUTLASS_DEVICE float
|
||||
_reduce_amax(float *reg, float *smem_ptr) {
|
||||
float amax = 1e-8;
|
||||
// thread reduction
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i)
|
||||
amax = max(amax, fabs(reg[i]));
|
||||
|
||||
__syncwarp();
|
||||
|
||||
// warp reduction
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 16; i >= 1; i /= 2) {
|
||||
amax = max(
|
||||
__shfl_xor_sync(0xffffffff, amax, i, 32),
|
||||
amax
|
||||
);
|
||||
}
|
||||
|
||||
// cta reduction
|
||||
if (threadIdx.x == 0) {
|
||||
*smem_ptr = 0;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
atomicMax((uint32_t*)smem_ptr, reinterpret_cast<const uint32_t&>(amax));
|
||||
|
||||
__syncthreads();
|
||||
|
||||
amax = *smem_ptr;
|
||||
|
||||
return amax;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
_quantization(float *float_reg, OutputDtype *out_reg, float scale) {
|
||||
FPConverter converter;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < NumElementPerThread; ++i) {
|
||||
out_reg[i] = converter(float_reg[i] * scale);
|
||||
}
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
template <
|
||||
class InputDtype,
|
||||
int BlockSize,
|
||||
int NumThrPerCta
|
||||
>
|
||||
bool quantization(
|
||||
void const *Iptr, void *Optr, void *OSptr,
|
||||
int64_t m, int64_t n,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
BOOL_SWITCH(m % BlockSize == 0, IsEvenM, [&] {
|
||||
BOOL_SWITCH(n % BlockSize == 0, IsEvenN, [&] {
|
||||
using Kernel = Quantization<
|
||||
InputDtype, NumThrPerCta, IsEvenM, IsEvenN>;
|
||||
using Arguments = typename Kernel::Arguments;
|
||||
Arguments args = {
|
||||
Iptr, Optr, OSptr,
|
||||
m, n
|
||||
};
|
||||
auto params = Kernel::to_underlying_arguments(args);
|
||||
auto grid_shape = Kernel::get_grid_size(m, n);
|
||||
auto cta_shape = Kernel::get_cta_size(m, n);
|
||||
static constexpr size_t ShmSize = Kernel::ShmSize;
|
||||
launch_kernel<Kernel>(params, grid_shape, cta_shape, ShmSize, stream);
|
||||
});
|
||||
});
|
||||
|
||||
return true;
|
||||
}
|
||||
Submodule fastvideo-kernel/include/cutlass deleted from e67e63c331
@@ -1,35 +0,0 @@
|
||||
[build-system]
|
||||
requires = [
|
||||
"scikit-build-core>=0.10",
|
||||
"torch>=2.5.0",
|
||||
"setuptools>=61.0.0",
|
||||
"wheel"
|
||||
]
|
||||
build-backend = "scikit_build_core.build"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.2.1"
|
||||
description = "Unified CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
license = { file = "LICENSE" }
|
||||
authors = [
|
||||
{ name = "Hao AI Lab", email = "contact@haoailab.com" }
|
||||
]
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA",
|
||||
]
|
||||
dependencies = [
|
||||
"torch>=2.5.0",
|
||||
"triton>=2.0.0",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
"Homepage" = "https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
[tool.scikit-build]
|
||||
cmake.build-type = "Release"
|
||||
minimum-version = "build-system.requires"
|
||||
wheel.packages = ["python/fastvideo_kernel"]
|
||||
@@ -1,34 +0,0 @@
|
||||
from .version import __version__
|
||||
|
||||
from fastvideo_kernel.ops import (
|
||||
sliding_tile_attention,
|
||||
video_sparse_attn,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.vmoba import (
|
||||
moba_attn_varlen,
|
||||
process_moba_input,
|
||||
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,112 +0,0 @@
|
||||
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)
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
def sliding_tile_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
window_size: list,
|
||||
text_length: int,
|
||||
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
|
||||
)
|
||||
|
||||
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
|
||||
if pad_size > 0:
|
||||
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],
|
||||
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]
|
||||
|
||||
|
||||
def video_sparse_attn(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
topk: int,
|
||||
block_size: int | tuple = 64,
|
||||
compress_attn_weight: torch.Tensor = None,
|
||||
) -> 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)
|
||||
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)
|
||||
|
||||
# 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)
|
||||
|
||||
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>
|
||||
else:
|
||||
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
|
||||
@@ -1,335 +0,0 @@
|
||||
import math
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
def is_cuda():
|
||||
return triton.runtime.driver.active.get_current_target().backend == "cuda"
|
||||
|
||||
|
||||
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 w in [4, 8]\
|
||||
]
|
||||
return configs
|
||||
|
||||
|
||||
def get_cuda_autotune_config():
|
||||
# cuda and hip can use differnt autotune configs
|
||||
return get_common_autotune_config()
|
||||
|
||||
|
||||
def get_hip_autotune_config():
|
||||
# cuda and hip can use differnt autotune configs
|
||||
return get_common_autotune_config()
|
||||
|
||||
|
||||
def get_autotune_config():
|
||||
if is_cuda():
|
||||
return get_cuda_autotune_config()
|
||||
else:
|
||||
return get_hip_autotune_config()
|
||||
|
||||
|
||||
@triton.jit
|
||||
def clamp_int(value, min_val, max_val):
|
||||
ret = tl.where(value > max_val, max_val, value)
|
||||
ret = tl.where(ret < min_val, min_val, ret)
|
||||
return ret
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_loop(
|
||||
q, k, v, kv_mask, m, l, acc, sm_scale,
|
||||
MASK_KV: tl.constexpr,
|
||||
):
|
||||
scores = tl.dot(q, k.T) #[BLOCK_Q, BLOCK_KV]
|
||||
scores = scores * sm_scale
|
||||
if MASK_KV:
|
||||
scores = tl.where(kv_mask[None, :], scores, -float('inf'))
|
||||
|
||||
current_m = tl.max(scores, axis=1)
|
||||
new_m = tl.maximum(m, current_m)
|
||||
exp_scores = tl.math.exp2(scores - new_m[:, None])
|
||||
current_l = tl.sum(exp_scores, axis=1)
|
||||
|
||||
# Update L <- L * exp(M - M') + L1, M <- M'
|
||||
alpha = tl.math.exp2(m - new_m)
|
||||
l = l * alpha + current_l
|
||||
m = new_m
|
||||
|
||||
# Update O <- O * exp(M - M') + P @ V
|
||||
acc = (acc * alpha[:, None] + tl.dot(exp_scores.to(v.type.element_ty), v))
|
||||
|
||||
return m, l, acc
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=get_autotune_config(),
|
||||
key=['head_dim'],
|
||||
)
|
||||
@triton.jit
|
||||
def triton_sta_kernel(
|
||||
Q, K, V, output,
|
||||
batch_size: int, num_heads: int, seq_len: int, head_dim: int,
|
||||
img_seq_len: int,
|
||||
text_length: int,
|
||||
canvas_t: int, canvas_h: int, canvas_w: int,
|
||||
kernel_t: int, kernel_h: int, kernel_w: int,
|
||||
tile_t: int, tile_h: int, tile_w: int,
|
||||
scale: float,
|
||||
has_text: tl.constexpr,
|
||||
text_q: tl.constexpr,
|
||||
BLOCK_Q: tl.constexpr,
|
||||
BLOCK_KV: tl.constexpr,
|
||||
BLOCK_DIM: tl.constexpr,
|
||||
):
|
||||
total_tile_size = tile_t * tile_h * tile_w
|
||||
q_block_per_tile = (total_tile_size + BLOCK_Q - 1) // BLOCK_Q
|
||||
|
||||
batch_idx = tl.program_id(0)
|
||||
head_idx = tl.program_id(1)
|
||||
if text_q:
|
||||
q_block_idx = tl.program_id(2)
|
||||
else:
|
||||
q_tile_flat = tl.program_id(2) // q_block_per_tile
|
||||
q_block_idx = tl.program_id(2) % q_block_per_tile
|
||||
|
||||
m = tl.full((BLOCK_Q,), -float('inf'), dtype=tl.float32)
|
||||
l = tl.zeros((BLOCK_Q,), dtype=tl.float32)
|
||||
acc = tl.zeros((BLOCK_Q, BLOCK_DIM), dtype=tl.float32)
|
||||
|
||||
q_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
if text_q:
|
||||
q_base_idx = img_seq_len + q_block_idx * BLOCK_Q
|
||||
else:
|
||||
q_base_idx = q_tile_flat * total_tile_size + q_block_idx * BLOCK_Q
|
||||
|
||||
q_offset_in_tile = tl.arange(0, BLOCK_Q)
|
||||
q_idx = q_base_idx + q_offset_in_tile
|
||||
q_mask = (q_block_idx * BLOCK_Q + tl.arange(0, BLOCK_Q)) < total_tile_size
|
||||
|
||||
q = tl.load(
|
||||
Q + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=q_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_Q, BLOCK_DIM]
|
||||
|
||||
# Scale sm_scale by log_2(e) and use 2^x instead of exp
|
||||
sm_scale = scale * 1.4426950408889634
|
||||
|
||||
num_tiles_t = canvas_t // tile_t
|
||||
num_tiles_h = canvas_h // tile_h
|
||||
num_tiles_w = canvas_w // tile_w
|
||||
tiles_per_hw = num_tiles_h * num_tiles_w
|
||||
|
||||
if text_q:
|
||||
kv_tile_start_t = 0
|
||||
kv_tile_end_t = num_tiles_t
|
||||
|
||||
kv_tile_start_h = 0
|
||||
kv_tile_end_h = num_tiles_h
|
||||
|
||||
kv_tile_start_w = 0
|
||||
kv_tile_end_w = num_tiles_w
|
||||
|
||||
else:
|
||||
q_tile_t = q_tile_flat // tiles_per_hw
|
||||
remaining = q_tile_flat % tiles_per_hw
|
||||
q_tile_h = remaining // num_tiles_w
|
||||
q_tile_w = remaining % num_tiles_w
|
||||
|
||||
kernel_center_t = clamp_int(q_tile_t, kernel_t // 2, (num_tiles_t - 1) - kernel_t // 2)
|
||||
kernel_center_h = clamp_int(q_tile_h, kernel_h // 2, (num_tiles_h - 1) - kernel_h // 2)
|
||||
kernel_center_w = clamp_int(q_tile_w, kernel_w // 2, (num_tiles_w - 1) - kernel_w // 2)
|
||||
|
||||
kv_tile_start_t = kernel_center_t - kernel_t // 2
|
||||
kv_tile_end_t = kernel_center_t + kernel_t // 2 + 1
|
||||
kv_tile_end_t = tl.where(kv_tile_end_t > num_tiles_t, num_tiles_t, kv_tile_end_t)
|
||||
|
||||
kv_tile_start_h = kernel_center_h - kernel_h // 2
|
||||
kv_tile_end_h = kernel_center_h + kernel_h // 2 + 1
|
||||
kv_tile_end_h = tl.where(kv_tile_end_h > num_tiles_h, num_tiles_h, kv_tile_end_h)
|
||||
|
||||
kv_tile_start_w = kernel_center_w - kernel_w // 2
|
||||
kv_tile_end_w = kernel_center_w + kernel_w // 2 + 1
|
||||
kv_tile_end_w = tl.where(kv_tile_end_w > num_tiles_w, num_tiles_w, kv_tile_end_w)
|
||||
|
||||
# for kv_img
|
||||
for kv_tile_t in tl.range(kv_tile_start_t, kv_tile_end_t):
|
||||
for kv_tile_h in tl.range(kv_tile_start_h, kv_tile_end_h):
|
||||
for kv_tile_w in tl.range(kv_tile_start_w, kv_tile_end_w):
|
||||
kv_base_idx = (kv_tile_t * num_tiles_h * num_tiles_w + kv_tile_h * num_tiles_w + kv_tile_w) * total_tile_size
|
||||
|
||||
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
|
||||
kv_offset_in_block = tl.arange(0, BLOCK_KV)
|
||||
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
|
||||
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < total_tile_size
|
||||
|
||||
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
|
||||
k = tl.load(
|
||||
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
v = tl.load(
|
||||
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
|
||||
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, False)
|
||||
|
||||
|
||||
# for kv_text
|
||||
if has_text:
|
||||
kv_base_idx = img_seq_len
|
||||
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
|
||||
kv_offset_in_block = tl.arange(0, BLOCK_KV)
|
||||
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
|
||||
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < text_length
|
||||
|
||||
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
|
||||
k = tl.load(
|
||||
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
v = tl.load(
|
||||
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
|
||||
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, True)
|
||||
|
||||
|
||||
output_acc = acc / l[:, None]
|
||||
tl.store(
|
||||
output + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
output_acc,
|
||||
mask=q_mask[:, None]
|
||||
) # [BLOCK_Q, BLOCK_DIM]
|
||||
|
||||
|
||||
def sliding_tile_attention_triton(
|
||||
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
window_size, text_length: int,
|
||||
has_text=True, dit_seq_shape='30x48x80') -> torch.Tensor:
|
||||
seq_length = q.shape[2]
|
||||
if has_text:
|
||||
assert q.shape[2] >= 115200 and q.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
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)
|
||||
else:
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo
|
||||
assert q.shape[2] == 82944
|
||||
elif dit_seq_shape == '18x48x80': # Wan
|
||||
assert q.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
assert q.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
if dit_seq_shape == '30x48x80': # Hunyuan
|
||||
canvas_t, canvas_h, canvas_w = 30, 48, 80
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
elif dit_seq_shape == '36x48x48': # Stepvideo
|
||||
canvas_t, canvas_h, canvas_w = 36, 48, 48
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
elif dit_seq_shape == '18x48x80': # Wan
|
||||
canvas_t, canvas_h, canvas_w = 18, 48, 80
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
num_tiles_t = canvas_t // tile_t
|
||||
num_tiles_h = canvas_h // tile_h
|
||||
num_tiles_w = canvas_w // tile_w
|
||||
num_tiles = num_tiles_t * num_tiles_h * num_tiles_w
|
||||
|
||||
total_tile_size = tile_t * tile_h * tile_w
|
||||
|
||||
# BLOCK_Q=128
|
||||
# BLOCK_KV=128
|
||||
BLOCK_DIM = head_dim
|
||||
|
||||
output = torch.empty_like(q)
|
||||
|
||||
# for q_img
|
||||
# kernel_size maybe different for different head
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (kernel_t, kernel_h, kernel_w) in enumerate(window_size):
|
||||
for batch in range(batch_size):
|
||||
q_head, k_head, v_head, o_head = (q[batch:batch + 1, head_index:head_index + 1],
|
||||
k[batch:batch + 1, head_index:head_index + 1],
|
||||
v[batch:batch + 1, head_index:head_index + 1],
|
||||
output[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
# triton_sta_kernel[(1, 1, num_tiles * triton.cdiv(total_tile_size, BLOCK_Q))](
|
||||
grid = lambda META: (1, 1, num_tiles * triton.cdiv(total_tile_size, META['BLOCK_Q']))
|
||||
triton_sta_kernel[grid](
|
||||
q_head, k_head, v_head, o_head,
|
||||
1, 1, seq_len, head_dim,
|
||||
img_seq_len,
|
||||
text_length,
|
||||
canvas_t, canvas_h, canvas_w,
|
||||
kernel_t, kernel_h, kernel_w,
|
||||
tile_t, tile_h, tile_w,
|
||||
scale=1.0 / (head_dim ** 0.5),
|
||||
has_text=has_text,
|
||||
text_q=False,
|
||||
# BLOCK_Q=BLOCK_Q,
|
||||
# BLOCK_KV=BLOCK_KV,
|
||||
BLOCK_DIM=BLOCK_DIM,
|
||||
)
|
||||
|
||||
# for q_text
|
||||
# kernel_t, kernel_h, kernel_w is not used, set to (3, 3, 3)
|
||||
if has_text:
|
||||
# triton_sta_kernel[(batch_size, num_heads, triton.cdiv(total_tile_size, BLOCK_Q))](
|
||||
grid = lambda META: (batch_size, num_heads, triton.cdiv(total_tile_size, META['BLOCK_Q']))
|
||||
triton_sta_kernel[grid](
|
||||
q, k, v, output,
|
||||
batch_size, num_heads, seq_len, head_dim,
|
||||
img_seq_len,
|
||||
text_length,
|
||||
canvas_t, canvas_h, canvas_w,
|
||||
3, 3, 3,
|
||||
#kernel_t, kernel_h, kernel_w,
|
||||
tile_t, tile_h, tile_w,
|
||||
scale=1.0 / (head_dim ** 0.5),
|
||||
has_text=has_text,
|
||||
text_q=True,
|
||||
# BLOCK_Q=BLOCK_Q,
|
||||
# BLOCK_KV=BLOCK_KV,
|
||||
BLOCK_DIM=BLOCK_DIM,
|
||||
)
|
||||
|
||||
if has_text:
|
||||
if pad_size > 0:
|
||||
output = output[:, :, :seq_length]
|
||||
|
||||
return output
|
||||
@@ -1,716 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional, Tuple
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# Try to load the C++ extension
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops
|
||||
quant_cuda = getattr(fastvideo_kernel_ops, "quant_cuda", None)
|
||||
gemm_cuda = getattr(fastvideo_kernel_ops, "gemm_cuda", None)
|
||||
rms_norm_cuda = getattr(fastvideo_kernel_ops, "rms_norm_cuda", None)
|
||||
layer_norm_cuda = getattr(fastvideo_kernel_ops, "layer_norm_cuda", None)
|
||||
except ImportError:
|
||||
quant_cuda = None
|
||||
gemm_cuda = None
|
||||
rms_norm_cuda = None
|
||||
layer_norm_cuda = None
|
||||
|
||||
def int8_quant(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Quantize a floating-point tensor to int8 using a custom CUDA kernel.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor of type float16/bfloat16.
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]:
|
||||
- x_q: Quantized int8 tensor.
|
||||
- x_scale: Per-block scale tensor used for quantization.
|
||||
"""
|
||||
x_q, x_scale = quant_cuda(x, None, None)
|
||||
return x_q, x_scale
|
||||
|
||||
|
||||
def int8_linear(
|
||||
x: torch.Tensor,
|
||||
w_q: torch.Tensor,
|
||||
w_s: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Perform an int8 GEMM (matrix multiplication) using quantized weights and a
|
||||
quantized version of the input. The underlying compute is performed by a
|
||||
custom CUDA kernel.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input activation of shape (M, K) in float32.
|
||||
w_q (torch.Tensor): Quantized int8 weight tensor of shape (N, K).
|
||||
w_s (torch.Tensor): Scale tensor associated with w_q.
|
||||
**kwargs: Additional options (reserved for future use).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor of shape (M, N) in float32.
|
||||
"""
|
||||
assert w_q.dtype == torch.int8, "Weight tensor must be int8."
|
||||
shape = x.shape
|
||||
x = x.reshape(-1, shape[-1])
|
||||
m = x.shape[0]
|
||||
n = w_q.shape[0]
|
||||
y = torch.zeros(m, n, dtype=x.dtype, device=x.device)
|
||||
|
||||
x_q, x_s = int8_quant(x)
|
||||
gemm_cuda(x_q, x_s, w_q, w_s, y)
|
||||
return y.reshape(*shape[:-1], n)
|
||||
|
||||
def flatten_if_batched(*tensors):
|
||||
"""
|
||||
Flattens all input tensors from (B, N, D_i) to (B * N, D_i) if they are batched (3D).
|
||||
|
||||
Args:
|
||||
*tensors: Any number of input tensors, each must have shape (B, N, D_i) or (N, D_i)
|
||||
|
||||
Returns:
|
||||
flat_tensors: List of flattened tensors
|
||||
batched: Boolean flag indicating whether inputs were batched
|
||||
batch_size: Batch size if batched, else None
|
||||
"""
|
||||
if not tensors:
|
||||
raise ValueError("At least one tensor must be provided.")
|
||||
|
||||
first = tensors[0]
|
||||
assert len(first.shape) in [
|
||||
2,
|
||||
3,
|
||||
], "Input tensors must be batched (3D) or not batched (2D)"
|
||||
|
||||
if len(first.shape) == 3: # batched
|
||||
batched = True
|
||||
batch_size = first.shape[0]
|
||||
assert all(t.shape[0] == batch_size for t in tensors), "All input tensors must have the same batch size"
|
||||
assert all(
|
||||
t.shape[1] == first.shape[1] for t in tensors
|
||||
), "All input tensors must have the same sequence length"
|
||||
flat_tensors = [t.reshape(-1, t.shape[-1]) for t in tensors]
|
||||
else:
|
||||
batched = False
|
||||
batch_size = None
|
||||
flat_tensors = list(tensors)
|
||||
|
||||
return flat_tensors, batched, batch_size
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _rms_norm_fwd_fused(
|
||||
X,
|
||||
Y,
|
||||
W,
|
||||
Rstd,
|
||||
x_stride,
|
||||
y_stride,
|
||||
N: tl.constexpr, # number of columns in X,
|
||||
N2: tl.constexpr,
|
||||
eps, # epsilon to avoid division by zero
|
||||
BLOCK_M: tl.constexpr,
|
||||
):
|
||||
# Map the program id to the row of X and Y it should compute.
|
||||
pid = tl.program_id(0)
|
||||
rows = pid * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
cols = tl.arange(0, N2)
|
||||
mask = cols < N
|
||||
|
||||
x_ptr = X + rows[:, None] * x_stride + cols[None, :]
|
||||
y_ptr = Y + rows[:, None] * y_stride + cols[None, :]
|
||||
|
||||
x = tl.load(x_ptr, mask=mask[None, :], other=0.0).to(tl.float32)
|
||||
|
||||
# Compute variance
|
||||
_var = x * x
|
||||
var = tl.sum(_var, axis=1) / N
|
||||
rstd = 1 / tl.sqrt(var + eps)
|
||||
|
||||
# Write mean / rstd
|
||||
tl.store(Rstd + rows, rstd)
|
||||
rstd = tl.reshape(rstd, (BLOCK_M, 1))
|
||||
|
||||
# Normalize and apply linear transformation
|
||||
w = tl.load(W + cols)
|
||||
x_hat = x * rstd
|
||||
y = x_hat * w
|
||||
|
||||
# Write output
|
||||
y = y.to(Y.type.element_ty)
|
||||
tl.store(y_ptr, y, mask=mask[None, :])
|
||||
|
||||
|
||||
def rmsnorm(x, w, eps):
|
||||
"""
|
||||
Forward pass of the RMSNorm.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor, High precision.
|
||||
w (torch.Tensor): RMSNorm weight tensor.
|
||||
eps (float): RMSNorm epsilon value.
|
||||
|
||||
Returns:
|
||||
y (torch.Tensor): Output tensor, High precision.
|
||||
rstd (torch.Tensor): Inverse standard deviation, needed for backward.
|
||||
"""
|
||||
assert x.is_contiguous(), "Input must be contiguous"
|
||||
# Change batched 3D input to 2D
|
||||
[x], batched, BS = flatten_if_batched(x)
|
||||
|
||||
# allocate output
|
||||
M, N = x.shape
|
||||
y = torch.empty_like(x, dtype=x.dtype)
|
||||
rstd = torch.empty((M,), dtype=torch.float32, device=x.device)
|
||||
|
||||
# heuristics for number of warps
|
||||
num_warps = 8
|
||||
|
||||
# Avoid illegal memory access
|
||||
N2 = triton.next_power_of_2(N)
|
||||
|
||||
if N <= 512:
|
||||
BLOCK_M = 32
|
||||
else:
|
||||
BLOCK_M = 1
|
||||
|
||||
# Call the triton kernel
|
||||
_rms_norm_fwd_fused[(triton.cdiv(M, BLOCK_M),)]( #
|
||||
x,
|
||||
y,
|
||||
w,
|
||||
rstd, #
|
||||
x.stride(0),
|
||||
y.stride(0),
|
||||
N,
|
||||
N2,
|
||||
eps,
|
||||
num_warps=num_warps,
|
||||
BLOCK_M=BLOCK_M,
|
||||
)
|
||||
|
||||
# Recover 2D to 3D
|
||||
if batched:
|
||||
y = y.reshape(BS, -1, y.shape[-1])
|
||||
|
||||
return y, rstd
|
||||
|
||||
@triton.jit
|
||||
def _layer_norm_param_fwd_fused(
|
||||
X, # pointer to the input
|
||||
Y, # pointer to the output
|
||||
W, # pointer to the weights
|
||||
B, # pointer to the biases
|
||||
Mean, # pointer to the mean
|
||||
Rstd, # pointer to the 1/std
|
||||
x_stride, # how much to increase the pointer when moving by 1 row
|
||||
y_stride, # how much to increase the pointer when moving by 1 row
|
||||
N: tl.constexpr, # number of columns in X,
|
||||
N2: tl.constexpr, # number of columns in X,
|
||||
eps, # epsilon to avoid division by zero
|
||||
BLOCK_M: tl.constexpr,
|
||||
):
|
||||
# Map the program id to the row of X and Y it should compute.
|
||||
pid = tl.program_id(0)
|
||||
rows = pid * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
cols = tl.arange(0, N2)
|
||||
mask = cols < N
|
||||
|
||||
x_ptr = X + rows[:, None] * x_stride + cols[None, :]
|
||||
y_ptr = Y + rows[:, None] * y_stride + cols[None, :]
|
||||
|
||||
x = tl.load(x_ptr, mask=mask[None, :], other=0.0).to(tl.float32)
|
||||
|
||||
# Compute mean and Variance
|
||||
mean = tl.sum(x, axis=1, keep_dims=True) / N
|
||||
# Compute variance
|
||||
_var = (x - mean) * (x - mean)
|
||||
var = tl.sum(_var, axis=1, keep_dims=True) / N
|
||||
rstd = 1 / tl.sqrt(var + eps)
|
||||
|
||||
# Write mean / rstd
|
||||
_mean = tl.reshape(mean, (BLOCK_M))
|
||||
_rstd = tl.reshape(rstd, (BLOCK_M))
|
||||
tl.store(Mean + rows, _mean)
|
||||
tl.store(Rstd + rows, _rstd)
|
||||
|
||||
# Normalize and apply linear transformation
|
||||
x_hat = (x - mean) * rstd
|
||||
|
||||
w = tl.load(W + cols)
|
||||
b = tl.load(B + cols)
|
||||
|
||||
x_hat = x_hat * w + b
|
||||
|
||||
# Write output
|
||||
x_hat = x_hat.to(Y.type.element_ty)
|
||||
tl.store(y_ptr, x_hat, mask=mask[None, :])
|
||||
|
||||
|
||||
def layernorm_param(x, w, b, eps):
|
||||
# Change batched 3D input to 2D
|
||||
[x], batched, BS = flatten_if_batched(x)
|
||||
|
||||
# allocate output
|
||||
M, N = x.shape
|
||||
y = torch.empty_like(x, dtype=torch.float32)
|
||||
mean = torch.empty((M,), dtype=torch.float32, device=x.device)
|
||||
rstd = torch.empty((M,), dtype=torch.float32, device=x.device)
|
||||
# heuristics for number of warps
|
||||
num_warps = 8
|
||||
|
||||
N2 = triton.next_power_of_2(N)
|
||||
|
||||
if N <= 512:
|
||||
BLOCK_M = 32
|
||||
else:
|
||||
BLOCK_M = 1
|
||||
|
||||
# enqueue kernel
|
||||
_layer_norm_param_fwd_fused[(triton.cdiv(M, BLOCK_M),)]( #
|
||||
x,
|
||||
y,
|
||||
w,
|
||||
b,
|
||||
mean,
|
||||
rstd, #
|
||||
x.stride(0),
|
||||
y.stride(0),
|
||||
N,
|
||||
N2,
|
||||
eps,
|
||||
num_warps=num_warps,
|
||||
BLOCK_M=BLOCK_M,
|
||||
)
|
||||
|
||||
# Recover 2D to 3D
|
||||
if batched:
|
||||
y = y.reshape(BS, -1, y.shape[-1])
|
||||
|
||||
return y, mean, rstd
|
||||
|
||||
|
||||
########################################################
|
||||
# Elementwise_affine=False
|
||||
########################################################
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _layer_norm_noparam_fwd_fused(
|
||||
X, # pointer to the input
|
||||
Y, # pointer to the output
|
||||
Mean, # pointer to the mean
|
||||
Rstd, # pointer to the 1/std
|
||||
x_stride, # how much to increase the pointer when moving by 1 row
|
||||
y_stride, # how much to increase the pointer when moving by 1 row
|
||||
N: tl.constexpr, # number of columns in X,
|
||||
N2: tl.constexpr, # number of columns in X,
|
||||
eps, # epsilon to avoid division by zero
|
||||
BLOCK_M: tl.constexpr,
|
||||
):
|
||||
# Map the program id to the row of X and Y it should compute.
|
||||
pid = tl.program_id(0)
|
||||
rows = pid * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
cols = tl.arange(0, N2)
|
||||
mask = cols < N
|
||||
|
||||
x_ptr = X + rows[:, None] * x_stride + cols[None, :]
|
||||
y_ptr = Y + rows[:, None] * y_stride + cols[None, :]
|
||||
|
||||
x = tl.load(x_ptr, mask=mask[None, :], other=0.0).to(tl.float32)
|
||||
|
||||
# Compute mean and Variance
|
||||
mean = tl.sum(x, axis=1, keep_dims=True) / N
|
||||
# Compute variance
|
||||
_var = (x - mean) * (x - mean)
|
||||
var = tl.sum(_var, axis=1, keep_dims=True) / N
|
||||
rstd = 1 / tl.sqrt(var + eps)
|
||||
|
||||
# Write mean / rstd
|
||||
_mean = tl.reshape(mean, (BLOCK_M))
|
||||
_rstd = tl.reshape(rstd, (BLOCK_M))
|
||||
tl.store(Mean + rows, _mean)
|
||||
tl.store(Rstd + rows, _rstd)
|
||||
|
||||
# Normalize and apply linear transformation
|
||||
x_hat = (x - mean) * rstd
|
||||
|
||||
# Write output
|
||||
x_hat = x_hat.to(Y.type.element_ty)
|
||||
tl.store(y_ptr, x_hat, mask=mask[None, :])
|
||||
|
||||
|
||||
def layernorm_noparam(x, eps):
|
||||
assert x.is_contiguous(), "Input must be contiguous"
|
||||
|
||||
# Change batched 3D input to 2D
|
||||
[x], batched, BS = flatten_if_batched(x)
|
||||
|
||||
# allocate output
|
||||
M, N = x.shape
|
||||
y = torch.empty_like(x, dtype=torch.float32)
|
||||
mean = torch.empty((M,), dtype=torch.float32, device=x.device)
|
||||
rstd = torch.empty((M,), dtype=torch.float32, device=x.device)
|
||||
# heuristics for number of warps
|
||||
num_warps = 8
|
||||
|
||||
N2 = triton.next_power_of_2(N)
|
||||
|
||||
if N <= 512:
|
||||
BLOCK_M = 32
|
||||
else:
|
||||
BLOCK_M = 1
|
||||
|
||||
# enqueue kernel
|
||||
_layer_norm_noparam_fwd_fused[(triton.cdiv(M, BLOCK_M),)]( #
|
||||
x,
|
||||
y,
|
||||
mean,
|
||||
rstd, #
|
||||
x.stride(0),
|
||||
y.stride(0),
|
||||
N,
|
||||
N2,
|
||||
eps,
|
||||
num_warps=num_warps,
|
||||
BLOCK_M=BLOCK_M,
|
||||
)
|
||||
|
||||
# Recover 2D to 3D
|
||||
if batched:
|
||||
y = y.reshape(BS, -1, y.shape[-1])
|
||||
|
||||
return y, mean, rstd
|
||||
|
||||
def layernorm(x, w, b, eps, elementwise_affine=True):
|
||||
if elementwise_affine:
|
||||
assert w is not None and b is not None
|
||||
return layernorm_param(x, w, b, eps)
|
||||
else:
|
||||
assert w is None and b is None
|
||||
return layernorm_noparam(x, eps)
|
||||
|
||||
def cdiv(a: int, b: int):
|
||||
return (a + b - 1) // b
|
||||
|
||||
def dequantize_weight(w_q, w_s, dtype):
|
||||
# w_q: (N, K) int8
|
||||
# w_s: (NB, KB) float32
|
||||
# Block size 128
|
||||
BLOCK = 128
|
||||
N, K = w_q.shape
|
||||
# Expand w_s
|
||||
# Repeat interleave
|
||||
w_s_exp = w_s.repeat_interleave(BLOCK, dim=0).repeat_interleave(BLOCK, dim=1)
|
||||
# Crop
|
||||
w_s_exp = w_s_exp[:N, :K]
|
||||
return w_q.to(dtype) * w_s_exp.to(dtype)
|
||||
|
||||
class Int8LinearFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x, w_q, w_s, bias=None):
|
||||
ctx.save_for_backward(x, w_q, w_s, bias)
|
||||
ctx.bias_requires_grad = bias.requires_grad if bias is not None else False
|
||||
|
||||
out = int8_linear(x, w_q, w_s)
|
||||
if bias is not None:
|
||||
out = out + bias
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
x, w_q, w_s, bias = ctx.saved_tensors
|
||||
|
||||
grad_input = None
|
||||
grad_weight = None
|
||||
grad_scale = None
|
||||
grad_bias = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
# grad_input = grad_output @ W
|
||||
w_float = dequantize_weight(w_q, w_s, grad_output.dtype)
|
||||
grad_input = torch.matmul(grad_output, w_float)
|
||||
|
||||
if ctx.bias_requires_grad:
|
||||
dim_to_sum = list(range(grad_output.dim() - 1))
|
||||
grad_bias = grad_output.sum(dim=dim_to_sum)
|
||||
|
||||
return grad_input, grad_weight, grad_scale, grad_bias
|
||||
|
||||
class RMSNormFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x, w, eps):
|
||||
y, rstd = rmsnorm(x, w, eps)
|
||||
ctx.save_for_backward(x, w, rstd)
|
||||
ctx.eps = eps
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
# x, w, rstd are saved
|
||||
# grad_output is dL/dy
|
||||
# We can implement backward using torch ops for correctness
|
||||
x, w, rstd = ctx.saved_tensors
|
||||
eps = ctx.eps
|
||||
N = x.shape[-1]
|
||||
|
||||
# dL/dy * w
|
||||
dx = grad_output * w
|
||||
|
||||
# Expand rstd
|
||||
rstd = rstd.unsqueeze(-1) # (B, 1) or (M, 1)
|
||||
if x.dim() == 3:
|
||||
# Flatten if necessary or handle dimensions.
|
||||
# forward flattens x to (M, N) internally but saved x might be 3D?
|
||||
# rmsnorm takes x, flattens it. But does it modify x in place? No.
|
||||
# But ctx.save_for_backward saves the original x (3D if input was 3D).
|
||||
# rstd is (M,).
|
||||
# We should flatten x and grad_output to match logic
|
||||
x_flat = x.reshape(-1, N)
|
||||
grad_output_flat = grad_output.reshape(-1, N)
|
||||
dx = dx.reshape(-1, N)
|
||||
else:
|
||||
x_flat = x
|
||||
grad_output_flat = grad_output
|
||||
dx = dx
|
||||
|
||||
# Standard RMSNorm backward
|
||||
# dy = grad_output_flat
|
||||
# x_hat = x * rstd
|
||||
# w * dy
|
||||
# c1 = mean(dy * w * x) * rstd^2
|
||||
# dx = (dy * w - c1 * x) * rstd
|
||||
|
||||
# More precise:
|
||||
# y = x * rstd * w
|
||||
# dy = grad_output
|
||||
# dw = sum(dy * x * rstd, dim=0)
|
||||
|
||||
grad_w = None
|
||||
if ctx.needs_input_grad[1]: # w
|
||||
# grad_w = (grad_output_flat * x_flat * rstd).sum(dim=0)
|
||||
# More accurate gradient for weight when considering rstd was computed from x
|
||||
# Actually, for Affine part (w), it is just sum(dL/dy * x_hat).
|
||||
# y = x_hat * w
|
||||
# dL/dw = sum(dL/dy * x_hat)
|
||||
x_hat = x_flat * rstd
|
||||
grad_w = (grad_output_flat * x_hat).sum(dim=0)
|
||||
|
||||
grad_input = None
|
||||
if ctx.needs_input_grad[0]: # x
|
||||
# x_hat = x * rstd
|
||||
# dx = rstd * (w * dy - mean(w * dy * x_hat) * x_hat)
|
||||
# but for RMSNorm:
|
||||
# dx = rstd * (w * dy - (x * rstd^2) * mean(w * dy * x)) -> check formula
|
||||
|
||||
# Using PyTorch autograd for reference logic:
|
||||
# y = x / sqrt(mean(x^2) + eps) * w
|
||||
# Let sigma = sqrt(...)
|
||||
# dL/dx = dL/dy * w * (1/sigma) + dL/dsigma * dsigma/dx
|
||||
# dsigma/dx = 1/(2*sigma) * 2x/N = x / (N * sigma)
|
||||
# dL/dsigma = sum(dL/dy * w * x * (-1/sigma^2))
|
||||
# dL/dx = (dL/dy * w)/sigma - sum(dL/dy * w * x) * x / (N * sigma^3)
|
||||
# = (1/sigma) * [ (dL/dy * w) - x * sum(dL/dy * w * x) / (N * sigma^2) ]
|
||||
# = rstd * [ (dL/dy * w) - x * rstd^2 * mean(dL/dy * w * x) ] # Wait, sum/N is mean
|
||||
|
||||
# Let's compute:
|
||||
dy_w = grad_output_flat * w
|
||||
# term2 = (dy_w * x_flat).sum(dim=-1, keepdim=True) / N # mean(dy*w*x)
|
||||
# grad_input = rstd * (dy_w - x_flat * (rstd ** 2) * term2)
|
||||
|
||||
# Re-derivation:
|
||||
# x_hat = x * rstd
|
||||
# y = x_hat * w
|
||||
# dL/dx = dL/dy * dy/dx
|
||||
# dy/dx = w * dx_hat/dx
|
||||
# dx_hat/dx = rstd * (I - x * x^T * rstd^2 / N) ?? No.
|
||||
# dx_hat_i / dx_j = rstd * (delta_ij - x_i * x_j * rstd^2 / N)
|
||||
|
||||
# dL/dx_hat = dL/dy * w
|
||||
# dL/dx = rstd * (dL/dx_hat - x * rstd^2 * mean(dL/dx_hat * x))
|
||||
|
||||
dx_hat = grad_output_flat * w
|
||||
dx_hat_x_mean = (dx_hat * x_flat).mean(dim=-1, keepdim=True)
|
||||
grad_input = rstd * (dx_hat - x_flat * (rstd ** 2) * dx_hat_x_mean)
|
||||
|
||||
if x.dim() == 3:
|
||||
grad_input = grad_input.reshape(x.shape)
|
||||
|
||||
return grad_input, grad_w, None
|
||||
|
||||
class LayerNormFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x, w, b, eps, elementwise_affine):
|
||||
if elementwise_affine:
|
||||
y, mean, rstd = layernorm_param(x, w, b, eps)
|
||||
ctx.save_for_backward(x, w, b, mean, rstd)
|
||||
else:
|
||||
y, mean, rstd = layernorm_noparam(x, eps)
|
||||
ctx.save_for_backward(x, None, None, mean, rstd)
|
||||
|
||||
ctx.eps = eps
|
||||
ctx.elementwise_affine = elementwise_affine
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
x, w, b, mean, rstd = ctx.saved_tensors
|
||||
N = x.shape[-1]
|
||||
|
||||
if x.dim() == 3:
|
||||
x_flat = x.reshape(-1, N)
|
||||
grad_output_flat = grad_output.reshape(-1, N)
|
||||
else:
|
||||
x_flat = x
|
||||
grad_output_flat = grad_output
|
||||
|
||||
if w is None:
|
||||
w_eff = 1.0
|
||||
else:
|
||||
w_eff = w
|
||||
|
||||
# dx calculation
|
||||
# x_hat = (x - mean) * rstd
|
||||
# y = x_hat * w + b
|
||||
# dL/dx_hat = dL/dy * w
|
||||
|
||||
dy = grad_output_flat
|
||||
dx_hat = dy * w_eff
|
||||
|
||||
# dL/dvar = sum(dL/dx_hat * (x-mean) * (-0.5) * (var+eps)^(-1.5))
|
||||
# = -0.5 * rstd^3 * sum(dx_hat * (x-mean))
|
||||
# dL/dmean = sum(dL/dx_hat * (-rstd)) + dL/dvar * (-2/N) * sum(x-mean)
|
||||
# term sum(x-mean) is 0. So second part vanishes.
|
||||
# dL/dmean = -rstd * sum(dx_hat)
|
||||
|
||||
# dL/dx = dL/dx_hat * rstd + dL/dvar * 2(x-mean)/N + dL/dmean * 1/N
|
||||
# = dx_hat * rstd + (-0.5 * rstd^3 * sum(dx_hat * (x-mean))) * 2(x-mean)/N + (-rstd * sum(dx_hat))/N
|
||||
# = rstd * [ dx_hat - mean(dx_hat) - (x-mean)*rstd^2 * mean(dx_hat * (x-mean)) ]
|
||||
|
||||
x_centered = x_flat - mean.unsqueeze(-1)
|
||||
dx_hat_mean = dx_hat.mean(dim=-1, keepdim=True)
|
||||
dx_hat_x_centered_mean = (dx_hat * x_centered).mean(dim=-1, keepdim=True)
|
||||
|
||||
grad_input = rstd.unsqueeze(-1) * (dx_hat - dx_hat_mean - x_centered * (rstd.unsqueeze(-1)**2) * dx_hat_x_centered_mean)
|
||||
|
||||
if x.dim() == 3:
|
||||
grad_input = grad_input.reshape(x.shape)
|
||||
|
||||
grad_w = None
|
||||
grad_b = None
|
||||
|
||||
if ctx.elementwise_affine:
|
||||
if ctx.needs_input_grad[1]:
|
||||
x_hat = x_centered * rstd.unsqueeze(-1)
|
||||
grad_w = (dy * x_hat).sum(dim=0)
|
||||
if ctx.needs_input_grad[2]:
|
||||
grad_b = dy.sum(dim=0)
|
||||
|
||||
return grad_input, grad_w, grad_b, None, None
|
||||
|
||||
|
||||
class Int8Linear(nn.Module):
|
||||
def __init__(self, in_features, out_features, bias=True, dtype=torch.bfloat16):
|
||||
super().__init__()
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
|
||||
row_blocks = cdiv(out_features, b=128)
|
||||
col_blocks = cdiv(in_features, b=128)
|
||||
|
||||
self.register_buffer("int8_weight", torch.empty((out_features, in_features), dtype=torch.int8))
|
||||
self.register_buffer("scale", torch.empty((row_blocks, col_blocks), dtype=torch.float32))
|
||||
if bias:
|
||||
self.register_buffer("bias", torch.empty(out_features, dtype=dtype))
|
||||
else:
|
||||
self.bias = None
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
return Int8LinearFunction.apply(x, self.int8_weight, self.scale, self.bias)
|
||||
|
||||
@classmethod
|
||||
def from_linear(cls, original_linear: nn.Linear, quantize: bool = True):
|
||||
|
||||
int8_layer = cls(
|
||||
original_linear.in_features,
|
||||
original_linear.out_features,
|
||||
bias=original_linear.bias is not None,
|
||||
dtype=original_linear.weight.dtype
|
||||
)
|
||||
if quantize:
|
||||
w_data = original_linear.weight.data.cuda()
|
||||
int8_w, scale = int8_quant(w_data)
|
||||
|
||||
int8_layer.int8_weight.copy_(int8_w)
|
||||
int8_layer.scale.copy_(scale)
|
||||
if original_linear.bias is not None:
|
||||
int8_layer.bias.data.copy_(original_linear.bias.data.cuda())
|
||||
|
||||
return int8_layer
|
||||
|
||||
class FastRMSNorm(nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.register_buffer("weight", torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return RMSNormFunction.apply(x.float(), self.weight, self.eps).to(x.dtype)
|
||||
|
||||
@classmethod
|
||||
def from_rmsnorm(cls, original_rmsnorm):
|
||||
rmsnorm_layer = cls(
|
||||
dim=original_rmsnorm.dim,
|
||||
eps=original_rmsnorm.eps
|
||||
)
|
||||
if original_rmsnorm.weight.device != torch.device('meta'):
|
||||
rmsnorm_layer.weight.data.copy_(original_rmsnorm.weight.float().data)
|
||||
return rmsnorm_layer
|
||||
|
||||
class FastLayerNorm(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
eps: float = 1e-5,
|
||||
elementwise_affine: bool = False,
|
||||
bias: bool = True
|
||||
) :
|
||||
super().__init__()
|
||||
self.dim = dim # type: ignore[arg-type]
|
||||
self.eps = eps
|
||||
self.elementwise_affine = elementwise_affine
|
||||
if self.elementwise_affine:
|
||||
self.register_buffer("weight", torch.empty(self.dim))
|
||||
if bias:
|
||||
self.register_buffer("bias", torch.empty(self.dim))
|
||||
else:
|
||||
self.bias = None
|
||||
else:
|
||||
self.register_parameter("weight", None)
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def forward(self, x):
|
||||
return LayerNormFunction.apply(x.float(), self.weight, self.bias, self.eps, self.elementwise_affine).to(x.dtype)
|
||||
|
||||
@classmethod
|
||||
def from_layernorm(cls, original_layernorm):
|
||||
layernorm_layer = cls(
|
||||
dim=original_layernorm.normalized_shape[0],
|
||||
eps=original_layernorm.eps,
|
||||
elementwise_affine=False if original_layernorm.weight is None else True,
|
||||
bias=original_layernorm.bias is not None
|
||||
)
|
||||
if original_layernorm.weight is not None and original_layernorm.weight.device != torch.device('meta'):
|
||||
layernorm_layer.weight.data.copy_(original_layernorm.weight.data)
|
||||
if original_layernorm.bias is not None and original_layernorm.bias.device != torch.device('meta'):
|
||||
layernorm_layer.bias.data.copy_(original_layernorm.bias.data)
|
||||
return layernorm_layer
|
||||
@@ -1 +0,0 @@
|
||||
__version__ = "0.2.1"
|
||||
@@ -1,294 +0,0 @@
|
||||
import torch
|
||||
import pytest
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Ensure local package is imported
|
||||
# sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../python")))
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops
|
||||
except ImportError:
|
||||
fastvideo_kernel_ops = None
|
||||
|
||||
from fastvideo_kernel import turbodiffusion_ops
|
||||
|
||||
# Helper for RMS Norm reference
|
||||
def rms_norm_ref(x, w, eps=1e-6):
|
||||
dtype = x.dtype
|
||||
x = x.float()
|
||||
variance = x.pow(2).mean(-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + eps)
|
||||
return (x * w.float()).to(dtype)
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
|
||||
class TestTurboDiffusion:
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("shape", [(16, 128), (32, 256), (1, 1024)])
|
||||
def test_quant_correctness(self, dtype, shape):
|
||||
if turbodiffusion_ops.quant_cuda is None:
|
||||
pytest.skip("quant_cuda not available")
|
||||
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
x_q, x_scale = turbodiffusion_ops.int8_quant(x)
|
||||
|
||||
assert x_q.dtype == torch.int8
|
||||
assert x_scale.dtype == torch.float32
|
||||
|
||||
# Simple check: dequantize and compute error
|
||||
# Note: The quantization scheme details matter here (per block? per tensor?).
|
||||
# Looking at quant.cu, it seems to be block-based but the output scale shape isn't immediately obvious from python signature
|
||||
# without looking at C++ code deeper.
|
||||
# But let's check shapes at least.
|
||||
|
||||
# If we can't easily dequantize without knowing block size logic in python,
|
||||
# checking that it runs and produces valid shapes is a good start.
|
||||
assert x_q.shape == shape
|
||||
# x_scale shape depends on block size, usually smaller than x
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_gemm_correctness(self, dtype):
|
||||
if turbodiffusion_ops.gemm_cuda is None:
|
||||
pytest.skip("gemm_cuda not available")
|
||||
|
||||
M, N, K = 32, 64, 128
|
||||
x = torch.randn(M, K, dtype=dtype, device="cuda")
|
||||
|
||||
# Create weights
|
||||
# For simplicity in testing, let's create random int8 weights and scales
|
||||
w_q = torch.randint(-127, 127, (N, K), dtype=torch.int8, device="cuda")
|
||||
|
||||
# Scale shape: The Int8Linear class uses:
|
||||
# row_blocks = cdiv(out_features, b=128)
|
||||
# col_blocks = cdiv(in_features, b=128)
|
||||
# scale shape: (row_blocks, col_blocks)
|
||||
|
||||
row_blocks = (N + 127) // 128
|
||||
col_blocks = (K + 127) // 128
|
||||
w_s = torch.randn(row_blocks, col_blocks, dtype=torch.float32, device="cuda").abs()
|
||||
|
||||
# Run int8_linear
|
||||
output = turbodiffusion_ops.int8_linear(x, w_q, w_s)
|
||||
|
||||
assert output.shape == (M, N)
|
||||
assert output.dtype == dtype
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_gemm_backward(self, dtype):
|
||||
if turbodiffusion_ops.gemm_cuda is None:
|
||||
pytest.skip("gemm_cuda not available")
|
||||
|
||||
M, N, K = 32, 64, 128
|
||||
x = torch.randn(M, K, dtype=dtype, device="cuda", requires_grad=True)
|
||||
|
||||
# Weights (frozen)
|
||||
w_q = torch.randint(-127, 127, (N, K), dtype=torch.int8, device="cuda")
|
||||
row_blocks = (N + 127) // 128
|
||||
col_blocks = (K + 127) // 128
|
||||
w_s = torch.randn(row_blocks, col_blocks, dtype=torch.float32, device="cuda").abs()
|
||||
|
||||
bias = torch.randn(N, dtype=dtype, device="cuda", requires_grad=True)
|
||||
|
||||
# Use Int8LinearFunction
|
||||
output = turbodiffusion_ops.Int8LinearFunction.apply(x, w_q, w_s, bias)
|
||||
loss = output.sum()
|
||||
loss.backward()
|
||||
|
||||
assert x.grad is not None
|
||||
assert bias.grad is not None
|
||||
assert x.grad.shape == (M, K)
|
||||
|
||||
# Check correctness against dequantized weight
|
||||
w_float = turbodiffusion_ops.dequantize_weight(w_q, w_s, dtype)
|
||||
|
||||
# Reference
|
||||
x_ref = x.detach().clone().requires_grad_()
|
||||
bias_ref = bias.detach().clone().requires_grad_()
|
||||
|
||||
# Note: Int8Linear forward uses quantized input, so output won't match exactly reference with float input.
|
||||
# But backward gradients should be consistent with the logic we implemented:
|
||||
# grad_input = grad_output @ w_float
|
||||
|
||||
# Let's verify our backward logic matches standard matmul backward with dequantized weights
|
||||
# We manually compute expected gradient given grad_output = ones
|
||||
grad_output = torch.ones_like(output)
|
||||
expected_x_grad = grad_output @ w_float
|
||||
expected_bias_grad = grad_output.sum(0)
|
||||
|
||||
torch.testing.assert_close(x.grad, expected_x_grad, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(bias.grad, expected_bias_grad, atol=1e-2, rtol=1e-2)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("shape", [(2, 16, 128), (4, 32, 256)])
|
||||
def test_rms_norm_triton(self, dtype, shape):
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
dim = shape[-1]
|
||||
w = torch.randn(dim, dtype=dtype, device="cuda")
|
||||
eps = 1e-5
|
||||
|
||||
# Triton implementation
|
||||
# Note: rmsnorm returns tuple now
|
||||
res = turbodiffusion_ops.rmsnorm(x, w, eps)
|
||||
if isinstance(res, tuple):
|
||||
out_triton = res[0]
|
||||
else:
|
||||
out_triton = res
|
||||
|
||||
# Reference
|
||||
out_ref = rms_norm_ref(x, w, eps)
|
||||
|
||||
torch.testing.assert_close(out_triton, out_ref, atol=1e-2, rtol=1e-2)
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("shape", [(2, 16, 128)])
|
||||
def test_rms_norm_backward(self, dtype, shape):
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda", requires_grad=True)
|
||||
dim = shape[-1]
|
||||
w = torch.randn(dim, dtype=dtype, device="cuda", requires_grad=True)
|
||||
eps = 1e-5
|
||||
|
||||
# Forward via Function
|
||||
y = turbodiffusion_ops.RMSNormFunction.apply(x, w, eps)
|
||||
loss = y.sum()
|
||||
loss.backward()
|
||||
|
||||
x_grad = x.grad
|
||||
w_grad = w.grad
|
||||
|
||||
# Reference
|
||||
x_ref = x.detach().clone().requires_grad_()
|
||||
w_ref = w.detach().clone().requires_grad_()
|
||||
# Custom RMSNorm ref in pytorch
|
||||
def rms_norm_ref_grad(x, w, eps):
|
||||
x_float = x.float()
|
||||
var = x_float.pow(2).mean(-1, keepdim=True)
|
||||
rstd = torch.rsqrt(var + eps)
|
||||
return (x_float * rstd * w.float()).to(x.dtype)
|
||||
|
||||
y_ref = rms_norm_ref_grad(x_ref, w_ref, eps)
|
||||
loss_ref = y_ref.sum()
|
||||
loss_ref.backward()
|
||||
|
||||
torch.testing.assert_close(x_grad, x_ref.grad, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(w_grad, w_ref.grad, atol=2e-2, rtol=2e-2)
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_rms_norm_cuda(self, dtype):
|
||||
if fastvideo_kernel_ops is None or not hasattr(fastvideo_kernel_ops, "rms_norm_cuda"):
|
||||
pytest.skip("rms_norm_cuda not available")
|
||||
|
||||
shape = (16, 128)
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
dim = shape[-1]
|
||||
w = torch.randn(dim, dtype=dtype, device="cuda")
|
||||
eps = 1e-5
|
||||
|
||||
# C++ implementation
|
||||
# Signature: rms_norm_cuda(Input, eps, Weight, Output) -> Output
|
||||
out_cuda = torch.empty_like(x)
|
||||
fastvideo_kernel_ops.rms_norm_cuda(x, eps, w, out_cuda)
|
||||
|
||||
# Reference
|
||||
out_ref = rms_norm_ref(x, w, eps)
|
||||
|
||||
torch.testing.assert_close(out_cuda, out_ref, atol=1e-2, rtol=1e-2)
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("shape", [(2, 16, 128), (4, 32, 256)])
|
||||
def test_layer_norm_triton(self, dtype, shape):
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
dim = shape[-1]
|
||||
eps = 1e-5
|
||||
|
||||
# With affine
|
||||
w = torch.randn(dim, dtype=dtype, device="cuda")
|
||||
b = torch.randn(dim, dtype=dtype, device="cuda")
|
||||
|
||||
# Triton implementation
|
||||
# Note: layernorm returns tuple now
|
||||
res = turbodiffusion_ops.layernorm(x, w, b, eps, elementwise_affine=True)
|
||||
if isinstance(res, tuple):
|
||||
out_triton = res[0]
|
||||
else:
|
||||
out_triton = res
|
||||
out_triton = out_triton.to(dtype)
|
||||
|
||||
# Reference
|
||||
ln = torch.nn.LayerNorm(dim, eps=eps, elementwise_affine=True, dtype=dtype).cuda()
|
||||
ln.weight.data.copy_(w)
|
||||
ln.bias.data.copy_(b)
|
||||
out_ref = ln(x)
|
||||
|
||||
torch.testing.assert_close(out_triton, out_ref, atol=1e-2, rtol=1e-2)
|
||||
|
||||
# Without affine
|
||||
res_no_affine = turbodiffusion_ops.layernorm(x, None, None, eps, elementwise_affine=False)
|
||||
if isinstance(res_no_affine, tuple):
|
||||
out_triton_no_affine = res_no_affine[0]
|
||||
else:
|
||||
out_triton_no_affine = res_no_affine
|
||||
out_triton_no_affine = out_triton_no_affine.to(dtype)
|
||||
ln_no_affine = torch.nn.LayerNorm(dim, eps=eps, elementwise_affine=False, dtype=dtype).cuda()
|
||||
out_ref_no_affine = ln_no_affine(x)
|
||||
|
||||
torch.testing.assert_close(out_triton_no_affine, out_ref_no_affine, atol=1e-2, rtol=1e-2)
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@pytest.mark.parametrize("shape", [(2, 16, 128)])
|
||||
def test_layer_norm_backward(self, dtype, shape):
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda", requires_grad=True)
|
||||
dim = shape[-1]
|
||||
eps = 1e-5
|
||||
w = torch.randn(dim, dtype=dtype, device="cuda", requires_grad=True)
|
||||
b = torch.randn(dim, dtype=dtype, device="cuda", requires_grad=True)
|
||||
|
||||
# Triton implementation via Function
|
||||
y = turbodiffusion_ops.LayerNormFunction.apply(x, w, b, eps, True)
|
||||
loss = y.sum()
|
||||
loss.backward()
|
||||
|
||||
x_grad = x.grad
|
||||
w_grad = w.grad
|
||||
b_grad = b.grad
|
||||
|
||||
# Reference
|
||||
x_ref = x.detach().clone().requires_grad_()
|
||||
ln = torch.nn.LayerNorm(dim, eps=eps, elementwise_affine=True, dtype=dtype).cuda()
|
||||
ln.weight.data.copy_(w.detach())
|
||||
ln.bias.data.copy_(b.detach())
|
||||
|
||||
y_ref = ln(x_ref)
|
||||
loss_ref = y_ref.sum()
|
||||
loss_ref.backward()
|
||||
|
||||
torch.testing.assert_close(x_grad, x_ref.grad, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(w_grad, ln.weight.grad, atol=1e-2, rtol=1e-2)
|
||||
torch.testing.assert_close(b_grad, ln.bias.grad, atol=1e-2, rtol=1e-2)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_layer_norm_cuda(self, dtype):
|
||||
if fastvideo_kernel_ops is None or not hasattr(fastvideo_kernel_ops, "layer_norm_cuda"):
|
||||
pytest.skip("layer_norm_cuda not available")
|
||||
|
||||
shape = (16, 128)
|
||||
x = torch.randn(shape, dtype=dtype, device="cuda")
|
||||
dim = shape[-1]
|
||||
eps = 1e-5
|
||||
w = torch.randn(dim, dtype=dtype, device="cuda")
|
||||
b = torch.randn(dim, dtype=dtype, device="cuda")
|
||||
|
||||
# C++ implementation
|
||||
# Signature: layer_norm_cuda(Input, eps, W, B, Output) -> Output
|
||||
out_cuda = torch.empty_like(x)
|
||||
fastvideo_kernel_ops.layer_norm_cuda(x, eps, w, b, out_cuda)
|
||||
|
||||
# Reference
|
||||
ln = torch.nn.LayerNorm(dim, eps=eps, elementwise_affine=True, dtype=dtype).cuda()
|
||||
ln.weight.data.copy_(w)
|
||||
ln.bias.data.copy_(b)
|
||||
out_ref = ln(x)
|
||||
|
||||
torch.testing.assert_close(out_cuda, out_ref, atol=1e-2, rtol=1e-2)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user