Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
59c63b4c75 | ||
|
|
bb208e9bab | ||
|
|
18b2c2673d | ||
|
|
65d47c5d29 | ||
|
|
ed02c87a4e | ||
|
|
d14cd2b07a | ||
|
|
86fde63d5d |
@@ -222,14 +222,3 @@ 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"
|
||||
|
||||
@@ -75,7 +75,7 @@ case "$TEST_TYPE" in
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
@@ -126,10 +126,6 @@ case "$TEST_TYPE" in
|
||||
log "Running unit tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
|
||||
;;
|
||||
"lora_extraction")
|
||||
log "Running LoRA extraction tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
|
||||
@@ -3,18 +3,8 @@ name: Deploy Documentation
|
||||
on:
|
||||
push:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/docs.yml'
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/docs.yml'
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
@@ -1,236 +0,0 @@
|
||||
name: Publish FastVideo Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/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 csrc/fastvideo_kernel
|
||||
# Get current commit's version from pyproject.toml
|
||||
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
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', '3.13']
|
||||
torch-cuda:
|
||||
- torch-version: '2.5.1'
|
||||
cuda-version: '12.4.1'
|
||||
torch-cuda-short: 'cu124'
|
||||
- torch-version: '2.6.0'
|
||||
cuda-version: '12.6.3'
|
||||
torch-cuda-short: 'cu126'
|
||||
- torch-version: '2.7.1'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
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
|
||||
|
||||
cd csrc/fastvideo_kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/fastvideo_kernel
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}-py${{ matrix.python-version }}
|
||||
path: csrc/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: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton
|
||||
|
||||
cd csrc/fastvideo_kernel
|
||||
git submodule update --init --recursive
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/fastvideo_kernel/dist/
|
||||
@@ -30,8 +30,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/c7g1qdD" target="_blank"> <b> WeChat </b> </a> |
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/TM8JyJCd" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -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)
|
||||
|
||||
@@ -1,103 +0,0 @@
|
||||
# FVD (Fréchet Video Distance) Benchmark
|
||||
|
||||
Evaluate generated video quality using FVD with the I3D feature extractor.
|
||||
|
||||
## Quick Start
|
||||
|
||||
**Run the benchmark:**
|
||||
|
||||
```bash
|
||||
bash benchmarks/scripts/run.sh
|
||||
```
|
||||
|
||||
That's it! The script auto-installs dependencies and runs the benchmark.
|
||||
|
||||
**To customize:** Edit `benchmarks/fvd/run_fvd.py` to change:
|
||||
- Video paths (`real_dir`, `gen_dir`)
|
||||
- Number of videos, frames, sampling strategy
|
||||
- Device, batch size, caching, etc.
|
||||
|
||||
## Advanced Usage (CLI)
|
||||
|
||||
For more control without editing Python files, use the CLI.
|
||||
|
||||
**First-time setup** (one-time per pod/environment):
|
||||
|
||||
```bash
|
||||
bash benchmarks/scripts/setup_fvd.sh
|
||||
```
|
||||
|
||||
Then run any configuration you want:
|
||||
|
||||
```bash
|
||||
# Custom configuration
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--num-videos 1024 \
|
||||
--num-frames 32 \
|
||||
--clip-strategy random \
|
||||
--batch-size 32 \
|
||||
--seed 42
|
||||
```
|
||||
|
||||
**Standard protocols:**
|
||||
|
||||
```bash
|
||||
# Use predefined protocols
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
|
||||
```
|
||||
|
||||
**Feature caching** (speed up repeated evaluations):
|
||||
|
||||
```bash
|
||||
python -m benchmarks.fvd.cli \
|
||||
--real-path data/real/ \
|
||||
--gen-path outputs/gen/ \
|
||||
--protocol fvd2048_16f \
|
||||
--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.
|
||||
|
||||
## Available Protocols
|
||||
|
||||
- `fvd2048_16f` - Standard (2048 videos, 16 frames)
|
||||
- `fvd2048_128f` - Long videos (128 frames)
|
||||
- `fvd2048_128f_subsample8` - Subsampled long videos
|
||||
- `quick_test` - Fast testing (10 videos)
|
||||
|
||||
## Configuration Options
|
||||
|
||||
Key options in `FVDConfig`:
|
||||
|
||||
```python
|
||||
num_videos=2048, # Videos to evaluate
|
||||
num_frames_per_clip=16, # Frames per clip
|
||||
clip_strategy='beginning', # beginning|random|uniform|middle|sliding
|
||||
frame_stride=1, # Frame subsampling
|
||||
batch_size=32, # GPU batch size
|
||||
device='cuda', # cuda|cpu
|
||||
cache_real_features=None, # Cache path for speed
|
||||
seed=42, # Reproducibility
|
||||
```
|
||||
|
||||
## Programmatic Usage
|
||||
|
||||
```python
|
||||
from benchmarks.fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
config = FVDConfig.fvd2048_16f() # or custom config
|
||||
results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
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,35 +0,0 @@
|
||||
"""
|
||||
FastVideo Frechet Video Distance (FVD) Benchmark Module.
|
||||
>>> from fastvideo.benchmarks.fvd import compute_fvd_with_config, FVDConfig
|
||||
>>> config = FVDConfig.fvd2048_16f() # Standard protocol
|
||||
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
|
||||
>>> print(f"FVD: {results['fvd']:.2f}")
|
||||
"""
|
||||
|
||||
from .fvd import (
|
||||
compute_fvd,
|
||||
compute_fvd_with_config,
|
||||
compute_frechet_distance,
|
||||
compute_statistics,
|
||||
FVDConfig,
|
||||
)
|
||||
from .i3d_model import I3DFeatureExtractor
|
||||
from .video_utils import (
|
||||
load_video_auto,
|
||||
sample_clips_from_video,
|
||||
load_video_clips_streaming,
|
||||
ClipSamplingStrategy,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
'compute_fvd',
|
||||
'compute_fvd_with_config',
|
||||
'compute_frechet_distance',
|
||||
'compute_statistics',
|
||||
'FVDConfig',
|
||||
'I3DFeatureExtractor',
|
||||
'load_video_auto',
|
||||
'sample_clips_from_video',
|
||||
'load_video_clips_streaming',
|
||||
'ClipSamplingStrategy',
|
||||
]
|
||||
@@ -1,185 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
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)',
|
||||
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 directory')
|
||||
parser.add_argument('--gen-path',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to generated videos directory')
|
||||
|
||||
# Reproducibility
|
||||
parser.add_argument(
|
||||
'--seed',
|
||||
type=int,
|
||||
default=None,
|
||||
help='Random seed for reproducibility (np.random, random, torch)')
|
||||
|
||||
# Protocol presets
|
||||
parser.add_argument('--protocol',
|
||||
type=str,
|
||||
default=None,
|
||||
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 (default: 2048)')
|
||||
|
||||
# Clip sampling
|
||||
parser.add_argument('--num-frames',
|
||||
type=int,
|
||||
default=16,
|
||||
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 (default: 32)')
|
||||
parser.add_argument('--device',
|
||||
type=str,
|
||||
default='cuda',
|
||||
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')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create config
|
||||
if args.protocol:
|
||||
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]()
|
||||
|
||||
# Override device and caching from args
|
||||
config.device = args.device
|
||||
config.cache_real_features = args.cache_real_features
|
||||
config.i3d_model_path = args.i3d_model_path
|
||||
config.batch_size = args.batch_size
|
||||
config.seed = args.seed
|
||||
else:
|
||||
# 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:
|
||||
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)
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
sys.exit(main())
|
||||
@@ -1,447 +0,0 @@
|
||||
import numpy as np
|
||||
import scipy.linalg
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from collections.abc import Iterator
|
||||
import pickle
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from .i3d_model import I3DFeatureExtractor
|
||||
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
|
||||
|
||||
|
||||
def compute_statistics(features: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Compute mean and covariance."""
|
||||
mu = np.mean(features, axis=0)
|
||||
sigma = np.cov(features, rowvar=False)
|
||||
return mu, sigma
|
||||
|
||||
|
||||
def compute_frechet_distance(mu1: np.ndarray,
|
||||
sigma1: np.ndarray,
|
||||
mu2: np.ndarray,
|
||||
sigma2: np.ndarray,
|
||||
eps: float = 1e-6) -> float:
|
||||
"""
|
||||
Compute Fréchet distance between two Gaussians.
|
||||
"""
|
||||
sigma1 = sigma1 + eps * np.eye(sigma1.shape[0])
|
||||
sigma2 = sigma2 + eps * np.eye(sigma2.shape[0])
|
||||
|
||||
diff = mu1 - mu2
|
||||
mean_distance = np.sum(diff**2)
|
||||
|
||||
trace_sum = np.trace(sigma1 + sigma2)
|
||||
|
||||
covmean = scipy.linalg.sqrtm(sigma1 @ sigma2)
|
||||
|
||||
if np.iscomplexobj(covmean):
|
||||
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
|
||||
print(
|
||||
f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
|
||||
covmean = covmean.real
|
||||
|
||||
trace_product = np.trace(covmean)
|
||||
|
||||
fvd = mean_distance + trace_sum - 2 * trace_product
|
||||
|
||||
return float(fvd)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FVDConfig:
|
||||
# default configuration for FVD computation:
|
||||
|
||||
# Video selection
|
||||
num_videos: int = 2048
|
||||
|
||||
# Clip sampling
|
||||
num_frames_per_clip: int = 16
|
||||
num_clips_per_video: int = 1
|
||||
clip_strategy: str | ClipSamplingStrategy = 'beginning'
|
||||
|
||||
# Temporal subsampling
|
||||
frame_stride: int = 1 # 1=no subsampling, 2=every 2nd, 8=every 8th
|
||||
temporal_stride: int = 1 # For sliding window clips
|
||||
|
||||
# Data processing
|
||||
video_extensions: list[str] = field(
|
||||
default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
|
||||
support_frame_dirs: bool = True
|
||||
|
||||
# Computation
|
||||
batch_size: int = 32
|
||||
device: str = 'cuda'
|
||||
|
||||
use_streaming: bool = True
|
||||
resize_before_extraction: bool = True
|
||||
|
||||
# Caching
|
||||
cache_real_features: str | None = None
|
||||
i3d_model_path: str | None = None
|
||||
|
||||
# Reproducibility
|
||||
seed: int | None = None
|
||||
|
||||
@classmethod
|
||||
def fvd2048_16f(cls) -> 'FVDConfig':
|
||||
"""
|
||||
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',
|
||||
use_streaming=True)
|
||||
|
||||
@classmethod
|
||||
def fvd2048_128f(cls) -> 'FVDConfig':
|
||||
"""Long video protocol: 2048 videos, 128 frames."""
|
||||
return cls(num_videos=2048,
|
||||
num_frames_per_clip=128,
|
||||
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."""
|
||||
return cls(num_videos=100,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning')
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Export config to dict for logging"""
|
||||
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.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:
|
||||
desc += f"_{self.num_clips_per_video}clips"
|
||||
if self.clip_strategy != 'beginning':
|
||||
desc += f"_{self.clip_strategy}"
|
||||
return desc
|
||||
|
||||
|
||||
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
|
||||
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}...")
|
||||
|
||||
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_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")
|
||||
|
||||
features = np.concatenate(all_features, axis=0)
|
||||
|
||||
if verbose:
|
||||
print(f"Extracted {len(features)} feature vectors")
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def load_or_compute_features(videos: str | Path | torch.Tensor,
|
||||
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:
|
||||
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:
|
||||
features = pickle.load(f)
|
||||
|
||||
# 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("Recomputing features...")
|
||||
elif len(features) > max_features:
|
||||
features = features[:max_features]
|
||||
return features
|
||||
else:
|
||||
return features
|
||||
|
||||
# Compute features
|
||||
if isinstance(videos, str | Path):
|
||||
target_size = (224, 224) if config.resize_before_extraction else None
|
||||
|
||||
video_generator = load_video_clips_streaming(
|
||||
videos,
|
||||
num_frames=config.num_frames_per_clip,
|
||||
max_videos=config.num_videos,
|
||||
clip_strategy=config.clip_strategy,
|
||||
frame_stride=config.frame_stride,
|
||||
num_clips_per_video=config.num_clips_per_video,
|
||||
video_extensions=config.video_extensions,
|
||||
support_frame_dirs=config.support_frame_dirs,
|
||||
target_size=target_size,
|
||||
verbose=True)
|
||||
|
||||
max_clips = config.num_videos * config.num_clips_per_video
|
||||
features = extract_features_streaming(video_generator,
|
||||
extractor,
|
||||
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,
|
||||
verbose=True)
|
||||
features = features.numpy()
|
||||
|
||||
# Validate feature count
|
||||
expected_count = config.num_videos * config.num_clips_per_video
|
||||
if len(features) < expected_count:
|
||||
raise ValueError(
|
||||
f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
|
||||
f"Found fewer videos than expected. Check your video directory.")
|
||||
elif len(features) > expected_count:
|
||||
print(f"Truncating {len(features)} features to {expected_count}")
|
||||
features = features[:expected_count]
|
||||
|
||||
# Cache features if requested
|
||||
if cache_path is not None:
|
||||
cache_dir = Path(cache_path)
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
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)
|
||||
|
||||
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)
|
||||
- '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
|
||||
_rnd.seed(config.seed)
|
||||
np.random.seed(config.seed)
|
||||
torch.manual_seed(config.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(config.seed)
|
||||
|
||||
if verbose:
|
||||
print("=" * 70)
|
||||
print(f"Computing FVD with protocol: {config}")
|
||||
print("=" * 70)
|
||||
print("\nConfiguration:")
|
||||
for key, value in config.to_dict().items():
|
||||
print(f" {key}: {value}")
|
||||
print()
|
||||
|
||||
# Initialize I3D
|
||||
if verbose:
|
||||
print(f"\nInitializing I3D model on {config.device}...")
|
||||
|
||||
extractor = I3DFeatureExtractor(device=config.device,
|
||||
cache_dir=config.i3d_model_path)
|
||||
|
||||
# Extract features
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Extracting REAL video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
real_features = load_or_compute_features(
|
||||
videos=real_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=config.cache_real_features,
|
||||
cache_name="real_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Extracting GENERATED video features...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
gen_features = load_or_compute_features(videos=gen_videos,
|
||||
extractor=extractor,
|
||||
config=config,
|
||||
cache_path=None,
|
||||
cache_name="gen_features")
|
||||
|
||||
if verbose:
|
||||
print(f"\nReal videos/clips: {len(real_features)}")
|
||||
print(f"Generated videos/clips: {len(gen_features)}")
|
||||
print(f"\n{'='*70}")
|
||||
print("Computing statistics...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
mu_real, sigma_real = compute_statistics(real_features)
|
||||
mu_gen, sigma_gen = compute_statistics(gen_features)
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print("Computing Fréchet distance...")
|
||||
print(f"{'='*70}")
|
||||
|
||||
fvd = compute_frechet_distance(mu_real, sigma_real, mu_gen, sigma_gen)
|
||||
|
||||
if verbose:
|
||||
print(f"\n{'='*70}")
|
||||
print(f"FVD Score: {fvd:.4f}")
|
||||
print(f"Protocol: {config}")
|
||||
print(f"{'='*70}\n")
|
||||
|
||||
results = {
|
||||
'fvd': fvd,
|
||||
'protocol': str(config),
|
||||
'config': config.to_dict(),
|
||||
}
|
||||
|
||||
return results
|
||||
@@ -1,142 +0,0 @@
|
||||
"""I3D Feature Extractor for FVD Computation"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pathlib import Path
|
||||
from huggingface_hub import hf_hub_download
|
||||
from tqdm import tqdm
|
||||
from contextlib import suppress
|
||||
|
||||
|
||||
class I3DFeatureExtractor(nn.Module):
|
||||
"""
|
||||
I3D feature extractor for FVD computation.
|
||||
Extracts 400-dimensional features from videos using I3D model
|
||||
trained on Kinetics-400.
|
||||
"""
|
||||
|
||||
REPO_ID = 'flateon/FVD-I3D-torchscript'
|
||||
MODEL_FILENAME = 'i3d_torchscript.pt'
|
||||
|
||||
def __init__(self,
|
||||
device: str = 'cuda',
|
||||
cache_dir: str | Path | None = None):
|
||||
super().__init__()
|
||||
|
||||
self.device_str = device
|
||||
if device == 'cuda' and not torch.cuda.is_available():
|
||||
print(
|
||||
"Warning: CUDA requested but not available – falling back to CPU"
|
||||
)
|
||||
self.device = torch.device('cpu')
|
||||
else:
|
||||
self.device = torch.device(device)
|
||||
|
||||
self.cache_dir: str | None
|
||||
if cache_dir is not None:
|
||||
self.cache_dir = str(Path(cache_dir).resolve())
|
||||
else:
|
||||
self.cache_dir = None # Use HF default cache
|
||||
|
||||
self.model = self._load_model()
|
||||
self.model.eval()
|
||||
|
||||
with suppress(Exception):
|
||||
self.model.to(self.device)
|
||||
|
||||
def _load_model(self) -> torch.nn.Module:
|
||||
"""Download and load I3D TorchScript model from Hugging Face Hub."""
|
||||
print(f"Loading I3D model from Hugging Face Hub ({self.REPO_ID})...")
|
||||
|
||||
try:
|
||||
# Download model from Hugging Face Hub
|
||||
model_path = hf_hub_download(repo_id=self.REPO_ID,
|
||||
filename=self.MODEL_FILENAME,
|
||||
cache_dir=self.cache_dir)
|
||||
|
||||
# Load directly to chosen device
|
||||
model = torch.jit.load(model_path, map_location=self.device)
|
||||
print("I3D model loaded successfully")
|
||||
return model
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Preprocess videos for I3D.
|
||||
|
||||
Args:
|
||||
videos: [B, T, C, H, W], values in [0, 255]
|
||||
|
||||
Returns:
|
||||
Preprocessed videos [B, C, T, 224, 224] (normalized and resized)
|
||||
"""
|
||||
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 needed
|
||||
if videos.max() > 1.0:
|
||||
videos = videos / 255.0
|
||||
|
||||
# Resize to 224x224 if needed
|
||||
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)
|
||||
|
||||
# Convert to [B, C, T, H, W] format
|
||||
videos = videos.permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
return videos
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_features(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32,
|
||||
verbose: bool = True) -> torch.Tensor:
|
||||
"""
|
||||
Extract I3D features
|
||||
|
||||
Args:
|
||||
videos: [N, T, C, H, W], values in [0, 255]
|
||||
batch_size: Batch size for processing
|
||||
verbose: Show progress bar
|
||||
|
||||
Returns:
|
||||
Features [N, 400]
|
||||
"""
|
||||
N = len(videos)
|
||||
all_features = []
|
||||
|
||||
iterator = range(0, N, batch_size)
|
||||
if verbose:
|
||||
iterator = tqdm(iterator, desc="Extracting I3D features")
|
||||
|
||||
for i in iterator:
|
||||
batch = videos[i:i + batch_size].to(self.device)
|
||||
batch = self.preprocess(batch) # Now returns [B, C, T, H, W]
|
||||
|
||||
# Use the HF model without rescale/resize (we handle it in preprocess)
|
||||
features = self.model(batch,
|
||||
rescale=False,
|
||||
resize=False,
|
||||
return_features=True)
|
||||
|
||||
all_features.append(features.cpu())
|
||||
|
||||
return torch.cat(all_features, dim=0)
|
||||
|
||||
def __call__(self,
|
||||
videos: torch.Tensor,
|
||||
batch_size: int = 32) -> torch.Tensor:
|
||||
return self.extract_features(videos, batch_size=batch_size)
|
||||
@@ -1,34 +0,0 @@
|
||||
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))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# Get script directory
|
||||
script_dir = Path(__file__).parent.resolve()
|
||||
|
||||
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"
|
||||
|
||||
results = compute_fvd_with_config(real_dir, gen_dir, cfg, verbose=True)
|
||||
print(f"FVD = {results['fvd']:.2f}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,97 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
import sys
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import random
|
||||
from fvd import compute_fvd_with_config, FVDConfig
|
||||
|
||||
script_path = Path(__file__).resolve()
|
||||
fastvideo_root = script_path.parent.parent.parent
|
||||
sys.path.insert(0, str(fastvideo_root))
|
||||
|
||||
|
||||
def split_videos(video_dir: Path, n_per_subset: int = 128, seed: int = 42):
|
||||
subset_a = video_dir.parent / 'bair_full_subset_A'
|
||||
subset_b = video_dir.parent / 'bair_full_subset_B'
|
||||
|
||||
if subset_a.exists():
|
||||
shutil.rmtree(subset_a)
|
||||
if subset_b.exists():
|
||||
shutil.rmtree(subset_b)
|
||||
|
||||
subset_a.mkdir(parents=True)
|
||||
subset_b.mkdir(parents=True)
|
||||
|
||||
videos = sorted(video_dir.glob('*.mp4'))
|
||||
|
||||
random.seed(seed)
|
||||
shuffled = list(videos)
|
||||
random.shuffle(shuffled)
|
||||
|
||||
needed = n_per_subset * 2
|
||||
if len(shuffled) > needed:
|
||||
shuffled = shuffled[:needed]
|
||||
|
||||
mid = len(shuffled) // 2
|
||||
|
||||
print(f"\nSplitting {len(shuffled)} BAIR FULL videos:")
|
||||
print(f" Subset A: {mid} videos")
|
||||
print(f" Subset B: {len(shuffled) - mid} videos")
|
||||
|
||||
for v in shuffled[:mid]:
|
||||
shutil.copy2(v, subset_a / v.name)
|
||||
|
||||
for v in shuffled[mid:]:
|
||||
shutil.copy2(v, subset_b / v.name)
|
||||
|
||||
return subset_a, subset_b, mid
|
||||
|
||||
|
||||
def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
|
||||
config = FVDConfig(num_videos=num_videos,
|
||||
num_frames_per_clip=16,
|
||||
clip_strategy='beginning',
|
||||
batch_size=8,
|
||||
device='cuda',
|
||||
seed=42)
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST 1: Identity Test")
|
||||
print("=" * 70)
|
||||
|
||||
result1 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_a),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_identity = result1['fvd']
|
||||
print(f"\nIdentity FVD: {fvd_identity:.2f}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST 2: Real vs Real")
|
||||
print("=" * 70)
|
||||
|
||||
result2 = compute_fvd_with_config(real_videos=str(subset_a),
|
||||
gen_videos=str(subset_b),
|
||||
config=config,
|
||||
verbose=False)
|
||||
fvd_real = result2['fvd']
|
||||
print(f"\nReal vs Real FVD: {fvd_real:.2f}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("RESULTS")
|
||||
print("=" * 70)
|
||||
print(f"Identity: {fvd_identity:.2f}")
|
||||
print(f"Real vs Real: {fvd_real:.2f}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
bair_dir = Path('benchmarks/data/bair_full_videos')
|
||||
|
||||
subset_a, subset_b, count = split_videos(bair_dir,
|
||||
n_per_subset=128,
|
||||
seed=42)
|
||||
validate_fvd(subset_a, subset_b, count)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,490 +0,0 @@
|
||||
import torch
|
||||
import cv2
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from collections.abc import Iterator
|
||||
from tqdm import tqdm
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ClipSamplingStrategy(Enum):
|
||||
"""Clip sampling strategies for FVD evaluation."""
|
||||
BEGINNING = 'beginning' # Take first N frames (most common)
|
||||
RANDOM = 'random' # Random N consecutive frames
|
||||
UNIFORM = 'uniform' # Uniformly spaced frames across video
|
||||
MIDDLE = 'middle' # Middle N frames
|
||||
SLIDING = 'sliding' # Multiple sliding windows
|
||||
ALL = 'all' # All possible clips
|
||||
|
||||
|
||||
def _load_video_cv2(video_path: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform') -> torch.Tensor:
|
||||
"""
|
||||
Load video from video file using OpenCV.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file (MP4, AVI, MOV, MKV)
|
||||
num_frames: Number of frames to extract
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
video_path = str(video_path)
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
|
||||
if not cap.isOpened():
|
||||
raise RuntimeError(f"Cannot open video: {video_path}")
|
||||
|
||||
frames = []
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
|
||||
if num_frames is None:
|
||||
# Read all available frames
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
cap.release()
|
||||
if len(frames) == 0:
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
return frames
|
||||
|
||||
if total_frames == 0:
|
||||
raise RuntimeError(f"Video has 0 frames: {video_path}")
|
||||
|
||||
# Determine frame indices for sampling
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(
|
||||
total_frames)) + [total_frames - 1] * (num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0, total_frames - 1, num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
# Extract frames
|
||||
for idx in frame_indices:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
||||
ret, frame = cap.read()
|
||||
|
||||
if not ret:
|
||||
if len(frames) > 0:
|
||||
frames.append(frames[-1].copy())
|
||||
else:
|
||||
h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
frames.append(np.zeros((h, w, 3), dtype=np.uint8))
|
||||
continue
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
cap.release()
|
||||
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _load_video_from_frames(
|
||||
frame_dir: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform',
|
||||
frame_extensions: list[str] | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Load video from directory of frame images.
|
||||
|
||||
Args:
|
||||
frame_dir: Directory containing frames
|
||||
num_frames: Number of frames to sample
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
frame_extensions: Image file extensions to look for
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
if frame_extensions is None:
|
||||
frame_extensions = ['.jpg', '.png', '.jpeg', '.bmp']
|
||||
|
||||
frame_dir = Path(frame_dir)
|
||||
|
||||
if not frame_dir.exists():
|
||||
raise FileNotFoundError(f"Frame directory not found: {frame_dir}")
|
||||
|
||||
# Find all frames
|
||||
frame_files: list[Path] = []
|
||||
for ext in frame_extensions:
|
||||
frame_files.extend(frame_dir.glob(f"*{ext}"))
|
||||
|
||||
if len(frame_files) == 0:
|
||||
raise ValueError(
|
||||
f"No frames found in {frame_dir} with extensions {frame_extensions}"
|
||||
)
|
||||
|
||||
frame_files = sorted(frame_files, key=lambda x: x.name)
|
||||
total_frames = len(frame_files)
|
||||
|
||||
# Determine frame indices
|
||||
if num_frames is None:
|
||||
frame_indices = list(range(total_frames))
|
||||
else:
|
||||
if total_frames < num_frames:
|
||||
frame_indices = list(range(total_frames)) + [total_frames - 1] * (
|
||||
num_frames - total_frames)
|
||||
elif sample_strategy == 'uniform':
|
||||
frame_indices = np.linspace(0,
|
||||
total_frames - 1,
|
||||
num_frames,
|
||||
dtype=int).tolist()
|
||||
elif sample_strategy == 'random':
|
||||
frame_indices = sorted(
|
||||
np.random.choice(total_frames, num_frames, replace=False))
|
||||
else:
|
||||
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
|
||||
|
||||
# Load frames
|
||||
frames = []
|
||||
for idx in frame_indices:
|
||||
frame_path = frame_files[idx]
|
||||
frame = cv2.imread(str(frame_path))
|
||||
|
||||
if frame is None:
|
||||
raise RuntimeError(f"Failed to load frame: {frame_path}")
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame)
|
||||
|
||||
# Stack and convert to tensor
|
||||
frames = np.stack(frames) # [T, H, W, C]
|
||||
frames = torch.from_numpy(frames).permute(0, 3, 1,
|
||||
2).float() # [T, C, H, W]
|
||||
|
||||
return frames
|
||||
|
||||
|
||||
def _detect_video_format(path: str | Path) -> str:
|
||||
"""
|
||||
Detect if path is a video file or frame directory.
|
||||
|
||||
Returns:
|
||||
'video_file', 'frame_directory', or 'unknown'
|
||||
"""
|
||||
path = Path(path)
|
||||
|
||||
if path.is_file():
|
||||
return 'video_file'
|
||||
elif path.is_dir():
|
||||
# Check if contains image files
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
|
||||
for ext in image_extensions:
|
||||
if list(path.glob(f"*{ext}")):
|
||||
return 'frame_directory'
|
||||
return 'unknown'
|
||||
else:
|
||||
raise ValueError(f"Path does not exist: {path}")
|
||||
|
||||
|
||||
def load_video_auto(video_path: str | Path,
|
||||
num_frames: int | None = 16,
|
||||
sample_strategy: str = 'uniform') -> torch.Tensor:
|
||||
"""
|
||||
Automatically detect format and load video.
|
||||
|
||||
Supports:
|
||||
- Video files (MP4, AVI, MOV, MKV)
|
||||
- Frame directories (JPG, PNG)
|
||||
|
||||
Args:
|
||||
video_path: Path to video file or frame directory
|
||||
num_frames: Number of frames to extract
|
||||
sample_strategy: 'uniform' or 'random'
|
||||
|
||||
Returns:
|
||||
video: [T, C, H, W]
|
||||
"""
|
||||
format_type = _detect_video_format(video_path)
|
||||
|
||||
if format_type == 'video_file':
|
||||
return _load_video_cv2(video_path, num_frames, sample_strategy)
|
||||
elif format_type == 'frame_directory':
|
||||
return _load_video_from_frames(video_path, num_frames, sample_strategy)
|
||||
else:
|
||||
raise ValueError(f"Unknown video format at {video_path}")
|
||||
|
||||
|
||||
def sample_clips_from_video(
|
||||
video: torch.Tensor,
|
||||
num_frames_per_clip: int = 16,
|
||||
num_clips: int = 1,
|
||||
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
|
||||
frame_stride: int = 1,
|
||||
temporal_stride: int = 1) -> list[torch.Tensor]:
|
||||
"""
|
||||
Sample clips from a video with various strategies.
|
||||
|
||||
Args:
|
||||
video: [T, C, H, W] full video
|
||||
num_frames_per_clip: Frames per clip
|
||||
num_clips: Number of clips to extract
|
||||
strategy: ClipSamplingStrategy or string ('beginning', 'random', etc.)
|
||||
frame_stride: Skip frames (FPS control: 1=all, 2=every 2nd, 8=every 8th)
|
||||
temporal_stride: Stride between clips for sliding window
|
||||
|
||||
Returns:
|
||||
List of clips, each [num_frames_per_clip, C, H, W]
|
||||
|
||||
Examples:
|
||||
>>> # Beginning clip (most common for FVD)
|
||||
>>> clips = sample_clips_from_video(video, 16, strategy='beginning')
|
||||
|
||||
>>> # Multiple random clips
|
||||
>>> clips = sample_clips_from_video(video, 16, num_clips=4, strategy='random')
|
||||
|
||||
>>> # Subsample FPS by 2x (every 2nd frame)
|
||||
>>> clips = sample_clips_from_video(video, 16, frame_stride=2)
|
||||
|
||||
>>> # Sliding window with overlap
|
||||
>>> clips = sample_clips_from_video(video, 16, strategy='sliding', temporal_stride=8)
|
||||
"""
|
||||
# Convert string to enum if needed
|
||||
if isinstance(strategy, str):
|
||||
strategy = ClipSamplingStrategy(strategy)
|
||||
|
||||
T, C, H, W = video.shape
|
||||
|
||||
# Apply frame stride (FPS subsampling)
|
||||
if frame_stride > 1:
|
||||
video = video[::frame_stride]
|
||||
T = len(video)
|
||||
|
||||
effective_clip_length = num_frames_per_clip
|
||||
|
||||
# Handle videos shorter than clip length
|
||||
if effective_clip_length > T:
|
||||
pad_length = effective_clip_length - T
|
||||
last_frame = video[-1:].repeat(pad_length, 1, 1, 1)
|
||||
video = torch.cat([video, last_frame], dim=0)
|
||||
T = len(video)
|
||||
|
||||
clips = []
|
||||
|
||||
if strategy == ClipSamplingStrategy.BEGINNING:
|
||||
# Take first clip (most common for FVD evaluation)
|
||||
clip = video[:effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.MIDDLE:
|
||||
# Take middle clip
|
||||
start = (T - effective_clip_length) // 2
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.RANDOM:
|
||||
# Sample N random clips
|
||||
for _ in range(num_clips):
|
||||
if effective_clip_length == T:
|
||||
start = 0
|
||||
else:
|
||||
start = np.random.randint(0, T - effective_clip_length + 1)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.UNIFORM:
|
||||
# Uniformly spaced clips
|
||||
if num_clips == 1:
|
||||
# Single clip from middle
|
||||
start = (T - effective_clip_length) // 2
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
else:
|
||||
# Multiple uniformly spaced clips
|
||||
step = (T - effective_clip_length) / (num_clips -
|
||||
1) if num_clips > 1 else 0
|
||||
for i in range(num_clips):
|
||||
start = int(i * step)
|
||||
start = min(start, T - effective_clip_length)
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
elif strategy == ClipSamplingStrategy.SLIDING:
|
||||
# Sliding window with stride
|
||||
for start in range(0, T - effective_clip_length + 1, temporal_stride):
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
if len(clips) >= num_clips:
|
||||
break
|
||||
|
||||
elif strategy == ClipSamplingStrategy.ALL:
|
||||
# All possible clips (overlapping)
|
||||
for start in range(T - effective_clip_length + 1):
|
||||
clip = video[start:start + effective_clip_length]
|
||||
clips.append(clip)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {strategy}")
|
||||
|
||||
return clips
|
||||
|
||||
|
||||
def load_video_clips_streaming(directory: str | Path,
|
||||
num_frames: int = 16,
|
||||
max_videos: int | None = None,
|
||||
clip_strategy: str
|
||||
| ClipSamplingStrategy = 'beginning',
|
||||
frame_stride: int = 1,
|
||||
num_clips_per_video: int = 1,
|
||||
video_extensions: list[str] | None = None,
|
||||
support_frame_dirs: bool = True,
|
||||
target_size: tuple[int, int] | None = (224, 224),
|
||||
verbose: bool = True) -> Iterator[torch.Tensor]:
|
||||
"""
|
||||
This generator yields clips one-by-one instead of loading all videos into RAM.
|
||||
Perfect for large datasets where memory is limited.
|
||||
|
||||
Args:
|
||||
directory: Path to directory with videos
|
||||
num_frames: Frames per clip
|
||||
max_videos: Max videos to load
|
||||
clip_strategy: 'beginning', 'random', 'uniform', etc.
|
||||
frame_stride: Frame skip (1=all, 2=every 2nd, 8=every 8th)
|
||||
num_clips_per_video: Number of clips per video
|
||||
video_extensions: Video file extensions
|
||||
support_frame_dirs: Also load frame directories
|
||||
target_size: Resize clips to (H, W). If None, keep original size.
|
||||
verbose: Show progress
|
||||
|
||||
Yields:
|
||||
clip: [T, C, H, W] individual clips
|
||||
|
||||
Example:
|
||||
>>> for clip in load_video_clips_streaming('data/videos/', num_frames=16):
|
||||
>>> features = model.extract_features(clip.unsqueeze(0))
|
||||
>>> # Process one clip at a time - low memory usage!
|
||||
"""
|
||||
if video_extensions is None:
|
||||
video_extensions = ['.mp4', '.avi', '.mov', '.mkv']
|
||||
|
||||
directory = Path(directory)
|
||||
|
||||
if not directory.exists():
|
||||
raise FileNotFoundError(f"Directory not found: {directory}")
|
||||
|
||||
# Find video paths
|
||||
video_paths: list[Path] = []
|
||||
|
||||
# Find video files
|
||||
for ext in video_extensions:
|
||||
video_paths.extend(directory.glob(f"**/*{ext}"))
|
||||
|
||||
# Find frame directories if enabled
|
||||
if support_frame_dirs:
|
||||
for subdir in directory.iterdir():
|
||||
if subdir.is_dir():
|
||||
# Check if it contains frames
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
|
||||
for ext in image_extensions:
|
||||
if list(subdir.glob(f"*{ext}")):
|
||||
video_paths.append(subdir)
|
||||
break
|
||||
|
||||
if len(video_paths) == 0:
|
||||
raise ValueError(f"No videos found in {directory}")
|
||||
|
||||
video_paths = sorted(video_paths)
|
||||
|
||||
if max_videos is not None:
|
||||
video_paths = video_paths[:max_videos]
|
||||
|
||||
if verbose:
|
||||
print(f"Found {len(video_paths)} videos in {directory}")
|
||||
if num_clips_per_video > 1:
|
||||
print(f"Extracting {num_clips_per_video} clips per video...")
|
||||
if frame_stride > 1:
|
||||
print(f"Subsampling frames with stride {frame_stride}...")
|
||||
if target_size:
|
||||
print(f"Resizing clips to {target_size}...")
|
||||
|
||||
# Track statistics
|
||||
failed_count = 0
|
||||
total_clips = 0
|
||||
|
||||
iterator = tqdm(video_paths,
|
||||
desc="Loading videos") if verbose else video_paths
|
||||
|
||||
for video_path in iterator:
|
||||
try:
|
||||
# Load full video
|
||||
video = load_video_auto(video_path,
|
||||
num_frames=None,
|
||||
sample_strategy='uniform')
|
||||
|
||||
# Sample clips from video
|
||||
clips = sample_clips_from_video(video,
|
||||
num_frames_per_clip=num_frames,
|
||||
num_clips=num_clips_per_video,
|
||||
strategy=clip_strategy,
|
||||
frame_stride=frame_stride)
|
||||
|
||||
if target_size is not None:
|
||||
resized_clips = []
|
||||
for clip in clips:
|
||||
T, C, H, W = clip.shape
|
||||
if target_size != (H, W):
|
||||
# Resize to target size
|
||||
clip = clip.contiguous(
|
||||
) # Fix non-contiguous tensors first
|
||||
clip_flat = clip.view(T * C, H,
|
||||
W).unsqueeze(0) # [1, T*C, H, W]
|
||||
clip_resized = torch.nn.functional.interpolate(
|
||||
clip_flat,
|
||||
size=target_size,
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
clip = clip_resized.squeeze(0).view(
|
||||
T, C, target_size[0],
|
||||
target_size[1]) # Back to [T, C, H, W]
|
||||
resized_clips.append(clip)
|
||||
clips = resized_clips
|
||||
|
||||
# Yield clips one by one
|
||||
for clip in clips:
|
||||
yield clip
|
||||
total_clips += 1
|
||||
|
||||
# Free memory
|
||||
del video, clips
|
||||
|
||||
except Exception as e:
|
||||
failed_count += 1
|
||||
if verbose:
|
||||
print(f"\nWarning: Failed to load {video_path}: {e}")
|
||||
continue
|
||||
|
||||
# Validate
|
||||
if total_clips == 0:
|
||||
raise RuntimeError(f"Failed to load any videos from {directory}")
|
||||
|
||||
failure_rate = failed_count / len(video_paths)
|
||||
if failure_rate > 0.1: # More than 10% failed
|
||||
print(
|
||||
f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})"
|
||||
)
|
||||
|
||||
if verbose:
|
||||
print(
|
||||
f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos"
|
||||
)
|
||||
@@ -1,7 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
@@ -1,4 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
@@ -2,12 +2,12 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We support H100 (via TK) and any other GPU (via triton) for STA.
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
@@ -16,14 +16,6 @@ git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you want to skip the compilation of the TK kernel and only use the Triton version, try below:
|
||||
|
||||
```bash
|
||||
SKIP_SM90_EXT=1 python setup.py install
|
||||
or
|
||||
SKIP_SM90_EXT=1 pip install --no-build-isolation .
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
@@ -38,7 +30,7 @@ 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 PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
@@ -51,7 +43,7 @@ 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.
|
||||
# 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.
|
||||
@@ -66,6 +58,7 @@ 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
|
||||
@@ -74,7 +67,7 @@ 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.
|
||||
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
|
||||
@@ -89,7 +82,7 @@ Here is a diagram of how the window is configured and passed through the FastVid
|
||||
|
||||
|
||||
## 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.
|
||||
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.
|
||||
|
||||
|
||||
@@ -51,28 +51,21 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
ext_modules = []
|
||||
|
||||
if os.environ.get("SKIP_SM90_EXT", "0") != "1":
|
||||
ext_modules.append(
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
)
|
||||
else:
|
||||
print("ENV SKIP_SM90_EXT=1, skip st_attn_cuda compile")
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=ext_modules,
|
||||
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",
|
||||
|
||||
@@ -7,17 +7,12 @@ try:
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
try:
|
||||
from st_attn.st_attn_triton import sliding_tile_attention_triton
|
||||
except ImportError:
|
||||
sliding_tile_attention_triton = None
|
||||
|
||||
def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
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,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
@@ -51,13 +46,4 @@ def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, h
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
major, minor = torch.cuda.get_device_capability(q_all.device)
|
||||
if major == 9 and minor == 0 and sta_fwd is not None:
|
||||
return sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
|
||||
elif sliding_tile_attention_triton is not None:
|
||||
return sliding_tile_attention_triton(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
|
||||
else:
|
||||
raise ImportError("No suitable sliding tile attention implementation found.")
|
||||
return hidden_states[:, :, :seq_length]
|
||||
@@ -1,327 +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 get_common_autotune_config():
|
||||
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 [1, 2, 3, 4]\
|
||||
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
|
||||
+10
-99
@@ -34,16 +34,16 @@ def pytorch_test(Q, K, V, block_sparse_mask, dO):
|
||||
)
|
||||
|
||||
|
||||
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, q_non_pad_index, kv_non_pad_index, q_num_blocks, kv_num_blocks, dO):
|
||||
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, q_non_pad_index, q_num_blocks, BLOCK_M)
|
||||
k_padded = vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
v_padded = vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
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[:, :, q_non_pad_index, :]
|
||||
output = output[:, :, non_pad_index, :]
|
||||
output.backward(dO)
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
@@ -64,7 +64,7 @@ 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=16, max_size=64, device="cuda"):
|
||||
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)
|
||||
|
||||
|
||||
@@ -86,21 +86,19 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
|
||||
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, num_blocks, k, device)
|
||||
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, variable_block_sizes, device)
|
||||
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)
|
||||
|
||||
# print(Q.shape, K.shape, V.shape, dO.shape)
|
||||
|
||||
# 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, non_pad_index, num_blocks, num_blocks, 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
|
||||
@@ -120,60 +118,6 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
|
||||
|
||||
return results
|
||||
|
||||
def check_correctness_qkdiff(h, d, num_q_blocks, num_kv_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"
|
||||
q_variable_block_sizes = generate_variable_block_sizes(num_q_blocks, device=device)
|
||||
kv_variable_block_sizes = generate_variable_block_sizes(num_kv_blocks, device=device)
|
||||
|
||||
S_q = int(q_variable_block_sizes.sum().item())
|
||||
S_kv = int(kv_variable_block_sizes.sum().item())
|
||||
|
||||
q_non_pad_index = get_non_pad_index(q_variable_block_sizes, num_q_blocks, BLOCK_M)
|
||||
kv_non_pad_index = get_non_pad_index(kv_variable_block_sizes, num_kv_blocks, BLOCK_M)
|
||||
|
||||
block_mask = generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device)
|
||||
full_mask = create_full_mask_from_block_mask(block_mask, q_variable_block_sizes, kv_variable_block_sizes, device)
|
||||
|
||||
for _ in range(num_iterations):
|
||||
Q = generate_tensor((1, h, S_q, d), torch.bfloat16, device)
|
||||
K = generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
|
||||
V = generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
|
||||
dO = generate_tensor((1, h, S_q, d), torch.bfloat16, device)
|
||||
|
||||
# print(Q.shape, K.shape, V.shape, dO.shape)
|
||||
|
||||
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), kv_variable_block_sizes, q_non_pad_index, kv_non_pad_index, num_q_blocks, num_kv_blocks, 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_q = h * S_q * d * num_iterations
|
||||
total_elements_kv = h * S_kv * d * num_iterations
|
||||
|
||||
for name, data in results.items():
|
||||
total_elements = total_elements_q if name in ['gQ', 'gO'] else total_elements_kv
|
||||
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"},
|
||||
@@ -203,43 +147,10 @@ def generate_error_graphs(h, d, error_mode='all'):
|
||||
|
||||
print("-" * 150)
|
||||
|
||||
def generate_error_graphs_qkdiff(h, d, error_mode='all'):
|
||||
test_configs = [
|
||||
{"num_q_blocks": 16, "num_kv_blocks": 32, "k": 2, "description": "Small Q, Med KV"},
|
||||
{"num_q_blocks": 32, "num_kv_blocks": 16, "k": 4, "description": "Med Q, Small KV"},
|
||||
{"num_q_blocks": 53, "num_kv_blocks": 32, "k": 6, "description": "Large Q, Med KV"},
|
||||
{"num_q_blocks": 16, "num_kv_blocks": 48, "k": 2, "description": "Small Q, Large KV"},
|
||||
{"num_q_blocks": 48, "num_kv_blocks": 16, "k": 2, "description": "Large Q, Small KV"},
|
||||
]
|
||||
|
||||
print(f"\nError Analysis (QK Diff) for h={h}, d={d}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Q Blks':<8} {'KV Blks':<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_q_blocks = config["num_q_blocks"]
|
||||
num_kv_blocks = config["num_kv_blocks"]
|
||||
k = config["k"]
|
||||
description = config["description"]
|
||||
results = check_correctness_qkdiff(h, d, num_q_blocks, num_kv_blocks, k, error_mode=error_mode)
|
||||
print(f"{description:<20} {num_q_blocks:<8} {num_kv_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)
|
||||
generate_error_graphs_qkdiff(h, d, error_mode=mode)
|
||||
print("\nAnalysis completed for all modes.")
|
||||
print("\nAnalysis completed for all modes.")
|
||||
|
||||
@@ -1,236 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
# Make sure we can import from the project root (`vsa`, `tests.utils`, etc.)
|
||||
CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
PROJECT_ROOT = os.path.dirname(CURRENT_DIR)
|
||||
if PROJECT_ROOT not in sys.path:
|
||||
sys.path.append(PROJECT_ROOT)
|
||||
if CURRENT_DIR not in sys.path:
|
||||
sys.path.append(CURRENT_DIR)
|
||||
|
||||
from tests.utils import (
|
||||
generate_block_sparse_mask_for_function,
|
||||
create_full_mask_from_block_mask,
|
||||
)
|
||||
from vsa import block_sparse_attn, BLOCK_M
|
||||
import test_vsa as ref # reuse helper functions from backward test
|
||||
|
||||
|
||||
def pytorch_forward(
|
||||
Q: torch.Tensor,
|
||||
K: torch.Tensor,
|
||||
V: torch.Tensor,
|
||||
block_sparse_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Dense PyTorch reference forward:
|
||||
- Q: [1, h, S_q, d]
|
||||
- K,V: [1, h, S_kv, d]
|
||||
- block_sparse_mask: [h, S_q, S_kv] bool
|
||||
"""
|
||||
q = Q.clone().float()
|
||||
k = K.clone().float()
|
||||
v = V.clone().float()
|
||||
|
||||
attn = torch.matmul(q, k.transpose(-2, -1)) # [1, h, S_q, S_kv]
|
||||
attn = attn / (q.size(-1) ** 0.5)
|
||||
attn = attn.masked_fill(~block_sparse_mask.unsqueeze(0), float("-inf"))
|
||||
attn = torch.nn.functional.softmax(attn, dim=-1)
|
||||
out = torch.matmul(attn, v) # [1, h, S_q, d]
|
||||
return out.to(torch.bfloat16)
|
||||
|
||||
|
||||
def block_sparse_forward_test(
|
||||
Q: torch.Tensor,
|
||||
K: torch.Tensor,
|
||||
V: torch.Tensor,
|
||||
block_sparse_mask: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
q_non_pad_index: torch.Tensor,
|
||||
kv_non_pad_index: torch.Tensor,
|
||||
q_num_blocks: int,
|
||||
kv_num_blocks: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward-only wrapper around `block_sparse_attn`, mirroring `block_sparse_kernel_test`
|
||||
but without any backward / grad logic.
|
||||
"""
|
||||
Q = Q.detach()
|
||||
K = K.detach()
|
||||
V = V.detach()
|
||||
|
||||
q_padded = ref.vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
|
||||
k_padded = ref.vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
v_padded = ref.vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
|
||||
|
||||
out_padded, _ = block_sparse_attn(
|
||||
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
|
||||
)
|
||||
# Remove padding on the query side
|
||||
out = out_padded[:, :, q_non_pad_index, :]
|
||||
return out
|
||||
|
||||
|
||||
def run_forward_equal_qk(
|
||||
h: int = 16,
|
||||
d: int = 128,
|
||||
num_blocks: int = 16,
|
||||
k: int = 2,
|
||||
num_iterations: int = 5,
|
||||
) -> Tuple[float, float]:
|
||||
"""
|
||||
Forward-only correctness test for the case S_q == S_kv.
|
||||
Mirrors `check_correctness` but only compares forward outputs.
|
||||
"""
|
||||
assert torch.cuda.is_available(), "VSA kernels require CUDA"
|
||||
device = "cuda"
|
||||
|
||||
variable_block_sizes = ref.generate_variable_block_sizes(
|
||||
num_blocks, device=device
|
||||
)
|
||||
S = int(variable_block_sizes.sum().item())
|
||||
non_pad_index = ref.get_non_pad_index(
|
||||
variable_block_sizes, num_blocks, BLOCK_M
|
||||
)
|
||||
|
||||
block_mask = generate_block_sparse_mask_for_function(
|
||||
h, num_blocks, num_blocks, k, device
|
||||
)
|
||||
full_mask = create_full_mask_from_block_mask(
|
||||
block_mask, variable_block_sizes, variable_block_sizes, device
|
||||
)
|
||||
print(f"[qkequal] h: {h}, d: {d}, num_blocks: {num_blocks}, k: {k}")
|
||||
print(f"[qkequal] variable_block_sizes: {variable_block_sizes}, non_pad_index: {non_pad_index.shape}, block_mask: {block_mask.shape}, full_mask: {full_mask.shape}")
|
||||
sum_diff = 0.0
|
||||
sum_abs = 0.0
|
||||
max_rel_diff = 0.0
|
||||
|
||||
for i in range(num_iterations):
|
||||
Q = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
K = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
V = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
|
||||
if i == 0: print(f"[qkequal] Q: {Q.shape}, K: {K.shape}, V: {V.shape}, full_mask: {full_mask.shape}")
|
||||
if i == 0: print(f"[qkequal] block_mask: {block_mask.shape}")
|
||||
|
||||
pt_o = pytorch_forward(Q, K, V, full_mask)
|
||||
bs_o = block_sparse_forward_test(
|
||||
Q,
|
||||
K,
|
||||
V,
|
||||
block_mask.unsqueeze(0),
|
||||
variable_block_sizes,
|
||||
non_pad_index,
|
||||
non_pad_index,
|
||||
num_blocks,
|
||||
num_blocks,
|
||||
)
|
||||
|
||||
diff = (pt_o - bs_o).abs()
|
||||
sum_diff += diff.sum().item()
|
||||
sum_abs += pt_o.abs().sum().item()
|
||||
rel_max = diff.max() / (pt_o.abs().mean() + 1e-6)
|
||||
max_rel_diff = max(max_rel_diff, rel_max.item())
|
||||
|
||||
total_elems = h * S * d * num_iterations
|
||||
avg_abs_err = sum_diff / total_elems
|
||||
return avg_abs_err, max_rel_diff
|
||||
|
||||
|
||||
def run_forward_qk_diff(
|
||||
h: int = 16,
|
||||
d: int = 128,
|
||||
num_q_blocks: int = 16,
|
||||
num_kv_blocks: int = 32,
|
||||
k: int = 2,
|
||||
num_iterations: int = 5,
|
||||
) -> Tuple[float, float]:
|
||||
"""
|
||||
Forward-only correctness test for the case S_q != S_kv.
|
||||
|
||||
NOTE:
|
||||
- The Triton backend supports different Q/KV logical lengths via padding.
|
||||
- The SM90 (H100) CUDA backend currently assumes the same number of blocks
|
||||
for Q and KV, so we skip this test there.
|
||||
"""
|
||||
assert torch.cuda.is_available(), "VSA kernels require CUDA"
|
||||
|
||||
device = "cuda"
|
||||
|
||||
q_variable_block_sizes = ref.generate_variable_block_sizes(
|
||||
num_q_blocks, device=device
|
||||
)
|
||||
kv_variable_block_sizes = ref.generate_variable_block_sizes(
|
||||
num_kv_blocks, device=device
|
||||
)
|
||||
|
||||
S_q = int(q_variable_block_sizes.sum().item())
|
||||
S_kv = int(kv_variable_block_sizes.sum().item())
|
||||
|
||||
q_non_pad_index = ref.get_non_pad_index(
|
||||
q_variable_block_sizes, num_q_blocks, BLOCK_M
|
||||
)
|
||||
kv_non_pad_index = ref.get_non_pad_index(
|
||||
kv_variable_block_sizes, num_kv_blocks, BLOCK_M
|
||||
)
|
||||
|
||||
block_mask = generate_block_sparse_mask_for_function(
|
||||
h, num_q_blocks, num_kv_blocks, k, device
|
||||
)
|
||||
full_mask = create_full_mask_from_block_mask(
|
||||
block_mask, q_variable_block_sizes, kv_variable_block_sizes, device
|
||||
)
|
||||
|
||||
sum_diff = 0.0
|
||||
sum_abs = 0.0
|
||||
max_rel_diff = 0.0
|
||||
|
||||
for _ in range(num_iterations):
|
||||
Q = ref.generate_tensor((1, h, S_q, d), torch.bfloat16, device)
|
||||
K = ref.generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
|
||||
V = ref.generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
|
||||
|
||||
pt_o = pytorch_forward(Q, K, V, full_mask)
|
||||
bs_o = block_sparse_forward_test(
|
||||
Q,
|
||||
K,
|
||||
V,
|
||||
block_mask.unsqueeze(0),
|
||||
kv_variable_block_sizes,
|
||||
q_non_pad_index,
|
||||
kv_non_pad_index,
|
||||
num_q_blocks,
|
||||
num_kv_blocks,
|
||||
)
|
||||
|
||||
diff = (pt_o - bs_o).abs()
|
||||
sum_diff += diff.sum().item()
|
||||
sum_abs += pt_o.abs().sum().item()
|
||||
rel_max = diff.max() / (pt_o.abs().mean() + 1e-6)
|
||||
max_rel_diff = max(max_rel_diff, rel_max.item())
|
||||
|
||||
total_elems = h * S_q * d * num_iterations
|
||||
avg_abs_err = sum_diff / total_elems
|
||||
return avg_abs_err, max_rel_diff
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
h, d = 16, 128
|
||||
print("Forward Block Sparse Attention Check (QK Equal)")
|
||||
print("=" * 80)
|
||||
avg_err_eq, max_rel_eq = run_forward_equal_qk(h, d, num_blocks=32, k=2)
|
||||
print(f"QK equal: avg |ΔO| = {avg_err_eq:.6e}, max rel ΔO = {max_rel_eq:.6e}")
|
||||
|
||||
print("\nForward Block Sparse Attention Check (QK Different)")
|
||||
print("=" * 80)
|
||||
avg_err_diff, max_rel_diff = run_forward_qk_diff(
|
||||
h, d, num_q_blocks=32, num_kv_blocks=48, k=2
|
||||
)
|
||||
print(
|
||||
f"QK diff: avg |ΔO| = {avg_err_diff:.6e}, max rel ΔO = {max_rel_diff:.6e}"
|
||||
)
|
||||
|
||||
+21
-27
@@ -1,60 +1,54 @@
|
||||
import torch
|
||||
|
||||
def generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate block sparse mask of shape [h, num_q_blocks, num_kv_blocks].
|
||||
Generate block sparse mask of shape [h, num_blocks, num_blocks].
|
||||
|
||||
Args:
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key/value blocks
|
||||
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_q_blocks, num_kv_blocks] bool tensor
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
"""
|
||||
k = min(k, num_kv_blocks)
|
||||
scores = torch.rand(h, num_q_blocks, num_kv_blocks, device=device)
|
||||
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_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
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, q_variable_block_sizes,
|
||||
kv_variable_block_sizes, device="cuda"):
|
||||
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_q_blocks, num_kv_blocks] bool tensor
|
||||
q_variable_block_sizes: [num_q_blocks] tensor
|
||||
kv_variable_block_sizes: [num_kv_blocks] tensor
|
||||
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_q, S_kv] bool tensor where S = total sequence length
|
||||
full_mask: [h, S, S] bool tensor where S = total sequence length
|
||||
"""
|
||||
h, num_q_blocks, num_kv_blocks = block_sparse_mask.shape
|
||||
total_q_seq_len = q_variable_block_sizes.sum().item()
|
||||
total_kv_seq_len = kv_variable_block_sizes.sum().item()
|
||||
|
||||
q_cumsum = torch.cat([torch.tensor([0], device=device), q_variable_block_sizes.cumsum(dim=0)[:-1]])
|
||||
kv_cumsum = torch.cat([torch.tensor([0], device=device), kv_variable_block_sizes.cumsum(dim=0)[:-1]])
|
||||
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_q_seq_len, total_kv_seq_len, dtype=torch.bool, device=device)
|
||||
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_q_blocks):
|
||||
q_start = q_cumsum[q_block]
|
||||
q_end = q_start + q_variable_block_sizes[q_block]
|
||||
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_kv_blocks):
|
||||
for kv_block in range(num_blocks):
|
||||
if block_sparse_mask[head, q_block, kv_block]:
|
||||
kv_start = kv_cumsum[kv_block]
|
||||
kv_end = kv_start + kv_variable_block_sizes[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
|
||||
@@ -247,12 +247,11 @@ def _attn_bwd_dq(dq, q, K, V, #
|
||||
|
||||
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):
|
||||
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
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
|
||||
@@ -672,32 +672,23 @@ block_sparse_attention_forward(
|
||||
torch::Tensor v,
|
||||
torch::Tensor q2k_block_sparse_index,
|
||||
torch::Tensor q2k_block_sparse_num,
|
||||
torch::Tensor kv_block_size
|
||||
torch::Tensor block_size
|
||||
)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
// q shape: (batch, qo_heads, q_seq_len, head_dim)
|
||||
// k shape: (batch, kv_heads, kv_seq_len, head_dim)
|
||||
// v shape: (batch, kv_heads, kv_seq_len, head_dim)
|
||||
// q2k_block_sparse_index shape: (batch, qo_heads, num_q_blocks, max_kv_blocks_per_q)
|
||||
// q2k_block_sparse_num shape: (batch, qo_heads, num_q_blocks)
|
||||
// kv_block_size shape: (num_kv_blocks) This does not need other dimensions because across all batch/heads the padding is the same.
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto q_seq_len = q.size(2);
|
||||
auto kv_seq_len = k.size(2);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
|
||||
auto num_q_blocks = q2k_block_sparse_index.size(2);
|
||||
auto num_kv_blocks = kv_block_size.size(0);
|
||||
auto num_q_blocks = block_size.size(0);
|
||||
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
|
||||
TORCH_CHECK(num_q_blocks * BLOCK_M == q_seq_len, "This kernel supports variable q block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_kv_blocks * BLOCK_M == kv_seq_len, "This kernel supports variable kv block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -705,8 +696,11 @@ block_sparse_attention_forward(
|
||||
TORCH_CHECK(q2k_block_sparse_index.size(0) == batch, "q2k_block_sparse_index batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(0) == batch, "q2k_block_sparse_num batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K inputs");
|
||||
TORCH_CHECK(q2k_block_sparse_num.size(2) == num_q_blocks, "q2k_block_sparse_num idx 2 - must match num_q_blocks");
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(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");
|
||||
@@ -733,12 +727,12 @@ block_sparse_attention_forward(
|
||||
// for the returned outputs
|
||||
torch::Tensor o = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, v.options());
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
@@ -768,11 +762,11 @@ block_sparse_attention_forward(
|
||||
|
||||
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>(q_seq_len), 64U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_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>(q_seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
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,
|
||||
@@ -780,17 +774,17 @@ block_sparse_attention_forward(
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(q_seq_len),
|
||||
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*>(kv_block_size.data_ptr())
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<64>,
|
||||
@@ -819,11 +813,11 @@ block_sparse_attention_forward(
|
||||
|
||||
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>(q_seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_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>(q_seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
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,
|
||||
@@ -831,17 +825,17 @@ block_sparse_attention_forward(
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(q_seq_len),
|
||||
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*>(kv_block_size.data_ptr())
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128>,
|
||||
@@ -868,7 +862,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
torch::Tensor og,
|
||||
torch::Tensor k2q_block_sparse_index,
|
||||
torch::Tensor k2q_block_sparse_num,
|
||||
torch::Tensor kv_block_size)
|
||||
torch::Tensor block_size)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
@@ -877,23 +871,11 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
CHECK_INPUT(o);
|
||||
CHECK_INPUT(og);
|
||||
|
||||
// q: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// k: [batch, kv_heads, kv_seq_len, head_dim]
|
||||
// v: [batch, kv_heads, kv_seq_len, head_dim]
|
||||
// o: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// l_vec: [batch, qo_heads, q_seq_len, 1]
|
||||
// og: [batch, qo_heads, q_seq_len, head_dim]
|
||||
// k2q_block_sparse_index: [batch, kv_heads, num_kv_blocks, max_num_q_blocks]
|
||||
// k2q_block_sparse_num: [batch, kv_heads, num_kv_blocks]
|
||||
// kv_block_size: [num_kv_blocks]
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto q_seq_len = q.size(2);
|
||||
auto kv_seq_len = k.size(2);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
|
||||
auto num_kv_blocks = kv_block_size.size(0);
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index.size(2) must match num_kv_blocks (kv_block_size.size(0))");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -904,18 +886,23 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(0) == batch, "k2q_block_sparse_index batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(0) == batch, "k2q_block_sparse_num batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K sequence length");
|
||||
TORCH_CHECK(l_vec.size(2) == q_seq_len, "L sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(o.size(2) == q_seq_len, "O sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(og.size(2) == q_seq_len, "OG sequence length dimension - idx 2 - must match Q sequence length");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(2) == num_kv_blocks, "k2q_block_sparse_num idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(l_vec.size(2) == seq_len, "L sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(o.size(2) == seq_len, "O sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(og.size(2) == seq_len, "OG sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_index idx 2 - must match seq_len / BLOCK_N");
|
||||
TORCH_CHECK(k2q_block_sparse_num.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_num idx 2 - must match seq_len / BLOCK_N");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(o.size(3) == head_dim, "O head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(og.size(3) == head_dim, "OG head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
|
||||
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
@@ -942,20 +929,20 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
torch::Tensor qg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
torch::Tensor kg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(kv_heads),
|
||||
static_cast<const uint>(kv_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
torch::Tensor vg = torch::zeros({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(kv_heads),
|
||||
static_cast<const uint>(kv_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(head_dim)}, l_vec.options());
|
||||
|
||||
torch::Tensor d_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(q_seq_len),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)}, l_vec.options());
|
||||
|
||||
float* qg_ptr = qg.data_ptr<float>();
|
||||
@@ -984,7 +971,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
// cudaStreamSynchronize(stream);
|
||||
|
||||
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
|
||||
dim3 grid_bwd(q_seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
if (head_dim == 64) {
|
||||
using og_tile = st_bf<4*16, 64>;
|
||||
@@ -997,9 +984,9 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_prep_globals = bwd_prep_globals<64>;
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
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)};
|
||||
|
||||
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
@@ -1036,15 +1023,15 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_global_args = bwd_globals<64>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_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>(q_seq_len)};
|
||||
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), 64U};
|
||||
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), 64U};
|
||||
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), 64U};
|
||||
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), 64U};
|
||||
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), 64U};
|
||||
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), 64U};
|
||||
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), 64U};
|
||||
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,
|
||||
@@ -1055,14 +1042,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_vg_arg,
|
||||
bwd_l_arg,
|
||||
bwd_d_arg,
|
||||
static_cast<int>(kv_seq_len), // N is not used in the kernel
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
@@ -1101,9 +1088,9 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_prep_globals = bwd_prep_globals<128>;
|
||||
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
|
||||
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
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)};
|
||||
|
||||
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
|
||||
|
||||
@@ -1140,15 +1127,15 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
|
||||
using bwd_global_args = bwd_globals<128>;
|
||||
|
||||
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
|
||||
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
|
||||
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_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>(q_seq_len)};
|
||||
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), 128U};
|
||||
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), 128U};
|
||||
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), 128U};
|
||||
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), 128U};
|
||||
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), 128U};
|
||||
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), 128U};
|
||||
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), 128U};
|
||||
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,
|
||||
@@ -1159,14 +1146,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_vg_arg,
|
||||
bwd_l_arg,
|
||||
bwd_d_arg,
|
||||
static_cast<int>(kv_seq_len), // N is not used in the kernel
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
__pycache__/
|
||||
*.so
|
||||
*.pyc
|
||||
.ipynb_checkpoints/
|
||||
@@ -1,187 +0,0 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
@@ -1,6 +0,0 @@
|
||||
include LICENSE
|
||||
include README.md
|
||||
include pyproject.toml
|
||||
recursive-include src/fastvideo_kernel *.cu *.cuh *.cpp *.h
|
||||
recursive-include csrc *.cu *.cuh *.cpp *.h
|
||||
recursive-include tk *.cu *.cuh *.cpp *.h
|
||||
@@ -1,31 +0,0 @@
|
||||
# FastVideo Kernel
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
cd csrc/fastvideo_kernel
|
||||
pip install .
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_attn_varlen
|
||||
|
||||
# Example: Sliding Tile Attention
|
||||
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
|
||||
|
||||
# Example: Video Sparse Attention (with Triton fallback)
|
||||
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
|
||||
|
||||
# Example: VMoBA
|
||||
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
|
||||
```
|
||||
|
||||
## Requirements
|
||||
|
||||
- H100 GPU (sm_90a) for CUDA kernels
|
||||
- Triton for non-H100 fallback
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,23 +0,0 @@
|
||||
#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
|
||||
}
|
||||
@@ -1,573 +0,0 @@
|
||||
// # Define TORCH_COMPILE macro
|
||||
|
||||
#include "kittens.cuh"
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <stdio.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
|
||||
return (value < min) ? min : ((value > max) ? max : value);
|
||||
}
|
||||
// #define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
__device__ __forceinline__ int abs_int(int value) {
|
||||
return (value < 0) ? -value : value;
|
||||
}
|
||||
|
||||
|
||||
constexpr int CONSUMER_WARPGROUPS = (3);
|
||||
constexpr int PRODUCER_WARPGROUPS = (1);
|
||||
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
|
||||
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
|
||||
|
||||
using namespace kittens;
|
||||
namespace cg = cooperative_groups;
|
||||
|
||||
template<int D> struct fwd_attend_ker_tile_dims {};
|
||||
template<> struct fwd_attend_ker_tile_dims<64> {
|
||||
constexpr static int tile_width = (64);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (4);
|
||||
};
|
||||
template<> struct fwd_attend_ker_tile_dims<128> {
|
||||
constexpr static int tile_width = (128);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (2);
|
||||
};
|
||||
|
||||
template<int D> struct fwd_globals {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
|
||||
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
q_gl q;
|
||||
k_gl k;
|
||||
v_gl v;
|
||||
l_gl l;
|
||||
o_gl o;
|
||||
|
||||
const int N;
|
||||
const int text_L;
|
||||
const int hr;
|
||||
};
|
||||
|
||||
|
||||
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
|
||||
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
|
||||
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
extern __shared__ int __shm[];
|
||||
tma_swizzle_allocator al((int*)&__shm[0]);
|
||||
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
|
||||
|
||||
using K = fwd_attend_ker_tile_dims<D>;
|
||||
|
||||
using q_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
using k_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using v_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
|
||||
using o_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
|
||||
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
|
||||
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
|
||||
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
|
||||
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
|
||||
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
|
||||
int img_kv_blocks;
|
||||
int kv_blocks = g.N / (K::kv_height);
|
||||
if constexpr (text_kv) {
|
||||
img_kv_blocks = kv_blocks - 3;
|
||||
} else {
|
||||
img_kv_blocks = kv_blocks;
|
||||
}
|
||||
int kv_head_idx = blockIdx.y / g.hr;
|
||||
int seq_idx;
|
||||
if constexpr (text_q) {
|
||||
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
} else {
|
||||
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
}
|
||||
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
|
||||
if (threadIdx.x == 0) {
|
||||
init_semaphore(qsmem_semaphore, 0, 1);
|
||||
for(int j = 0; j < K::stages; j++) {
|
||||
init_semaphore(k_smem_arrived[j], 0, 1);
|
||||
init_semaphore(v_smem_arrived[j], 0, 1);
|
||||
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
|
||||
}
|
||||
|
||||
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
|
||||
|
||||
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
|
||||
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
|
||||
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
|
||||
}
|
||||
|
||||
if constexpr (text_q){
|
||||
for (int j = 0; j < K::stages - 1; j++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
|
||||
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
|
||||
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
|
||||
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int count = 0;
|
||||
int j = 0;
|
||||
while (count < K::stages - 1) {
|
||||
int kt = j / 3 / (CH * CW);
|
||||
int kh = (j / 3) % (CH * CW) / CW;
|
||||
int kw = (j / 3) % CW;
|
||||
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
|
||||
if (mask){
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
|
||||
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
|
||||
count += 1;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int pipe_idx = K::stages - 1;
|
||||
|
||||
if(warpgroupid == NUM_WARPGROUPS-1) {
|
||||
warpgroup::decrease_registers<32>();
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
|
||||
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
|
||||
}
|
||||
else { kv_iters = kv_blocks-2;}
|
||||
|
||||
if(warpid == NUM_WORKERS-4) {
|
||||
if constexpr (text_q){
|
||||
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
|
||||
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int k_t_min = clamp_int(qt-DT, 0, CT-1);
|
||||
int k_t_max = clamp_int(qt+DT, 0, CT-1);
|
||||
int k_h_min = clamp_int(qh-DH, 0, CH-1);
|
||||
int k_h_max = clamp_int(qh+DH, 0, CH-1);
|
||||
int k_w_min = clamp_int(qw-DW, 0, CW-1);
|
||||
int k_w_max = clamp_int(qw+DW, 0, CW-1);
|
||||
int count = 0;
|
||||
for (int kt = k_t_min; kt <= k_t_max; kt++) {
|
||||
for (int kh = k_h_min; kh <= k_h_max; kh++) {
|
||||
for (int kw = k_w_min; kw <= k_w_max; kw++) {
|
||||
for (int j = 0; j <= 2; j++){
|
||||
if (count >= K::stages - 1) {
|
||||
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
} else {
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// for text
|
||||
for (int index = img_kv_blocks; index < kv_blocks; index++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
warpgroup::increase_registers<160>();
|
||||
|
||||
rt_fl<16, K::kv_height> att_block;
|
||||
rt_bf<16, K::kv_height> att_block_mma;
|
||||
rt_fl<16, K::tile_width> o_reg;
|
||||
|
||||
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
|
||||
|
||||
neg_infty(max_vec);
|
||||
zero(norm_vec);
|
||||
zero(o_reg);
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
|
||||
kv_iters = (kv_iters/8);
|
||||
}
|
||||
else if constexpr (text_q){
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
kv_iters = img_kv_blocks - 1;
|
||||
} else {
|
||||
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
}
|
||||
|
||||
kittens::wait(qsmem_semaphore, 0);
|
||||
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
if constexpr(text_kv) {
|
||||
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
// apply non-pad mask
|
||||
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
|
||||
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
|
||||
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
|
||||
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
}
|
||||
|
||||
div_row(o_reg, o_reg, norm_vec);
|
||||
warpgroup::store(o_smem[warpgroupid], o_reg);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
|
||||
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
|
||||
}
|
||||
|
||||
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
|
||||
log(norm_vec, norm_vec);
|
||||
add(norm_vec, norm_vec, max_vec_scaled);
|
||||
|
||||
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
|
||||
else { mul(norm_vec, norm_vec, -11.313708499f); }
|
||||
|
||||
warpgroup::store(l_smem[warpgroupid], norm_vec);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
|
||||
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
|
||||
}
|
||||
tma::store_async_wait();
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
#include "pyutils/torch_helpers.cuh"
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <iostream>
|
||||
|
||||
torch::Tensor
|
||||
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
|
||||
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
|
||||
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
|
||||
auto hr = qo_heads / kv_heads;
|
||||
|
||||
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
|
||||
|
||||
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
|
||||
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
|
||||
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
|
||||
|
||||
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
|
||||
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
|
||||
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
|
||||
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
|
||||
if (head_dim == 128) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<128>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
// 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) {
|
||||
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
dim3 grid_text(2, qo_heads, batch);
|
||||
if (!process_text) {
|
||||
#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) { 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);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
} else {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
}
|
||||
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
if (kernel_aspect_ratio_flag == 2){
|
||||
#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) { 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);
|
||||
}
|
||||
#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) { 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);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
|
||||
else {
|
||||
TORCH_CHECK(false, "Unsupported kernel_aspect_ratio_flag: ", kernel_aspect_ratio_flag);
|
||||
}
|
||||
|
||||
}
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return o;
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
#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
|
||||
}
|
||||
@@ -1,27 +0,0 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=61.0", "torch>=2.5.0", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.1.0"
|
||||
description = "CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
license = {text = "Apache-2.0"}
|
||||
authors = [{name = "Hao AI Lab"}]
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
]
|
||||
dependencies = [
|
||||
"torch>=2.5.0",
|
||||
"triton>=2.0.0"
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -1,132 +0,0 @@
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
ROOT = Path(__file__).parent.absolute()
|
||||
CSRC_DIR = ROOT / "csrc"
|
||||
|
||||
# Path to ThunderKittens (TK)
|
||||
def get_tk_dir():
|
||||
tk_env = os.getenv("THUNDERKITTENS_ROOT")
|
||||
if tk_env:
|
||||
return tk_env
|
||||
|
||||
# Check common locations
|
||||
possible_paths = [
|
||||
ROOT / "tk",
|
||||
ROOT / "csrc" / "tk",
|
||||
ROOT.parent / "attn" / "sliding_tile_attn" / "tk",
|
||||
ROOT.parent / "attn" / "video_sparse_attn" / "tk",
|
||||
]
|
||||
for p in possible_paths:
|
||||
if (p / "include" / "kittens.cuh").exists():
|
||||
return str(p)
|
||||
|
||||
# Default fallback
|
||||
return str(ROOT.parent / "attn" / "sliding_tile_attn" / "tk")
|
||||
|
||||
TK_DIR = get_tk_dir()
|
||||
|
||||
def get_cuda_flags(tk_root: str) -> list:
|
||||
python_include = subprocess.check_output(
|
||||
["python", "-c", "import sysconfig; print(sysconfig.get_path('include'))"]
|
||||
).decode().strip()
|
||||
|
||||
torch_includes = subprocess.check_output([
|
||||
"python", "-c",
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; "
|
||||
"print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip().split()
|
||||
|
||||
return [
|
||||
"-DNDEBUG",
|
||||
"-Xcompiler=-Wno-psabi",
|
||||
"-Xcompiler=-fno-strict-aliasing",
|
||||
"--expt-extended-lambda",
|
||||
"--expt-relaxed-constexpr",
|
||||
"-forward-unknown-to-host-compiler",
|
||||
"--use_fast_math",
|
||||
"-std=c++20",
|
||||
"-O3",
|
||||
"-Xnvlink=--verbose",
|
||||
"-Xptxas=--verbose",
|
||||
"-Xptxas=--warn-on-spills",
|
||||
f"-I{tk_root}/include",
|
||||
f"-I{tk_root}/prototype",
|
||||
f"-I{python_include}",
|
||||
"-DTORCH_COMPILE",
|
||||
"-DKITTENS_HOPPER",
|
||||
"-arch=sm_90a",
|
||||
] + torch_includes
|
||||
|
||||
def get_extensions():
|
||||
if not torch.cuda.is_available():
|
||||
return []
|
||||
|
||||
extensions = []
|
||||
cpp_flags = ["-std=c++20", "-O3"]
|
||||
|
||||
# Check if TK is available
|
||||
if not os.path.exists(os.path.join(TK_DIR, "include", "kittens.cuh")):
|
||||
print(f"Warning: ThunderKittens not found at {TK_DIR}. CUDA kernels will not be built.")
|
||||
return []
|
||||
|
||||
cuda_flags = get_cuda_flags(TK_DIR)
|
||||
|
||||
# STA Extension
|
||||
extensions.append(CUDAExtension(
|
||||
"fastvideo_kernel._C.st_attn",
|
||||
sources=[
|
||||
"csrc/st_attn.cpp",
|
||||
"csrc/st_attn_h100.cu",
|
||||
],
|
||||
extra_compile_args={
|
||||
"cxx": cpp_flags + ["-DTK_COMPILE_ST_ATTN"],
|
||||
"nvcc": cuda_flags + ["-DTK_COMPILE_ST_ATTN"]
|
||||
},
|
||||
libraries=["cuda"],
|
||||
))
|
||||
|
||||
# VSA Extension
|
||||
extensions.append(CUDAExtension(
|
||||
"fastvideo_kernel._C.vsa",
|
||||
sources=[
|
||||
"csrc/vsa.cpp",
|
||||
"csrc/block_sparse_h100.cu",
|
||||
],
|
||||
extra_compile_args={
|
||||
"cxx": cpp_flags + ["-DTK_COMPILE_BLOCK_SPARSE"],
|
||||
"nvcc": cuda_flags + ["-DTK_COMPILE_BLOCK_SPARSE"]
|
||||
},
|
||||
libraries=["cuda"],
|
||||
))
|
||||
|
||||
return extensions
|
||||
|
||||
ext_modules = []
|
||||
if not any(arg in sys.argv for arg in ["clean", "egg_info", "--version"]):
|
||||
try:
|
||||
import torch
|
||||
ext_modules = get_extensions()
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to configure CUDA extensions: {e}")
|
||||
|
||||
setup(
|
||||
name="fastvideo-kernel",
|
||||
version="0.1.0",
|
||||
description="Unified CUDA kernels for FastVideo",
|
||||
long_description=open("README.md").read(),
|
||||
long_description_content_type="text/markdown",
|
||||
license="Apache-2.0",
|
||||
author="Hao AI Lab",
|
||||
url="https://github.com/hao-ai-lab/FastVideo",
|
||||
package_dir={"": "src"},
|
||||
packages=find_packages(where="src"),
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={"build_ext": BuildExtension} if ext_modules else {},
|
||||
python_requires=">=3.10",
|
||||
install_requires=["torch>=2.5.0", "triton>=2.0.0"],
|
||||
)
|
||||
@@ -1,21 +0,0 @@
|
||||
__version__ = "0.1.0"
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
"__version__",
|
||||
]
|
||||
@@ -1,103 +0,0 @@
|
||||
import math
|
||||
import torch
|
||||
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
|
||||
from .triton_kernels.index import map_to_index
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C.st_attn import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C.vsa import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
|
||||
def sliding_tile_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
window_size: list,
|
||||
text_length: int,
|
||||
has_text: bool = True,
|
||||
seq_shape: str = "30x48x80",
|
||||
) -> torch.Tensor:
|
||||
if sta_fwd is None:
|
||||
raise RuntimeError("STA kernel not compiled. Requires H100 and ThunderKittens at build time.")
|
||||
|
||||
seq_length = q.shape[2]
|
||||
shape_map = {"30x48x80": 1, "36x48x48": 2, "18x48x80": 3}
|
||||
|
||||
if has_text:
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
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)
|
||||
|
||||
if block_sparse_fwd is not None:
|
||||
idx, num = map_to_index(mask)
|
||||
out_s, _ = block_sparse_fwd(q, k, v, idx, num, variable_block_sizes.int())
|
||||
else:
|
||||
idx, num = map_to_index(mask)
|
||||
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
return out_c + out_s
|
||||
-449
@@ -1,449 +0,0 @@
|
||||
"""
|
||||
Fused Attention
|
||||
===============
|
||||
|
||||
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
(https://tridao.me/publications/flash2/flash2.pdf)
|
||||
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
import math # small utility needed by the sparse wrapper
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
|
||||
# the code below and commenting out the equivalent parameters is convenient for
|
||||
# re-tuning.
|
||||
configs = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64]\
|
||||
for BN in [64]\
|
||||
for s in [3, 4, 7]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(Q, K, V, sm_scale, #
|
||||
q2k_index, q2k_num, max_kv_blks, #
|
||||
variable_block_sizes,
|
||||
M, Out, #
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vk, stride_vn,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
Z, H, N_CTX, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr):
|
||||
"""
|
||||
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
|
||||
(32×64 and 64×32) – memory footprint unchanged.
|
||||
"""
|
||||
|
||||
# ----- program-id mapping -----
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(1) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
# ----- base pointers -----
|
||||
qvk_off = (b.to(tl.int64) * stride_qz +
|
||||
h.to(tl.int64) * stride_qh)
|
||||
|
||||
Q_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
K_base = tl.make_block_ptr(
|
||||
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
|
||||
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_base = tl.make_block_ptr(
|
||||
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
|
||||
|
||||
O_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
# ----- accumulators -----
|
||||
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
q = tl.load(Q_ptr)
|
||||
|
||||
# ----- sparse loop over valid K/V tiles -----
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
|
||||
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
|
||||
|
||||
k = tl.load(K_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
# mask out invalid columns
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptr)
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
# ----- epilogue -----
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(M + off_hz * N_CTX + offs_m, m_i)
|
||||
tl.store(O_ptr, acc.to(Out.type.element_ty))
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO, #
|
||||
Delta, #
|
||||
Z, H, N_CTX, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
|
||||
):
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
# Filled in by the wrapper.
|
||||
start_n, start_m, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
|
||||
step_m = BLOCK_M1
|
||||
kv_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_N1
|
||||
meta_base = ((b * H + h) * q_tiles + kv_blk)
|
||||
|
||||
q_blocks = tl.load(k2q_num + meta_base) # int32
|
||||
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + kv_blk)
|
||||
|
||||
|
||||
|
||||
for blk_idx in range(q_blocks*2):
|
||||
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
|
||||
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
|
||||
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(tl.bfloat16)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(tl.bfloat16)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
# Increment pointers.
|
||||
return dk, dv
|
||||
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr,
|
||||
# Filled in by the wrapper.
|
||||
start_m, start_n, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N2)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
|
||||
step_n = BLOCK_N2
|
||||
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M2
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
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
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
|
||||
p = tl.where(mask[None, :], p , 0.0)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd(Q, K, V, sm_scale, #
|
||||
DO, #
|
||||
DQ, DK, DV, #
|
||||
M, D,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_z, stride_h, stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr):
|
||||
LN2 = 0.6931471824645996 # = ln(2)
|
||||
|
||||
bhid = tl.program_id(2)
|
||||
off_chz = (bhid * N_CTX).to(tl.int64)
|
||||
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
|
||||
pid = tl.program_id(0)
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += adj
|
||||
K += adj
|
||||
V += adj
|
||||
DO += adj
|
||||
DQ += adj
|
||||
DK += adj
|
||||
DV += adj
|
||||
M += off_chz
|
||||
D += off_chz
|
||||
|
||||
# load scales
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
start_n = pid * BLOCK_N1
|
||||
start_m = 0
|
||||
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
|
||||
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
|
||||
num_steps = N_CTX // BLOCK_M1
|
||||
|
||||
dk, dv = _attn_bwd_dkdv( #
|
||||
dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
|
||||
start_n, start_m, num_steps #
|
||||
)
|
||||
|
||||
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dv_ptrs, dv)
|
||||
|
||||
# Write back dK.
|
||||
dk *= sm_scale
|
||||
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dk_ptrs, dk)
|
||||
|
||||
# THIS BLOCK DOES DQ:
|
||||
start_m = pid * BLOCK_M2
|
||||
end_n = 0
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
|
||||
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
num_steps = N_CTX // BLOCK_N2
|
||||
dq = _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
|
||||
start_m, end_n, num_steps #
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
|
||||
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
|
||||
|
||||
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
q, k, v, sm_scale,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
M, o,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
B, H, T,
|
||||
HEAD_DIM=D, STAGE=3
|
||||
)
|
||||
|
||||
return o, M
|
||||
|
||||
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do, #
|
||||
delta, #
|
||||
BATCH, N_HEAD, N_CTX, #
|
||||
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
|
||||
)
|
||||
|
||||
|
||||
max_q_blks = k2q_index.shape[-1]
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
|
||||
N_HEAD, N_CTX, #
|
||||
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
|
||||
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
|
||||
HEAD_DIM=D #
|
||||
)
|
||||
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@@ -1,152 +0,0 @@
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
@@ -1,868 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
|
||||
|
||||
import random
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
except ImportError:
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
_flash_attn_varlen_backward = _unsupported
|
||||
flash_attn_varlen_func = _unsupported
|
||||
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def calc_chunks(cu_seqlen, moba_chunk_size):
|
||||
"""
|
||||
Calculate chunk boundaries.
|
||||
|
||||
For vision tasks we include all chunks (even the last one which might be shorter)
|
||||
so that every chunk can be selected.
|
||||
"""
|
||||
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
|
||||
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
|
||||
cu_num_chunk = torch.ones(
|
||||
batch_num_chunk.numel() + 1,
|
||||
device=cu_seqlen.device,
|
||||
dtype=batch_num_chunk.dtype,
|
||||
)
|
||||
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
|
||||
num_chunk = cu_num_chunk[-1]
|
||||
chunk_sizes = torch.full(
|
||||
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_sizes[0] = 0
|
||||
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
|
||||
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
|
||||
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
|
||||
chunk_to_batch = torch.zeros(
|
||||
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_to_batch[cu_num_chunk[1:-1]] = 1
|
||||
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
|
||||
|
||||
# Do not filter out any chunk
|
||||
filtered_chunk_indices = torch.arange(
|
||||
num_chunk, device=cu_seqlen.device, dtype=torch.int32
|
||||
)
|
||||
num_filtered_chunk = num_chunk
|
||||
|
||||
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
|
||||
|
||||
|
||||
# --- Threshold Selection Helper Functions ---
|
||||
|
||||
def _select_threshold_query_head(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects chunks for each <query, head> pair based on threshold.
|
||||
Normalization and sorting happen along the chunk dimension (dim=0).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per <head, query> (across chunks)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
row_min = gate_min_val.amin(dim=0) # (H, S)
|
||||
row_max = gate_masked.amax(dim=0) # (H, S)
|
||||
denom = row_max - row_min
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
|
||||
|
||||
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
|
||||
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
|
||||
|
||||
# 2) compute how much more normalized weight we need beyond self
|
||||
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
|
||||
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
|
||||
|
||||
# 3) zero out the self‐chunk in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0
|
||||
|
||||
# 4) sort the other chunks by descending norm, per <head,seq>
|
||||
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per <head,seq>
|
||||
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
|
||||
|
||||
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
|
||||
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
|
||||
any_cond = cond.any(dim=0) # (H, S)
|
||||
# Find the index of the first True value along dim 0. If none, use C-1.
|
||||
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff
|
||||
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
|
||||
|
||||
# 8) scatter it back to original chunk order
|
||||
others_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
others_mask.scatter_(0, sorted_idx, sorted_mask)
|
||||
|
||||
# 9) finally, include every self‐chunk plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_block(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <query, head> pairs for each block based on threshold.
|
||||
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
HS = H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per block (across heads and queries)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_denom = block_max - block_min
|
||||
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
|
||||
|
||||
gate_norm = (gate - block_min) / block_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights *per block*
|
||||
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
|
||||
|
||||
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
|
||||
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
|
||||
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort the other <head, seq> pairs by descending norm, per block
|
||||
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
|
||||
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per block
|
||||
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
|
||||
|
||||
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
|
||||
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
|
||||
any_cond = cond_flat.any(dim=1) # (C,)
|
||||
# Find the index of the first True value along dim 1. If none, use HS-1.
|
||||
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff per block
|
||||
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
|
||||
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
|
||||
|
||||
# 8) scatter it back to original <head, seq> order per block
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
|
||||
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_overall(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query, head> triplets globally based on threshold.
|
||||
Normalization and sorting happen across all valid entries.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
CHS = C * H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization globally across all valid entries
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
overall_max = gate_masked.max() # scalar
|
||||
overall_min = gate_min_val.min() # scalar
|
||||
overall_denom = overall_max - overall_min
|
||||
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
|
||||
|
||||
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights globally
|
||||
self_norm_sum_overall = self_norm_entries.sum() # scalar
|
||||
|
||||
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
|
||||
total_norm_sum_overall = gate_norm.sum() # scalar
|
||||
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort all other entries by descending norm, globally
|
||||
others_flat = others_norm.flatten() # (C*H*S,)
|
||||
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
|
||||
|
||||
# Only sort the valid 'other' entries
|
||||
valid_others_indices = torch.where(valid_others_mask_flat)[0]
|
||||
valid_others_values = others_flat[valid_others_indices]
|
||||
|
||||
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
|
||||
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
|
||||
|
||||
# 5) cumulative‑sum the sorted valid 'other' norms globally
|
||||
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
|
||||
|
||||
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
|
||||
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
|
||||
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
|
||||
any_cond = cond_values.any() # scalar
|
||||
|
||||
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
|
||||
cutoff_idx_in_sorted = torch.where(
|
||||
any_cond,
|
||||
cond_values.float().argmax(dim=0),
|
||||
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
|
||||
)
|
||||
|
||||
# 7) build a mask selecting the top-k others based on the cutoff
|
||||
# Select the original indices corresponding to the top entries in the sorted list
|
||||
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
|
||||
|
||||
# 8) create the mask in the original flat shape
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
|
||||
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
|
||||
others_mask_flat[selected_other_indices] = True
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_head_global(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query> globally for each head based on threshold.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# 1) LSE‐style normalization per head (across chunks and sequence dims)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
|
||||
|
||||
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
denom = max_per_head - min_per_head
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
|
||||
|
||||
gate_norm = (gate - min_per_head) / denom
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 2) sum normalized self‐chunk contributions per head
|
||||
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 3) total normalized sum per head
|
||||
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 4) how much more normalized weight needed per head
|
||||
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0)
|
||||
|
||||
# 5) zero out self‐chunk entries to focus on "others"
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
|
||||
|
||||
# 6) flatten chunk and sequence dims, per head
|
||||
CS = C * S
|
||||
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
|
||||
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
|
||||
# 7) vectorized selection of “others” per head
|
||||
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
|
||||
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
|
||||
|
||||
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
|
||||
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
|
||||
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
|
||||
|
||||
has_cutoff = cond.any(dim=1) # (H,)
|
||||
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
|
||||
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
|
||||
|
||||
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
|
||||
|
||||
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
|
||||
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
|
||||
|
||||
# 8) reshape selection mask back to (C, H, S)
|
||||
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
|
||||
|
||||
# 9) include self‐chunks plus selected others, and obey valid mask
|
||||
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
class MixedAttention(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
):
|
||||
ctx.max_seqlen = max_seqlen
|
||||
ctx.moba_chunk_size = moba_chunk_size
|
||||
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
|
||||
|
||||
# Non-causal self-attention branch
|
||||
# return out, softmax_lse, S_dmask, rng_state
|
||||
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
# MOBA attention branch (non-causal)
|
||||
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
|
||||
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
|
||||
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
|
||||
|
||||
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
output_2d = output.view(-1, q.shape[2])
|
||||
|
||||
max_lse_1d = self_attn_lse_sh.view(-1)
|
||||
max_lse_1d = max_lse_1d.index_reduce(
|
||||
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
|
||||
)
|
||||
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
|
||||
moba_attn_lse = (
|
||||
moba_attn_lse.view(-1)
|
||||
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
|
||||
.reshape_as(moba_attn_lse)
|
||||
)
|
||||
|
||||
mixed_attn_se_sh = self_attn_lse_sh.exp()
|
||||
moba_attn_se = moba_attn_lse.exp()
|
||||
|
||||
mixed_attn_se_sh.view(-1).index_add_(
|
||||
0, moba_q_sh_indices, moba_attn_se.view(-1)
|
||||
)
|
||||
mixed_attn_lse_sh = mixed_attn_se_sh.log()
|
||||
|
||||
# Combine self-attention output
|
||||
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
|
||||
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
|
||||
output_2d += self_attn_out_sh.reshape_as(output_2d)
|
||||
|
||||
# Combine MOBA attention output
|
||||
mixed_attn_lse = (
|
||||
mixed_attn_lse_sh.view(-1)
|
||||
.index_select(0, moba_q_sh_indices)
|
||||
.view_as(moba_attn_lse)
|
||||
)
|
||||
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
|
||||
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
|
||||
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
|
||||
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
|
||||
output = output.to(q.dtype)
|
||||
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
|
||||
ctx.save_for_backward(
|
||||
output,
|
||||
mixed_attn_lse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, d_output):
|
||||
|
||||
max_seqlen = ctx.max_seqlen
|
||||
moba_chunk_size = ctx.moba_chunk_size
|
||||
softmax_scale = ctx.softmax_scale
|
||||
|
||||
(
|
||||
output,
|
||||
mixed_attn_vlse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
) = ctx.saved_tensors
|
||||
|
||||
d_output = d_output.contiguous()
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_output,
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
out=output,
|
||||
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
|
||||
dq=dq,
|
||||
dk=dk,
|
||||
dv=dv,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
headdim = q.shape[-1]
|
||||
d_moba_output = (
|
||||
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
moba_output = (
|
||||
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
|
||||
mixed_attn_vlse = (
|
||||
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
|
||||
)
|
||||
|
||||
dmq = torch.empty_like(moba_q)
|
||||
dmkv = torch.empty_like(moba_kv)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_moba_output,
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
out=moba_output,
|
||||
softmax_lse=mixed_attn_vlse,
|
||||
dq=dmq,
|
||||
dk=dmkv[:,0],
|
||||
dv=dmkv[:,1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
|
||||
|
||||
|
||||
def moba_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int,
|
||||
moba_chunk_size: int,
|
||||
moba_topk: int,
|
||||
select_mode: str = 'threshold', # "topk" or "threshold"
|
||||
simsum_threshold: float = 0.25,
|
||||
threshold_type: str = 'query_head',
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Accelerated MOBA attention for vision tasks with proper LSE normalization.
|
||||
|
||||
This version:
|
||||
- Splits KV into chunks.
|
||||
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
|
||||
by amplifying the diagonal (self-chunk) logits.
|
||||
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
|
||||
reduction so that attending to each query over the selected chunks is equivalent
|
||||
to the original algorithm.
|
||||
"""
|
||||
# Stack keys and values.
|
||||
kv = torch.stack((k, v), dim=1)
|
||||
seqlen, num_head, head_dim = q.shape
|
||||
|
||||
# Compute chunk boundaries.
|
||||
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
|
||||
cu_seqlens, moba_chunk_size
|
||||
)
|
||||
|
||||
self_attn_cu_seqlen = cu_chunk
|
||||
|
||||
# Update top-k selection to include the self chunk.
|
||||
moba_topk = min(moba_topk, num_filtered_chunk)
|
||||
|
||||
# --- Build filtered KV from chunks ---
|
||||
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
|
||||
max_chunk_len = int(chunk_lengths.max().item())
|
||||
|
||||
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
|
||||
indices = chunk_starts.unsqueeze(1) + range_tensor
|
||||
indices = torch.clamp(indices, max=kv.shape[0] - 1)
|
||||
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
|
||||
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
|
||||
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
|
||||
|
||||
# Compute key_gate_weight over valid tokens.
|
||||
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
|
||||
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
|
||||
key_sum = (key_values * valid_mask_exp).sum(dim=1)
|
||||
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
|
||||
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
|
||||
|
||||
# Compute gate logits between key_gate_weight and queries.
|
||||
q_float = q.float()
|
||||
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
|
||||
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
|
||||
|
||||
# Amplify the diagonal (self chunk) contributions.
|
||||
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
|
||||
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
|
||||
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
|
||||
amplification_factor = 1e9 # Example factor; adjust as needed.
|
||||
origin_gate = gate.clone()
|
||||
gate = gate.clone()
|
||||
if select_mode == "topk":
|
||||
gate[gate_self_chunk_mask] += amplification_factor
|
||||
|
||||
# Exclude positions that are outside the valid batch boundaries.
|
||||
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
|
||||
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
|
||||
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
|
||||
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
|
||||
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
|
||||
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
|
||||
|
||||
if select_mode == 'topk':
|
||||
# We amplify self‐chunk in gate already, so self entries will rank highest.
|
||||
valid_gate_mask = gate != -float("inf")
|
||||
if threshold_type == 'query_head':
|
||||
# === per‐<head,seq> top-k across chunks (original behavior) ===
|
||||
# gate: (C, H, S)
|
||||
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
|
||||
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
gate_idx_mask.scatter_(0, gate_topk_idx, True)
|
||||
gate_mask = valid_gate_mask & gate_idx_mask
|
||||
elif threshold_type == 'overall':
|
||||
# === global top-k across all (chunk, head, seq) entries ===
|
||||
C, H, S = gate.shape
|
||||
flat_gate = gate.flatten()
|
||||
flat_mask = valid_gate_mask.flatten()
|
||||
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
|
||||
# pick topk global entries
|
||||
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
|
||||
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
|
||||
others_mask_flat[idx] = True
|
||||
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
|
||||
elif threshold_type == 'head_global':
|
||||
# per-head top-k across all chunks and sequence positions
|
||||
C, H, S = gate.shape
|
||||
CS = C * S
|
||||
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
|
||||
# pick top-k indices per head
|
||||
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
|
||||
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
|
||||
gate_idx_flat.scatter_(1, topk_idx, True)
|
||||
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid threshold_type for topk: {threshold_type}. "
|
||||
"Choose 'query_head', 'block', or 'overall'."
|
||||
)
|
||||
elif select_mode == 'threshold':
|
||||
# Delegate to the specific thresholding function
|
||||
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
|
||||
if threshold_type == 'query_head':
|
||||
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'block':
|
||||
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'overall':
|
||||
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'head_global':
|
||||
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
else:
|
||||
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
|
||||
else:
|
||||
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
|
||||
|
||||
# eliminate self_chunk in MoBA branch
|
||||
gate_mask = gate_mask & ~gate_self_chunk_mask
|
||||
# if gate_mask is all false, perform flash_attn instead
|
||||
if gate_mask.sum() == 0:
|
||||
return flash_attn_varlen_func(
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
|
||||
)
|
||||
|
||||
# Determine which query positions are selected.
|
||||
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
|
||||
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
|
||||
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
|
||||
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
|
||||
|
||||
# Build cumulative sequence lengths for the selected queries.
|
||||
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
|
||||
q_zero_mask = moba_seqlen_q == 0
|
||||
valid_expert_mask = ~q_zero_mask
|
||||
if q_zero_mask.sum() > 0:
|
||||
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
|
||||
moba_cu_seqlen_q = torch.cat(
|
||||
(
|
||||
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
|
||||
moba_seqlen_q.cumsum(dim=0),
|
||||
),
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
# Rearrange gathered KV for the MOBA branch.
|
||||
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
|
||||
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
|
||||
if q_zero_mask.sum() > 0:
|
||||
experts_tensor = experts_tensor[valid_expert_mask]
|
||||
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
|
||||
|
||||
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
|
||||
mask = seq_range < valid_expert_lengths.unsqueeze(1)
|
||||
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
|
||||
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
|
||||
|
||||
moba_cu_seqlen_kv = torch.cat(
|
||||
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
|
||||
valid_expert_lengths.cumsum(dim=0)],
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
assert (
|
||||
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
|
||||
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
|
||||
|
||||
return MixedAttention.apply(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
|
||||
def process_moba_input(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
"""
|
||||
Process inputs for the attention function.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
|
||||
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
|
||||
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed input tensor.
|
||||
"""
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
|
||||
else:
|
||||
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
|
||||
if len(chunk_size) == 2:
|
||||
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
|
||||
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
|
||||
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
|
||||
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
|
||||
elif len(chunk_size) == 3:
|
||||
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
|
||||
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
|
||||
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
|
||||
else:
|
||||
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
|
||||
|
||||
return x, moba_chunk_size
|
||||
|
||||
|
||||
def process_moba_output(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
pass
|
||||
elif len(chunk_size) == 2:
|
||||
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
|
||||
elif len(chunk_size) == 3:
|
||||
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# TEST
|
||||
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
|
||||
random.seed(0)
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
|
||||
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
|
||||
max_seqlen = q.shape[1]
|
||||
q = rearrange(q, "b s ... -> (b s) ...")
|
||||
k = rearrange(k, "b s ... -> (b s) ...")
|
||||
v = rearrange(v, "b s ... -> (b s) ...")
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
|
||||
"""Speed test comparing flash_attn vs moba_attention"""
|
||||
# Get data
|
||||
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
|
||||
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
|
||||
vo_grad = torch.randn_like(q)
|
||||
|
||||
# Warmup
|
||||
warmup_iters = 3
|
||||
perf_test_iters = 10
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_flash = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_moba = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
|
||||
|
||||
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
|
||||
print(f"Speedup: {time_flash / time_moba:.2f}x")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
CUDA_VISIBLE_DEVICES=1 \
|
||||
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
"""
|
||||
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
|
||||
@@ -1,71 +0,0 @@
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import BoolTensor, IntTensor
|
||||
from torch.nn.attention.flex_attention import create_block_mask
|
||||
|
||||
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
|
||||
torch._inductor.config.realize_opcount_threshold = 100
|
||||
|
||||
|
||||
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
|
||||
"""Generates a 3D NATTEN attention mask with a given kernel size.
|
||||
|
||||
Args:
|
||||
canvas_t: The time dimension of the canvas.
|
||||
canvas_h: The height of the canvas.
|
||||
canvas_w: The width of the canvas.
|
||||
kernel_t: The time dimension of the kernel.
|
||||
kernel_h: The height of the kernel.
|
||||
kernel_w: The width of the kernel.
|
||||
"""
|
||||
canvas_t, canvas_h, canvas_w = canvas_twh
|
||||
kernel_t, kernel_h, kernel_w = kernel_twh
|
||||
tile_t_size, tile_h_size, tile_w_size = tile_twh
|
||||
total_tile_size = tile_t_size * tile_h_size * tile_w_size
|
||||
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
|
||||
tile_id = idx // total_tile_size
|
||||
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
|
||||
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
|
||||
tile_w = tile_id % canvas_tile_w
|
||||
return tile_t, tile_h, tile_w
|
||||
|
||||
def sta_mask_3d(
|
||||
b: IntTensor,
|
||||
h: IntTensor,
|
||||
q_idx: IntTensor,
|
||||
kv_idx: IntTensor,
|
||||
) -> BoolTensor:
|
||||
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
|
||||
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
|
||||
# kernel nominally attempts to center itself on the query, but kernel center
|
||||
# is clamped to a fixed distance (kernel half-length) from the canvas edge
|
||||
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
|
||||
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
|
||||
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
|
||||
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
|
||||
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
|
||||
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
|
||||
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
|
||||
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
|
||||
|
||||
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return sta_mask_3d
|
||||
|
||||
|
||||
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
|
||||
img_seq_len = img_size[0] * img_size[1] * img_size[2]
|
||||
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
|
||||
mask = create_block_mask(image_mask,
|
||||
B=None,
|
||||
H=None,
|
||||
Q_LEN=img_seq_len + text_max_len,
|
||||
KV_LEN=img_seq_len + text_max_len,
|
||||
device=device,
|
||||
_compile=True)
|
||||
return mask
|
||||
@@ -1,63 +0,0 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
|
||||
# Local support import
|
||||
from .support_flex_sta import get_sliding_tile_attention_mask
|
||||
|
||||
# USE OUR NEW PACKAGE!
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
return output
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
# Using the same parameters as the original test
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
|
||||
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
|
||||
for kernel_size in kernel_size_ls:
|
||||
print(f"Testing kernel_size: {kernel_size}")
|
||||
for xi in tqdm(range(num_iterations)):
|
||||
torch.manual_seed(xi)
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
|
||||
pt_o = flex_test(Q, K, V, kernel_size)
|
||||
|
||||
diff = pt_o - tk_o
|
||||
abs_diff = torch.abs(diff)
|
||||
max_d = torch.max(abs_diff).item()
|
||||
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
|
||||
|
||||
if max_d > 0.1:
|
||||
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
|
||||
|
||||
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
|
||||
@@ -1,97 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from fastvideo_kernel.vmoba import moba_attn_varlen
|
||||
|
||||
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
|
||||
"""
|
||||
Generates random data for testing the variable-length attention function.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
random.seed(42)
|
||||
torch.cuda.manual_seed_all(42)
|
||||
|
||||
# Generate sequence lengths for each item in the batch
|
||||
if batch_size > 1:
|
||||
# Ensure sequence lengths are reasonably distributed
|
||||
avg_seqlen = total_seqlen // batch_size
|
||||
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
|
||||
remaining_len = total_seqlen - sum(seqlens)
|
||||
if remaining_len > 0:
|
||||
seqlens.append(remaining_len)
|
||||
else: # Adjust if sum exceeds total_seqlen
|
||||
seqlens.append(avg_seqlen)
|
||||
current_sum = sum(seqlens)
|
||||
seqlens[-1] -= (current_sum - total_seqlen)
|
||||
# Ensure all lengths are positive
|
||||
seqlens = [max(1, s) for s in seqlens]
|
||||
# Final adjustment to match total_seqlen
|
||||
seqlens[-1] += total_seqlen - sum(seqlens)
|
||||
|
||||
else:
|
||||
seqlens = [total_seqlen]
|
||||
|
||||
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
|
||||
max_seqlen = max(seqlens) if seqlens else 0
|
||||
|
||||
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [1, 2])
|
||||
@pytest.mark.parametrize("total_seqlen", [512, 1024])
|
||||
@pytest.mark.parametrize("num_heads", [8])
|
||||
@pytest.mark.parametrize("head_dim", [64])
|
||||
@pytest.mark.parametrize("moba_chunk_size", [64])
|
||||
@pytest.mark.parametrize("moba_topk", [2, 4])
|
||||
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
|
||||
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_moba_attn_varlen_forward(
|
||||
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
|
||||
):
|
||||
"""
|
||||
Tests the forward pass of moba_attn_varlen for basic correctness.
|
||||
It checks output shape, dtype, and for the presence of NaNs/Infs.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
pytest.skip("float32 is not supported in flash attention")
|
||||
|
||||
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
|
||||
batch_size, total_seqlen, num_heads, head_dim, dtype
|
||||
)
|
||||
|
||||
# Ensure chunk size is not larger than the smallest sequence length
|
||||
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
|
||||
if moba_chunk_size > min_seqlen:
|
||||
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
|
||||
|
||||
try:
|
||||
output = moba_attn_varlen(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
moba_chunk_size=moba_chunk_size,
|
||||
moba_topk=moba_topk,
|
||||
select_mode=select_mode,
|
||||
threshold_type=threshold_type,
|
||||
simsum_threshold=0.5, # A reasonable default for threshold mode
|
||||
)
|
||||
except Exception as e:
|
||||
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
|
||||
|
||||
# 1. Check output shape
|
||||
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
|
||||
|
||||
# 2. Check output dtype
|
||||
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
|
||||
|
||||
# 3. Check for NaNs or Infs in the output
|
||||
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
|
||||
@@ -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 . .
|
||||
|
||||
|
||||
@@ -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,10 +55,17 @@ 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 Kernels
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -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,10 +55,17 @@ 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 Kernels
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -55,10 +55,17 @@ 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 Kernels
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -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 Kernels
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r requirements-mkdocs.txt
|
||||
pip install -r docs/requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -70,7 +70,3 @@ uv pip install ninja
|
||||
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Testing
|
||||
|
||||
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Profiling FastVideo
|
||||
|
||||
!!! warning
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down inference.
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
|
||||
|
||||
## Profiling with PyTorch
|
||||
|
||||
@@ -49,5 +49,5 @@ Traces can be visualized using <https://ui.perfetto.dev/>.
|
||||
### Best Practices
|
||||
|
||||
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
|
||||
- After profiling, clean up trace directories to avoid filling disk storage.
|
||||
- After profiling, clean up trace directories to avoid filling disks.
|
||||
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
|
||||
|
||||
@@ -1,131 +0,0 @@
|
||||
# Testing in FastVideo
|
||||
|
||||
This guide explains how to add and run tests in FastVideo. The testing suite is divided into several categories to ensure correctness across components, training workflows, and inference quality.
|
||||
|
||||
## Test Types
|
||||
|
||||
* **Unit Tests**: Located in `fastvideo/tests/dataset`, `fastvideo/tests/entrypoints`, and `fastvideo/tests/workflow`. These test individual functions and classes.
|
||||
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
|
||||
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
|
||||
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
|
||||
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
|
||||
|
||||
For now, we will focus on **SSIM Tests**.
|
||||
|
||||
## SSIM Tests
|
||||
|
||||
SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos using specific models and parameters, and compare them against reference videos to ensure that changes in the codebase do not degrade generation quality or alter the output unexpectedly.
|
||||
|
||||
!!! note
|
||||
If you are adding an SSIM test, this serves as a safeguard. Any future code changes that break or cause errors with the specific arguments and configurations you defined will trigger a failure. Therefore, it is important to include multiple settings and arguments that cover the core features of your new pipeline to ensure robust regression testing.
|
||||
|
||||
### Directory Structure
|
||||
|
||||
```
|
||||
fastvideo/tests/ssim/
|
||||
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
|
||||
│ ├── <Model_Name>/
|
||||
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
|
||||
│ │ │ └── <Video_File>
|
||||
├── test_causal_similarity.py
|
||||
├── test_inference_similarity.py
|
||||
├── update_reference_videos.sh
|
||||
└── ...
|
||||
```
|
||||
|
||||
### Adding a New SSIM Test
|
||||
|
||||
To add a new SSIM test, follow these steps:
|
||||
|
||||
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
|
||||
|
||||
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
|
||||
|
||||
```python
|
||||
MY_MODEL_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": "organization/model-name",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
# ... other parameters
|
||||
}
|
||||
```
|
||||
|
||||
3. **Implement the Test Function**:
|
||||
* Use `pytest.mark.parametrize` to run the test with different prompts, backends, and models.
|
||||
* Set the attention backend environment variable.
|
||||
* Initialize the `VideoGenerator`.
|
||||
* Generate the video.
|
||||
* Compare the generated video with the reference video using `compute_video_ssim_torchvision`.
|
||||
|
||||
Example structure:
|
||||
|
||||
```python
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
def test_my_model_similarity(prompt, ATTENTION_BACKEND):
|
||||
# Setup output directories
|
||||
# ...
|
||||
|
||||
# Initialize Generator
|
||||
generator = VideoGenerator.from_pretrained(...)
|
||||
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
|
||||
```
|
||||
|
||||
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`.
|
||||
* Inspect the generated video to ensure it meets quality expectations.
|
||||
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
|
||||
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
|
||||
|
||||
### Running Tests Locally
|
||||
|
||||
To run the SSIM tests locally:
|
||||
|
||||
```bash
|
||||
pytest fastvideo/tests/ssim/ -vs
|
||||
```
|
||||
|
||||
Ensure you have the necessary GPUs available as defined in your test parameters.
|
||||
|
||||
## Modal Workflow
|
||||
|
||||
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
|
||||
|
||||
### `pr_test.py`
|
||||
|
||||
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
|
||||
|
||||
### Updating Modal Configuration
|
||||
|
||||
If you add a new test that requires:
|
||||
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
|
||||
* **Longer Execution Time**: Increase the `timeout` parameter.
|
||||
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
|
||||
|
||||
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
|
||||
|
||||
```python
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_ssim_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
```
|
||||
|
||||
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
|
||||
|
||||
### Workflow Scripts
|
||||
|
||||
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
|
||||
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
|
||||
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
|
||||
|
||||
!!! note
|
||||
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
|
||||
@@ -1,6 +1,6 @@
|
||||
# 🎯 Distillation
|
||||
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computations, enabling much faster video generation.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
@@ -13,7 +13,7 @@ We provide two distilled models:
|
||||
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
|
||||
|
||||
## ⚙️ Inference
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
|
||||
@@ -7,11 +7,6 @@ Get up and running with FastVideo in minutes!
|
||||
First, install FastVideo:
|
||||
|
||||
```bash
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
@@ -20,53 +15,36 @@ pip install fastvideo
|
||||
### Text-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo import FastVideoPipeline
|
||||
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-t2v-1.3B")
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
# Generate a video
|
||||
prompt = "A cat playing with a ball of yarn"
|
||||
video = pipe(prompt, num_frames=16, height=512, width=512)
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
from fastvideo import FastVideoPipeline
|
||||
from PIL import Image
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
# Load an image
|
||||
image = Image.open("input.jpg")
|
||||
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-i2v-14B-480p")
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
# Generate a video from the image
|
||||
video = pipe(image, num_frames=16, height=480, width=480)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
```
|
||||
|
||||
## Next Steps
|
||||
|
||||
+2
-3
@@ -24,13 +24,12 @@ FastVideo is an inference and post-training framework for diffusion models. It f
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- E2E post-training support
|
||||
- Data preprocessing pipeline for video data
|
||||
- Data preprocessing pipeline for video data.
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
@@ -43,7 +42,7 @@ Use the navigation menu on the left to explore different sections:
|
||||
|
||||
- **Getting Started**: Installation and quick start guides
|
||||
- **Inference**: Learn how to use FastVideo for video generation
|
||||
- **Training**: Data preprocessing and fine-tuning workflows
|
||||
- **Training**: Data preprocessing and fine-tuning workflows
|
||||
- **Distillation**: Post-training optimization techniques
|
||||
- **Sliding Tile Attention**: Advanced attention mechanisms
|
||||
- **Video Sparse Attention**: Efficient attention for video models
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# Seed Parameter Behavior in vLLM
|
||||
|
||||
## Overview
|
||||
|
||||
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
|
||||
|
||||
## Default Behavior
|
||||
|
||||
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
|
||||
|
||||
## Specifying a Seed
|
||||
|
||||
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
|
||||
|
||||
## Example Usage
|
||||
|
||||
### Without Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model without specifying a seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs different numbers across runs
|
||||
```
|
||||
|
||||
### Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model with a specific seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs the same number across runs
|
||||
```
|
||||
|
||||
## Important Notes
|
||||
|
||||
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
|
||||
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
|
||||
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
|
||||
|
||||
## Conclusion
|
||||
|
||||
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
|
||||
@@ -7,7 +7,7 @@ pip install st_attn
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently, we only have an implementation for H100s.
|
||||
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
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 98 KiB |
@@ -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
|
||||
|
||||
@@ -7,9 +7,9 @@ pip install vsa
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We support H100s (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
We support H100 (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
|
||||
First, install C++20 for ThunderKittens (if using an H100):
|
||||
First, install C++20 for ThunderKittens (if using H100):
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
|
||||
@@ -6,7 +6,7 @@ The `VideoGenerator` class provides the primary Python interface for doing offli
|
||||
- Python 3.10-3.12
|
||||
|
||||
## Installation
|
||||
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation) first.
|
||||
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) first.
|
||||
|
||||
## Usage
|
||||
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
|
||||
|
||||
@@ -1,42 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,75 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.wan import MatrixGameI2V480PConfig
|
||||
from fastvideo.models.dits.matrix_game.utils import create_action_presets
|
||||
|
||||
import torch
|
||||
|
||||
# Available variants: "base_distilled_model", "gta_distilled_model", "templerun_distilled_model"
|
||||
# Each variant has different keyboard_dim:
|
||||
# - base_distilled_model: keyboard_dim=4
|
||||
# - gta_distilled_model: keyboard_dim=2
|
||||
# - templerun_distilled_model: keyboard_dim=7 (keyboard only, no mouse)
|
||||
MODEL_VARIANT = "base_distilled_model"
|
||||
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
config = VARIANT_CONFIG[MODEL_VARIANT]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
num_frames = 597
|
||||
actions = create_action_presets(num_frames, keyboard_dim=config["keyboard_dim"])
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
|
||||
generator.generate_video(
|
||||
prompt="",
|
||||
image_path=config["image_url"],
|
||||
mouse_cond=actions["mouse"].unsqueeze(0),
|
||||
keyboard_cond=actions["keyboard"].unsqueeze(0),
|
||||
grid_sizes=grid_sizes,
|
||||
num_frames=num_frames,
|
||||
height=352,
|
||||
width=640,
|
||||
num_inference_steps=50,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -10,8 +10,9 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_id = "FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
model_id,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
@@ -25,7 +26,7 @@ def main():
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
|
||||
sampling_param = SamplingParam.from_pretrained(model_id)
|
||||
sampling_param.num_frames = 81
|
||||
sampling_param.width = 832
|
||||
sampling_param.height = 480
|
||||
|
||||
@@ -6,8 +6,10 @@ import time
|
||||
import json
|
||||
from pathlib import Path
|
||||
import tempfile
|
||||
from io import BytesIO
|
||||
|
||||
import gradio as gr
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
@@ -15,7 +17,7 @@ from fastvideo.configs.sample.base import SamplingParam
|
||||
MODEL_PATH_MAPPING = {
|
||||
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"CausalWan2.2-I2V-A14B-Preview": "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
"CausalWan2.2-I2V-A14B-Preview": "FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
}
|
||||
|
||||
|
||||
@@ -85,28 +87,39 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
|
||||
return None
|
||||
|
||||
|
||||
def encode_image_to_base64(image_path: str) -> str:
|
||||
"""Encode an image file to base64 string."""
|
||||
if not image_path or not os.path.exists(image_path):
|
||||
def encode_image_to_base64(image_input) -> str:
|
||||
"""Encode an image file path or in-memory image to a base64 string."""
|
||||
if image_input is None:
|
||||
return None
|
||||
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
|
||||
try:
|
||||
with open(image_path, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
if isinstance(image_input, str):
|
||||
if not os.path.exists(image_input):
|
||||
return None
|
||||
|
||||
with open(image_input, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
|
||||
ext = os.path.splitext(image_input)[1].lower()
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
elif isinstance(image_input, Image.Image):
|
||||
buffer = BytesIO()
|
||||
image_to_save = image_input.convert("RGB")
|
||||
image_to_save.save(buffer, format="PNG")
|
||||
image_bytes = buffer.getvalue()
|
||||
mime_type = 'image/png'
|
||||
else:
|
||||
return None
|
||||
|
||||
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
|
||||
|
||||
# Determine image type from extension
|
||||
ext = os.path.splitext(image_path)[1].lower()
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
|
||||
return f"data:{mime_type};base64,{image_base64}"
|
||||
|
||||
except Exception as e:
|
||||
@@ -426,7 +439,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
gr.Markdown("**Please make sure you upload a 480x832 image**")
|
||||
input_image = gr.Image(
|
||||
label="",
|
||||
type="filepath",
|
||||
type="pil",
|
||||
height=400,
|
||||
)
|
||||
|
||||
@@ -512,7 +525,17 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
if example_label and example_label in example_labels:
|
||||
index = example_labels.index(example_label)
|
||||
selected_prompt = examples[index]
|
||||
selected_image = example_images[index] if index < len(example_images) else None
|
||||
selected_image_path = example_images[index] if index < len(example_images) else None
|
||||
|
||||
if selected_image_path and os.path.exists(selected_image_path):
|
||||
try:
|
||||
with Image.open(selected_image_path) as img:
|
||||
selected_image = img.convert("RGB")
|
||||
except Exception:
|
||||
selected_image = None
|
||||
else:
|
||||
selected_image = None
|
||||
|
||||
return selected_prompt, selected_image
|
||||
return "", None
|
||||
|
||||
@@ -615,7 +638,7 @@ def main():
|
||||
default="",
|
||||
help="Comma separated list of paths to the T2V model(s)")
|
||||
parser.add_argument("--i2v_model_paths", type=str,
|
||||
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
default="FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
help="Comma separated list of paths to the I2V model(s)")
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
@@ -733,6 +756,7 @@ def main():
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
root_path="/gradio",
|
||||
allowed_paths=[
|
||||
os.path.abspath("outputs"),
|
||||
os.path.abspath("fastvideo-logos"),
|
||||
|
||||
@@ -12,4 +12,4 @@ Man dressed in 80's style dances very happily in his kitchen while listening to
|
||||
Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.
|
||||
A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.
|
||||
Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.
|
||||
Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.
|
||||
Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.
|
||||
|
||||
@@ -26,7 +26,7 @@ SEED_RANGE_MAX = 1_000_000
|
||||
SUPPORTED_MODELS = [
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
"FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
]
|
||||
|
||||
MODEL_CONFIGS = {
|
||||
@@ -272,7 +272,7 @@ class T2VModelDeployment(BaseModelDeployment):
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class T2V14BModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
|
||||
@@ -284,7 +284,7 @@ class T2V14BModelDeployment(BaseModelDeployment):
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class I2VModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, i2v_model_path: str, output_path: str = "outputs"):
|
||||
@@ -445,7 +445,7 @@ if __name__ == "__main__":
|
||||
help="Comma separated list of number of replicas for the T2V model(s)")
|
||||
parser.add_argument("--i2v_model_paths",
|
||||
type=str,
|
||||
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
default="FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
help="Comma separated list of paths to the I2V model(s)")
|
||||
parser.add_argument("--i2v_model_replicas",
|
||||
type=str,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
python examples/inference/gradio/serving/start_ray_serve_app.py \
|
||||
--t2v_model_paths "" \
|
||||
--t2v_model_replicas "" \
|
||||
--i2v_model_paths "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers" \
|
||||
--i2v_model_paths "FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers" \
|
||||
--i2v_model_replicas "1"
|
||||
|
||||
@@ -5,12 +5,12 @@ These are e2e example scripts for finetuning Wan2.1 T2V 1.3B on the crush-smol d
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/download_dataset.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
|
||||
|
||||
@@ -54,7 +54,7 @@ validation_args=(
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
|
||||
@@ -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[@]}"
|
||||
@@ -3,4 +3,4 @@ from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
@@ -1,9 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
@@ -47,29 +46,6 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlashAttnMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
attn_mask: torch.Tensor | None = None
|
||||
|
||||
|
||||
class FlashAttnMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> FlashAttnMetadata:
|
||||
return FlashAttnMetadata(current_timestep=current_timestep,
|
||||
attn_mask=attn_mask)
|
||||
|
||||
|
||||
class FlashAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -90,27 +66,12 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
attn_metadata: AttentionMetadata,
|
||||
):
|
||||
if attn_metadata is not None and hasattr(
|
||||
attn_metadata,
|
||||
"attn_mask") and attn_metadata.attn_mask is not None:
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import flash_attn_no_pad
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
attn_mask = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0),
|
||||
value=True)
|
||||
output = flash_attn_no_pad(qkv,
|
||||
attn_mask,
|
||||
causal=False,
|
||||
dropout_p=0,
|
||||
softmax_scale=None)
|
||||
else:
|
||||
output = flash_attn_func(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
output = flash_attn_func(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
return output
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -31,29 +30,6 @@ class SDPABackend(AttentionBackend):
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
@dataclass
|
||||
class SDPAMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
attn_mask: torch.Tensor | None = None
|
||||
|
||||
|
||||
class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
return SDPAMetadata(current_timestep=current_timestep,
|
||||
attn_mask=attn_mask)
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -75,15 +51,14 @@ class SDPAImpl(AttentionImpl):
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: SDPAMetadata,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
# transpose to bs, heads, seq_len, head_dim
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
attn_mask = attn_metadata.attn_mask if attn_metadata is not None else None
|
||||
attn_kwargs = {
|
||||
"attn_mask": attn_mask,
|
||||
"attn_mask": None,
|
||||
"dropout_p": self.dropout,
|
||||
"is_causal": self.causal,
|
||||
"scale": self.softmax_scale
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
|
||||
@@ -6,7 +6,7 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
|
||||
try:
|
||||
from fastvideo_kernel import video_sparse_attn
|
||||
from vsa import video_sparse_attn
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo_kernel import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
|
||||
@@ -11,7 +11,6 @@ from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
from fastvideo.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
@@ -65,8 +64,6 @@ class DistributedAttention(nn.Module):
|
||||
replicated_q: torch.Tensor | None = None,
|
||||
replicated_k: torch.Tensor | None = None,
|
||||
replicated_v: torch.Tensor | None = None,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
@@ -77,7 +74,6 @@ class DistributedAttention(nn.Module):
|
||||
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
|
||||
replicated_k (Optional[torch.Tensor]): Replicated key tensor
|
||||
replicated_v (Optional[torch.Tensor]): Replicated value tensor
|
||||
attention_mask (Optional[torch.Tensor]): Attention mask [batch_size, seq_len]
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
|
||||
@@ -95,30 +91,12 @@ class DistributedAttention(nn.Module):
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
# Stack QKV
|
||||
qkv = torch.cat([q, k, v],
|
||||
dim=0) # [3*batch, seq_len, num_heads, head_dim]
|
||||
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# After all-to-all, each rank has the full sequence but only a subset of heads
|
||||
# The attention mask should now apply to the full sequence length
|
||||
# Since mask is [batch, full_seq_len], it's already in the correct format
|
||||
|
||||
# LOAY TODO, instead of slicing repeatedly maintain an original qkv and rewrite into that
|
||||
valid_seq_len = None
|
||||
if attention_mask is not None:
|
||||
valid_seq_len = (attention_mask[0] == 1).sum().item()
|
||||
qkv = qkv[:, :valid_seq_len, :, :]
|
||||
|
||||
if freqs_cis is not None:
|
||||
cos, sin = freqs_cis
|
||||
qkv[:batch_size * 2] = _apply_rotary_emb(qkv[:batch_size * 2],
|
||||
cos,
|
||||
sin,
|
||||
is_neox_style=False)
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
@@ -141,23 +119,17 @@ class DistributedAttention(nn.Module):
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
if replicated_q is not None:
|
||||
split_idx = seq_len * world_size if valid_seq_len is None else valid_seq_len
|
||||
replicated_output = output[:, split_idx:]
|
||||
output = output[:, :split_idx]
|
||||
replicated_output = output[:, seq_len * world_size:]
|
||||
output = output[:, :seq_len * world_size]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
if attention_mask is not None:
|
||||
pad_len = (attention_mask[0] == 0).sum().item()
|
||||
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
@@ -175,8 +147,6 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
replicated_k: torch.Tensor | None = None,
|
||||
replicated_v: torch.Tensor | None = None,
|
||||
gate_compress: torch.Tensor | None = None,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
@@ -188,7 +158,6 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
|
||||
replicated_k (Optional[torch.Tensor]): Replicated key tensor
|
||||
replicated_v (Optional[torch.Tensor]): Replicated value tensor
|
||||
attention_mask (Optional[torch.Tensor]): Attention mask [batch_size, seq_len]
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
|
||||
@@ -204,32 +173,15 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
# Stack QKV
|
||||
qkvg = torch.cat([q, k, v, gate_compress],
|
||||
dim=0) # [4*batch, seq_len, num_heads, head_dim]
|
||||
dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
# Before: [4*batch, shard_seq_len, num_heads, head_dim]
|
||||
# After: [4*batch, full_seq_len, shard_num_heads, head_dim]
|
||||
qkvg = sequence_model_parallel_all_to_all_4D(qkvg,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# After all-to-all, each rank has the full sequence but only a subset of heads
|
||||
# The attention mask should now apply to the full sequence length
|
||||
|
||||
if attention_mask is not None:
|
||||
valid_seq_len = (attention_mask[0] == 1).sum().item()
|
||||
qkvg = qkvg[:, :valid_seq_len, :, :]
|
||||
|
||||
if freqs_cis is not None:
|
||||
cos, sin = freqs_cis
|
||||
qkvg[:batch_size * 2] = _apply_rotary_emb(qkvg[:batch_size * 2],
|
||||
cos,
|
||||
sin,
|
||||
is_neox_style=False)
|
||||
|
||||
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
|
||||
|
||||
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
|
||||
@@ -242,10 +194,6 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
if attention_mask is not None:
|
||||
pad_len = (attention_mask[0] == 0).sum().item()
|
||||
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
@@ -296,7 +244,6 @@ class LocalAttention(nn.Module):
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply local attention between query, key and value tensors.
|
||||
@@ -316,10 +263,5 @@ class LocalAttention(nn.Module):
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
if freqs_cis is not None:
|
||||
cos, sin = freqs_cis
|
||||
q = _apply_rotary_emb(q, cos, sin, is_neox_style=False)
|
||||
k = _apply_rotary_emb(k, cos, sin, is_neox_style=False)
|
||||
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
return output
|
||||
|
||||
@@ -1,99 +0,0 @@
|
||||
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
|
||||
#
|
||||
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
|
||||
# output and results there from are provided "AS IS" without any express or implied warranties of
|
||||
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
|
||||
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
|
||||
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
|
||||
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
|
||||
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
|
||||
# of rights and permissions under this agreement.
|
||||
# See the License for the specific language governing permissions and limitations under the License.
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
def flash_attn_no_pad(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
batch_size = qkv.shape[0]
|
||||
seqlen = qkv.shape[1]
|
||||
nheads = qkv.shape[-2]
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
|
||||
x, key_padding_mask)
|
||||
|
||||
x_unpad = rearrange(x_unpad,
|
||||
"nnz (three h d) -> nnz three h d",
|
||||
three=3,
|
||||
h=nheads)
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
output = rearrange(
|
||||
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices,
|
||||
batch_size, seqlen),
|
||||
"b s (h d) -> b s h d",
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def flash_attn_no_pad_v3(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
from flash_attn_interface import flash_attn_varlen_func as flash_attn_varlen_func_v3
|
||||
|
||||
if flash_attn_varlen_func_v3 is None:
|
||||
raise ImportError("FlashAttention V3 backend not available")
|
||||
|
||||
batch_size, seqlen, _, nheads, head_dim = qkv.shape
|
||||
query, key, value = qkv.unbind(dim=2)
|
||||
|
||||
query_unpad, indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
|
||||
rearrange(query, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
key_unpad, _, cu_seqlens_k, _, _ = unpad_input(
|
||||
rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
value_unpad, _, _, _, _ = unpad_input(
|
||||
rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
|
||||
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
|
||||
output_unpad = flash_attn_varlen_func_v3(query_unpad,
|
||||
key_unpad,
|
||||
value_unpad,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_q,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic)
|
||||
|
||||
output = rearrange(pad_input(
|
||||
rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size,
|
||||
seqlen),
|
||||
"b s (h d) -> b s h d",
|
||||
h=nheads)
|
||||
return output
|
||||
@@ -1,13 +1,9 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
|
||||
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
|
||||
"LongCatVideoConfig"
|
||||
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
|
||||
"CosmosVideoConfig"
|
||||
]
|
||||
|
||||
@@ -23,8 +23,6 @@ class DiTArchConfig(ArchConfig):
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
in_channels: int = 0
|
||||
out_channels: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
|
||||
@@ -1,181 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_transformer_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25ArchConfig(DiTArchConfig):
|
||||
"""Configuration for Cosmos 2.5 architecture (MiniTrainDIT)."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Remove "net." prefix and map official structure to FastVideo
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$":
|
||||
r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding: net.t_embedder.1.linear_1.weight -> time_embed.t_embedder.linear_1.weight
|
||||
r"^net\.t_embedder\.1\.linear_1\.(.*)$":
|
||||
r"time_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.t_embedder\.1\.linear_2\.(.*)$":
|
||||
r"time_embed.t_embedder.linear_2.\1",
|
||||
# Time embedding norm: net.t_embedding_norm.weight -> time_embed.norm.weight
|
||||
# Note: This also handles _extra_state if present
|
||||
r"^net\.t_embedding_norm\.(.*)$":
|
||||
r"time_embed.norm.\1",
|
||||
|
||||
# Cross-attention projection (optional): net.crossattn_proj.0.weight -> crossattn_proj.0.weight
|
||||
r"^net\.crossattn_proj\.0\.weight$":
|
||||
r"crossattn_proj.0.weight",
|
||||
r"^net\.crossattn_proj\.0\.bias$":
|
||||
r"crossattn_proj.0.bias",
|
||||
|
||||
# Transformer blocks: net.blocks.N -> transformer_blocks.N
|
||||
# Self-attention (self_attn -> attn1)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.weight",
|
||||
# RMSNorm _extra_state keys (internal PyTorch state, will be recomputed automatically)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_q._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_k._extra_state",
|
||||
|
||||
# Cross-attention (cross_attn -> attn2)
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.weight",
|
||||
# RMSNorm _extra_state keys for cross-attention
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_q._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_k._extra_state",
|
||||
|
||||
# MLP: net.blocks.N.mlp.layer1 -> transformer_blocks.N.mlp.fc_in
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
|
||||
# AdaLN-LoRA modulations: net.blocks.N.adaln_modulation_* -> transformer_blocks.N.adaln_modulation_*
|
||||
# These are now at the block level, not inside norm layers
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
|
||||
|
||||
# Layer norms: net.blocks.N.layer_norm_* -> transformer_blocks.N.norm*.norm
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_self_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm1.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_cross_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm2.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$":
|
||||
r"transformer_blocks.\1.norm3.norm._extra_state",
|
||||
|
||||
# Final layer: net.final_layer.linear -> final_layer.proj_out
|
||||
r"^net\.final_layer\.linear\.(.*)$":
|
||||
r"final_layer.proj_out.\1",
|
||||
# Final layer AdaLN-LoRA: net.final_layer.adaln_modulation -> final_layer.linear_*
|
||||
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$":
|
||||
r"final_layer.linear_1.\1",
|
||||
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$":
|
||||
r"final_layer.linear_2.\1",
|
||||
|
||||
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
|
||||
# - net.pos_embedder.* (seq, dim_spatial_range, dim_temporal_range) - These are computed dynamically
|
||||
# in FastVideo's Cosmos25RotaryPosEmbed forward() method, so they don't need to be loaded.
|
||||
# - net.accum_* keys (training metadata) - These are skipped during checkpoint loading.
|
||||
})
|
||||
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.\2",
|
||||
})
|
||||
|
||||
# Cosmos 2.5 specific config parameters
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128 # 2048 / 16
|
||||
num_layers: int = 28
|
||||
mlp_ratio: float = 4.0
|
||||
text_embed_dim: int = 1024
|
||||
adaln_lora_dim: int = 256
|
||||
use_adaln_lora: bool = True
|
||||
max_size: tuple[int, int, int] = (128, 240, 240)
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0) # T, H, W scaling
|
||||
concat_padding_mask: bool = True
|
||||
extra_pos_embed_type: str | None = None # "learnable" or None
|
||||
# Note: Official checkpoint has use_crossattn_projection=True with 100K-dim input from Qwen 7B.
|
||||
# When enabled, must provide 100,352-dim embeddings to match the projection layer in checkpoint.
|
||||
use_crossattn_projection: bool = False
|
||||
crossattn_proj_in_channels: int = 100352 # Qwen 7B embedding dimension
|
||||
rope_enable_fps_modulation: bool = True
|
||||
qk_norm: str = "rms_norm"
|
||||
eps: float = 1e-6
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25VideoConfig(DiTConfig):
|
||||
"""Configuration for Cosmos 2.5 video generation model."""
|
||||
arch_config: DiTArchConfig = field(default_factory=Cosmos25ArchConfig)
|
||||
prefix: str = "Cosmos25"
|
||||
@@ -1,157 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_double_block(n: str, m) -> bool:
|
||||
return "double_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_refiner_block(n: str, m) -> bool:
|
||||
return "refiner" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_txt_in(n: str, m) -> bool:
|
||||
return n.split(".")[-1] == "txt_in"
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideo15ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block])
|
||||
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^context_embedder\.proj_in\.(.*)$":
|
||||
r"txt_in.input_embedder.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
|
||||
|
||||
# 2. txt_in_2 mapping:
|
||||
r"^context_embedder_2\.(.*)$":
|
||||
r"txt_in_2.\1",
|
||||
|
||||
# 3. x_embedder mapping:
|
||||
r"^x_embedder\.proj\.(.*)$":
|
||||
r"img_in.proj.\1",
|
||||
|
||||
# 4. Top-level time_text_embed mappings:
|
||||
r"^time_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"time_in.timestep_embedder.mlp.fc_in.\1",
|
||||
r"^time_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"time_in.timestep_embedder.mlp.fc_out.\1",
|
||||
r"^time_embed\.timestep_embedder_r\.linear_1\.(.*)$":
|
||||
r"time_in.timestep_embedder_r.mlp.fc_in.\1",
|
||||
r"^time_embed\.timestep_embedder_r\.linear_2\.(.*)$":
|
||||
r"time_in.timestep_embedder_r.mlp.fc_out.\1",
|
||||
|
||||
# 5. transformer_blocks mapping:
|
||||
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
|
||||
r"double_blocks.\1.img_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
|
||||
r"double_blocks.\1.txt_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 7. Final layers mapping:
|
||||
r"^norm_out\.linear\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^proj_out\.(.*)$":
|
||||
r"final_layer.linear.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: custom -> hf
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
in_channels: int = 65
|
||||
out_channels: int = 32
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128
|
||||
num_layers: int = 54
|
||||
num_refiner_layers: int = 2
|
||||
mlp_ratio: float = 4.0
|
||||
patch_size: int = 1
|
||||
patch_size_t: int = 1
|
||||
qk_norm: str = "rms_norm"
|
||||
text_embed_dim: int = 3584
|
||||
text_embed_2_dim: int = 1472
|
||||
image_embed_dim: int = 1152
|
||||
rope_theta: float = 256.0
|
||||
rope_axes_dim: tuple[int, ...] = (16, 56, 56)
|
||||
target_size: int = 640
|
||||
task_type: str = "i2v"
|
||||
use_meanflow: bool = False
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
|
||||
self.num_channels_latents: int = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideo15Config(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=HunyuanVideo15ArchConfig)
|
||||
|
||||
prefix: str = "Hunyuan15"
|
||||
@@ -1,149 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Video DiT configuration for native FastVideo implementation.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def is_longcat_blocks(n: str, m) -> bool:
|
||||
"""FSDP shard condition for LongCat transformer blocks."""
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatVideoArchConfig(DiTArchConfig):
|
||||
"""Architecture configuration for native LongCat Video DiT."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Enable torch.compile for transformer blocks (major speedup!)
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Parameter name mapping for weight conversion
|
||||
# Maps original LongCat third_party names -> native FastVideo names
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Embedders
|
||||
r"^x_embedder\.(.*)$": r"patch_embed.\1",
|
||||
r"^t_embedder\.mlp\.0\.(.*)$": r"time_embedder.linear_1.\1",
|
||||
r"^t_embedder\.mlp\.2\.(.*)$": r"time_embedder.linear_2.\1",
|
||||
r"^y_embedder\.y_proj\.0\.(.*)$": r"caption_embedder.linear_1.\1",
|
||||
r"^y_embedder\.y_proj\.2\.(.*)$": r"caption_embedder.linear_2.\1",
|
||||
|
||||
# Transformer blocks - AdaLN modulation
|
||||
r"^blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
|
||||
r"blocks.\1.adaln_linear_1.\2",
|
||||
|
||||
# Transformer blocks - Normalization
|
||||
r"^blocks\.(\d+)\.mod_norm_attn\.(.*)$": r"blocks.\1.norm_attn.\2",
|
||||
r"^blocks\.(\d+)\.mod_norm_ffn\.(.*)$": r"blocks.\1.norm_ffn.\2",
|
||||
r"^blocks\.(\d+)\.pre_crs_attn_norm\.(.*)$":
|
||||
r"blocks.\1.norm_cross.\2",
|
||||
|
||||
# Self-attention: QKV fused -> separate (will need splitting in converter)
|
||||
# Original has attn.qkv.weight -> need to split into to_q, to_k, to_v
|
||||
r"^blocks\.(\d+)\.attn\.qkv\.(.*)$":
|
||||
r"blocks.\1.self_attn.qkv_fused.\2", # Marker for splitting
|
||||
r"^blocks\.(\d+)\.attn\.proj\.(.*)$":
|
||||
r"blocks.\1.self_attn.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn\.q_norm\.(.*)$":
|
||||
r"blocks.\1.self_attn.q_norm.\2",
|
||||
r"^blocks\.(\d+)\.attn\.k_norm\.(.*)$":
|
||||
r"blocks.\1.self_attn.k_norm.\2",
|
||||
|
||||
# Cross-attention
|
||||
r"^blocks\.(\d+)\.cross_attn\.q_linear\.(.*)$":
|
||||
r"blocks.\1.cross_attn.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.kv_linear\.(.*)$":
|
||||
r"blocks.\1.cross_attn.kv_fused.\2", # Marker for splitting
|
||||
r"^blocks\.(\d+)\.cross_attn\.proj\.(.*)$":
|
||||
r"blocks.\1.cross_attn.to_out.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q_norm\.(.*)$":
|
||||
r"blocks.\1.cross_attn.q_norm.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k_norm\.(.*)$":
|
||||
r"blocks.\1.cross_attn.k_norm.\2",
|
||||
|
||||
# FFN (SwiGLU)
|
||||
r"^blocks\.(\d+)\.ffn\.w1\.(.*)$": r"blocks.\1.ffn.w1.\2", # gate
|
||||
r"^blocks\.(\d+)\.ffn\.w2\.(.*)$": r"blocks.\1.ffn.w2.\2", # down
|
||||
r"^blocks\.(\d+)\.ffn\.w3\.(.*)$": r"blocks.\1.ffn.w3.\2", # up
|
||||
|
||||
# Final layer
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
|
||||
r"final_layer.adaln_linear.\1",
|
||||
r"^final_layer\.norm_final\.(.*)$": r"final_layer.norm.\1",
|
||||
r"^final_layer\.linear\.(.*)$": r"final_layer.proj.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: custom -> hf
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# LoRA parameter name mapping
|
||||
lora_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Model architecture parameters
|
||||
hidden_size: int = 4096
|
||||
depth: int = 48 # Number of transformer blocks
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128 # hidden_size / num_attention_heads
|
||||
|
||||
in_channels: int = 16 # Latent space channels
|
||||
out_channels: int = 16
|
||||
num_channels_latents: int = 16
|
||||
|
||||
# Patch embedding
|
||||
patch_size: tuple[int, int,
|
||||
int] = (1, 2, 2) # [T, H, W] - no temporal compression
|
||||
|
||||
# Text/caption embedding
|
||||
caption_channels: int = 4096 # UMT5 d_model
|
||||
|
||||
# Timestep embedding
|
||||
adaln_tembed_dim: int = 512
|
||||
frequency_embedding_size: int = 256
|
||||
|
||||
# FFN
|
||||
mlp_ratio: int = 4
|
||||
|
||||
# Attention backend support
|
||||
_supported_attention_backends: tuple = field(default_factory=lambda: (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
))
|
||||
|
||||
# Text padding behavior
|
||||
text_tokens_zero_pad: bool = True
|
||||
|
||||
# Block Sparse Attention (BSA)
|
||||
enable_bsa: bool = False
|
||||
bsa_params: dict | None = field(
|
||||
default_factory=lambda: {
|
||||
"sparsity": 0.9375,
|
||||
"cdf_threshold": None,
|
||||
"chunk_3d_shape_q": [4, 4, 4],
|
||||
"chunk_3d_shape_k": [4, 4, 4],
|
||||
})
|
||||
|
||||
# LoRA exclusions
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
# Ensure attention_head_dim matches
|
||||
self.attention_head_dim = self.hidden_size // self.num_attention_heads
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatVideoConfig(DiTConfig):
|
||||
"""Main configuration for LongCat Video DiT."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=LongCatVideoArchConfig)
|
||||
|
||||
prefix: str = "longcat"
|
||||
@@ -1,83 +0,0 @@
|
||||
from dataclasses import dataclass, field
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
|
||||
# Override param_names_mapping to remove patch_embedding transformation
|
||||
# because MatrixGame checkpoints already have patch_embedding.proj format
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Removed: r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1"
|
||||
# because checkpoint already has correct format
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$":
|
||||
r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_out.\1",
|
||||
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
|
||||
r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
|
||||
r"blocks.\1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.attn2.to_out.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_out.\2",
|
||||
r"^blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
|
||||
action_config: dict = field(
|
||||
default_factory=lambda: {
|
||||
"blocks": list(range(15)),
|
||||
"enable_mouse": True,
|
||||
"enable_keyboard": True,
|
||||
"heads_num": 16,
|
||||
"hidden_size": 128,
|
||||
"img_hidden_size": 1536,
|
||||
"keyboard_dim_in": 4,
|
||||
"keyboard_hidden_dim": 1024,
|
||||
"mouse_dim_in": 2,
|
||||
"mouse_hidden_dim": 1024,
|
||||
"mouse_qk_dim_list": [8, 28, 28],
|
||||
"patch_size": [1, 2, 2],
|
||||
"qk_norm": True,
|
||||
"qkv_bias": False,
|
||||
"rope_dim_list": [8, 28, 28],
|
||||
"rope_theta": 256,
|
||||
"vae_time_compression_ratio": 4,
|
||||
"windows_size": 3,
|
||||
})
|
||||
|
||||
local_attn_size: int = -1
|
||||
sink_size: int = 0
|
||||
num_frames_per_block: int = 3
|
||||
text_len: int = 512
|
||||
text_dim: int = 0
|
||||
image_dim: int = 1280
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatrixGameWanVideoConfig(WanVideoConfig):
|
||||
arch_config: MatrixGameWanVideoArchConfig = field(
|
||||
default_factory=MatrixGameWanVideoArchConfig)
|
||||
prefix: str = "Wan"
|
||||
@@ -6,11 +6,9 @@ from fastvideo.configs.models.encoders.clip import (
|
||||
CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
|
||||
"Qwen2_5_VLConfig"
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig"
|
||||
]
|
||||
|
||||
@@ -72,7 +72,6 @@ class EncoderConfig(ModelConfig):
|
||||
@dataclass
|
||||
class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||
is_chat_model: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,93 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("embed_tokens")
|
||||
|
||||
|
||||
def _is_final_norm(n: str, m) -> bool:
|
||||
return n.endswith("norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Qwen2_5_VLArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 152064
|
||||
hidden_size: int = 8192
|
||||
intermediate_size: int = 29568
|
||||
num_hidden_layers: int = 80
|
||||
num_attention_heads: int = 64
|
||||
num_key_value_heads: int = 8
|
||||
hidden_act: str = "silu"
|
||||
max_position_embeddings: int = 32768
|
||||
initializer_range: float = 0.02
|
||||
rms_norm_eps: float = 1e-05
|
||||
use_cache: bool = True
|
||||
tie_word_embeddings: bool = False
|
||||
rope_theta: float = 1000000.0
|
||||
use_sliding_window: bool = False
|
||||
sliding_window: int | None = 4096
|
||||
max_window_layers: int = 80
|
||||
layer_types: list = field(default_factory=list)
|
||||
attention_dropout: float = 0.0
|
||||
rope_scaling: dict | None = None
|
||||
bos_token_id: int | None = None
|
||||
eos_token_id: int | None = None
|
||||
pad_token_id: int | None = None
|
||||
vision_token_id: int = 151654
|
||||
model_type: str = "qwen2_5_vl_text"
|
||||
dtype: str = "bfloat16"
|
||||
|
||||
stacked_params_mapping: list[tuple[str, str, str
|
||||
| int]] = field(default_factory=lambda: [
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0),
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_norm])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.sliding_window = self.sliding_window if self.use_sliding_window else None
|
||||
# for backward compatibility
|
||||
if self.num_key_value_heads is None:
|
||||
self.num_key_value_heads = self.num_attention_heads
|
||||
if self.layer_types is None:
|
||||
self.layer_types = [
|
||||
"sliding_attention" if self.sliding_window is not None
|
||||
and i >= self.max_window_layers else "full_attention"
|
||||
for i in range(self.num_hidden_layers)
|
||||
]
|
||||
if self.rope_scaling is not None and "type" in self.rope_scaling:
|
||||
if self.rope_scaling["type"] == "mrope":
|
||||
self.rope_scaling["type"] = "default"
|
||||
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
|
||||
|
||||
self.tokenizer_kwargs = {
|
||||
"add_generation_prompt": True,
|
||||
"tokenize": True,
|
||||
"return_dict": True,
|
||||
"padding": "max_length",
|
||||
"max_length": 1000 + 108,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Qwen2_5_VLConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=Qwen2_5_VLArchConfig)
|
||||
prefix: str = "qwen2_5_vl"
|
||||
is_chat_model: bool = True
|
||||
@@ -40,8 +40,6 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
eos_token_id: int = 1
|
||||
classifier_dropout: float = 0.0
|
||||
text_len: int = 512
|
||||
dtype: str | None = None
|
||||
gradient_checkpointing: bool = False
|
||||
stacked_params_mapping: list[tuple[str, str,
|
||||
str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
@@ -70,7 +68,6 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
self.hidden_size = self.d_model
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
@@ -9,5 +8,4 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"StepVideoVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
]
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15VAEArchConfig(VAEArchConfig):
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 32
|
||||
block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024)
|
||||
layers_per_block: int = 2
|
||||
spatial_compression_ratio: int = 16
|
||||
temporal_compression_ratio: int = 4
|
||||
downsample_match_channel: bool = True
|
||||
upsample_match_channel: bool = True
|
||||
scaling_factor: float = 1.03682
|
||||
|
||||
def __post_init__(self):
|
||||
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
|
||||
1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15VAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=Hunyuan15VAEArchConfig)
|
||||
@@ -2,7 +2,6 @@ from fastvideo.configs.pipelines.base import (PipelineConfig,
|
||||
SlidingTileAttnConfig)
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
@@ -12,8 +11,8 @@ from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "get_pipeline_config_cls_from_name"
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"SelfForcingWanT2V480PConfig", "CosmosConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -1,139 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
import re
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
|
||||
Qwen2_5_VLConfig, T5Config)
|
||||
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
PROMPT_TEMPLATE_TOKEN_LENGTH = 108
|
||||
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = "You are a helpful assistant. Describe the video by detailing the following aspects: \
|
||||
1. The main content and theme of the video. \
|
||||
2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects. \
|
||||
3. Actions, events, behaviors temporal relationships, physical movement changes of the objects. \
|
||||
4. background environment, light, style and atmosphere. \
|
||||
5. camera angles, movements, and transitions used in the video."
|
||||
|
||||
|
||||
def extract_glyph_texts(prompt: str) -> str | None:
|
||||
"""
|
||||
Extract glyph texts from prompt using regex pattern.
|
||||
|
||||
Args:
|
||||
prompt: Input prompt string
|
||||
|
||||
Returns:
|
||||
List of extracted glyph texts
|
||||
"""
|
||||
pattern = r"\"(.*?)\"|“(.*?)”"
|
||||
matches = re.findall(pattern, prompt)
|
||||
result = [match[0] or match[1] for match in matches]
|
||||
result = list(dict.fromkeys(result)) if len(result) > 1 else result
|
||||
|
||||
if result:
|
||||
formatted_result = ". ".join([f'Text "{text}"'
|
||||
for text in result]) + ". "
|
||||
else:
|
||||
formatted_result = None
|
||||
|
||||
return formatted_result
|
||||
|
||||
|
||||
def format_text_input(prompt: str, system_message: str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Apply text to template.
|
||||
|
||||
Args:
|
||||
prompt (List[str]): Input text.
|
||||
system_message (str): System message.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: List of chat conversation.
|
||||
"""
|
||||
|
||||
template = [{
|
||||
"role": "system",
|
||||
"content": system_message
|
||||
}, {
|
||||
"role": "user",
|
||||
"content": prompt if prompt else " "
|
||||
}]
|
||||
|
||||
return template
|
||||
|
||||
|
||||
def qwen_preprocess_text(prompt: str) -> list[dict[str, Any]]:
|
||||
output = format_text_input(prompt, PROMPT_TEMPLATE_ENCODE_VIDEO)
|
||||
return output
|
||||
|
||||
|
||||
def qwen_postprocess_text(
|
||||
outputs: BaseEncoderOutput,
|
||||
mask: torch.tensor) -> tuple[torch.tensor, torch.tensor]:
|
||||
assert outputs.hidden_states is not None
|
||||
output = outputs.hidden_states[-3]
|
||||
output = output[:, PROMPT_TEMPLATE_TOKEN_LENGTH:]
|
||||
mask = mask[:, PROMPT_TEMPLATE_TOKEN_LENGTH:]
|
||||
return output, mask
|
||||
|
||||
|
||||
def byt5_preprocess_text(prompt: str) -> str | None:
|
||||
prompts = [prompt] if isinstance(prompt, str) else prompt
|
||||
glyph_texts = [extract_glyph_texts(p) for p in prompts]
|
||||
return glyph_texts[0]
|
||||
|
||||
|
||||
def byt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15T2V480PConfig(PipelineConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=HunyuanVideo15Config)
|
||||
# VAE
|
||||
vae_config: VAEConfig = field(default_factory=Hunyuan15VAEConfig)
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (Qwen2_5_VLConfig(), T5Config()))
|
||||
preprocess_text_funcs: tuple[Callable[[Any], Any], ...] = field(
|
||||
default_factory=lambda: (qwen_preprocess_text, byt5_preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
|
||||
default_factory=lambda: (qwen_postprocess_text, byt5_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", "fp32"))
|
||||
text_encoder_crop_start: int = PROMPT_TEMPLATE_TOKEN_LENGTH
|
||||
text_encoder_max_lengths: tuple[int, ...] = field(
|
||||
default_factory=lambda: (1000 + PROMPT_TEMPLATE_TOKEN_LENGTH, 256))
|
||||
|
||||
vae_tiling: bool = True
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
flow_shift: int = 9
|
||||
@@ -1,355 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
import html
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields.
|
||||
|
||||
NOTE: This is for Phase 1 wrapper compatibility. For native model (Phase 2),
|
||||
use LongCatVideoConfig from fastvideo.configs.models.dits.longcat instead.
|
||||
"""
|
||||
# LongCat-specific architecture parameters
|
||||
adaln_tembed_dim: int = 512
|
||||
caption_channels: int = 4096
|
||||
depth: int = 48
|
||||
enable_bsa: bool = False
|
||||
enable_flashattn3: bool = False
|
||||
enable_flashattn2: bool = True
|
||||
enable_xformers: bool = False
|
||||
frequency_embedding_size: int = 256
|
||||
in_channels: int = 16
|
||||
mlp_ratio: int = 4
|
||||
num_heads: int = 32
|
||||
out_channels: int = 16
|
||||
text_tokens_zero_pad: bool = True
|
||||
patch_size: list[int] = field(default_factory=lambda: [1, 2, 2])
|
||||
cp_split_hw: list[int] | None = None
|
||||
bsa_params: dict | None = None
|
||||
|
||||
|
||||
def longcat_preprocess_text(prompt: str) -> str:
|
||||
"""Clean and preprocess text like original LongCat implementation.
|
||||
|
||||
This function applies the same text cleaning pipeline as the original
|
||||
LongCat-Video implementation to ensure identical tokenization results.
|
||||
|
||||
Steps:
|
||||
1. basic_clean: Fix unicode issues and unescape HTML entities
|
||||
2. whitespace_clean: Normalize whitespace to single spaces
|
||||
|
||||
Args:
|
||||
prompt: Raw input text prompt
|
||||
|
||||
Returns:
|
||||
Cleaned and normalized text prompt
|
||||
"""
|
||||
# basic_clean: fix unicode and HTML entities
|
||||
text = ftfy.fix_text(prompt)
|
||||
text = html.unescape(html.unescape(text))
|
||||
text = text.strip()
|
||||
|
||||
# whitespace_clean: normalize whitespace
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
text = text.strip()
|
||||
|
||||
return text
|
||||
|
||||
|
||||
def umt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""
|
||||
Postprocess UMT5/T5 encoder outputs to fixed length 512 embeddings.
|
||||
"""
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
assert torch.isnan(hidden_state).sum() == 0
|
||||
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT2V480PConfig(PipelineConfig):
|
||||
"""Configuration for LongCat pipeline (480p) aligned to LongCat-Video modules.
|
||||
|
||||
Components expected by loaders:
|
||||
- tokenizer: AutoTokenizer
|
||||
- text_encoder: UMT5EncoderModel
|
||||
- transformer: LongCatVideoTransformer3DModel (Phase 1 wrapper)
|
||||
OR LongCatTransformer3DModel (Phase 2 native)
|
||||
- vae: AutoencoderKLWan (Wan VAE, 4x8 compression)
|
||||
- scheduler: FlowMatchEulerDiscreteScheduler
|
||||
"""
|
||||
|
||||
# DiT config with LongCat-specific arch_config
|
||||
# NOTE: For Phase 1 wrapper, uses LongCatDiTArchConfig
|
||||
# For Phase 2 native model, can use LongCatVideoConfig directly
|
||||
dit_config: DiTConfig = field(
|
||||
default_factory=lambda: DiTConfig(arch_config=LongCatDiTArchConfig()))
|
||||
|
||||
# VAE config: Wan VAE with encoder+decoder enabled
|
||||
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Precision defaults
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", ))
|
||||
|
||||
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(
|
||||
default_factory=lambda: (T5Config(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (longcat_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(umt5_postprocess_text, ))
|
||||
|
||||
# LongCat-specific runtime toggles (consumed by pipeline/stages)
|
||||
enable_kv_cache: bool = True
|
||||
offload_kv_cache: bool = False
|
||||
enable_bsa: bool = False
|
||||
use_distill: bool = False
|
||||
enhance_hf: bool = False
|
||||
# Optional BSA parameter dict (kept for backward/phase-1 compatibility).
|
||||
# `LongCatPipeline.initialize_pipeline()` uses this as a base and then applies
|
||||
# CLI overrides (bsa_sparsity / bsa_chunk_{q,k} / bsa_cdf_threshold).
|
||||
bsa_params: dict | None = None
|
||||
# BSA runtime overrides (preferred over bsa_params if provided via CLI)
|
||||
bsa_sparsity: float | None = None
|
||||
bsa_cdf_threshold: float | None = None
|
||||
bsa_chunk_q: list[int] | None = None
|
||||
bsa_chunk_k: list[int] | None = None
|
||||
t_thresh: float | None = None # refine stage default controlled by sampling args
|
||||
|
||||
# LongCat does not need flow_shift
|
||||
flow_shift: float | None = None
|
||||
dmd_denoising_steps: list[int] | None = None
|
||||
|
||||
def __post_init__(self):
|
||||
# LongCat inference requires vae encoder and decoder
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT2V704PConfig(LongCatT2V480PConfig):
|
||||
"""Configuration for LongCat pipeline (704p) with BSA enabled by default.
|
||||
|
||||
Uses the same resolution and BSA parameters as original LongCat refinement stage.
|
||||
BSA parameters configured in transformer config.json with chunk_3d_shape=[4,4,4]:
|
||||
- Input: 704×1280×96
|
||||
- VAE (8x): 88×160×96
|
||||
- Patch [1,2,2]: 44×80×96
|
||||
- chunk [4,4,4]: 96%4=0, 44%4=0, 80%4=0 ✅
|
||||
|
||||
This configuration matches the original LongCat refinement stage parameters.
|
||||
"""
|
||||
|
||||
# Enable BSA by default for 704p
|
||||
enable_bsa: bool = True
|
||||
|
||||
|
||||
ASPECT_RATIO_627 = {
|
||||
'0.26': ([320, 1216], 1),
|
||||
'0.31': ([352, 1120], 1),
|
||||
'0.38': ([384, 1024], 1),
|
||||
'0.43': ([416, 960], 1),
|
||||
'0.52': ([448, 864], 1),
|
||||
'0.58': ([480, 832], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'0.74': ([544, 736], 1),
|
||||
'0.86': ([576, 672], 1),
|
||||
'0.95': ([608, 640], 1),
|
||||
'1.05': ([640, 608], 1),
|
||||
'1.17': ([672, 576], 1),
|
||||
'1.29': ([704, 544], 1),
|
||||
'1.35': ([736, 544], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'1.67': ([800, 480], 1),
|
||||
'1.73': ([832, 480], 1),
|
||||
'2.00': ([896, 448], 1),
|
||||
'2.31': ([960, 416], 1),
|
||||
'2.58': ([992, 384], 1),
|
||||
'2.75': ([1056, 384], 1),
|
||||
'3.09': ([1088, 352], 1),
|
||||
'3.70': ([1184, 320], 1),
|
||||
'3.80': ([1216, 320], 1),
|
||||
'3.90': ([1248, 320], 1),
|
||||
'4.00': ([1280, 320], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_627_F64 = {
|
||||
'0.26': ([320, 1216], 1),
|
||||
'0.38': ([384, 1024], 1),
|
||||
'0.50': ([448, 896], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'0.82': ([576, 704], 1),
|
||||
'1.00': ([640, 640], 1),
|
||||
'1.22': ([704, 576], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'1.86': ([832, 448], 1),
|
||||
'2.00': ([896, 448], 1),
|
||||
'2.50': ([960, 384], 1),
|
||||
'2.83': ([1088, 384], 1),
|
||||
'3.60': ([1152, 320], 1),
|
||||
'3.80': ([1216, 320], 1),
|
||||
'4.00': ([1280, 320], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_627_F128 = {
|
||||
'0.25': ([256, 1024], 1),
|
||||
'0.38': ([384, 1024], 1),
|
||||
'0.43': ([384, 896], 1),
|
||||
'0.57': ([512, 896], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'1.00': ([640, 640], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'1.75': ([896, 512], 1),
|
||||
'2.33': ([896, 384], 1),
|
||||
'2.67': ([1024, 384], 1),
|
||||
'4.00': ([1024, 256], 1),
|
||||
}
|
||||
|
||||
ASPECT_RATIO_627_F256 = {
|
||||
'0.25': ([256, 1024], 1),
|
||||
'0.33': ([256, 768], 1),
|
||||
'0.50': ([256, 512], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'1.00': ([512, 512], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'2.00': ([512, 256], 1),
|
||||
'3.00': ([768, 256], 1),
|
||||
'4.00': ([1024, 256], 1),
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960 = {
|
||||
'0.25': ([480, 1920], 1),
|
||||
'0.29': ([512, 1792], 1),
|
||||
'0.32': ([544, 1696], 1),
|
||||
'0.36': ([576, 1600], 1),
|
||||
'0.40': ([608, 1504], 1),
|
||||
'0.49': ([672, 1376], 1),
|
||||
'0.54': ([704, 1312], 1),
|
||||
'0.59': ([736, 1248], 1),
|
||||
'0.69': ([800, 1152], 1),
|
||||
'0.74': ([832, 1120], 1),
|
||||
'0.82': ([864, 1056], 1),
|
||||
'0.88': ([896, 1024], 1),
|
||||
'0.94': ([928, 992], 1),
|
||||
'1.00': ([960, 960], 1),
|
||||
'1.07': ([992, 928], 1),
|
||||
'1.14': ([1024, 896], 1),
|
||||
'1.22': ([1056, 864], 1),
|
||||
'1.31': ([1088, 832], 1),
|
||||
'1.35': ([1120, 832], 1),
|
||||
'1.44': ([1152, 800], 1),
|
||||
'1.70': ([1248, 736], 1),
|
||||
'2.00': ([1344, 672], 1),
|
||||
'2.05': ([1376, 672], 1),
|
||||
'2.47': ([1504, 608], 1),
|
||||
'2.53': ([1536, 608], 1),
|
||||
'2.83': ([1632, 576], 1),
|
||||
'3.06': ([1664, 544], 1),
|
||||
'3.12': ([1696, 544], 1),
|
||||
'3.62': ([1856, 512], 1),
|
||||
'3.93': ([1888, 480], 1),
|
||||
'4.00': ([1920, 480], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960_F64 = {
|
||||
'0.22': ([448, 2048], 1),
|
||||
'0.29': ([512, 1792], 1),
|
||||
'0.36': ([576, 1600], 1),
|
||||
'0.45': ([640, 1408], 1),
|
||||
'0.55': ([704, 1280], 1),
|
||||
'0.63': ([768, 1216], 1),
|
||||
'0.76': ([832, 1088], 1),
|
||||
'0.88': ([896, 1024], 1),
|
||||
'1.00': ([960, 960], 1),
|
||||
'1.14': ([1024, 896], 1),
|
||||
'1.31': ([1088, 832], 1),
|
||||
'1.50': ([1152, 768], 1),
|
||||
'1.58': ([1216, 768], 1),
|
||||
'1.82': ([1280, 704], 1),
|
||||
'1.91': ([1344, 704], 1),
|
||||
'2.20': ([1408, 640], 1),
|
||||
'2.30': ([1472, 640], 1),
|
||||
'2.67': ([1536, 576], 1),
|
||||
'2.89': ([1664, 576], 1),
|
||||
'3.62': ([1856, 512], 1),
|
||||
'3.75': ([1920, 512], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960_F128 = {
|
||||
'0.20': ([384, 1920], 1),
|
||||
'0.27': ([512, 1920], 1),
|
||||
'0.33': ([512, 1536], 1),
|
||||
'0.42': ([640, 1536], 1),
|
||||
'0.50': ([640, 1280], 1),
|
||||
'0.60': ([768, 1280], 1),
|
||||
'0.67': ([768, 1152], 1),
|
||||
'0.78': ([896, 1152], 1),
|
||||
'1.00': ([1024, 1024], 1),
|
||||
'1.29': ([1152, 896], 1),
|
||||
'1.50': ([1152, 768], 1),
|
||||
'1.67': ([1280, 768], 1),
|
||||
'2.00': ([1280, 640], 1),
|
||||
'2.40': ([1536, 640], 1),
|
||||
'3.00': ([1536, 512], 1),
|
||||
'3.75': ([1920, 512], 1),
|
||||
'5.00': ([1920, 384], 1),
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960_F256 = {
|
||||
'0.33': ([512, 1536], 1),
|
||||
'0.60': ([768, 1280], 1),
|
||||
'1.00': ([1024, 1024], 1),
|
||||
'1.67': ([1280, 768], 1),
|
||||
'3.00': ([1536, 512], 1),
|
||||
}
|
||||
|
||||
|
||||
def get_bucket_config(resolution, scale_factor_spatial):
|
||||
if resolution == '480p':
|
||||
if scale_factor_spatial == 16 or scale_factor_spatial == 32:
|
||||
return ASPECT_RATIO_627
|
||||
elif scale_factor_spatial == 64:
|
||||
return ASPECT_RATIO_627_F64
|
||||
elif scale_factor_spatial == 128:
|
||||
return ASPECT_RATIO_627_F128
|
||||
elif scale_factor_spatial == 256:
|
||||
return ASPECT_RATIO_627_F256
|
||||
elif resolution == '720p':
|
||||
if scale_factor_spatial == 16 or scale_factor_spatial == 32:
|
||||
return ASPECT_RATIO_960
|
||||
elif scale_factor_spatial == 64:
|
||||
return ASPECT_RATIO_960_F64
|
||||
elif scale_factor_spatial == 128:
|
||||
return ASPECT_RATIO_960_F128
|
||||
elif scale_factor_spatial == 256:
|
||||
return ASPECT_RATIO_960_F256
|
||||
|
||||
raise ValueError(
|
||||
f"Unsupported resolution '{resolution}' or scale_factor_spatial '{scale_factor_spatial}'"
|
||||
)
|
||||
@@ -7,17 +7,14 @@ from collections.abc import Callable
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig,
|
||||
MatrixGameI2V480PConfig)
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -29,10 +26,6 @@ logger = init_logger(__name__)
|
||||
PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
|
||||
Hunyuan15T2V480PConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
|
||||
Hunyuan15T2V720PConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
|
||||
@@ -52,45 +45,25 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World": CosmosConfig,
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGameI2V480PConfig,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGameI2V480PConfig,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGameI2V480PConfig,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan":
|
||||
lambda id: "hunyuan" in id.lower(),
|
||||
"hunyuan15":
|
||||
lambda id: "hunyuan15" in id.lower(),
|
||||
"matrixgame":
|
||||
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
|
||||
"wanpipeline":
|
||||
lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo":
|
||||
lambda id: "wanimagetovideo" in id.lower(),
|
||||
"wandmdpipeline":
|
||||
lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline":
|
||||
lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"stepvideo":
|
||||
lambda id: "stepvideo" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"cosmos": lambda id: "cosmos" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"longcat": LongCatT2V480PConfig,
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"matrixgame": MatrixGameI2V480PConfig,
|
||||
"hunyuan15":
|
||||
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
|
||||
"wanpipeline":
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
@@ -136,7 +109,6 @@ def get_pipeline_config_cls_from_name(
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
|
||||
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
|
||||
return pipeline_config_cls
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
|
||||
|
||||
@@ -6,7 +6,6 @@ import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPVisionConfig, T5Config,
|
||||
WAN2_1ControlCLIPVisionConfig)
|
||||
@@ -191,23 +190,3 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Matrix Game ===================
|
||||
# =============================================
|
||||
@dataclass
|
||||
class MatrixGameI2V480PConfig(WanI2V480PConfig):
|
||||
dit_config: DiTConfig = field(default_factory=MatrixGameWanVideoConfig)
|
||||
|
||||
image_encoder_config: EncoderConfig = field(
|
||||
default_factory=WAN2_1ControlCLIPVisionConfig)
|
||||
|
||||
is_causal: bool = True
|
||||
flow_shift: float | None = 5.0
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 666, 333])
|
||||
warp_denoising_step: bool = True
|
||||
context_noise: int = 0
|
||||
num_frames_per_block: int = 3
|
||||
# sliding_window_num_frames: int = 15
|
||||
|
||||
@@ -3,7 +3,6 @@ from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -18,27 +17,10 @@ class SamplingParam:
|
||||
|
||||
# Image inputs
|
||||
image_path: str | None = None
|
||||
pil_image: Any | None = None
|
||||
|
||||
# Video inputs
|
||||
video_path: str | None = None
|
||||
|
||||
# Action control inputs (Matrix-Game)
|
||||
mouse_cond: Any | None = None # Shape: (B, T, 2)
|
||||
keyboard_cond: Any | None = None # Shape: (B, T, K)
|
||||
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
|
||||
|
||||
# Refine inputs (LongCat 480p->720p upscaling)
|
||||
# Path-based refine (load stage1 video from disk, e.g. MP4)
|
||||
refine_from: str | None = None # Path to stage1 video (480p output from distill)
|
||||
t_thresh: float = 0.5 # Threshold for timestep scheduling in refinement
|
||||
spatial_refine_only: bool = False # If True, only spatial (no temporal doubling)
|
||||
num_cond_frames: int = 0 # Number of conditioning frames
|
||||
# In-memory refine input (for two-stage pipeline where stage1 frames are already in memory)
|
||||
# This mirrors LongCat's demo where a list of frames (e.g. np.ndarray or PIL.Image)
|
||||
# is passed directly to the refinement pipeline instead of reloading from disk.
|
||||
stage1_video: Any | None = None
|
||||
|
||||
# Text inputs
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
@@ -62,7 +44,6 @@ class SamplingParam:
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -228,31 +209,6 @@ class SamplingParam:
|
||||
default=SamplingParam.video_path,
|
||||
help="Path to input video for video-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--refine-from",
|
||||
type=str,
|
||||
default=SamplingParam.refine_from,
|
||||
help="Path to stage1 video for refinement (LongCat 480p->720p)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t-thresh",
|
||||
type=float,
|
||||
default=SamplingParam.t_thresh,
|
||||
help=
|
||||
"Threshold for timestep scheduling in refinement (default: 0.5)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial-refine-only",
|
||||
action=StoreBoolean,
|
||||
default=SamplingParam.spatial_refine_only,
|
||||
help="Only perform spatial super-resolution (no temporal doubling)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-cond-frames",
|
||||
type=int,
|
||||
default=SamplingParam.num_cond_frames,
|
||||
help="Number of conditioning frames for refinement",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--moba-config-path",
|
||||
type=str,
|
||||
|
||||
@@ -1,29 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import numpy as np
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_480P_SamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 121
|
||||
height: int = 480
|
||||
width: int = 848
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 6.0
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
sigmas: list[float] | None = field(
|
||||
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
|
||||
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
@@ -5,7 +5,6 @@ from typing import Any
|
||||
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
|
||||
@@ -24,7 +23,6 @@ from fastvideo.configs.sample.wan import (
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
MatrixGame2_SamplingParam,
|
||||
)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -34,36 +32,44 @@ from fastvideo.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
|
||||
Hunyuan15_480P_SamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
|
||||
Hunyuan15_720P_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
"FastVideo/FastHunyuan-diffusers":
|
||||
FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo":
|
||||
HunyuanSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers":
|
||||
StepVideoT2VSamplingParam,
|
||||
|
||||
# Wan2.1
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
|
||||
WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
|
||||
# Wan2.2
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
|
||||
# FastWan2.1
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480P_SamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
FastWanT2V480P_SamplingParam,
|
||||
|
||||
# FastWan2.2
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
|
||||
@@ -79,32 +85,17 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World":
|
||||
Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
|
||||
# MatrixGame2.0 models
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
|
||||
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan":
|
||||
lambda id: "hunyuan" in id.lower(),
|
||||
"hunyuan15":
|
||||
lambda id: "hunyuan15" in id.lower(),
|
||||
"wanpipeline":
|
||||
lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo":
|
||||
lambda id: "wanimagetovideo" in id.lower(),
|
||||
"stepvideo":
|
||||
lambda id: "stepvideo" in id.lower(),
|
||||
"wandmdpipeline":
|
||||
lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline":
|
||||
lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
"matrixgame":
|
||||
lambda id: "matrixgame" in id.lower() or "matrix-game" in id.lower(),
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -112,15 +103,12 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
"hunyuan":
|
||||
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"hunyuan15":
|
||||
Hunyuan15_480P_SamplingParam, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
|
||||
"wanpipeline":
|
||||
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
|
||||
"wandmdpipeline": FastWanT2V480P_SamplingParam,
|
||||
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam,
|
||||
"matrixgame": MatrixGame2_SamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
@@ -128,20 +116,6 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
|
||||
"""Get the appropriate sampling param for specific pretrained weights."""
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
|
||||
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
matrixgame_patterns = ["Matrix-Game", "Skywork--Matrix-Game", "matrixgame"]
|
||||
for pattern in matrixgame_patterns:
|
||||
if pattern.lower() in pipeline_name_or_path.lower():
|
||||
return MatrixGame2_SamplingParam
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
@@ -152,6 +126,15 @@ def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
|
||||
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
# Try to determine pipeline architecture for fallback
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user