Compare commits

..
Author SHA1 Message Date
SolitaryThinker 19c1d164c3 fix 2025-12-12 22:51:21 +00:00
233 changed files with 3958 additions and 20886 deletions
+58 -19
View File
@@ -22,7 +22,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- TEST_TYPE=encoder
@@ -35,7 +35,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- TEST_TYPE=vae
@@ -129,7 +129,11 @@ steps:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "csrc/attn/video_sparse_attn/**"
- "csrc/attn/video_sparse_attn/tk/**"
- "csrc/attn/video_sparse_attn/setup.py"
- "csrc/attn/video_sparse_attn/config_vsa.py"
- "csrc/attn/video_sparse_attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -141,7 +145,10 @@ steps:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "csrc/attn/sliding_tile_attn/**"
- "csrc/attn/sliding_tile_attn/setup.py"
- "csrc/attn/sliding_tile_attn/config_sta.py"
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -152,16 +159,48 @@ steps:
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "csrc/attn/sliding_tile_attn/**"
- "csrc/attn/sliding_tile_attn/setup.py"
- "csrc/attn/sliding_tile_attn/config_sta.py"
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Kernel Tests"
label: "Precision Tests STA"
env:
- TEST_TYPE=kernel_tests
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "csrc/attn/video_sparse_attn/**"
- "csrc/attn/video_sparse_attn/tk/**"
- "csrc/attn/tests/test_vsa.py"
- "csrc/attn/video_sparse_attn/setup.py"
- "csrc/attn/video_sparse_attn/config_vsa.py"
- "csrc/attn/video_sparse_attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- TEST_TYPE=precision_vsa
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VMoBA"
env:
- TEST_TYPE=precision_vmoba
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/vmoba/**"
- "fastvideo/attention/backends/vmoba.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
@@ -183,14 +222,14 @@ steps:
- TEST_TYPE=unit_test
agents:
queue: "default"
# - path:
# - "scripts/lora_extraction/**"
# - "pyproject.toml"
# - "docker/Dockerfile.python3.12"
# config:
# command: "timeout 90m .buildkite/scripts/pr_test.sh"
# label: "LoRA Extraction Tests"
# env:
# - TEST_TYPE=lora_extraction
# agents:
# queue: "default"
- path:
- "scripts/lora_extraction/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
+11 -3
View File
@@ -93,9 +93,13 @@ case "$TEST_TYPE" in
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"kernel_tests")
log "Running kernel tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_kernel_tests"
"precision_sta")
log "Running precision STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
;;
"precision_vsa")
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
;;
"inference_lora")
log "Running LoRA tests..."
@@ -114,6 +118,10 @@ case "$TEST_TYPE" in
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
;;
"precision_vmoba")
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
@@ -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/
-2
View File
@@ -30,8 +30,6 @@ env
**/build/
**.pyc
**.txt
*.log
weights/
# Distribution / packaging
build/
+6 -5
View File
@@ -1,6 +1,7 @@
[submodule "fastvideo-kernel/include/tk"]
path = fastvideo-kernel/include/tk
[submodule "csrc/attn/video_sparse_attn/tk"]
path = csrc/attn/video_sparse_attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
[submodule "csrc/attn/sliding_tile_attn/tk"]
path = csrc/attn/sliding_tile_attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
[submodule "fastvideo-kernel/include/cutlass"]
path = fastvideo-kernel/include/cutlass
url = https://github.com/NVIDIA/cutlass.git
+1 -1
View File
@@ -4,7 +4,7 @@ default_stages:
exclude: |
(?x)(
fastvideo/third_party/.*|
fastvideo-kernel/.*|
csrc/.*|
assets/.*|
tests/.*|
demo/.*|
+2 -2
View File
@@ -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)
+3 -6
View File
@@ -38,8 +38,7 @@ python -m benchmarks.fvd.cli \
--num-frames 32 \
--clip-strategy random \
--batch-size 32 \
--seed 42 \
--extractor clip
--seed 42
```
**Standard protocols:**
@@ -52,8 +51,6 @@ python -m benchmarks.fvd.cli \
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
```
This would use i3d model by default as the feature extractor
**Feature caching** (speed up repeated evaluations):
```bash
@@ -61,7 +58,7 @@ python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f \
--cache-real-features fvd-cache/extractor_name # Directory path (will save/load fvd-cache/extractor_name/extractor-name_real_features.pkl)
--cache-real-features cache/real # Directory path (will save/load cache/real/real_features.pkl)
```
Run `python -m benchmarks.fvd.cli --help` for all options.
@@ -86,7 +83,6 @@ batch_size=32, # GPU batch size
device='cuda', # cuda|cpu
cache_real_features=None, # Cache path for speed
seed=42, # Reproducibility
extractor='i3d', # i3d|clip|videomae
```
## Programmatic Usage
@@ -101,6 +97,7 @@ print(f"FVD: {results['fvd']:.2f}")
## Notes
- I3D model auto-downloads from Hugging Face on first run
- Requires minimum 10 frames per clip
- Supports both video files (.mp4, .avi, etc.) and frame directories
- `--cache-real-features` expects a **directory path** (e.g., `cache/real`), it will automatically create/load `real_features.pkl` inside that directory
+1 -4
View File
@@ -13,8 +13,7 @@ from .fvd import (
compute_statistics,
FVDConfig,
)
from .feature_extractors import (BaseFeatureExtractor, I3DFeatureExtractor,
load_extractor)
from .i3d_model import I3DFeatureExtractor
from .video_utils import (
load_video_auto,
sample_clips_from_video,
@@ -28,9 +27,7 @@ __all__ = [
'compute_frechet_distance',
'compute_statistics',
'FVDConfig',
'BaseFeatureExtractor',
'I3DFeatureExtractor',
'load_extractor',
'load_video_auto',
'sample_clips_from_video',
'load_video_clips_streaming',
+119 -41
View File
@@ -1,64 +1,122 @@
import argparse
import json
import sys
import traceback
from pathlib import Path
from .fvd import compute_fvd_with_config, FVDConfig
def main() -> int:
parser = argparse.ArgumentParser(
description='Compute Fréchet Video Distance (FVD)')
description='Compute Fréchet Video Distance (FVD)',
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Standard FVD2048_16f protocol
python -m fastvideo.benchmarks.fvd.cli \\
--real-path data/real/ \\
--gen-path outputs/gen/ \\
--protocol fvd2048_16f
# Custom configuration
python -m fastvideo.benchmarks.fvd.cli \\
--real-path data/real/ \\
--gen-path outputs/gen/ \\
--num-videos 1024 \\
--num-frames 32 \\
--clip-strategy random \\
--frame-stride 2
""")
# Required arguments
parser.add_argument('--real-path',
type=str,
required=True,
help='Path to real videos')
help='Path to real videos directory')
parser.add_argument('--gen-path',
type=str,
required=True,
help='Path to generated videos')
help='Path to generated videos directory')
# Extractor selection
parser.add_argument('--extractor',
type=str,
default='i3d',
choices=['i3d', 'clip', 'videomae'],
help='Feature extractor model to use (default: i3d)')
# Reproducibility
parser.add_argument(
'--seed',
type=int,
default=None,
help='Random seed for reproducibility (np.random, random, torch)')
# Standard args
parser.add_argument('--seed',
type=int,
default=None,
help='Random seed for reproducibility')
# Protocol presets
parser.add_argument('--protocol',
type=str,
default=None,
choices=['fvd2048_16f', 'fvd2048_128f', 'quick_test'],
choices=[
'fvd2048_16f', 'fvd2048_128f',
'fvd2048_128f_subsample8', 'quick_test'
],
help='Use standard protocol (overrides other settings)')
# Video selection
parser.add_argument('--num-videos',
type=int,
default=2048,
help='Number of videos to use')
help='Number of videos to use (default: 2048)')
# Clip sampling
parser.add_argument('--num-frames',
type=int,
default=16,
help='Number of frames per clip')
parser.add_argument('--clip-strategy',
type=str,
default='beginning',
help='Clip sampling strategy')
help='Number of frames per clip (default: 16)')
parser.add_argument('--num-clips',
type=int,
default=1,
help='Number of clips per video (default: 1)')
parser.add_argument(
'--clip-strategy',
type=str,
default='beginning',
choices=['beginning', 'random', 'uniform', 'middle', 'sliding', 'all'],
help='Clip sampling strategy (default: beginning)')
parser.add_argument(
'--frame-stride',
type=int,
default=1,
help='Frame stride for FPS subsampling (default: 1, no subsampling)')
parser.add_argument('--temporal-stride',
type=int,
default=1,
help='Temporal stride for sliding window (default: 1)')
# Data processing
parser.add_argument('--no-frame-dirs',
action='store_true',
help='Disable frame directory support')
# Computation
parser.add_argument('--batch-size',
type=int,
default=32,
help='Batch size for feature extraction')
help='Batch size for feature extraction (default: 32)')
parser.add_argument('--device',
type=str,
default='cuda',
help='Device to use (cuda or cpu)')
choices=['cuda', 'cpu'],
help='Device to use (default: cuda)')
# Caching
parser.add_argument('--cache-real-features',
type=str,
default=None,
help='Path to cache real video features')
parser.add_argument('--i3d-model-path',
type=str,
default=None,
help='Custom cache path for I3D model')
# Output
parser.add_argument('--output',
type=str,
default='fvd_results.json',
help='Output JSON file (default: fvd_results.json)')
parser.add_argument('--quiet',
action='store_true',
help='Suppress progress output')
@@ -70,36 +128,56 @@ def main() -> int:
protocol_map = {
'fvd2048_16f': FVDConfig.fvd2048_16f,
'fvd2048_128f': FVDConfig.fvd2048_128f,
'fvd2048_128f_subsample8': FVDConfig.fvd2048_128f_subsample8,
'quick_test': FVDConfig.quick_test,
}
config = protocol_map[args.protocol]()
# Apply overrides
# Override device and caching from args
config.device = args.device
config.cache_real_features = args.cache_real_features
config.extractor_model = args.extractor # Apply extractor arg
config.i3d_model_path = args.i3d_model_path
config.batch_size = args.batch_size
config.seed = args.seed
else:
config = FVDConfig(
num_videos=args.num_videos,
num_frames_per_clip=args.num_frames,
extractor_model=args.extractor, # Apply extractor arg
clip_strategy=args.clip_strategy,
batch_size=args.batch_size,
device=args.device,
cache_real_features=args.cache_real_features,
seed=args.seed)
# Custom config from args
config = FVDConfig(num_videos=args.num_videos,
num_frames_per_clip=args.num_frames,
num_clips_per_video=args.num_clips,
clip_strategy=args.clip_strategy,
frame_stride=args.frame_stride,
temporal_stride=args.temporal_stride,
support_frame_dirs=not args.no_frame_dirs,
batch_size=args.batch_size,
device=args.device,
cache_real_features=args.cache_real_features,
i3d_model_path=args.i3d_model_path,
seed=args.seed)
# Compute FVD
try:
_ = compute_fvd_with_config(
args.real_path, # noqa: F841
args.gen_path,
config,
verbose=not args.quiet)
results = compute_fvd_with_config(real_videos=args.real_path,
gen_videos=args.gen_path,
config=config,
verbose=not args.quiet)
# Save results
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w') as f:
json.dump(results, f, indent=2)
print(f"\nResults saved to {output_path}")
print(f"FVD: {results['fvd']:.2f}")
print(f"Protocol: {results['protocol']}")
return 0
except Exception as e:
print(f"Error: {e}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
import traceback
traceback.print_exc()
return 1
-264
View File
@@ -1,264 +0,0 @@
"""
Pluggable Feature Extractors for FVD Computation.
Supports I3D (standard), CLIP, and VideoMAE via a common interface.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from abc import ABC, abstractmethod
from huggingface_hub import hf_hub_download
from tqdm import tqdm
try:
from transformers import CLIPModel, CLIPProcessor, VideoMAEModel
TRANSFORMERS_AVAILABLE = True
except ImportError:
TRANSFORMERS_AVAILABLE = False
class BaseFeatureExtractor(ABC, nn.Module):
"""Abstract base class for all video feature extractors."""
def __init__(self, device: str = 'cuda'):
super().__init__()
self.device = torch.device(
device if torch.cuda.is_available() else 'cpu')
@property
@abstractmethod
def feature_dim(self) -> int:
"""Dimension of the output feature vector."""
pass
@abstractmethod
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Args:
videos: [B, T, C, H, W] in [0, 255] range.
Returns:
Preprocessed tensor ready for the model.
"""
pass
@abstractmethod
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
"""
Extract features for a single batch.
Args:
videos: [B, T, C, H, W] (raw input)
Returns:
Features: [B, feature_dim]
"""
pass
@torch.no_grad()
def extract_features(self,
videos: torch.Tensor,
batch_size: int = 32,
verbose: bool = True) -> torch.Tensor:
"""
Extract features for a large tensor of videos by batching.
"""
N = len(videos)
all_features = []
iterator = range(0, N, batch_size)
if verbose:
iterator = tqdm(
iterator,
desc=f"Extracting features ({self.__class__.__name__})")
for i in iterator:
batch = videos[i:i + batch_size].to(self.device)
features = self.extract_features_batch(batch)
all_features.append(features.cpu())
return torch.cat(all_features, dim=0)
# 1. I3D Extractor (The Standard FVD Metric)
class I3DFeatureExtractor(BaseFeatureExtractor):
REPO_ID = 'flateon/FVD-I3D-torchscript'
MODEL_FILENAME = 'i3d_torchscript.pt'
def __init__(self, device: str = 'cuda', cache_dir: str | None = None):
super().__init__(device)
self.cache_dir = cache_dir
self.model = self._load_model()
self.model.eval()
self.model.to(self.device)
@property
def feature_dim(self) -> int:
return 400
def _load_model(self) -> torch.nn.Module:
try:
model_path = hf_hub_download(repo_id=self.REPO_ID,
filename=self.MODEL_FILENAME,
cache_dir=self.cache_dir)
return torch.jit.load(model_path, map_location=self.device)
except Exception as e:
raise RuntimeError(f"Failed to load I3D model: {e}") from e
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""Standard I3D preprocessing: Resize to 224, Norm to [-1, 1]."""
B, T, C, H, W = videos.shape
if T < 10:
raise ValueError(f"I3D requires at least 10 frames, got {T}")
# Normalize to [0, 1]
if videos.max() > 1.0:
videos = videos / 255.0
# Scale to [-1, 1]
videos = videos * 2.0 - 1.0
# Resize to 224x224
if H != 224 or W != 224:
videos = videos.reshape(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.reshape(B, T, C, 224, 224)
# [B, T, C, H, W] -> [B, C, T, H, W]
return videos.permute(0, 2, 1, 3, 4).contiguous()
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
batch = self.preprocess(videos)
# TorchScript I3D returns raw logits when return_features=True
return self.model(batch,
rescale=False,
resize=False,
return_features=True)
# 2. CLIP Extractor (Semantic/Content Quality)
class CLIPFeatureExtractor(BaseFeatureExtractor):
def __init__(self,
device: str = 'cuda',
model_name: str = "openai/clip-vit-base-patch32"):
if not TRANSFORMERS_AVAILABLE:
raise ImportError(
"Please install transformers: pip install transformers")
super().__init__(device)
self.processor = CLIPProcessor.from_pretrained(model_name)
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
self.model.eval()
self._feature_dim = self.model.config.projection_dim
@property
def feature_dim(self) -> int:
return self._feature_dim
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
# Ensure values are [0, 255]
if videos.max() <= 1.0:
videos = videos * 255.0
return videos.to(torch.uint8)
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
# Input: [B, T, C, H, W]
B, T, C, H, W = videos.shape
videos = self.preprocess(videos)
# Flatten B*T to treat frames as images
images = videos.view(B * T, C, H, W)
# HF Processor
inputs = self.processor(images=images,
return_tensors="pt",
padding=True)
inputs = {k: v.to(self.device) for k, v in inputs.items()}
# Extract features [B*T, Dim]
outputs = self.model.get_image_features(**inputs)
# Reshape [B, T, Dim] and Average Pooling over time
outputs = outputs.view(B, T, -1)
return outputs.mean(dim=1)
# 3. VideoMAE Extractor (Structure/Motion Quality)
class VideoMAEFeatureExtractor(BaseFeatureExtractor):
def __init__(self,
device: str = 'cuda',
model_name: str = "MCG-NJU/videomae-base"):
if not TRANSFORMERS_AVAILABLE:
raise ImportError(
"Please install transformers: pip install transformers")
super().__init__(device)
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
self.model.eval()
self.register_buffer(
'mean',
torch.tensor([0.485, 0.456, 0.406],
device=self.device).view(1, 1, 3, 1, 1))
self.register_buffer(
'std',
torch.tensor([0.229, 0.224, 0.225],
device=self.device).view(1, 1, 3, 1, 1))
@property
def feature_dim(self) -> int:
return self.model.config.hidden_size
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Efficient GPU-based preprocessing.
Input: [B, T, C, H, W] in range [0, 255]
"""
B, T, C, H, W = videos.shape
# 1. Resize to 224x224
if H != 224 or W != 224:
videos = videos.view(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.view(B, T, C, 224, 224)
# 2. Normalize to [0, 1]
if videos.dtype != torch.float32:
videos = videos.float()
if videos.max() > 1.0:
videos = videos / 255.0
# 3. Apply ImageNet Mean/Std
return (videos - self.mean) / self.std
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
# Input: [B, T, C, H, W]
# Fast GPU Preprocessing
pixel_values = self.preprocess(videos)
# Forward pass
outputs = self.model(pixel_values)
# Global Average Pooling of last hidden state [B, T_patches, 768] -> [B, 768]
return outputs.last_hidden_state.mean(dim=1)
# Factory
def load_extractor(name: str, device: str = 'cuda') -> BaseFeatureExtractor:
name = name.lower()
if name == 'i3d':
return I3DFeatureExtractor(device)
elif name == 'clip':
return CLIPFeatureExtractor(device)
elif name == 'videomae':
return VideoMAEFeatureExtractor(device)
else:
raise ValueError(
f"Unknown extractor: {name}. Options: i3d, clip, videomae")
+150 -108
View File
@@ -5,7 +5,8 @@ from pathlib import Path
from collections.abc import Iterator
import pickle
from dataclasses import dataclass, field
from .feature_extractors import BaseFeatureExtractor, load_extractor
from .i3d_model import I3DFeatureExtractor
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
@@ -54,9 +55,6 @@ class FVDConfig:
# Video selection
num_videos: int = 2048
# Feature Extractor Selection
extractor_model: str = 'i3d' # Options: 'i3d', 'clip', 'videomae'
# Clip sampling
num_frames_per_clip: int = 16
num_clips_per_video: int = 1
@@ -87,7 +85,11 @@ class FVDConfig:
@classmethod
def fvd2048_16f(cls) -> 'FVDConfig':
"""Standard FVD protocol: 2048 videos, 16 frames, beginning clip."""
"""
Standard FVD protocol: 2048 videos, 16 frames, beginning clip.
most common FVD configuration used in papers
"""
return cls(num_videos=2048,
num_frames_per_clip=16,
clip_strategy='beginning',
@@ -101,6 +103,18 @@ class FVDConfig:
clip_strategy='beginning',
use_streaming=True)
@classmethod
def fvd2048_128f_subsample8(cls) -> 'FVDConfig':
"""
Long video with FPS subsampling: 2048 videos, 128 frames (every 8th).
Used for very long videos - samples every 8th frame
"""
return cls(num_videos=2048,
num_frames_per_clip=16,
frame_stride=8,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def quick_test(cls) -> 'FVDConfig':
"""Quick test config: 100 videos, 16 frames."""
@@ -110,13 +124,22 @@ class FVDConfig:
def to_dict(self) -> dict:
"""Export config to dict for logging"""
d = self.__dict__.copy()
d['clip_strategy'] = str(self.clip_strategy)
return d
return {
'num_videos': self.num_videos,
'num_frames_per_clip': self.num_frames_per_clip,
'num_clips_per_video': self.num_clips_per_video,
'clip_strategy': str(self.clip_strategy),
'frame_stride': self.frame_stride,
'temporal_stride': self.temporal_stride,
'batch_size': self.batch_size,
'device': self.device,
'seed': self.seed,
'use_streaming': self.use_streaming,
}
def __str__(self) -> str:
"""Human-readable protocol name"""
desc = f"FVD_{self.extractor_model.upper()}_{self.num_videos}_{self.num_frames_per_clip}f"
desc = f"FVD{self.num_videos}_{self.num_frames_per_clip}f"
if self.frame_stride > 1:
desc += f"_subsample{self.frame_stride}"
if self.num_clips_per_video > 1:
@@ -127,42 +150,57 @@ class FVDConfig:
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
extractor: BaseFeatureExtractor,
extractor: I3DFeatureExtractor,
batch_size: int = 32,
max_clips: int | None = None,
verbose: bool = True) -> np.ndarray:
"""
Extract features from a video clip generator using streaming.
Args:
video_generator: Iterator yielding clips [T, C, H, W]
extractor: I3D feature extractor
batch_size: Batch size for processing
max_clips: Maximum clips to process (for validation)
verbose: Show progress
Returns:
features: [N, 400] numpy array
"""
all_features = []
batch = []
clip_count = 0
if verbose:
print(f"Extracting features with batch_size={batch_size}...")
with torch.no_grad():
for clip_count, clip in enumerate(video_generator):
batch.append(clip)
for clip_count, clip in enumerate(video_generator):
batch.append(clip)
# Process batch when full
if len(batch) == batch_size:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features_batch(batch_tensor)
all_features.append(features.detach().cpu().numpy())
batch = []
if verbose and clip_count % (batch_size * 10) == 0:
print(f"Processed {clip_count} clips...")
if max_clips is not None and clip_count >= max_clips:
break
# Process remaining clips
if len(batch) > 0:
# Process batch when full
if len(batch) == batch_size:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features_batch(batch_tensor)
all_features.append(features.detach().cpu().numpy())
features = extractor.extract_features(batch_tensor,
batch_size=batch_size,
verbose=False)
all_features.append(features.cpu().numpy())
batch = [] # Clear batch
if verbose and clip_count % (batch_size * 10) == 0:
print(f"Processed {clip_count} clips...")
# Stop if we've reached max_clips
if max_clips is not None and clip_count >= max_clips:
break
# Process remaining clips
if len(batch) > 0:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features(batch_tensor,
batch_size=len(batch),
verbose=False)
all_features.append(features.cpu().numpy())
if len(all_features) == 0:
raise RuntimeError("No features extracted - check video loading")
@@ -176,17 +214,14 @@ def extract_features_streaming(video_generator: Iterator[torch.Tensor],
def load_or_compute_features(videos: str | Path | torch.Tensor,
extractor: BaseFeatureExtractor,
extractor: I3DFeatureExtractor,
config: FVDConfig,
cache_path: str | None = None,
cache_name: str = "real_features") -> np.ndarray:
"""Load features from cache or compute (with streaming support)"""
if cache_path is not None:
script_dir = Path(__file__).parent
cache_dir = script_dir / cache_path
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
cache_file = Path(cache_path) / f"{cache_name}.pkl"
if cache_file.exists():
print(f"Loading cached features from {cache_file}")
with open(cache_file, 'rb') as f:
@@ -194,25 +229,19 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
# Validate and limit based on config
max_features = config.num_videos * config.num_clips_per_video
if len(features) < max_features:
print(
f"WARNING: Cache has {len(features)} features but need {max_features}"
)
print("Cached features insufficient - will recompute...")
print("Recomputing features...")
elif len(features) > max_features:
print(
f"Using {max_features} features from cache (truncated from {len(features)})"
)
features = features[:max_features]
return features
else:
print(f"Using all {len(features)} cached features")
return features
print("Computing features from scratch...")
if isinstance(videos, (str | Path)):
# Compute features
if isinstance(videos, str | Path):
target_size = (224, 224) if config.resize_before_extraction else None
video_generator = load_video_clips_streaming(
@@ -233,7 +262,9 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
batch_size=config.batch_size,
max_clips=max_clips,
verbose=True)
else:
# Already a tensor
print(f"Extracting features from {len(videos)} video tensors...")
features = extractor.extract_features(videos,
batch_size=config.batch_size,
@@ -252,10 +283,9 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
# Cache features if requested
if cache_path is not None:
script_dir = Path(__file__).parent
cache_dir = script_dir / cache_path
cache_dir = Path(cache_path)
cache_dir.mkdir(parents=True, exist_ok=True)
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
cache_file = cache_dir / f"{cache_name}.pkl"
print(f"Caching features to {cache_file}")
with open(cache_file, 'wb') as f:
pickle.dump(features, f)
@@ -263,34 +293,79 @@ def load_or_compute_features(videos: str | Path | torch.Tensor,
return features
def compute_fvd(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
num_frames: int = 16,
batch_size: int = 32,
device: str = 'cuda',
num_videos: int | None = 2048,
cache_real_features: str | None = None,
i3d_model_path: str | None = None,
seed: int | None = None,
verbose: bool = True) -> float:
"""
Compute Fréchet Video Distance (FVD)
For advanced control, use compute_fvd_with_config() instead.
Args:
real_videos: Path to real videos or tensor [N, T, C, H, W]
gen_videos: Path to generated videos or tensor [N, T, C, H, W]
num_frames: Frames per video (default: 16)
batch_size: Batch size (default: 32)
device: 'cuda' or 'cpu' (default: 'cuda')
num_videos: Max videos (default: 2048)
cache_real_features: Cache path for real features
i3d_model_path: Custom I3D model cache path
seed: Random seed for reproducibility
verbose: Print progress
Returns:
FVD score (float). Lower is better.
"""
num_videos = num_videos if num_videos is not None else 2048
config = FVDConfig(
num_videos=num_videos,
num_frames_per_clip=num_frames,
batch_size=batch_size,
device=device,
cache_real_features=cache_real_features,
i3d_model_path=i3d_model_path,
seed=seed,
)
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
return result['fvd']
def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
config: FVDConfig,
verbose: bool = True) -> dict:
"""
Compute FVD using a standardized configuration.
This is the recommended way to compute FVD for reproducibility.
Args:
real_videos: Path or tensors
gen_videos: Path or tensors
config: FVDConfig specifying protocol
verbose: Print progress
Returns:
results: Dictionary with:
- 'fvd': FVD score (float)
- 'protocol': Protocol name (str)
- 'model': Feature extractor model name (str)
- 'config': Configuration dict
Example:
>>> config = FVDConfig.fvd2048_16f()
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
Compute FVD using a standardized configuration.
This is the recommended way to compute FVD for reproducibility.
Args:
real_videos: Path or tensors
gen_videos: Path or tensors
config: FVDConfig specifying protocol
verbose: Print progress
Returns:
results: Dictionary with:
- 'fvd': FVD score (float)
- 'protocol': Protocol name (str)
- 'config': Configuration dict
Example:
>>> config = FVDConfig.fvd2048_16f()
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
>>> print(f"Protocol: {results['protocol']}") # "FVD2048_16f"
"""
# Seed for reproducibility
if config.seed is not None:
import random as _rnd
@@ -303,20 +378,18 @@ def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
if verbose:
print("=" * 70)
print(f"Computing FVD with protocol: {config}")
print(f"Model: {config.extractor_model.upper()}")
print("=" * 70)
print("\nConfiguration:")
for key, value in config.to_dict().items():
print(f" {key}: {value}")
print()
# Initialize Extractor using Factory
# Initialize I3D
if verbose:
print(
f"\nInitializing {config.extractor_model.upper()} model on {config.device}..."
)
print(f"\nInitializing I3D model on {config.device}...")
extractor = load_extractor(config.extractor_model, device=config.device)
extractor = I3DFeatureExtractor(device=config.device,
cache_dir=config.i3d_model_path)
# Extract features
if verbose:
@@ -361,45 +434,14 @@ def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
if verbose:
print(f"\n{'='*70}")
print(f"FVD Score ({config.extractor_model.upper()}): {fvd:.4f}")
print(f"FVD Score: {fvd:.4f}")
print(f"Protocol: {config}")
print(f"{'='*70}\n")
results = {
'fvd': fvd,
'protocol': str(config),
'model': config.extractor_model,
'config': config.to_dict(),
}
return results
def compute_fvd(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
num_frames: int = 16,
batch_size: int = 32,
device: str = 'cuda',
num_videos: int | None = 2048,
cache_real_features: str | None = None,
i3d_model_path: str | None = None,
seed: int | None = None,
verbose: bool = True) -> float:
"""
Backward compatibility wrapper for computing FVD (defaults to I3D).
"""
num_videos = num_videos if num_videos is not None else 2048
config = FVDConfig(
num_videos=num_videos,
num_frames_per_clip=num_frames,
extractor_model='i3d', # Default to I3D
batch_size=batch_size,
device=device,
cache_real_features=cache_real_features,
i3d_model_path=i3d_model_path,
seed=seed,
)
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
return result['fvd']
+17 -37
View File
@@ -1,53 +1,33 @@
import sys
from pathlib import Path
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config
root_dir = Path(__file__).parent.parent.parent
sys.path.insert(0, str(root_dir))
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config # noqa: E402
def main() -> None:
# Get script directory
script_dir = Path(__file__).parent.resolve()
# Define directories
clip_strategy = 'beginning' # Options: 'uniform', 'random', 'beginning', 'end', 'all'
cfg = FVDConfig(
num_videos=650,
num_frames_per_clip=16,
num_clips_per_video=1,
clip_strategy=clip_strategy,
frame_stride=1,
batch_size=32,
device='cuda',
seed=42,
cache_real_features=str(script_dir / f'fvd-cache/{clip_strategy}'),
)
real_dir = "benchmarks/data/real_videos"
gen_dir = "benchmarks/data/generated_videos"
# Compare all 3 models
models_to_test = ['i3d', 'clip', 'videomae']
print(f"\n{'='*60}")
print("STARTING COMPARISON BENCHMARK")
print(f"{'='*60}")
for model_name in models_to_test:
print(f"\n>>> Running evaluation with {model_name.upper()}...")
try:
cfg = FVDConfig(
num_videos=650,
num_frames_per_clip=16,
extractor_model=model_name,
clip_strategy='beginning',
device='cuda',
seed=42,
# Use separate cache folders for each model to avoid conflicts
cache_real_features=str(script_dir / f'fvd-cache/{model_name}'),
)
results = compute_fvd_with_config(real_dir,
gen_dir,
cfg,
verbose=False)
print(f"FVD: {results['fvd']}\nModel: {results['model']}")
except Exception as e:
print(f"{model_name.upper()} Failed: {e}")
print(f"\n{'='*60}")
print("BENCHMARK COMPLETE")
print(f"{'='*60}")
results = compute_fvd_with_config(real_dir, gen_dir, cfg, verbose=True)
print(f"FVD = {results['fvd']:.2f}")
if __name__ == '__main__':
+1 -1
View File
@@ -1,7 +1,7 @@
#!/bin/bash
# 1. Install missing dependency
pip install -q opencv-python-headless transformers huggingface_hub
pip install -q opencv-python-headless
# 2. Run FVD script
python benchmarks/fvd/run_fvd.py
+113
View File
@@ -0,0 +1,113 @@
# Attention Kernel Used in FastVideo
## Video Sparse Attention (VSA)
### Installation
We support H100 (via TK) and any other GPU (via triton) for VSA.
```bash
git submodule update --init --recursive
python setup_vsa.py install
```
If you encounter error during installation, try below:
Install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
(If you use CUDA12.8)
```bash
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
### Verify if you have successfully installed
```bash
# test numerical
python tests/test_vsa.py
# (For H100) test speed
python benchmarks/bench_vsa_hopper.py
```
bench_vsa_hopper.py should print something like this:
```bash
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
=== BLOCK SPARSE ATTENTION BENCHMARK ===
Block Sparse Forward - TFLOPS: 5622.26
Block Sparse Backward - TFLOPS: 3865.68
```
## Sliding Tile Attention (STA)
We only support H100 for STA.
```bash
git submodule update --init --recursive
python setup_sta.py install
```
### Usage
End-2-end inference with FastVideo:
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
If you want to use sliding tile attention in your custom model:
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
# a tile is a cube of size (6, 8, 8)
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
# text_length: int ranging from 0 to 256
# If your attention contains text token (Hunyuan)
out = sliding_tile_attention(q, k, v, window_size, text_length)
# If your attention does not contain text token (StepVideo)
out = sliding_tile_attention(q, k, v, window_size, 0, False)
```
### Test
```bash
python tests/test_sta.py # test STA
python tests/test_vsa.py # test VSA
```
### Benchmark
```bash
python benchmarks/bench_sta.py
```
### How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
## Why is STA Fast?
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
STA removes mixed blocks.
<div align="center">
<img src=../../assets/sliding_tile_attn_map.png width="80%"/>
</div>
## Acknowledgement
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
+145
View File
@@ -0,0 +1,145 @@
import os
from collections import defaultdict
import matplotlib.pyplot as plt
import numpy as np
import torch
from st_attn import sliding_tile_attention
from triton.testing import do_bench
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
assert mode in ["fwd", "bwd", "fwd_bwd"]
f = 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
def compute_TFLOPS(flops, ms):
flops = flops / 1e12
ms = ms / 1e3
return flops / ms
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
# # Warmup for forward pass
# for _ in range(10):
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# # Time the forward pass
# for i in range(10):
# start_events_fwd[i].record()
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# end_events_fwd[i].record()
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
# time_us_fwd = np.mean(times_fwd) * 1000
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass (ms): {ms:.2f}")
print(f"Average TFLOPS: {tflops_fwd}")
print("-" * 60)
# torch.cuda.empty_cache()
# torch.cuda.synchronize()
# # Prepare for timing backward pass
# start_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# end_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# # Warmup for backward pass
# for _ in range(10):
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# # Time the backward pass
# for i in range(10):
# start_events_bwd[i].record()
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# end_events_bwd[i].record()
# torch.cuda.synchronize()
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
# time_us_bwd = np.mean(times_bwd) * 1000
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass(ms): {ms:.2f}")
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
return results
def plot_results(results):
os.makedirs('benchmark_results', exist_ok=True)
for mode in ['fwd', 'bwd']:
for (D, causal), values in results[mode].items():
seq_lens = [x[0] for x in values]
tflops = [x[1] for x in values]
plt.figure(figsize=(10, 6))
bars = plt.bar(range(len(seq_lens)), tflops, tick_label=seq_lens)
plt.xlabel('Sequence Length')
plt.ylabel('TFLOPS')
plt.title(f'{mode.upper()} Pass - Head Dim: {D}, Causal: {causal}')
plt.grid(True)
# Adding the numerical y value on top of each bar
for bar in bars:
yval = bar.get_height()
plt.text(bar.get_x() + bar.get_width() / 2, yval, round(yval, 2), ha='center', va='bottom')
filename = f'benchmark_results/{mode}_D{D}_causal{causal}.png'
plt.savefig(filename)
plt.close()
# Example list of configurations to test
configurations = [
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
# (16, 16, 768*8, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768, 128, True),
# (16, 16, 768*2, 128, True),
# (16, 16, 768*4, 128, True),
# (16, 16, 768*8, 128, True),
# (16, 16, 768*16, 128, True),
# (16, 32, 768, 64, False),
# (16, 32, 768*2, 64, False),
# (16, 32, 768*4, 64, False),
# (16, 32, 768*8, 64, False),
# (16, 32, 768*16, 64, False),
# (16, 32, 768, 64, True),
# (16, 32, 768*2, 64, True),
# (16, 32, 768*4, 64, True),
# (16, 32, 768*8, 64, True),
# (16, 32, 768*16, 64, True),
]
results = benchmark_attention(configurations)
# plot_results(results)
+224
View File
@@ -0,0 +1,224 @@
import torch
import argparse
from triton.testing import do_bench
from vsa import block_sparse_fwd, block_sparse_bwd
from vsa import BLOCK_M, BLOCK_N
import triton
import numpy as np
import random
def set_seed(seed: int = 42):
# Python random module
random.seed(seed)
# NumPy
np.random.seed(seed)
# PyTorch
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
return parser.parse_args()
def create_input_tensors(batch, head, seq_len, headdim):
"""Create random input tensors for attention."""
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
return q, k, v
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
Args:
bs: batch size
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key-value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to k).
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
Binary mask where 1 indicates attention connection.
"""
# Ensure k is not larger than num_kv_blocks
k = min(k, num_kv_blocks)
# Create random scores for sampling
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
# Get top-k indices for each q block
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
# sort q2k_block_sparse_index
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
# All q blocks attend to exactly k kv blocks
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
# Create the corresponding mask
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
# Fill in the mask based on the indices
for b in range(bs):
for head in range(h):
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx]
block_sparse_mask[b, head, q_idx, kv_indices] = True
# Create the reverse mapping (k2q)
# First, initialize lists to collect q indices for each kv block
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
# Populate the lists based on q2k mapping
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
for kv_idx in kv_indices:
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
# Find the maximum number of q blocks that attend to any kv block
max_q_per_kv = 0
for flat_idx in range(bs * h):
for kv_idx in range(num_kv_blocks):
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
# Create tensors for k2q mapping
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
dtype=torch.int32, device=device)
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
dtype=torch.int32, device=device)
# Fill the tensors
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for kv_idx in range(num_kv_blocks):
q_indices = k2q_indices_list[flat_idx][kv_idx]
num_q = len(q_indices)
k2q_block_sparse_num[b, head, kv_idx] = num_q
if num_q > 0:
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
q_indices, dtype=torch.int32, device=device)
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
"""Benchmark block sparse attention forward and backward passes."""
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
# Forward pass
# Warm-up run
variable_block_sizes = torch.ones(q2k_block_sparse_index.shape[2], device=q.device).int() * BLOCK_M
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
torch.cuda.synchronize()
# Benchmark forward
fwd_time = do_bench(
lambda: block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes),
warmup=5,
rep=20,
quantiles=None
)
sparse_tflops = flops / fwd_time * 1e-12 * 1e3
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
# Backward pass
grad_output = torch.randn_like(o)
# Warm-up runs
for _ in range(5):
block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
torch.cuda.synchronize()
# Benchmark backward
bwd_time = do_bench(
lambda: block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes),
warmup=5,
rep=20,
quantiles=None
)
bwd_flops = 2.5 * flops # Approximation
sparse_bwd_tflops = bwd_flops / bwd_time * 1e-12 * 1e3
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
return sparse_tflops, sparse_bwd_tflops
def main():
args = parse_arguments()
set_seed(42)
# Extract parameters
batch = args.batch_size
head = args.num_heads
headdim = args.head_dim
print(f"Block Sparse Attention Benchmark")
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
# Test with different sequence lengths
for seq_len in args.seq_lengths:
# Skip very long sequences if they might cause OOM
if seq_len > 16384 and batch > 1:
continue
print("="*100)
print(f"\nSequence length: {seq_len}")
# Calculate theoretical FLOPs for attention
flops = 4 * batch * head * headdim * seq_len * seq_len
# Create input tensors
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
# Setup block sparse parameters
num_q_blocks = seq_len // BLOCK_M
num_kv_blocks = seq_len // BLOCK_N
# Determine k value (number of kv blocks per q block)
topk = args.topk
if topk is None:
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
topk = max(1, topk)
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
# Generate block sparse pattern
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
# Benchmark block sparse attention
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
)
# Print results
print("\n=== PERFORMANCE RESULTS ===")
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
if __name__ == "__main__":
main()
+217
View File
@@ -0,0 +1,217 @@
import torch
import argparse
import triton.testing
from vsa import block_sparse_attn
from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
def set_seed(seed: int = 42):
# Python random module
random.seed(seed)
# NumPy
np.random.seed(seed)
# PyTorch
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
return parser.parse_args()
def create_input_tensors(batch, head, seq_len, headdim):
"""Create random input tensors for attention."""
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
return q, k, v
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
Args:
bs: batch size
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key-value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to k).
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
Binary mask where 1 indicates attention connection.
"""
# Ensure k is not larger than num_kv_blocks
k = min(k, num_kv_blocks)
# Create random scores for sampling
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
# Get top-k indices for each q block
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
# sort q2k_block_sparse_index
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
# All q blocks attend to exactly k kv blocks
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
# Create the corresponding mask
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
# Fill in the mask based on the indices
for b in range(bs):
for head in range(h):
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx]
block_sparse_mask[b, head, q_idx, kv_indices] = True
# Create the reverse mapping (k2q)
# First, initialize lists to collect q indices for each kv block
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
# Populate the lists based on q2k mapping
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
for kv_idx in kv_indices:
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
# Find the maximum number of q blocks that attend to any kv block
max_q_per_kv = 0
for flat_idx in range(bs * h):
for kv_idx in range(num_kv_blocks):
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
# Create tensors for k2q mapping
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
dtype=torch.int32, device=device)
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
dtype=torch.int32, device=device)
# Fill the tensors
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for kv_idx in range(num_kv_blocks):
q_indices = k2q_indices_list[flat_idx][kv_idx]
num_q = len(q_indices)
k2q_block_sparse_num[b, head, kv_idx] = num_q
if num_q > 0:
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
q_indices, dtype=torch.int32, device=device)
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
"""Benchmark block sparse attention forward+backward pass."""
print("\n=== BLOCK SPARSE ATTENTION FORWARD+BACKWARD BENCHMARK ===")
# Combined forward+backward pass
# Warm-up run
q_fwd = q.clone().requires_grad_(True)
k_fwd = k.clone().requires_grad_(True)
v_fwd = v.clone().requires_grad_(True)
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
grad_output = torch.randn_like(o)
o.backward(grad_output)
torch.cuda.synchronize()
# Benchmark forward+backward
def forward_backward_fn():
q_fwd = q.clone().requires_grad_(True)
k_fwd = k.clone().requires_grad_(True)
v_fwd = v.clone().requires_grad_(True)
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
grad_output = torch.randn_like(o)
o.backward(grad_output)
total_time = triton.testing.do_bench(
forward_backward_fn,
warmup=25,
rep=100,
return_mode='mean'
)
# Total flops for forward + backward (forward + 2.5x backward approximation)
total_flops = flops + 2.5 * flops # 3.5x the forward flops
sparse_tflops = total_flops / total_time * 1e-12 * 1e3
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_tflops:.2f}")
return sparse_tflops
def main():
args = parse_arguments()
set_seed(42)
# Extract parameters
batch = args.batch_size
head = args.num_heads
headdim = args.head_dim
print(f"Block Sparse Attention Benchmark")
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
# Test with different sequence lengths
for seq_len in args.seq_lengths:
# Skip very long sequences if they might cause OOM
if seq_len > 16384 and batch > 1:
continue
print("="*100)
print(f"\nSequence length: {seq_len}")
# Calculate theoretical FLOPs for attention
flops = 4 * batch * head * headdim * seq_len * seq_len
# Create input tensors
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
# Setup block sparse parameters
num_q_blocks = seq_len // BLOCK_M
num_kv_blocks = seq_len // BLOCK_N
# Determine k value (number of kv blocks per q block)
topk = args.topk
if topk is None:
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
topk = max(1, topk)
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
# Generate block sparse pattern
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
# Benchmark block sparse attention
sparse_fwd = benchmark_block_sparse_attention(
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
)
# Print results
print("\n=== PERFORMANCE RESULTS ===")
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_fwd:.2f}")
if __name__ == "__main__":
main()
+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config_sta.py
+96
View File
@@ -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.
+15
View File
@@ -0,0 +1,15 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'st_attn': {
'source_files': {
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['st_attn']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
+76
View File
@@ -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"])
+23
View File
@@ -0,0 +1,23 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ST_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
@@ -0,0 +1,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]
@@ -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']}")
+156
View File
@@ -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.")
+54
View File
@@ -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
+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config_vsa.py
+61
View File
@@ -0,0 +1,61 @@
# Attention Kernel Used in FastVideo
## Video Sparse Attention (VSA)
### Installation
We support H100 (via TK) and any other GPU (via triton) for VSA.
```bash
pip install vsa
```
Install from source:
```bash
git submodule update --init --recursive
python setup.py install
```
If you encounter error during installation, try below:
Install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
(If you use CUDA12.8)
```bash
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
### Verify if you have successfully installed
```bash
# test numerical
python ../tests/test_vsa.py
# (For H100) test speed
python ../benchmarks/bench_vsa_hopper.py
```
bench_vsa_hopper.py should print something like this:
```bash
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
=== BLOCK SPARSE ATTENTION BENCHMARK ===
Block Sparse Forward - TFLOPS: 5622.26
Block Sparse Backward - TFLOPS: 3865.68
```
## Acknowledgement
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
+15
View File
@@ -0,0 +1,15 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'block_sparse': {
'source_files': {
'h100': 'vsa/block_sparse_h100.cu'
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['block_sparse']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
+81
View File
@@ -0,0 +1,81 @@
import os
import subprocess
from config_vsa import kernels, sources, target
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
target = target.lower()
# Package metadata
PACKAGE_NAME = "vsa"
VERSION = "0.0.3"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
# Set environment variables
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
python_include = subprocess.check_output(['python', '-c',
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
torch_include = subprocess.check_output([
'python', '-c',
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip()
print('vsa root:', tk_root)
print('Python include:', python_include)
print('Torch include directories:', torch_include)
# CUDA flags
cuda_flags = [
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
] + torch_include.split()
cpp_flags = ['-std=c++20', '-O3']
if target == 'h100':
cuda_flags.append('-DKITTENS_HOPPER')
cuda_flags.append('-arch=sm_90a')
else:
raise ValueError(f'Target {target} not supported')
source_files = ['vsa.cpp']
for k in kernels:
if target not in sources[k]['source_files']:
raise KeyError(f'Target {target} not found in source files for kernel {k}')
if isinstance(sources[k]['source_files'][target], list):
source_files.extend(sources[k]['source_files'][target])
else:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
ext_modules = [
CUDAExtension('vsa_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
]
setup(name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
ext_modules=ext_modules,
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.10',
install_requires=["torch>=2.5.0"])
+27
View File
@@ -0,0 +1,27 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_BLOCK_SPARSE
extern std::vector<torch::Tensor> block_sparse_attention_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
);
extern std::vector<torch::Tensor> block_sparse_attention_backward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_BLOCK_SPARSE
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
#endif
}
@@ -0,0 +1,80 @@
import torch
from typing import Tuple
block_sparse_attn=None
import torch
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
block_sparse_attn = block_sparse_attn_SM90
else:
from vsa.block_sparse_wrapper import block_sparse_attn_triton
block_sparse_fwd = None
block_sparse_bwd = None
block_sparse_attn = block_sparse_attn_triton
BLOCK_M = 64
BLOCK_N = 64
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight=None):
"""
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
NOTE: We assume q, k, v is zero padded!!
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements == 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = (q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
k_compress = (k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
v_compress = (v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
topK_indices = torch.topk(block_attn_score, topk, dim=-1).indices
block_mask = torch.zeros_like(block_attn_score, dtype=torch.bool).scatter_(-1, topK_indices, True)
output_select, _ = block_sparse_attn(q, k, v, block_mask, variable_block_sizes)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
@@ -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
@@ -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,9 +1,9 @@
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
@@ -26,7 +26,6 @@ def topk_index_to_map_kernel(
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
@@ -60,7 +59,6 @@ def map_to_index_kernel(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
@@ -108,7 +106,6 @@ def topk_index_to_map(index: torch.Tensor,
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
+32
View File
@@ -0,0 +1,32 @@
# Attention Kernel Used in FastVideo
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
### Installation
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
### Usage
You can use `moba_attn_varlen` in the following ways:
**Install from source:**
```bash
python setup.py install
```
**Import after installation:**
```python
from vmoba import moba_attn_varlen
```
**Or import directly from the project root:**
```python
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
```
### Verify if you have successfully installed
```bash
python csrc/attn/vmoba_attn/vmoba/vmoba.py
```
+26
View File
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
from setuptools import find_packages, setup
PACKAGE_NAME = "vmoba"
VERSION = "0.0.0"
AUTHOR = "JianzongWu"
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
URL = "https://github.com/KwaiVGI/VMoBA"
setup(
name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
classifiers=[
"Programming Language :: Python :: 3",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[
"flash-attn >= 2.7.1",
]
)
@@ -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
):
+2
View File
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
+11 -5
View File
@@ -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
+11 -5
View File
@@ -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
+11 -4
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir 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
+9 -3
View File
@@ -55,12 +55,18 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
./build.sh
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
EXPOSE 22
-58
View File
@@ -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
View File
@@ -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
+3 -5
View File
@@ -2,11 +2,9 @@
writing-mode: sideways-lr;
white-space: nowrap;
max-width: 0;
}
/* Keep header cell paragraph content tight (avoid CSS nesting for compatibility) */
.vertical-table-header th.head:not(.stub) p {
margin: 0;
p {
margin: 0;
}
}
/* Image sizing classes */
-179
View File
@@ -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`.
-53
View File
@@ -1,53 +0,0 @@
# FastVideo Attention Kernels
FastVideo provides highly optimized custom attention kernels to accelerate video generation.
## Supported Kernels
* **[Video Sparse Attention (VSA)](vsa/index.md)**: Sparse attention mechanism selecting top-k blocks.
* **[Sliding Tile Attention (STA)](sta/index.md)**: Optimized attention for window-based video generation.
## General Build Instructions
These instructions apply to building the `fastvideo-kernel` package from source, which includes both STA and VSA kernels.
### Prerequisites
* **PyTorch**: 2.5.0+
* **CUDA**: 12.4+ (12.8 recommended for best performance)
* **C++ Compiler**: GCC 11+ (C++20 support required for ThunderKittens)
Install system dependencies:
```bash
sudo apt update
sudo apt install -y gcc-11 g++-11 clang-11 ninja-build
# Set gcc-11 as default
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
```
Set up your CUDA environment variables (adjust version as needed):
```bash
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
### Compile and Install
Clone the repository and build the kernel:
```bash
# Clone recursively to get ThunderKittens submodule
git clone --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.
-36
View File
@@ -1,36 +0,0 @@
# Sliding Tile Attention (STA)
Optimized attention for window-based video generation (e.g., HunyuanVideo).
## Installation
STA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
## Usage
```python
from fastvideo_kernel import sliding_tile_attention
# q, k, v: [batch_size, num_heads, seq_length, head_dim]
# window_size: List of (t, h, w) tiles. Tile size is (6, 8, 8).
# text_length: Number of text tokens (0-256)
out = sliding_tile_attention(
q, k, v,
window_size=[(3, 3, 3)], # Example window
text_length=256
)
```
## Citation
If you use Sliding Tile Attention in your research, please cite:
```bibtex
@article{zhang2025fast,
title={Fast video generation with sliding tile attention},
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
journal={arXiv preprint arXiv:2502.04507},
year={2025}
}
```
-36
View File
@@ -1,36 +0,0 @@
# Video Sparse Attention (VSA)
Sparse attention mechanism selecting top-k blocks.
## Installation
VSA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
## Usage
```python
from fastvideo_kernel import video_sparse_attn
# q, k, v: [batch_size, num_heads, seq_len, head_dim]
# variable_block_sizes: Number of valid tokens per block
# topk: Number of blocks to attend
output = video_sparse_attn(
q, k, v,
variable_block_sizes=block_sizes,
topk=32
)
```
## Citation
If you use Video Sparse Attention in your research, please cite:
```bibtex
@article{zhang2025vsa,
title={Vsa: Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
```
+1 -1
View File
@@ -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:
+2 -2
View File
@@ -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.
+3 -5
View File
@@ -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`.
+2 -2
View File
@@ -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
+18 -41
View File
@@ -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.)
+1 -1
View File
@@ -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
+16
View File
@@ -0,0 +1,16 @@
# 🔍 Demo
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
<div style="text-align: center;">
<video controls width="800">
<source src="https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747" type="video/mp4">
Your browser does not support the video tag.
</video>
</div>
You can run STA using the following command:
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
@@ -0,0 +1,65 @@
# 🔧 Installation
You can install the Sliding Tile Attention package using
```
pip install st_attn
```
# Building from Source
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have 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
```
+1 -1
View File
@@ -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
```
+1 -1
View File
@@ -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.
-42
View File
@@ -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[@]}"
-159
View File
@@ -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()
-187
View File
@@ -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.
-6
View File
@@ -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
-69
View File
@@ -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}
}
```
-38
View File
@@ -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 &params,
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;
}
-35
View File
@@ -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"
View File
@@ -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