Compare commits

...
Author SHA1 Message Date
SolitaryThinker 104a539a22 update docker 2025-12-24 09:32:42 +00:00
Shreejith SGandWilliam Lin f8bfc76015 feat: consolidate attention kernels into unified fastvideo-kernel package (#946)
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-12-24 01:39:51 -06:00
alexzmsandShao Duan 8f1e6c3336 Add LongCat T2V (Base, Distillation and Refinement) Support to FastVideo (#883)
Co-authored-by: Shao Duan <shaoxiongduan@gmail.com>
2025-12-23 01:11:18 -06:00
William Lin 8e7d2e7879 [bugfix] [dmd2] allow dmd2 simulate_student_forward to use text-only dataset (#951) 2025-12-23 00:41:31 -06:00
William Lin 6ab2870942 [rocm] Add rocm fastvideo docker image (#952) 2025-12-22 18:21:04 -06:00
RoyWang e0ad145152 [feat] add sliding_tile attention triton kernel and ROCM support (#916) 2025-12-22 18:02:51 -06:00
Matthew Noto da04d08426 [docs] small fixes (#947) 2025-12-22 15:23:55 -06:00
Wei Zhou 1f70032af5 [New Model] Hunyuan1.5 (#943) 2025-12-21 00:57:52 -06:00
William Lin 7f71994653 [misc] Allow manual override of Pipeline class through override_pipeline_cls_name (#945) 2025-12-20 14:39:17 -06:00
Loay Rashid 2bb3349da1 [bugfix] Added VSA Padding logic (#944) 2025-12-20 14:29:11 -06:00
Kaiqin Kong 8fe1689968 [feat] Add Matrix-Game 2.0 (#938) 2025-12-20 14:09:12 -06:00
Loay Rashid e53730f324 [docs] Minor Fixes (#942) 2025-12-19 16:48:16 -06:00
Loay Rashid 7a4fe9086a [feat] Support sequence packing and shard after pachification for USP (#894) 2025-12-19 16:19:46 -06:00
Ohm-Rishabh d277361aae [misc] add schedule configurations to pytorch profiler (#934) 2025-12-18 01:45:23 -06:00
alexzms 734a54e7a9 [ci]: Use pre-built docker image & skip VSA compilation (#939) 2025-12-16 23:14:11 -08:00
alexzms 91364982df [Feature] Support for Variable Q/KV Sequence Lengths in VSA ThunderKittens kernel (#911) 2025-12-16 20:08:15 -08:00
William Lin 50145e4fcb [CI] Fix CI tests (#935) 2025-12-16 04:59:43 -08:00
William Lin 4112507e99 [misc] upgrade pytorch version to 2.9.0 (#928) 2025-12-15 04:12:43 -08:00
William Lin 424fc2b4ae [bugfix] [lora] [distillation] Fix lora distillation bug (#933) 2025-12-15 04:12:02 -08:00
William Lin e6066223e6 [bugfix] [VSA] [distillation] Various bugfixes for VSA and distillation and nightly tests (#932) 2025-12-12 16:51:54 -08:00
175 changed files with 19266 additions and 971 deletions
@@ -0,0 +1,236 @@
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/
+2
View File
@@ -30,6 +30,8 @@ env
**/build/
**.pyc
**.txt
*.log
weights/
# Distribution / packaging
build/
+2 -2
View File
@@ -15,7 +15,7 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
## NEWS
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) 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) 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/468).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
## Acknowledgement
We learned and reused code from the following projects:
- [Wan-Video](https://github.com/Wan-Video)
+14 -7
View File
@@ -2,12 +2,12 @@
# Attention Kernel Used in FastVideo
## Sliding Tile Attention (STA)
We only support H100 for STA.
We support H100 (via TK) and any other GPU (via triton) for STA.
### Installation
```bash
pip install st_attn
```
```
Install from source:
@@ -16,6 +16,14 @@ 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
@@ -30,7 +38,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
```
@@ -43,7 +51,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.
@@ -58,7 +66,6 @@ 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
@@ -67,7 +74,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
@@ -82,7 +89,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.
+16 -9
View File
@@ -51,21 +51,28 @@ 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=[
CUDAExtension('st_attn_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
],
ext_modules=ext_modules,
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
@@ -7,12 +7,17 @@ try:
except ImportError:
sta_fwd = None
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
try:
from st_attn.st_attn_triton import sliding_tile_attention_triton
except ImportError:
sliding_tile_attention_triton = None
def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
seq_length = q_all.shape[2]
dit_seq_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
@@ -46,4 +51,13 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
_ = 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]
return hidden_states[:, :, :seq_length]
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
major, minor = torch.cuda.get_device_capability(q_all.device)
if major == 9 and minor == 0 and sta_fwd is not None:
return sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
elif sliding_tile_attention_triton is not None:
return sliding_tile_attention_triton(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
else:
raise ImportError("No suitable sliding tile attention implementation found.")
@@ -0,0 +1,327 @@
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
+99 -10
View File
@@ -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, non_pad_index, 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):
Q = Q.detach().requires_grad_()
K = K.detach().requires_grad_()
V = V.detach().requires_grad_()
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
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)
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
output = output[:, :, non_pad_index, :]
output = output[:, :, q_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=32, max_size=64, device="cuda"):
def generate_variable_block_sizes(num_blocks, min_size=16, max_size=64, device="cuda"):
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
@@ -86,19 +86,21 @@ 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, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
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)
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, 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)
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
@@ -118,6 +120,60 @@ 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"},
@@ -147,10 +203,43 @@ 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)
print("\nAnalysis completed for all modes.")
generate_error_graphs_qkdiff(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
+236
View File
@@ -0,0 +1,236 @@
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}"
)
+27 -21
View File
@@ -1,54 +1,60 @@
import torch
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
def generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate block sparse mask of shape [h, num_blocks, num_blocks].
Generate block sparse mask of shape [h, num_q_blocks, num_kv_blocks].
Args:
h: number of heads
num_blocks: number of blocks
num_q_blocks: number of query blocks
num_kv_blocks: number of key/value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
block_sparse_mask: [h, num_q_blocks, num_kv_blocks] bool tensor
"""
k = min(k, num_blocks)
scores = torch.rand(h, num_blocks, num_blocks, device=device)
k = min(k, num_kv_blocks)
scores = torch.rand(h, num_q_blocks, num_kv_blocks, device=device)
_, indices = torch.topk(scores, k, dim=-1)
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
block_sparse_mask = torch.zeros(h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
return block_sparse_mask
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
def create_full_mask_from_block_mask(block_sparse_mask, q_variable_block_sizes,
kv_variable_block_sizes, device="cuda"):
"""
Convert block-level sparse mask to full attention mask.
Args:
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
variable_block_sizes: [num_blocks] tensor
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
device: device to create tensors on
Returns:
full_mask: [h, S, S] bool tensor where S = total sequence length
full_mask: [h, S_q, S_kv] bool tensor where S = total sequence length
"""
h, num_blocks, _ = block_sparse_mask.shape
total_seq_len = variable_block_sizes.sum().item()
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
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]])
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
full_mask = torch.zeros(h, total_q_seq_len, total_kv_seq_len, dtype=torch.bool, device=device)
for head in range(h):
for q_block in range(num_blocks):
q_start = cumsum[q_block]
q_end = q_start + variable_block_sizes[q_block]
for q_block in range(num_q_blocks):
q_start = q_cumsum[q_block]
q_end = q_start + q_variable_block_sizes[q_block]
for kv_block in range(num_blocks):
for kv_block in range(num_kv_blocks):
if block_sparse_mask[head, q_block, kv_block]:
kv_start = cumsum[kv_block]
kv_end = kv_start + variable_block_sizes[kv_block]
kv_start = kv_cumsum[kv_block]
kv_end = kv_start + kv_variable_block_sizes[kv_block]
full_mask[head, q_start:q_end, kv_start:kv_end] = True
return full_mask
@@ -672,23 +672,32 @@ block_sparse_attention_forward(
torch::Tensor v,
torch::Tensor q2k_block_sparse_index,
torch::Tensor q2k_block_sparse_num,
torch::Tensor block_size
torch::Tensor kv_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 seq_len = q.size(2);
auto q_seq_len = q.size(2);
auto kv_seq_len = k.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 = block_size.size(0);
auto num_q_blocks = q2k_block_sparse_index.size(2);
auto num_kv_blocks = kv_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 * 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");
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.");
// 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");
@@ -696,11 +705,8 @@ 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(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(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(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
@@ -727,12 +733,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>(seq_len),
static_cast<const uint>(q_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>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
@@ -762,11 +768,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>(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};
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};
globals g{
qg_arg,
@@ -774,17 +780,17 @@ block_sparse_attention_forward(
vg_arg,
lg_arg,
og_arg,
static_cast<int>(seq_len),
static_cast<int>(q_seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
};
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<64>,
@@ -813,11 +819,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>(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};
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};
globals g{
qg_arg,
@@ -825,17 +831,17 @@ block_sparse_attention_forward(
vg_arg,
lg_arg,
og_arg,
static_cast<int>(seq_len),
static_cast<int>(q_seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
};
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<128>,
@@ -862,7 +868,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 block_size)
torch::Tensor kv_block_size)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -871,11 +877,23 @@ 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 seq_len = q.size(2);
auto q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto head_dim = q.size(3);
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
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))");
// 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");
@@ -886,23 +904,18 @@ 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(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(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(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);
@@ -929,20 +942,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>(seq_len),
static_cast<const uint>(q_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>(seq_len),
static_cast<const uint>(kv_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>(seq_len),
static_cast<const uint>(kv_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>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(1)}, l_vec.options());
float* qg_ptr = qg.data_ptr<float>();
@@ -971,7 +984,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(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
dim3 grid_bwd(q_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>;
@@ -984,9 +997,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>(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)};
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)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1023,15 +1036,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>(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_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_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1042,14 +1055,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(seq_len),
static_cast<int>(kv_seq_len), // N is not used in the kernel
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*>(block_size.data_ptr())};
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
@@ -1088,9 +1101,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>(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)};
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)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1127,15 +1140,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>(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_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_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1146,14 +1159,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(seq_len),
static_cast<int>(kv_seq_len), // N is not used in the kernel
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*>(block_size.data_ptr())};
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
+7
View File
@@ -0,0 +1,7 @@
build/
dist/
*.egg-info/
__pycache__/
*.so
*.pyc
.ipynb_checkpoints/
+187
View File
@@ -0,0 +1,187 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
+6
View File
@@ -0,0 +1,6 @@
include LICENSE
include README.md
include pyproject.toml
recursive-include src/fastvideo_kernel *.cu *.cuh *.cpp *.h
recursive-include csrc *.cu *.cuh *.cpp *.h
recursive-include tk *.cu *.cuh *.cpp *.h
+31
View File
@@ -0,0 +1,31 @@
# FastVideo Kernel
CUDA kernels for FastVideo video generation.
## Installation
```bash
git submodule update --init --recursive
cd csrc/fastvideo_kernel
pip install .
```
## Usage
```python
from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_attn_varlen
# Example: Sliding Tile Attention
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
# Example: Video Sparse Attention (with Triton fallback)
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
# Example: VMoBA
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
```
## Requirements
- H100 GPU (sm_90a) for CUDA kernels
- Triton for non-H100 fallback
File diff suppressed because it is too large Load Diff
+23
View File
@@ -0,0 +1,23 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ST_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
+573
View File
@@ -0,0 +1,573 @@
// # 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();
}
+27
View File
@@ -0,0 +1,27 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_BLOCK_SPARSE
extern std::vector<torch::Tensor> block_sparse_attention_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
);
extern std::vector<torch::Tensor> block_sparse_attention_backward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_BLOCK_SPARSE
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
#endif
}
+27
View File
@@ -0,0 +1,27 @@
[build-system]
requires = ["setuptools>=61.0", "torch>=2.5.0", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "fastvideo-kernel"
version = "0.1.0"
description = "CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
license = {text = "Apache-2.0"}
authors = [{name = "Hao AI Lab"}]
classifiers = [
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
]
dependencies = [
"torch>=2.5.0",
"triton>=2.0.0"
]
[project.urls]
Repository = "https://github.com/hao-ai-lab/FastVideo"
[tool.setuptools.packages.find]
where = ["src"]
+132
View File
@@ -0,0 +1,132 @@
import os
import subprocess
import sys
from pathlib import Path
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
ROOT = Path(__file__).parent.absolute()
CSRC_DIR = ROOT / "csrc"
# Path to ThunderKittens (TK)
def get_tk_dir():
tk_env = os.getenv("THUNDERKITTENS_ROOT")
if tk_env:
return tk_env
# Check common locations
possible_paths = [
ROOT / "tk",
ROOT / "csrc" / "tk",
ROOT.parent / "attn" / "sliding_tile_attn" / "tk",
ROOT.parent / "attn" / "video_sparse_attn" / "tk",
]
for p in possible_paths:
if (p / "include" / "kittens.cuh").exists():
return str(p)
# Default fallback
return str(ROOT.parent / "attn" / "sliding_tile_attn" / "tk")
TK_DIR = get_tk_dir()
def get_cuda_flags(tk_root: str) -> list:
python_include = subprocess.check_output(
["python", "-c", "import sysconfig; print(sysconfig.get_path('include'))"]
).decode().strip()
torch_includes = subprocess.check_output([
"python", "-c",
"import torch; from torch.utils.cpp_extension import include_paths; "
"print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip().split()
return [
"-DNDEBUG",
"-Xcompiler=-Wno-psabi",
"-Xcompiler=-fno-strict-aliasing",
"--expt-extended-lambda",
"--expt-relaxed-constexpr",
"-forward-unknown-to-host-compiler",
"--use_fast_math",
"-std=c++20",
"-O3",
"-Xnvlink=--verbose",
"-Xptxas=--verbose",
"-Xptxas=--warn-on-spills",
f"-I{tk_root}/include",
f"-I{tk_root}/prototype",
f"-I{python_include}",
"-DTORCH_COMPILE",
"-DKITTENS_HOPPER",
"-arch=sm_90a",
] + torch_includes
def get_extensions():
if not torch.cuda.is_available():
return []
extensions = []
cpp_flags = ["-std=c++20", "-O3"]
# Check if TK is available
if not os.path.exists(os.path.join(TK_DIR, "include", "kittens.cuh")):
print(f"Warning: ThunderKittens not found at {TK_DIR}. CUDA kernels will not be built.")
return []
cuda_flags = get_cuda_flags(TK_DIR)
# STA Extension
extensions.append(CUDAExtension(
"fastvideo_kernel._C.st_attn",
sources=[
"csrc/st_attn.cpp",
"csrc/st_attn_h100.cu",
],
extra_compile_args={
"cxx": cpp_flags + ["-DTK_COMPILE_ST_ATTN"],
"nvcc": cuda_flags + ["-DTK_COMPILE_ST_ATTN"]
},
libraries=["cuda"],
))
# VSA Extension
extensions.append(CUDAExtension(
"fastvideo_kernel._C.vsa",
sources=[
"csrc/vsa.cpp",
"csrc/block_sparse_h100.cu",
],
extra_compile_args={
"cxx": cpp_flags + ["-DTK_COMPILE_BLOCK_SPARSE"],
"nvcc": cuda_flags + ["-DTK_COMPILE_BLOCK_SPARSE"]
},
libraries=["cuda"],
))
return extensions
ext_modules = []
if not any(arg in sys.argv for arg in ["clean", "egg_info", "--version"]):
try:
import torch
ext_modules = get_extensions()
except Exception as e:
print(f"Warning: Failed to configure CUDA extensions: {e}")
setup(
name="fastvideo-kernel",
version="0.1.0",
description="Unified CUDA kernels for FastVideo",
long_description=open("README.md").read(),
long_description_content_type="text/markdown",
license="Apache-2.0",
author="Hao AI Lab",
url="https://github.com/hao-ai-lab/FastVideo",
package_dir={"": "src"},
packages=find_packages(where="src"),
ext_modules=ext_modules,
cmdclass={"build_ext": BuildExtension} if ext_modules else {},
python_requires=">=3.10",
install_requires=["torch>=2.5.0", "triton>=2.0.0"],
)
@@ -0,0 +1,21 @@
__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__",
]
@@ -0,0 +1,103 @@
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
@@ -0,0 +1,449 @@
"""
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
@@ -0,0 +1,152 @@
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
@@ -0,0 +1,868 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
import random
import time
import os
import torch
from typing import Tuple
try:
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
from functools import lru_cache
from einops import rearrange
@lru_cache(maxsize=16)
def calc_chunks(cu_seqlen, moba_chunk_size):
"""
Calculate chunk boundaries.
For vision tasks we include all chunks (even the last one which might be shorter)
so that every chunk can be selected.
"""
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
cu_num_chunk = torch.ones(
batch_num_chunk.numel() + 1,
device=cu_seqlen.device,
dtype=batch_num_chunk.dtype,
)
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
num_chunk = cu_num_chunk[-1]
chunk_sizes = torch.full(
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
)
chunk_sizes[0] = 0
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
chunk_to_batch = torch.zeros(
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
)
chunk_to_batch[cu_num_chunk[1:-1]] = 1
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
# Do not filter out any chunk
filtered_chunk_indices = torch.arange(
num_chunk, device=cu_seqlen.device, dtype=torch.int32
)
num_filtered_chunk = num_chunk
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
# --- Threshold Selection Helper Functions ---
def _select_threshold_query_head(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects chunks for each <query, head> pair based on threshold.
Normalization and sorting happen along the chunk dimension (dim=0).
"""
C, H, S = gate.shape
eps = 1e-6
# LSE‐style normalization per <head, query> (across chunks)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
row_min = gate_min_val.amin(dim=0) # (H, S)
row_max = gate_masked.amax(dim=0) # (H, S)
denom = row_max - row_min
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
# 2) compute how much more normalized weight we need beyond self
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
# 3) zero out the self‐chunk in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0
# 4) sort the other chunks by descending norm, per <head,seq>
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
# 5) cumulative‑sum the sorted norms per <head,seq>
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
any_cond = cond.any(dim=0) # (H, S)
# Find the index of the first True value along dim 0. If none, use C-1.
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
# 7) build a mask in sorted order up to that cutoff
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
# 8) scatter it back to original chunk order
others_mask = torch.zeros_like(gate, dtype=torch.bool)
others_mask.scatter_(0, sorted_idx, sorted_mask)
# 9) finally, include every self‐chunk plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_block(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <query, head> pairs for each block based on threshold.
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
"""
C, H, S = gate.shape
HS = H * S
eps = 1e-6
# LSE‐style normalization per block (across heads and queries)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_denom = block_max - block_min
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
gate_norm = (gate - block_min) / block_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights *per block*
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort the other <head, seq> pairs by descending norm, per block
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
# 5) cumulative‑sum the sorted norms per block
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
any_cond = cond_flat.any(dim=1) # (C,)
# Find the index of the first True value along dim 1. If none, use HS-1.
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
# 7) build a mask in sorted order up to that cutoff per block
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
# 8) scatter it back to original <head, seq> order per block
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_overall(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query, head> triplets globally based on threshold.
Normalization and sorting happen across all valid entries.
"""
C, H, S = gate.shape
CHS = C * H * S
eps = 1e-6
# LSE‐style normalization globally across all valid entries
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
overall_max = gate_masked.max() # scalar
overall_min = gate_min_val.min() # scalar
overall_denom = overall_max - overall_min
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights globally
self_norm_sum_overall = self_norm_entries.sum() # scalar
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
total_norm_sum_overall = gate_norm.sum() # scalar
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort all other entries by descending norm, globally
others_flat = others_norm.flatten() # (C*H*S,)
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
# Only sort the valid 'other' entries
valid_others_indices = torch.where(valid_others_mask_flat)[0]
valid_others_values = others_flat[valid_others_indices]
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
# 5) cumulative‑sum the sorted valid 'other' norms globally
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
any_cond = cond_values.any() # scalar
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
cutoff_idx_in_sorted = torch.where(
any_cond,
cond_values.float().argmax(dim=0),
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
)
# 7) build a mask selecting the top-k others based on the cutoff
# Select the original indices corresponding to the top entries in the sorted list
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
# 8) create the mask in the original flat shape
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
others_mask_flat[selected_other_indices] = True
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_head_global(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query> globally for each head based on threshold.
"""
C, H, S = gate.shape
eps = 1e-6
# 1) LSE‐style normalization per head (across chunks and sequence dims)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
denom = max_per_head - min_per_head
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
gate_norm = (gate - min_per_head) / denom
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 2) sum normalized self‐chunk contributions per head
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
# 3) total normalized sum per head
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
# 4) how much more normalized weight needed per head
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
remain_ratio = torch.clamp(remain_ratio, min=0.0)
# 5) zero out self‐chunk entries to focus on "others"
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
# 6) flatten chunk and sequence dims, per head
CS = C * S
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
# 7) vectorized selection of “others” per head
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
has_cutoff = cond.any(dim=1) # (H,)
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
# 8) reshape selection mask back to (C, H, S)
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
# 9) include self‐chunks plus selected others, and obey valid mask
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
return final_gate_mask
class MixedAttention(torch.autograd.Function):
@staticmethod
def forward(
ctx,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
):
ctx.max_seqlen = max_seqlen
ctx.moba_chunk_size = moba_chunk_size
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
# Non-causal self-attention branch
# return out, softmax_lse, S_dmask, rng_state
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=q,
k=k,
v=v,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
# MOBA attention branch (non-causal)
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
output_2d = output.view(-1, q.shape[2])
max_lse_1d = self_attn_lse_sh.view(-1)
max_lse_1d = max_lse_1d.index_reduce(
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
)
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
moba_attn_lse = (
moba_attn_lse.view(-1)
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
.reshape_as(moba_attn_lse)
)
mixed_attn_se_sh = self_attn_lse_sh.exp()
moba_attn_se = moba_attn_lse.exp()
mixed_attn_se_sh.view(-1).index_add_(
0, moba_q_sh_indices, moba_attn_se.view(-1)
)
mixed_attn_lse_sh = mixed_attn_se_sh.log()
# Combine self-attention output
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
output_2d += self_attn_out_sh.reshape_as(output_2d)
# Combine MOBA attention output
mixed_attn_lse = (
mixed_attn_lse_sh.view(-1)
.index_select(0, moba_q_sh_indices)
.view_as(moba_attn_lse)
)
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
output = output.to(q.dtype)
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
ctx.save_for_backward(
output,
mixed_attn_lse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
)
return output
@staticmethod
def backward(ctx, d_output):
max_seqlen = ctx.max_seqlen
moba_chunk_size = ctx.moba_chunk_size
softmax_scale = ctx.softmax_scale
(
output,
mixed_attn_vlse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
) = ctx.saved_tensors
d_output = d_output.contiguous()
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
_ = _flash_attn_varlen_backward(
dout=d_output,
q=q,
k=k,
v=v,
out=output,
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
dq=dq,
dk=dk,
dv=dv,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
headdim = q.shape[-1]
d_moba_output = (
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
moba_output = (
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
mixed_attn_vlse = (
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
)
dmq = torch.empty_like(moba_q)
dmkv = torch.empty_like(moba_kv)
_ = _flash_attn_varlen_backward(
dout=d_moba_output,
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
out=moba_output,
softmax_lse=mixed_attn_vlse,
dq=dmq,
dk=dmkv[:,0],
dv=dmkv[:,1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
def moba_attn_varlen(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens: torch.Tensor,
max_seqlen: int,
moba_chunk_size: int,
moba_topk: int,
select_mode: str = 'threshold', # "topk" or "threshold"
simsum_threshold: float = 0.25,
threshold_type: str = 'query_head',
) -> torch.Tensor:
"""
Accelerated MOBA attention for vision tasks with proper LSE normalization.
This version:
- Splits KV into chunks.
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
by amplifying the diagonal (self-chunk) logits.
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
reduction so that attending to each query over the selected chunks is equivalent
to the original algorithm.
"""
# Stack keys and values.
kv = torch.stack((k, v), dim=1)
seqlen, num_head, head_dim = q.shape
# Compute chunk boundaries.
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
cu_seqlens, moba_chunk_size
)
self_attn_cu_seqlen = cu_chunk
# Update top-k selection to include the self chunk.
moba_topk = min(moba_topk, num_filtered_chunk)
# --- Build filtered KV from chunks ---
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
max_chunk_len = int(chunk_lengths.max().item())
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
indices = chunk_starts.unsqueeze(1) + range_tensor
indices = torch.clamp(indices, max=kv.shape[0] - 1)
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
# Compute key_gate_weight over valid tokens.
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
key_sum = (key_values * valid_mask_exp).sum(dim=1)
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
# Compute gate logits between key_gate_weight and queries.
q_float = q.float()
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
# Amplify the diagonal (self chunk) contributions.
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
amplification_factor = 1e9 # Example factor; adjust as needed.
origin_gate = gate.clone()
gate = gate.clone()
if select_mode == "topk":
gate[gate_self_chunk_mask] += amplification_factor
# Exclude positions that are outside the valid batch boundaries.
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
if select_mode == 'topk':
# We amplify self‐chunk in gate already, so self entries will rank highest.
valid_gate_mask = gate != -float("inf")
if threshold_type == 'query_head':
# === per‐<head,seq> top-k across chunks (original behavior) ===
# gate: (C, H, S)
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
gate_idx_mask.scatter_(0, gate_topk_idx, True)
gate_mask = valid_gate_mask & gate_idx_mask
elif threshold_type == 'overall':
# === global top-k across all (chunk, head, seq) entries ===
C, H, S = gate.shape
flat_gate = gate.flatten()
flat_mask = valid_gate_mask.flatten()
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
# pick topk global entries
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
others_mask_flat[idx] = True
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
elif threshold_type == 'head_global':
# per-head top-k across all chunks and sequence positions
C, H, S = gate.shape
CS = C * S
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
# pick top-k indices per head
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
gate_idx_flat.scatter_(1, topk_idx, True)
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
else:
raise ValueError(
f"Invalid threshold_type for topk: {threshold_type}. "
"Choose 'query_head', 'block', or 'overall'."
)
elif select_mode == 'threshold':
# Delegate to the specific thresholding function
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
if threshold_type == 'query_head':
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'block':
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'overall':
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'head_global':
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
else:
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
else:
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
# eliminate self_chunk in MoBA branch
gate_mask = gate_mask & ~gate_self_chunk_mask
# if gate_mask is all false, perform flash_attn instead
if gate_mask.sum() == 0:
return flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
)
# Determine which query positions are selected.
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
# Build cumulative sequence lengths for the selected queries.
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
q_zero_mask = moba_seqlen_q == 0
valid_expert_mask = ~q_zero_mask
if q_zero_mask.sum() > 0:
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
moba_cu_seqlen_q = torch.cat(
(
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
moba_seqlen_q.cumsum(dim=0),
),
dim=0,
).to(torch.int32)
# Rearrange gathered KV for the MOBA branch.
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
if q_zero_mask.sum() > 0:
experts_tensor = experts_tensor[valid_expert_mask]
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
mask = seq_range < valid_expert_lengths.unsqueeze(1)
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
moba_cu_seqlen_kv = torch.cat(
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
valid_expert_lengths.cumsum(dim=0)],
dim=0,
).to(torch.int32)
assert (
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
return MixedAttention.apply(
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
)
def process_moba_input(
x,
patch_resolution,
chunk_size,
):
"""
Process inputs for the attention function.
Args:
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
Returns:
torch.Tensor: Processed input tensor.
"""
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
else:
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
if len(chunk_size) == 2:
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
elif len(chunk_size) == 3:
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
else:
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
return x, moba_chunk_size
def process_moba_output(
x,
patch_resolution,
chunk_size,
):
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
pass
elif len(chunk_size) == 2:
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
elif len(chunk_size) == 3:
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
return x
# TEST
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
random.seed(0)
torch.manual_seed(0)
torch.cuda.manual_seed(0)
device = torch.cuda.current_device()
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
max_seqlen = q.shape[1]
q = rearrange(q, "b s ... -> (b s) ...")
k = rearrange(k, "b s ... -> (b s) ...")
v = rearrange(v, "b s ... -> (b s) ...")
return q, k, v, cu_seqlens, max_seqlen
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
"""Speed test comparing flash_attn vs moba_attention"""
# Get data
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
vo_grad = torch.randn_like(q)
# Warmup
warmup_iters = 3
perf_test_iters = 10
# Warmup
for _ in range(warmup_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
start_flash = time.perf_counter()
for _ in range(perf_test_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
# Warmup
for _ in range(warmup_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
start_moba = time.perf_counter()
for _ in range(perf_test_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
print(f"Speedup: {time_flash / time_moba:.2f}x")
if __name__ == "__main__":
"""
CUDA_VISIBLE_DEVICES=1 \
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
"""
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
@@ -0,0 +1,71 @@
from typing import Tuple
import torch
from torch import BoolTensor, IntTensor
from torch.nn.attention.flex_attention import create_block_mask
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
torch._inductor.config.realize_opcount_threshold = 100
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
"""Generates a 3D NATTEN attention mask with a given kernel size.
Args:
canvas_t: The time dimension of the canvas.
canvas_h: The height of the canvas.
canvas_w: The width of the canvas.
kernel_t: The time dimension of the kernel.
kernel_h: The height of the kernel.
kernel_w: The width of the kernel.
"""
canvas_t, canvas_h, canvas_w = canvas_twh
kernel_t, kernel_h, kernel_w = kernel_twh
tile_t_size, tile_h_size, tile_w_size = tile_twh
total_tile_size = tile_t_size * tile_h_size * tile_w_size
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
img_seq_len = canvas_t * canvas_h * canvas_w
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
tile_id = idx // total_tile_size
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
tile_w = tile_id % canvas_tile_w
return tile_t, tile_h, tile_w
def sta_mask_3d(
b: IntTensor,
h: IntTensor,
q_idx: IntTensor,
kv_idx: IntTensor,
) -> BoolTensor:
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
# kernel nominally attempts to center itself on the query, but kernel center
# is clamped to a fixed distance (kernel half-length) from the canvas edge
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
return sta_mask_3d
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
img_seq_len = img_size[0] * img_size[1] * img_size[2]
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
mask = create_block_mask(image_mask,
B=None,
H=None,
Q_LEN=img_seq_len + text_max_len,
KV_LEN=img_seq_len + text_max_len,
device=device,
_compile=True)
return mask
@@ -0,0 +1,63 @@
import torch
import sys
import os
from tqdm import tqdm
# Local support import
from .support_flex_sta import get_sliding_tile_attention_mask
# USE OUR NEW PACKAGE!
from fastvideo_kernel import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
# Using the same parameters as the original test
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
for kernel_size in kernel_size_ls:
print(f"Testing kernel_size: {kernel_size}")
for xi in tqdm(range(num_iterations)):
torch.manual_seed(xi)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
pt_o = flex_test(Q, K, V, kernel_size)
diff = pt_o - tk_o
abs_diff = torch.abs(diff)
max_d = torch.max(abs_diff).item()
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
if max_d > 0.1:
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
if __name__ == "__main__":
b, h, d = 2, 24, 128
n = 69120
causal = False
mean = 1e-1
std = 10
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
@@ -0,0 +1,97 @@
# 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"
+2 -2
View File
@@ -1,4 +1,4 @@
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
FROM nvidia/cuda:12.8.0-cudnn-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 flash-attn==2.8.3 --no-build-isolation
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
COPY . .
+4 -11
View File
@@ -1,4 +1,4 @@
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
FROM nvidia/cuda:12.8.0-cudnn-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 flash-attn==2.8.3 --no-build-isolation
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
COPY . .
@@ -55,17 +55,10 @@ 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 STA (Sliding Tile Attention)
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
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 && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
python setup.py install
+3 -10
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
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
COPY . .
@@ -55,17 +55,10 @@ 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 STA (Sliding Tile Attention)
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
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 && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
python setup.py install
+2 -9
View File
@@ -55,17 +55,10 @@ 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 STA (Sliding Tile Attention)
# Install FastVideo Kernels
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
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 && \
cd csrc/fastvideo_kernel && \
git submodule update --init --recursive && \
python setup.py install
+58
View File
@@ -0,0 +1,58 @@
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
View File
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
```bash
# Install dependencies
pip install -r docs/requirements-mkdocs.txt
pip install -r requirements-mkdocs.txt
# Serve docs with live reload (recommended for development)
mkdocs serve
+1 -1
View File
@@ -6,7 +6,7 @@ Thank you for your interest in contributing to FastVideo. We want to make the pr
Our community is open to everyone and welcomes any contributions no matter how large or small.
# Developer Environment:
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support 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 supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
+2 -2
View File
@@ -1,7 +1,7 @@
# Profiling FastVideo
!!! warning
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the 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 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 disks.
- After profiling, clean up trace directories to avoid filling disk storage.
- 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.
+5 -3
View File
@@ -74,9 +74,11 @@ To add a new SSIM test, follow these steps:
generator.generate_video(prompt, ...)
# Compare with Reference
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
assert ssim_values[0] >= 0.98 # Threshold
```
ssim_values = compute_video_ssim_torchvision(
reference_path, generated_path, use_ms_ssim=True
)
assert ssim_values[0] >= 0.98 # Threshold
```
4. **Reference Videos**:
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
+2 -2
View File
@@ -1,6 +1,6 @@
# 🎯 Distillation
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computations, 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 computation, 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.html). Set `MODEL_BASE` to your own model path and run:
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
```bash
bash scripts/inference/v1_inference_wan_dmd.sh
+40 -18
View File
@@ -7,6 +7,11 @@ 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
```
@@ -15,36 +20,53 @@ pip install fastvideo
### Text-to-Video Generation
```python
from fastvideo import FastVideoPipeline
from fastvideo import VideoGenerator
# Initialize the pipeline
pipe = FastVideoPipeline.from_pretrained("wan2.1-t2v-1.3B")
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
)
# Generate a video
prompt = "A cat playing with a ball of yarn"
video = pipe(prompt, num_frames=16, height=512, width=512)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Save the video
video.save("output.mp4")
# 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()
```
### Image-to-Video Generation
```python
from fastvideo import FastVideoPipeline
from PIL import Image
from fastvideo import VideoGenerator, SamplingParam
# Load an image
image = Image.open("input.jpg")
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Initialize the pipeline
pipe = FastVideoPipeline.from_pretrained("wan2.1-i2v-14B-480p")
# 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
# Generate a video from the image
video = pipe(image, num_frames=16, height=480, width=480)
# 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)
# Save the video
video.save("output.mp4")
if __name__ == '__main__':
main()
```
## Next Steps
+1 -1
View File
@@ -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 implementation on H100.
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently, we only have an implementation for H100s.
First, install C++20 for ThunderKittens:
```bash
+1 -1
View File
@@ -30,7 +30,7 @@ path_to_your_dataset_folder/
└── prompt.txt
```
To geranate the `videos2caption.json` and `merge.txt`, run
To generate 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
+2 -2
View File
@@ -7,9 +7,9 @@ pip install vsa
```
# Building from Source
We support H100 (via ThunderKittens) and any other GPU (via Triton) for VSA.
We support H100s (via ThunderKittens) and any other GPU (via Triton) for VSA.
First, install C++20 for ThunderKittens (if using H100):
First, install C++20 for ThunderKittens (if using an H100):
```bash
sudo apt update
+1 -1
View File
@@ -6,7 +6,7 @@ The `VideoGenerator` class provides the primary Python interface for doing offli
- Python 3.10-3.12
## Installation
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) first.
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation) first.
## Usage
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
+42
View File
@@ -0,0 +1,42 @@
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()
@@ -0,0 +1,75 @@
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()
@@ -0,0 +1,107 @@
#!/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[@]}"
+47 -8
View File
@@ -1,8 +1,9 @@
# 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
@@ -46,6 +47,29 @@ 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__(
@@ -66,12 +90,27 @@ class FlashAttentionImpl(AttentionImpl):
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
attn_metadata: FlashAttnMetadata,
):
output = flash_attn_func(
query, # type: ignore[no-untyped-call]
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal)
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)
return output
+29 -4
View File
@@ -1,9 +1,10 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from dataclasses import dataclass
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
AttentionBackend, AttentionImpl, AttentionMetadata)
AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@@ -30,6 +31,29 @@ 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__(
@@ -51,14 +75,15 @@ class SDPAImpl(AttentionImpl):
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
attn_metadata: SDPAMetadata,
) -> 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": None,
"attn_mask": attn_mask,
"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 st_attn import sliding_tile_attention
from fastvideo_kernel 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 vsa import video_sparse_attn
from fastvideo_kernel import video_sparse_attn
except ImportError:
video_sparse_attn = None
+2 -2
View File
@@ -6,8 +6,8 @@ from dataclasses import dataclass
import torch
from einops import rearrange
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
from fastvideo_kernel import (moba_attn_varlen, process_moba_input,
process_moba_output)
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
+62 -4
View File
@@ -11,6 +11,7 @@ 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):
@@ -64,6 +65,8 @@ 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.
@@ -74,6 +77,7 @@ 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:
@@ -91,12 +95,30 @@ class DistributedAttention(nn.Module):
ctx_attn_metadata = forward_context.attn_metadata
# Stack QKV
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
qkv = torch.cat([q, k, v],
dim=0) # [3*batch, 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)
@@ -119,17 +141,23 @@ class DistributedAttention(nn.Module):
# Redistribute back if using sequence parallelism
replicated_output = None
if replicated_q is not None:
replicated_output = output[:, seq_len * world_size:]
output = output[:, :seq_len * world_size]
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]
# 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
@@ -147,6 +175,8 @@ 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.
@@ -158,6 +188,7 @@ 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:
@@ -173,15 +204,32 @@ 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) # [3, seq_len, num_heads, head_dim]
dim=0) # [4*batch, 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)
@@ -194,6 +242,10 @@ 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)
@@ -244,6 +296,7 @@ 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.
@@ -263,5 +316,10 @@ 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
@@ -0,0 +1,99 @@
# 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
+5 -2
View File
@@ -1,10 +1,13 @@
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", "WanVideoConfig", "StepVideoConfig",
"CosmosVideoConfig", "Cosmos25VideoConfig"
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig"
]
+2
View File
@@ -23,6 +23,8 @@ 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
@@ -0,0 +1,157 @@
# 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"
+149
View File
@@ -0,0 +1,149 @@
# 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"
@@ -0,0 +1,83 @@
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,9 +6,11 @@ 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"
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig"
]
@@ -72,6 +72,7 @@ class EncoderConfig(ModelConfig):
@dataclass
class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
is_chat_model: bool = False
@dataclass
@@ -0,0 +1,93 @@
# 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
+3
View File
@@ -40,6 +40,8 @@ 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)
@@ -68,6 +70,7 @@ class T5ArchConfig(TextEncoderArchConfig):
"return_attention_mask": True,
"return_tensors": "pt",
}
self.hidden_size = self.d_model
@dataclass
@@ -1,5 +1,6 @@
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
@@ -8,4 +9,5 @@ __all__ = [
"WanVAEConfig",
"StepVideoVAEConfig",
"CosmosVAEConfig",
"Hunyuan15VAEConfig",
]
@@ -0,0 +1,27 @@
# 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)
+5 -4
View File
@@ -2,6 +2,7 @@ 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
@@ -11,8 +12,8 @@ from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "CosmosConfig",
"get_pipeline_config_cls_from_name"
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "get_pipeline_config_cls_from_name"
]
+139
View File
@@ -0,0 +1,139 @@
# 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
+355
View File
@@ -0,0 +1,355 @@
# 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}'"
)
+36 -8
View File
@@ -7,14 +7,17 @@ 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)
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig,
MatrixGameI2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -26,6 +29,10 @@ 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,
@@ -45,25 +52,45 @@ 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(),
"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(),
"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(),
# 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,
@@ -109,6 +136,7 @@ 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():
+21
View File
@@ -6,6 +6,7 @@ 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)
@@ -190,3 +191,23 @@ 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
+44
View File
@@ -3,6 +3,7 @@ from dataclasses import dataclass
from typing import Any
from fastvideo.logger import init_logger
from fastvideo.utils import StoreBoolean
logger = init_logger(__name__)
@@ -17,10 +18,27 @@ 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"
@@ -44,6 +62,7 @@ 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
@@ -209,6 +228,31 @@ 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,
+29
View File
@@ -0,0 +1,29 @@
# 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
+56 -39
View File
@@ -5,6 +5,7 @@ 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
@@ -23,6 +24,7 @@ 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
@@ -32,44 +34,36 @@ 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,
"FastVideo/stepvideo-t2v-diffusers":
StepVideoT2VSamplingParam,
"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,
# 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":
@@ -85,17 +79,32 @@ 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(),
"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(),
"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(),
# Add other pipeline architecture detectors
}
@@ -103,12 +112,15 @@ 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
}
@@ -116,6 +128,20 @@ 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(
@@ -126,15 +152,6 @@ 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
+11
View File
@@ -196,3 +196,14 @@ class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
height: int = 448
width: int = 832
fps: int = 16
@dataclass
class MatrixGame2_SamplingParam(SamplingParam):
height: int = 352
width: int = 640
num_frames: int = 57
fps: int = 25
guidance_scale: float = 1.0
num_inference_steps: int = 3
negative_prompt: str | None = None
+72 -1
View File
@@ -4,7 +4,13 @@
import torch
import torch.distributed
from fastvideo.distributed.parallel_state import get_sp_group, get_tp_group
from fastvideo.distributed.parallel_state import (get_sp_group,
get_sp_parallel_rank,
get_sp_world_size,
get_tp_group)
from fastvideo.distributed.utils import (unpad_sequence_tensor,
compute_padding_for_sp,
pad_sequence_tensor)
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
@@ -30,3 +36,68 @@ def sequence_model_parallel_all_gather(input_: torch.Tensor,
dim: int = -1) -> torch.Tensor:
"""All-gather the input tensor across model parallel group."""
return get_sp_group().all_gather(input_, dim)
def sequence_model_parallel_all_gather_with_unpad(
input_: torch.Tensor,
original_seq_len: int,
dim: int = -1) -> torch.Tensor:
"""All-gather the input tensor and remove padding.
Args:
input_: Sharded (and possibly padded) tensor to gather
original_seq_len: Original sequence length before padding
dim: Dimension to gather along (default: -1)
Returns:
Tensor: Gathered and unpadded tensor
"""
# First gather across all ranks
gathered = get_sp_group().all_gather(input_, dim)
current_seq_len = gathered.shape[dim]
if current_seq_len > original_seq_len:
gathered = unpad_sequence_tensor(gathered,
original_seq_len,
seq_dim=dim)
return gathered
def sequence_model_parallel_shard(input_: torch.Tensor,
dim: int = 1) -> tuple[torch.Tensor, int]:
"""Shard the input tensor across model parallel group with optional padding.
Args:
input_: Input tensor to shard
dim: Dimension to shard along (default: 1)
Returns:
tuple: (sharded_tensor, original_seq_len)
- sharded_tensor: The sharded (and possibly padded) tensor
- original_seq_len: Original sequence length before padding
"""
sp_rank = get_sp_parallel_rank()
sp_world_size = get_sp_world_size()
original_seq_len = input_.shape[dim]
# Compute padding if needed
padded_seq_len, padding_amount = compute_padding_for_sp(
original_seq_len, sp_world_size)
# Pad if necessary
if padding_amount > 0:
input_ = pad_sequence_tensor(input_, padded_seq_len, seq_dim=dim)
elements_per_rank = padded_seq_len // sp_world_size
# Sharding along dim
input_ = input_.movedim(dim, 0)
input_ = input_[sp_rank * elements_per_rank:(sp_rank + 1) *
elements_per_rank]
input_ = input_.movedim(0, dim)
return input_, original_seq_len
+123
View File
@@ -61,6 +61,129 @@ def split_tensor_along_last_dim(
return tuple(tensor_list)
def compute_padding_for_sp(seq_len: int, sp_world_size: int) -> tuple[int, int]:
"""
Compute padding needed for sequence parallel.
Args:
seq_len: Original sequence length
sp_world_size: Sequence parallel world size
Returns:
tuple: (padded_seq_len, padding_amount)
"""
if seq_len % sp_world_size == 0:
return seq_len, 0
padding_amount = sp_world_size - (seq_len % sp_world_size)
padded_seq_len = seq_len + padding_amount
return padded_seq_len, padding_amount
def create_attention_mask_for_padding(
seq_len: int,
padded_seq_len: int,
batch_size: int,
device: torch.device,
dtype: torch.dtype = torch.bool,
) -> torch.Tensor | None:
"""
Create attention mask to ignore padded tokens.
Args:
seq_len: Original sequence length (before padding)
padded_seq_len: Padded sequence length
batch_size: Batch size
device: Device to create mask on
dtype: Data type for the mask (default: bool)
Returns:
Tensor: Boolean mask [B, padded_seq_len] where True = valid token,
or None if no padding is needed
"""
if seq_len == padded_seq_len:
return None
# Create mask: True for valid tokens, False for padding
attention_mask = torch.ones(
(batch_size, padded_seq_len),
dtype=dtype,
device=device,
)
# Mask out padding tokens
attention_mask[:, seq_len:] = 0
return attention_mask
def pad_sequence_tensor(
tensor: torch.Tensor,
target_seq_len: int,
seq_dim: int = 1,
pad_value: float = 0.0,
) -> torch.Tensor:
"""
Pad a tensor along the sequence dimension.
Args:
tensor: Input tensor to pad
target_seq_len: Target sequence length after padding
seq_dim: Dimension to pad along (default: 1)
pad_value: Value to use for padding (default: 0.0)
Returns:
Tensor: Padded tensor
"""
current_seq_len = tensor.shape[seq_dim]
if current_seq_len >= target_seq_len:
return tensor
padding_amount = target_seq_len - current_seq_len
# Create padding shape
pad_shape = list(tensor.shape)
pad_shape[seq_dim] = padding_amount
# Create padding tensor
padding = torch.full(
pad_shape,
pad_value,
dtype=tensor.dtype,
device=tensor.device,
)
# Concatenate along sequence dimension
padded_tensor = torch.cat([tensor, padding], dim=seq_dim)
return padded_tensor
def unpad_sequence_tensor(
tensor: torch.Tensor,
original_seq_len: int,
seq_dim: int = 1,
) -> torch.Tensor:
"""
Remove padding from a tensor along the sequence dimension.
Args:
tensor: Padded tensor
original_seq_len: Original sequence length (before padding)
seq_dim: Dimension to unpad along (default: 1)
Returns:
Tensor: Unpadded tensor
"""
# Use slice to remove padding
indices = [slice(None)] * tensor.dim()
indices[seq_dim] = slice(0, original_seq_len)
return tensor[tuple(indices)]
@dataclasses.dataclass
class StatelessProcessGroup:
"""A dataclass to hold a metadata store, and the rank, world_size of the
+18 -12
View File
@@ -102,6 +102,11 @@ class VideoGenerator:
self,
prompt: str | None = None,
sampling_param: SamplingParam | None = None,
# Action control inputs (Matrix-Game)
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
grid_sizes: tuple[int, int, int] | list[int] | torch.Tensor
| None = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]]:
"""
@@ -131,6 +136,15 @@ class VideoGenerator:
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
self.fastvideo_args.model_path)
# Add action control inputs to kwargs if provided
if mouse_cond is not None:
kwargs['mouse_cond'] = mouse_cond
if keyboard_cond is not None:
kwargs['keyboard_cond'] = keyboard_cond
if grid_sizes is not None:
kwargs['grid_sizes'] = grid_sizes
sampling_param.update(kwargs)
if self.fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
@@ -297,27 +311,19 @@ class VideoGenerator:
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
# Adjust latent frames to be divisible by number of GPUs
if sampling_param.num_frames_round_down:
# Ensure we have at least 1 batch per GPU
new_latent_num_frames = max(
1, (orig_latent_num_frames // num_gpus)) * num_gpus
else:
new_latent_num_frames = math.ceil(
orig_latent_num_frames / num_gpus) * num_gpus
if use_temporal_scaling_frames:
# Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor
new_num_frames = (new_latent_num_frames -
new_num_frames = (orig_latent_num_frames -
1) * temporal_scale_factor + 1
else: # stepvideo only
# Find the least common multiple of 3 and num_gpus
divisor = math.lcm(3, num_gpus)
# Round up to the nearest multiple of this LCM
new_latent_num_frames = (
(new_latent_num_frames + divisor - 1) // divisor) * divisor
orig_latent_num_frames = (
(orig_latent_num_frames + divisor - 1) // divisor) * divisor
# Convert back to actual frames using the StepVideo formula
new_num_frames = new_latent_num_frames // 3 * 17
new_num_frames = orig_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
+15
View File
@@ -32,6 +32,9 @@ if TYPE_CHECKING:
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
FASTVIDEO_TORCH_PROFILER_WAIT_STEPS: int = 2
FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS: int = 1
FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS: int = 2
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
FASTVIDEO_SERVER_DEV_MODE: bool = False
FASTVIDEO_STAGE_LOGGING: bool = False
@@ -247,6 +250,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
# not profile flops.
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
# Wait steps per profiling cycle (torch.profiler.schedule wait parameter)
# Defaults to 2 if not set.
"FASTVIDEO_TORCH_PROFILER_WAIT_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WAIT_STEPS", "2")),
# Warmup steps per profiling cycle (torch.profiler.schedule warmup parameter)
# Defaults to 1 if not set.
"FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS", "1")),
# Active steps per profiling cycle (torch.profiler.schedule active parameter)
# Defaults to 2 if not set.
"FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS", "2")),
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
+64
View File
@@ -174,6 +174,8 @@ class FastVideoArgs:
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
override_pipeline_cls_name: str | None = None
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
@@ -317,6 +319,62 @@ class FastVideoArgs:
"Path to a text file containing prompts (one per line) for batch processing",
)
# LoRA parameters (inference-time adapter loading)
parser.add_argument(
"--lora-path",
type=str,
default=FastVideoArgs.lora_path,
help=
"Path to a LoRA adapter (directory or HF repo id). If set, LoRA will be applied at inference.",
)
parser.add_argument(
"--lora-nickname",
type=str,
default=FastVideoArgs.lora_nickname,
help=
"Nickname to refer to the loaded LoRA adapter (useful for swapping).",
)
parser.add_argument(
"--lora-target-modules",
nargs="+",
type=str,
default=FastVideoArgs.lora_target_modules,
help=
"Optional list of module name substrings to restrict LoRA injection (e.g. q_proj k_proj v_proj).",
)
# BSA runtime control (LongCat)
parser.add_argument(
"--enable-bsa",
action=StoreBoolean,
help=
"Enable Block Sparse Attention (BSA) at runtime (overrides config).",
)
parser.add_argument(
"--bsa-sparsity",
type=float,
help="BSA sparsity (e.g., 0.9375).",
)
parser.add_argument(
"--bsa-cdf-threshold",
type=float,
help="BSA CDF threshold (optional).",
)
parser.add_argument(
"--bsa-chunk-q",
nargs=3,
type=int,
metavar=("T", "H", "W"),
help="BSA chunk_3d_shape_q as three ints, e.g., 4 4 4.",
)
parser.add_argument(
"--bsa-chunk-k",
nargs=3,
type=int,
metavar=("T", "H", "W"),
help="BSA chunk_3d_shape_k as three ints, e.g., 4 4 4.",
)
# STA (Sliding Tile Attention) parameters
parser.add_argument(
"--STA-mode",
@@ -424,6 +482,12 @@ class FastVideoArgs:
default=FastVideoArgs.override_transformer_cls_name,
help="Override transformer cls name",
)
parser.add_argument(
"--override-pipeline-cls-name",
type=str,
default=FastVideoArgs.override_pipeline_cls_name,
help="Override pipeline cls name",
)
parser.add_argument(
"--init-weights-from-safetensors",
type=str,
+1
View File
@@ -87,6 +87,7 @@ _ACTIVATION_REGISTRY = {
"gelu_pytorch_tanh": lambda: nn.GELU(approximate="tanh"),
"relu": nn.ReLU,
"silu": nn.SiLU,
"swish": nn.SiLU,
"quick_gelu": QuickGELU,
}
+9 -3
View File
@@ -429,6 +429,7 @@ def get_rotary_pos_embed(
theta_rescale_factor=1.0,
interpolation_factor=1.0,
shard_dim: int = 0,
do_sp_sharding: bool = False,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
@@ -444,6 +445,7 @@ def get_rotary_pos_embed(
theta_rescale_factor: Rescale factor for theta. Defaults to 1.0
interpolation_factor: Factor to scale positions. Defaults to 1.0
shard_dim: Which dimension to shard for sequence parallelism. Defaults to 0.
do_sp_sharding: Whether to shard the positional embeddings for sequence parallelism. Defaults to False.
Returns:
Tuple of (cos, sin) tensors for rotary embeddings
@@ -460,9 +462,13 @@ def get_rotary_pos_embed(
) == head_dim, "sum(rope_dim_list) should equal to head_dim of attention layer"
# Get SP info
sp_group = get_sp_group()
sp_rank = sp_group.rank_in_group
sp_world_size = sp_group.world_size
if do_sp_sharding:
sp_group = get_sp_group()
sp_rank = sp_group.rank_in_group
sp_world_size = sp_group.world_size
else:
sp_rank = 0
sp_world_size = 1
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
+209
View File
@@ -0,0 +1,209 @@
# SPDX-License-Identifier: Apache-2.0
"""
3D Rotary Position Embedding (RoPE) for video transformers.
Reference: https://arxiv.org/pdf/2104.09864.pdf
"""
import torch
import torch.nn as nn
from einops import rearrange, repeat
def broadcast(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor:
"""
Broadcast and concatenate tensors along a dimension.
"""
num_tensors = len(tensors)
shape_lens = set(len(t.shape) for t in tensors)
assert len(
shape_lens) == 1, "tensors must all have the same number of dimensions"
shape_len = list(shape_lens)[0]
dim = (dim + shape_len) if dim < 0 else dim
dims = list(zip(*[list(t.shape) for t in tensors], strict=False))
expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim]
assert all(
len(set(t[1])) <= 2 for t in
expandable_dims), "invalid dimensions for broadcastable concatenation"
max_dims = [(t[0], max(t[1])) for t in expandable_dims]
expanded_dims = [(t[0], (t[1], ) * num_tensors) for t in max_dims]
expanded_dims.insert(dim, (dim, dims[dim]))
expandable_shapes = list(zip(*[t[1] for t in expanded_dims], strict=False))
tensors = [
t[0].expand(*t[1])
for t in zip(tensors, expandable_shapes, strict=False)
]
return torch.cat(tensors, dim=dim)
def rotate_half(x: torch.Tensor) -> torch.Tensor:
"""
Rotate half the hidden dims of the input.
"""
x = rearrange(x, "... (d r) -> ... d r", r=2)
x1, x2 = x.unbind(dim=-1)
x = torch.stack((-x2, x1), dim=-1)
return rearrange(x, "... d r -> ... (d r)")
class RotaryPositionalEmbedding3D(nn.Module):
"""
3D Rotary Positional Embedding for video transformers.
Splits the head dimension across temporal, height, and width dimensions,
computing separate rotary embeddings for each and concatenating them.
"""
def __init__(
self,
head_dim: int,
base: float = 10000.0,
):
"""
Args:
head_dim: Dimension of each attention head
base: Base value for exponential frequency
"""
super().__init__()
self.head_dim = head_dim
assert self.head_dim % 8 == 0, "head_dim must be a multiple of 8 for 3D RoPE"
self.base = base
# Cache for precomputed frequencies
self.freqs_dict: dict[tuple, torch.Tensor] = {}
def register_grid_size(self, grid_size: tuple[int, int, int]) -> None:
"""
Precompute and register frequencies for a given grid size.
Args:
grid_size: (T, H, W) tuple of grid dimensions
"""
if grid_size not in self.freqs_dict:
self.freqs_dict[grid_size] = self.precompute_freqs_3d(grid_size)
def precompute_freqs_3d(self, grid_size: tuple[int, int,
int]) -> torch.Tensor:
"""
Precompute 3D rotary frequencies.
Args:
grid_size: (num_frames, height, width)
Returns:
freqs: [T*H*W, head_dim] tensor of frequencies
"""
num_frames, height, width = grid_size
# Split head_dim across 3 dimensions
# Temporal gets the remainder to ensure exact division
dim_t = self.head_dim - 4 * (self.head_dim // 6)
dim_h = 2 * (self.head_dim // 6)
dim_w = 2 * (self.head_dim // 6)
# Compute frequency bands for each dimension
freqs_t = 1.0 / (self.base**(
torch.arange(0, dim_t, 2)[:(dim_t // 2)].float() / dim_t))
freqs_h = 1.0 / (self.base**(
torch.arange(0, dim_h, 2)[:(dim_h // 2)].float() / dim_h))
freqs_w = 1.0 / (self.base**(
torch.arange(0, dim_w, 2)[:(dim_w // 2)].float() / dim_w))
# Create position grids
grid_t = torch.arange(num_frames, dtype=torch.float32)
grid_h = torch.arange(height, dtype=torch.float32)
grid_w = torch.arange(width, dtype=torch.float32)
# Compute frequencies for each position
freqs_t = torch.einsum("..., f -> ... f", grid_t, freqs_t)
freqs_h = torch.einsum("..., f -> ... f", grid_h, freqs_h)
freqs_w = torch.einsum("..., f -> ... f", grid_w, freqs_w)
# Duplicate for complex pair representation
freqs_t = repeat(freqs_t, "... n -> ... (n r)", r=2)
freqs_h = repeat(freqs_h, "... n -> ... (n r)", r=2)
freqs_w = repeat(freqs_w, "... n -> ... (n r)", r=2)
# Broadcast and concatenate across all 3 dimensions
freqs = broadcast(
[
freqs_t[:, None, None, :], # [T, 1, 1, dim_t]
freqs_h[None, :, None, :], # [1, H, 1, dim_h]
freqs_w[None, None, :, :], # [1, 1, W, dim_w]
],
dim=-1,
)
# Flatten spatial dimensions: [T, H, W, head_dim] -> [T*H*W, head_dim]
freqs = rearrange(freqs, "T H W D -> (T H W) D")
return freqs
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
grid_size: tuple[int, int, int],
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply 3D rotary positional embedding to queries and keys.
Args:
q: Query tensor [B, num_heads, seq_len, head_dim]
k: Key tensor [B, num_heads, seq_len, head_dim]
grid_size: (T, H, W) tuple of grid dimensions
Returns:
(q_rotated, k_rotated): Rotated query and key tensors
"""
# Register grid size if not cached
if grid_size not in self.freqs_dict:
self.register_grid_size(grid_size)
# Get cached frequencies
freqs_cis = self.freqs_dict[grid_size].to(q.device)
# Cast to float32 for precision
q_, k_ = q.float(), k.float()
freqs_cis = freqs_cis.float()
# Compute cos and sin
cos = freqs_cis.cos()
sin = freqs_cis.sin()
# Reshape for broadcasting: [1, 1, seq_len, head_dim]
cos = rearrange(cos, "n d -> 1 1 n d")
sin = rearrange(sin, "n d -> 1 1 n d")
# Apply rotation
q_ = (q_ * cos) + (rotate_half(q_) * sin)
k_ = (k_ * cos) + (rotate_half(k_) * sin)
# Cast back to original dtype
return q_.type_as(q), k_.type_as(k)
def apply_rotary_emb_3d(
q: torch.Tensor,
k: torch.Tensor,
rope_module: RotaryPositionalEmbedding3D,
grid_size: tuple[int, int, int],
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Convenience function to apply 3D RoPE.
Args:
q: Query tensor [B, num_heads, seq_len, head_dim]
k: Key tensor [B, num_heads, seq_len, head_dim]
rope_module: RotaryPositionalEmbedding3D module
grid_size: (T, H, W) grid dimensions
Returns:
(q_rotated, k_rotated): Rotated tensors
"""
return rope_module(q, k, grid_size)
+8 -19
View File
@@ -23,6 +23,7 @@ from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
from fastvideo.models.dits.base import CachableDiT
from fastvideo.models.utils import modulate
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.distributed.communication_op import sequence_model_parallel_shard, sequence_model_parallel_all_gather
class HunyuanRMSNorm(nn.Module):
@@ -239,14 +240,7 @@ class MMDoubleStreamBlock(nn.Module):
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings
cos, sin = freqs_cis
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
@@ -265,7 +259,7 @@ class MMDoubleStreamBlock(nn.Module):
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# Run distributed attention
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis)
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
@@ -395,18 +389,11 @@ class MMSingleStreamBlock(nn.Module):
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
# Apply rotary embeddings to image parts
cos, sin = freqs_cis
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Run distributed attention
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q,
txt_k, txt_v)
txt_k, txt_v, freqs_cis = freqs_cis)
attn_output = torch.cat((img_attn_output, txt_attn_output),
dim=1).view(batch_size, seq_len, -1)
# Process MLP activation
@@ -593,7 +580,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(tt * get_sp_world_size(), th, tw), self.hidden_size,
(tt, th, tw), self.hidden_size,
self.num_attention_heads, self.rope_dim_list, self.rope_theta)
freqs_cos = freqs_cos.to(x.device)
freqs_sin = freqs_sin.to(x.device)
@@ -608,6 +595,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
vec = vec + self.guidance_in(guidance)
# Embed image and text
img = self.img_in(img)
img, _ = sequence_model_parallel_shard(img, dim=1)
txt = self.txt_in(txt, t)
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
@@ -648,6 +636,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
self.maybe_cache_states(img, original_img)
# Final layer processing
img = sequence_model_parallel_all_gather(img, dim=1)
img = self.final_layer(img, vec)
# Unpatchify to get original shape
img = unpatchify(img, tt, th, tw, self.patch_size, self.out_channels)
+853
View File
@@ -0,0 +1,853 @@
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Any, Dict, Optional, List
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.attention import DistributedAttention, LocalAttention
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.configs.models.dits import HunyuanVideo15Config
from fastvideo.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
# TODO(will-PY-refactor): RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
TimestepEmbedder, unpatchify)
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.forward_context import set_forward_context
from fastvideo.attention.backends.abstract import AttentionMetadata
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
logger = init_logger(__name__)
class HunyuanRMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x) -> torch.Tensor:
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
class HunyuanVideo15TimeEmbedding(nn.Module):
r"""
Time embedding for HunyuanVideo 1.5.
Supports standard timestep embedding and optional reference timestep embedding for MeanFlow-based super-resolution
models.
Args:
embedding_dim (`int`):
The dimension of the output embedding.
"""
def __init__(self, embedding_dim: int, use_meanflow: bool = False):
super().__init__()
self.timestep_embedder = TimestepEmbedder(hidden_size=embedding_dim)
self.use_meanflow = use_meanflow
self.time_proj_r = None
self.timestep_embedder_r = None
if use_meanflow:
self.timestep_embedder_r = TimestepEmbedder(hidden_size=embedding_dim)
def forward(
self,
timestep: torch.Tensor,
timestep_r: Optional[torch.Tensor] = None,
) -> torch.Tensor:
timesteps_emb = self.timestep_embedder(timestep)
if timestep_r is not None:
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
timesteps_emb = timesteps_emb + timesteps_emb_r
return timesteps_emb
class HunyuanVideo15ByT5TextProjection(nn.Module):
def __init__(self, in_features: int, hidden_size: int, out_features: int):
super().__init__()
self.norm = nn.LayerNorm(in_features)
self.linear_1 = nn.Linear(in_features, hidden_size)
self.linear_2 = nn.Linear(hidden_size, hidden_size)
self.linear_3 = nn.Linear(hidden_size, out_features)
self.act_fn = nn.GELU()
def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm(encoder_hidden_states)
hidden_states = self.linear_1(hidden_states)
hidden_states = self.act_fn(hidden_states)
hidden_states = self.linear_2(hidden_states)
hidden_states = self.act_fn(hidden_states)
hidden_states = self.linear_3(hidden_states)
return hidden_states
class HunyuanVideo15ImageProjection(nn.Module):
def __init__(self, in_channels: int, hidden_size: int):
super().__init__()
self.norm_in = nn.LayerNorm(in_channels)
self.linear_1 = nn.Linear(in_channels, in_channels)
self.act_fn = nn.GELU()
self.linear_2 = nn.Linear(in_channels, hidden_size)
self.norm_out = nn.LayerNorm(hidden_size)
def forward(self, image_embeds: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm_in(image_embeds)
hidden_states = self.linear_1(hidden_states)
hidden_states = self.act_fn(hidden_states)
hidden_states = self.linear_2(hidden_states)
hidden_states = self.norm_out(hidden_states)
return hidden_states
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
encoder_attention_mask: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple,
seq_attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# seq_len = txt_q.shape[1] + img_q.shape[1]
# attention_mask = F.pad(encoder_attention_mask, (seq_len - encoder_attention_mask.shape[1], 0), value=True)
# attention_mask = attention_mask.bool()
# self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1)
# self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)
# attention_mask = (self_attn_mask_1 & self_attn_mask_2).bool()
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis, attention_mask=seq_attention_mask)
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class HunyuanVideo15Transformer3DModel(CachableDiT):
r"""
A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5).
"""
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = HunyuanVideo15Config()._fsdp_shard_conditions
_compile_conditions = HunyuanVideo15Config()._compile_conditions
_supported_attention_backends = HunyuanVideo15Config(
)._supported_attention_backends
param_names_mapping = HunyuanVideo15Config().param_names_mapping
reverse_param_names_mapping = HunyuanVideo15Config(
).reverse_param_names_mapping
lora_param_names_mapping = HunyuanVideo15Config().lora_param_names_mapping
def __init__(
self,
config: HunyuanVideo15Config,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
self.out_channels = config.out_channels or config.in_channels
self.patch_size = (config.patch_size_t, config.patch_size, config.patch_size)
# 1. Latent and condition embedders
self.img_in = PatchEmbed(self.patch_size,
config.in_channels,
self.hidden_size,
prefix=f"{config.prefix}.img_in")
self.image_embedder = HunyuanVideo15ImageProjection(config.image_embed_dim, self.hidden_size)
self.txt_in = SingleTokenRefiner(config.text_embed_dim,
self.hidden_size,
config.num_attention_heads,
depth=config.num_refiner_layers,
dtype=None,
prefix=f"{config.prefix}.txt_in")
self.txt_in_2 = HunyuanVideo15ByT5TextProjection(config.text_embed_2_dim, 2048, self.hidden_size)
self.time_in = HunyuanVideo15TimeEmbedding(self.hidden_size, use_meanflow=config.use_meanflow)
self.cond_type_embed = nn.Embedding(3, self.hidden_size)
# 3. Dual stream transformer blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.num_layers)
]
)
# 5. Output projection
self.final_layer = FinalLayer(self.hidden_size,
self.patch_size,
self.out_channels,
prefix=f"{config.prefix}.final_layer")
self.gradient_checkpointing = False
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: List[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: List[torch.Tensor],
encoder_attention_mask: List[torch.Tensor],
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
):
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width), self.hidden_size,
self.num_attention_heads, self.config.rope_axes_dim, self.config.rope_theta)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
seq_attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
seq_attention_mask = None
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
freqs_cis,
seq_attention_mask
)
else:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
freqs_cis,
seq_attention_mask
)
# Final layer processing
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states
class SingleTokenRefiner(nn.Module):
"""
A token refiner that processes text embeddings with attention to improve
their representation for cross-attention with image features.
"""
def __init__(
self,
in_channels,
hidden_size,
num_attention_heads,
depth=2,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
# Input projection
# self.input_embedder = ReplicatedLinear(
# in_channels,
# hidden_size,
# bias=True,
# params_dtype=dtype,
# prefix=f"{prefix}.input_embedder")
self.input_embedder = nn.Linear(in_channels, hidden_size, bias=True)
# Timestep embedding
self.t_embedder = TimestepEmbedder(hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.t_embedder")
# Context embedding
self.c_embedder = MLP(in_channels,
hidden_size,
hidden_size,
act_type="silu",
dtype=dtype,
prefix=f"{prefix}.c_embedder")
# Refiner blocks
self.refiner_blocks = nn.ModuleList([
IndividualTokenRefinerBlock(
hidden_size,
num_attention_heads,
qkv_bias=qkv_bias,
dtype=dtype,
prefix=f"{prefix}.refiner_blocks.{i}",
) for i in range(depth)
])
def forward(self, x, t, mask=None):
# Get timestep embeddings
timestep_aware_representations = self.t_embedder(t)
# Get context-aware representations
original_dtype = x.dtype
if mask is None:
context_aware_representations = x.mean(dim=1)
else:
mask_float = mask.float().unsqueeze(-1) # [B, L, 1]
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(
context_aware_representations)
c = timestep_aware_representations + context_aware_representations
# Project input
x = self.input_embedder(x)
# Process through refiner blocks
for block in self.refiner_blocks:
x = block(x, c, mask)
return x
class IndividualTokenRefinerBlock(nn.Module):
"""
A transformer block for refining individual tokens with self-attention.
"""
def __init__(
self,
hidden_size,
num_attention_heads,
mlp_ratio=4.0,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Normalization and attention
self.norm1 = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=True,
dtype=dtype)
self.self_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
params_dtype=dtype,
prefix=f"{prefix}.self_attn_qkv")
self.self_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype,
prefix=f"{prefix}.self_attn_proj")
# MLP
self.norm2 = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=True,
dtype=dtype)
self.mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
act_type="silu",
dtype=dtype,
prefix=f"{prefix}.mlp")
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
# Scaled dot product attention
self.attn = LocalAttention(
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA),
)
def forward(self, x, c, mask=None):
if mask is not None:
mask = mask.clone().bool()
mask[:, 0] = True # Prevent attention weights from becoming NaN
# Get modulation parameters
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=-1)
# Self-attention
norm_x = self.norm1(x)
qkv, _ = self.self_attn_qkv(norm_x)
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
# Run scaled dot product attention
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
attn_output = self.attn(q, k, v) # [B, L, H, D]
attn_output = attn_output.reshape(batch_size, seq_len,
-1) # [B, L, H*D]
# Project and apply residual connection with gating
attn_out, _ = self.self_attn_proj(attn_output)
x = x + attn_out * gate_msa.unsqueeze(1)
# MLP
mlp_out = self.mlp(self.norm2(x))
x = x + mlp_out * gate_mlp.unsqueeze(1)
return x
class FinalLayer(nn.Module):
"""
The final layer of DiT that projects features to pixel space.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
# Normalization
self.norm_final = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
+866
View File
@@ -0,0 +1,866 @@
# SPDX-License-Identifier: Apache-2.0
"""
Native LongCat Video DiT implementation using FastVideo conventions.
This is a Phase 2 reimplementation that replaces the third_party wrapper
with native FastVideo layers for better performance and integration.
"""
from typing import Any
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from fastvideo.configs.models.dits import LongCatVideoConfig
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.layernorm import RMSNorm, FP32LayerNorm
from fastvideo.layers.activation import get_act_fn
from fastvideo.layers.rotary_embedding_3d import RotaryPositionalEmbedding3D
from fastvideo.attention.layer import DistributedAttention, LocalAttention
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.third_party.longcat_video.block_sparse_attention.bsa_interface import flash_attn_bsa_3d
# ============================================================================
# Embeddings
# ============================================================================
class PatchEmbed3D(nn.Module):
"""
3D patch embedding using Conv3d.
"""
def __init__(
self,
patch_size: tuple[int, int, int] = (1, 2, 2),
in_channels: int = 16,
embed_dim: int = 4096,
):
super().__init__()
self.patch_size = patch_size
self.in_channels = in_channels
self.embed_dim = embed_dim
self.proj = nn.Conv3d(
in_channels,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=True,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: [B, C, T, H, W]
Returns:
[B, N, C] where N = (T/pt) * (H/ph) * (W/pw)
"""
# Padding if needed
_, _, T, H, W = x.shape
if W % self.patch_size[2] != 0:
x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
if H % self.patch_size[1] != 0:
x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
if T % self.patch_size[0] != 0:
x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - T % self.patch_size[0]))
x = self.proj(x) # [B, C, T', H', W']
x = x.flatten(2).transpose(1, 2) # [B, N, C]
return x
class TimestepEmbedder(nn.Module):
"""
Sinusoidal timestep embedding + MLP projection.
"""
def __init__(
self,
frequency_embedding_size: int = 256,
adaln_tembed_dim: int = 512,
dtype: torch.dtype | None = None,
):
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
# Use FastVideo's ReplicatedLinear
self.linear_1 = ReplicatedLinear(
frequency_embedding_size,
adaln_tembed_dim,
bias=True,
params_dtype=dtype,
)
self.act = nn.SiLU()
self.linear_2 = ReplicatedLinear(
adaln_tembed_dim,
adaln_tembed_dim,
bias=True,
params_dtype=dtype,
)
@staticmethod
def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000) -> torch.Tensor:
"""
Create sinusoidal timestep embeddings.
"""
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half
)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, t: torch.Tensor, latent_shape: tuple | None = None) -> torch.Tensor:
"""
Args:
t: [B] or [B, T] timesteps
latent_shape: (T, H, W) for temporal expansion
Returns:
[B, T, C]
"""
# Sinusoidal embedding in FP32
t_freq = self.timestep_embedding(t.flatten(), self.frequency_embedding_size)
# Cast to model dtype before MLP
# Handle LoRA wrapper if present
linear_layer = self.linear_1.base_layer if hasattr(self.linear_1, 'base_layer') else self.linear_1
target_dtype = linear_layer.weight.dtype
if t_freq.dtype != target_dtype:
t_freq = t_freq.to(target_dtype)
# MLP projection
t_emb, _ = self.linear_1(t_freq)
t_emb = self.act(t_emb)
t_emb, _ = self.linear_2(t_emb)
# Reshape if needed
if latent_shape is not None and len(t.shape) > 1:
B = t.shape[0]
T = latent_shape[0]
t_emb = t_emb.reshape(B, T, -1)
return t_emb
class CaptionEmbedder(nn.Module):
"""
Caption embedding with MLP projection and optional text compaction.
"""
def __init__(
self,
caption_channels: int = 4096,
hidden_size: int = 4096,
text_tokens_zero_pad: bool = True,
dtype: torch.dtype | None = None,
):
super().__init__()
self.text_tokens_zero_pad = text_tokens_zero_pad
# Two-layer MLP using ReplicatedLinear
self.linear_1 = ReplicatedLinear(
caption_channels,
hidden_size,
bias=True,
params_dtype=dtype,
)
self.act = nn.SiLU()
self.linear_2 = ReplicatedLinear(
hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
)
def forward(
self,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Args:
encoder_hidden_states: [B, N_text, C_text] or [B, 1, N_text, C_text]
encoder_attention_mask: [B, N_text] or [B, 1, 1, N_text]
Returns:
y: [B, N_text, C] - standard padded representation (like other models)
"""
# Handle extra dimension from wrapper
if len(encoder_hidden_states.shape) == 4:
encoder_hidden_states = encoder_hidden_states.squeeze(1)
# Project
y, _ = self.linear_1(encoder_hidden_states)
y = self.act(y)
y, _ = self.linear_2(y) # [B, N_text, C]
# Handle attention masking - just zero out padded tokens if requested
if encoder_attention_mask is not None:
# Remove extra dimensions
if len(encoder_attention_mask.shape) == 4:
encoder_attention_mask = encoder_attention_mask.squeeze(1).squeeze(1)
elif len(encoder_attention_mask.shape) == 3:
encoder_attention_mask = encoder_attention_mask.squeeze(1)
# Zero out padded tokens if requested
if self.text_tokens_zero_pad:
y = y * encoder_attention_mask.unsqueeze(-1)
# Return standard format [B, N_text, C] - no compaction!
return y
# ============================================================================
# Attention Modules (Placeholders for now)
# ============================================================================
class LongCatSelfAttention(nn.Module):
"""
Self-attention with 3D RoPE support and optional BSA.
"""
def __init__(
self,
dim: int,
num_heads: int,
config: LongCatVideoConfig,
dtype: torch.dtype | None = None,
):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
# Separate Q/K/V projections (not fused like original)
self.to_q = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
self.to_k = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
self.to_v = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
# Per-head RMS normalization
self.q_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
self.k_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
# Output projection
self.to_out = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
# 3D RoPE
self.rope_3d = RotaryPositionalEmbedding3D(head_dim=self.head_dim)
# BSA configuration
self.enable_bsa = getattr(config, 'enable_bsa', False)
self.bsa_params = getattr(config, 'bsa_params', None)
# FastVideo attention backend (used when BSA is disabled)
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
supported_attention_backends=config._supported_attention_backends,
)
def forward(
self,
x: torch.Tensor, # [B, N, C]
latent_shape: tuple, # (T, H, W)
**kwargs
) -> torch.Tensor:
"""
Forward pass with 3D RoPE and optional BSA.
"""
B, N, C = x.shape
T, H, W = latent_shape
# Project to Q/K/V
q, _ = self.to_q(x)
k, _ = self.to_k(x)
v, _ = self.to_v(x)
# Reshape to heads: [B, N, num_heads, head_dim]
q = q.view(B, N, self.num_heads, self.head_dim)
k = k.view(B, N, self.num_heads, self.head_dim)
v = v.view(B, N, self.num_heads, self.head_dim)
# Per-head RMS normalization
q = self.q_norm(q)
k = self.k_norm(k)
# For RoPE: need [B, num_heads, N, head_dim]
q_rope = q.transpose(1, 2)
k_rope = k.transpose(1, 2)
# Apply 3D RoPE
q_rope, k_rope = self.rope_3d(q_rope, k_rope, grid_size=latent_shape)
# Transpose back: [B, N, num_heads, head_dim] or [B, H, N, D] for BSA
q = q_rope.transpose(1, 2)
k = k_rope.transpose(1, 2)
# === Attention: BSA or standard ===
if self.enable_bsa and T > 1: # Only use BSA for multi-frame videos
# BSA expects [B, H, S, D] format
q_bsa = q.transpose(1, 2).contiguous() # [B, num_heads, N, head_dim]
k_bsa = k.transpose(1, 2).contiguous()
v_bsa = v.transpose(1, 2).contiguous()
# Handle SP split: BSA operates on per-rank spatial dimensions
# Replicate LongCat's cp_split_hw logic exactly
from fastvideo.distributed.parallel_state import get_sp_world_size
sp_size = get_sp_world_size()
if sp_size > 1:
# Calculate optimal 2D split (same as LongCat's get_optimal_split)
factors = []
for i in range(1, int(sp_size**0.5) + 1):
if sp_size % i == 0:
factors.append([i, sp_size // i])
cp_split_hw = min(factors, key=lambda x: abs(x[0] - x[1]))
# Split H and W dimensions by their respective factors
T_bsa, H_bsa, W_bsa = latent_shape
assert H_bsa % cp_split_hw[0] == 0 and W_bsa % cp_split_hw[1] == 0, \
f"H {H_bsa} must be divisible by {cp_split_hw[0]}, W {W_bsa} must be divisible by {cp_split_hw[1]}"
H_bsa = H_bsa // cp_split_hw[0]
W_bsa = W_bsa // cp_split_hw[1]
latent_shape_bsa = (T_bsa, H_bsa, W_bsa)
else:
latent_shape_bsa = latent_shape
# Call BSA with per-rank latent shape
out = flash_attn_bsa_3d(
q_bsa, k_bsa, v_bsa,
latent_shape_q=latent_shape_bsa,
latent_shape_k=latent_shape_bsa,
**self.bsa_params
) # [B, num_heads, N, head_dim]
# Transpose back: [B, N, num_heads, head_dim]
out = out.transpose(1, 2)
else:
# Standard attention: [B, N, num_heads, head_dim]
out, _ = self.attn(q, k, v)
# Reshape and project out
out = out.reshape(B, N, C)
out, _ = self.to_out(out)
return out
class LongCatCrossAttention(nn.Module):
"""
Cross-attention for text conditioning (standard implementation like other models).
"""
def __init__(
self,
dim: int,
num_heads: int,
config: LongCatVideoConfig,
dtype: torch.dtype | None = None,
):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
# Separate Q/K/V projections
self.to_q = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
self.to_k = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
self.to_v = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
# Per-head RMS normalization
self.q_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
self.k_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
# Output projection
self.to_out = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
# Cross-attention uses LocalAttention (FastVideo standard)
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=config.arch_config._supported_attention_backends,
)
def forward(
self,
x: torch.Tensor, # [B, N_img, C]
context: torch.Tensor, # [B, N_text, C]
**kwargs
) -> torch.Tensor:
"""
Forward pass for cross-attention (standard implementation).
Args:
x: Image tokens [B, N_img, C]
context: Text tokens [B, N_text, C] (standard padded format)
"""
B, N_img, C = x.shape
# Project Q, K, V (standard cross-attention like WanVideo/StepVideo/Cosmos)
q, _ = self.to_q(x)
k, _ = self.to_k(context)
v, _ = self.to_v(context)
N_text = context.shape[1]
# Reshape to heads
q = q.view(B, N_img, self.num_heads, self.head_dim)
k = k.view(B, N_text, self.num_heads, self.head_dim)
v = v.view(B, N_text, self.num_heads, self.head_dim)
# Per-head RMS normalization
q = self.q_norm(q)
k = self.k_norm(k)
# Run cross-attention using FastVideo's LocalAttention
# LocalAttention handles different q and k/v sequence lengths automatically
out = self.attn(q, k, v) # [B, N_img, num_heads, head_dim]
# Reshape and project out
out = out.reshape(B, N_img, C)
out, _ = self.to_out(out)
return out
# ============================================================================
# Feed-Forward Network
# ============================================================================
class LongCatSwiGLUFFN(nn.Module):
"""
SwiGLU feed-forward network using FastVideo's ReplicatedLinear.
FFN(x) = down(gate(x) * SiLU(up(x)))
"""
def __init__(
self,
dim: int,
hidden_dim: int,
dtype: torch.dtype | None = None,
):
super().__init__()
# Three projections for SwiGLU (no bias as per original)
self.w1 = ReplicatedLinear(dim, hidden_dim, bias=False, params_dtype=dtype) # gate
self.w3 = ReplicatedLinear(dim, hidden_dim, bias=False, params_dtype=dtype) # up
self.w2 = ReplicatedLinear(hidden_dim, dim, bias=False, params_dtype=dtype) # down
self.act = nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass: SiLU(w1(x)) * w3(x) -> w2 (matching original LongCat)
"""
w1_out, _ = self.w1(x)
w3_out, _ = self.w3(x)
combined = self.act(w1_out) * w3_out
out, _ = self.w2(combined)
return out
# ============================================================================
# Modulation Utilities
# ============================================================================
def modulate_fp32(norm: nn.Module, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
"""
Apply modulation in FP32 for numerical stability (matching original LongCat).
shift and scale should already be FP32 from torch.amp.autocast context.
"""
# Ensure modulation params are FP32 (should be from autocast)
assert shift.dtype == torch.float32 and scale.dtype == torch.float32, \
f"shift and scale must be FP32, got {shift.dtype} and {scale.dtype}"
orig_dtype = x.dtype
# Normalize and modulate in FP32
x_norm = norm(x.to(torch.float32))
x_mod = x_norm * (scale + 1) + shift
return x_mod.to(orig_dtype)
# ============================================================================
# Transformer Block
# ============================================================================
class LongCatTransformerBlock(nn.Module):
"""
Single-stream transformer block with:
- AdaLN modulation (FP32)
- Self-attention
- Cross-attention
- SwiGLU FFN
"""
def __init__(
self,
hidden_size: int,
num_heads: int,
mlp_ratio: int,
adaln_tembed_dim: int,
config: LongCatVideoConfig,
dtype: torch.dtype | None = None,
):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
# AdaLN modulation (6 parameters: scale/shift for attn & ffn, gate for residual)
self.adaln_linear_1 = ReplicatedLinear(
adaln_tembed_dim,
6 * hidden_size,
bias=True,
params_dtype=dtype,
)
self.adaln_act = nn.SiLU()
# Normalization layers (CRITICAL: Use LayerNorm not RMSNorm like original!)
# Original LongCat uses LayerNorm_FP32 with elementwise_affine=False
self.norm_attn = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
self.norm_ffn = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
# Cross-attention norm has elementwise_affine=True (has weight and bias)
self.norm_cross = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=True)
# Self-attention
self.self_attn = LongCatSelfAttention(
dim=hidden_size,
num_heads=num_heads,
config=config,
dtype=dtype,
)
# Cross-attention
self.cross_attn = LongCatCrossAttention(
dim=hidden_size,
num_heads=num_heads,
config=config,
dtype=dtype,
)
# SwiGLU FFN
ffn_hidden_dim = int(hidden_size * mlp_ratio * 2 / 3)
# Round up to nearest multiple of 256
ffn_hidden_dim = 256 * ((ffn_hidden_dim + 255) // 256)
self.ffn = LongCatSwiGLUFFN(
dim=hidden_size,
hidden_dim=ffn_hidden_dim,
dtype=dtype,
)
def forward(
self,
x: torch.Tensor, # [B, N, C]
context: torch.Tensor, # [B, N_text, C]
t: torch.Tensor, # [B, T, C_t]
latent_shape: tuple, # (T, H, W)
**kwargs
) -> torch.Tensor:
"""
Forward pass with AdaLN modulation.
"""
B, N, C = x.shape
T, H, W = latent_shape
x_orig_dtype = x.dtype # Save for later casting
# === AdaLN Modulation (CRITICAL: FP32 for stability like original) ===
# Use autocast to compute modulation params in FP32
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
t_mod = self.adaln_act(t)
mod_params, _ = self.adaln_linear_1(t_mod)
# Ensure FP32 output (needed when LoRA is applied)
if mod_params.dtype != torch.float32:
mod_params = mod_params.float()
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = \
mod_params.unsqueeze(2).chunk(6, dim=-1) # [B, T, 1, C]
# === Self-Attention ===
x_norm = modulate_fp32(self.norm_attn, x.view(B, T, -1, C), shift_msa, scale_msa)
x_norm = x_norm.view(B, N, C)
attn_out = self.self_attn(x_norm, latent_shape=latent_shape)
# Residual with gating (CRITICAL: FP32 like original, then cast back)
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
x = x + (gate_msa * attn_out.view(B, T, -1, C)).view(B, N, C)
x = x.to(x_orig_dtype)
# === Cross-Attention ===
x_norm_cross = self.norm_cross(x)
cross_out = self.cross_attn(x_norm_cross, context)
x = x + cross_out
# === FFN ===
x_norm_ffn = modulate_fp32(self.norm_ffn, x.view(B, T, -1, C), shift_mlp, scale_mlp)
x_norm_ffn = x_norm_ffn.view(B, N, C)
ffn_out = self.ffn(x_norm_ffn)
# Residual with gating (CRITICAL: FP32 like original, then cast back)
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
x = x + (gate_mlp * ffn_out.view(B, T, -1, C)).view(B, N, C)
x = x.to(x_orig_dtype)
return x
# ============================================================================
# Final Layer
# ============================================================================
class FinalLayer(nn.Module):
"""
Final output projection with AdaLN modulation.
"""
def __init__(
self,
hidden_size: int,
out_channels: int,
adaln_tembed_dim: int,
patch_size: tuple[int, int, int],
dtype: torch.dtype | None = None,
):
super().__init__()
# AdaLN for final layer (2 parameters: scale and shift)
self.adaln_linear = ReplicatedLinear(
adaln_tembed_dim,
2 * hidden_size,
bias=True,
params_dtype=dtype,
)
self.adaln_act = nn.SiLU()
# CRITICAL: Use LayerNorm not RMSNorm! (matches original)
self.norm = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
# Output projection
num_patch = patch_size[0] * patch_size[1] * patch_size[2]
self.proj = ReplicatedLinear(
hidden_size,
num_patch * out_channels,
bias=True,
params_dtype=dtype,
)
def forward(
self,
x: torch.Tensor, # [B, N, C]
t: torch.Tensor, # [B, T, C_t]
latent_shape: tuple,
) -> torch.Tensor:
"""
Returns: [B, N, out_channels * patch_size^3]
"""
B, N, C = x.shape
T, _, _ = latent_shape
# AdaLN modulation (FP32 for stability like original)
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
t_mod = self.adaln_act(t)
mod_params, _ = self.adaln_linear(t_mod)
# Ensure FP32 output (needed when LoRA is applied)
if mod_params.dtype != torch.float32:
mod_params = mod_params.float()
shift, scale = mod_params.unsqueeze(2).chunk(2, dim=-1)
# Modulate
x = modulate_fp32(self.norm, x.view(B, T, -1, C), shift, scale)
x = x.reshape(B, N, C)
# Project
x, _ = self.proj(x)
return x
# ============================================================================
# Main Model
# ============================================================================
class LongCatTransformer3DModel(CachableDiT):
"""
Native LongCat Video Transformer using FastVideo layers.
This is a Phase 2 implementation that replaces third_party dependencies.
"""
# FSDP sharding: shard at each transformer block
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and n.split(".")[-1].isdigit(),
]
# torch.compile optimization: compile each transformer block for speedup
_compile_conditions = [
lambda n, m: "blocks" in n and n.split(".")[-1].isdigit(),
]
# Parameter name mapping (for weight conversion)
param_names_mapping = {} # Will be defined in config
reverse_param_names_mapping = {}
lora_param_names_mapping = {}
# Supported attention backends
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
def __init__(self, config: LongCatVideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
# Extract architecture parameters
self.hidden_size = config.hidden_size # 4096
self.num_attention_heads = config.num_attention_heads # 32
self.depth = config.depth # 48
self.mlp_ratio = config.mlp_ratio # 4
self.in_channels = config.in_channels # 16
self.out_channels = config.out_channels # 16
self.num_channels_latents = config.in_channels
self.patch_size = config.patch_size # [1, 2, 2]
# Embeddings
self.patch_embed = PatchEmbed3D(
patch_size=self.patch_size,
in_channels=self.in_channels,
embed_dim=self.hidden_size,
)
self.time_embedder = TimestepEmbedder(
frequency_embedding_size=config.frequency_embedding_size,
adaln_tembed_dim=config.adaln_tembed_dim,
)
self.caption_embedder = CaptionEmbedder(
caption_channels=config.caption_channels,
hidden_size=self.hidden_size,
text_tokens_zero_pad=getattr(config, 'text_tokens_zero_pad', True),
)
# Transformer blocks (48 blocks)
self.blocks = nn.ModuleList([
LongCatTransformerBlock(
hidden_size=self.hidden_size,
num_heads=self.num_attention_heads,
mlp_ratio=self.mlp_ratio,
adaln_tembed_dim=config.adaln_tembed_dim,
config=config,
)
for _ in range(self.depth)
])
# Output projection
self.final_layer = FinalLayer(
hidden_size=self.hidden_size,
out_channels=self.out_channels,
adaln_tembed_dim=config.adaln_tembed_dim,
patch_size=self.patch_size,
)
def enable_bsa(self):
"""Enable BSA for all self-attention layers."""
for block in self.blocks:
block.self_attn.enable_bsa = True
def disable_bsa(self):
"""Disable BSA for all self-attention layers."""
for block in self.blocks:
block.self_attn.enable_bsa = False
def forward(
self,
hidden_states: torch.Tensor, # [B, C, T, H, W]
encoder_hidden_states: torch.Tensor | list[torch.Tensor], # [B, N_text, C_text]
timestep: torch.LongTensor, # [B] or [B, T]
encoder_attention_mask: torch.Tensor | None = None, # [B, N_text]
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance: float | None = None, # Unused, for API compatibility
**kwargs
) -> torch.Tensor:
"""
Forward pass with FastVideo parameter ordering.
NOTE: This follows FastVideo convention:
(hidden_states, encoder_hidden_states, timestep)
"""
B, _, T, H, W = hidden_states.shape
N_t = T // self.patch_size[0]
N_h = H // self.patch_size[1]
N_w = W // self.patch_size[2]
# Handle list of encoder outputs (take first one)
if isinstance(encoder_hidden_states, list):
encoder_hidden_states = encoder_hidden_states[0]
# 1. Patch embedding
x = self.patch_embed(hidden_states) # [B, N, C]
# 2. Timestep embedding
# Expand timestep from [B] to [B, T] if needed
if timestep.ndim == 1:
timestep = timestep.unsqueeze(1).expand(-1, N_t) # [B, T]
t = self.time_embedder(timestep.flatten(), latent_shape=(N_t, N_h, N_w))
if t.ndim == 2:
t = t.reshape(B, N_t, -1) # [B, T, C_t]
# 3. Caption embedding (standard format, no compaction)
context = self.caption_embedder(
encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask
) # [B, N_text, C]
# 4. Transformer blocks
for i, block in enumerate(self.blocks):
x = block(
x, context, t,
latent_shape=(N_t, N_h, N_w)
)
# 5. Output projection
output = self.final_layer(x, t, latent_shape=(N_t, N_h, N_w))
# Reshape to [B, C_out, T, H, W]
output = self.unpatchify(output, N_t, N_h, N_w)
# Cast to float32 for better accuracy (as per original)
output = output.to(torch.float32)
return output
def unpatchify(self, x: torch.Tensor, N_t: int, N_h: int, N_w: int) -> torch.Tensor:
"""
Args:
x: [B, N, C] where C = T_p * H_p * W_p * C_out
Returns:
[B, C_out, T, H, W]
"""
T_p, H_p, W_p = self.patch_size
x = rearrange(
x,
"B (N_t N_h N_w) (T_p H_p W_p C_out) -> B C_out (N_t T_p) (N_h H_p) (N_w W_p)",
N_t=N_t,
N_h=N_h,
N_w=N_w,
T_p=T_p,
H_p=H_p,
W_p=W_p,
C_out=self.out_channels,
)
return x
@@ -0,0 +1,11 @@
from .model import MatrixGameWanModel, MatrixGameTransformerBlock
from .causal_model import CausalMatrixGameWanModel, CausalMatrixGameTransformerBlock
from .action_module import ActionModule
__all__ = [
"MatrixGameWanModel",
"MatrixGameTransformerBlock",
"CausalMatrixGameWanModel",
"CausalMatrixGameTransformerBlock",
"ActionModule",
]
@@ -0,0 +1,567 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from Matrix-Game: https://github.com/SkyworkAI/Matrix-Game/blob/main/Matrix-Game-2/wan/modules/action_module.py
from einops import rearrange
import torch
import torch.nn as nn
import math
from torch.nn.attention.flex_attention import flex_attention
from fastvideo.attention import LocalAttention
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.layernorm import FP32LayerNorm, RMSNorm
from fastvideo.layers.rotary_embedding import (
get_nd_rotary_pos_embed as _fv_get_nd_rotary_pos_embed,
_apply_rotary_emb,
)
from fastvideo.platforms import AttentionBackendEnum
DISABLE_COMPILE = False
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
def _get_nd_rotary_pos_embed_matrixgame(
rope_dim_list,
rope_sizes,
theta: float = 10000.0,
theta_rescale_factor: float = 1.0,
):
cos, sin = _fv_get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
theta=theta,
theta_rescale_factor=theta_rescale_factor,
dtype=torch.float32,
)
# convert from [S, D/2] to [S, D] format
cos = cos.repeat_interleave(2, dim=1)
sin = sin.repeat_interleave(2, dim=1)
return cos, sin
def _apply_rotary_emb_qk(
xq: torch.Tensor,
xk: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor,
start_offset: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
seq_len = xq.shape[1]
# Slice frequencies based on offset
cos = freqs_cos[start_offset:start_offset + seq_len] # [S, D]
sin = freqs_sin[start_offset:start_offset + seq_len] # [S, D]
# Move to device
cos = cos.to(xq.device)
sin = sin.to(xq.device)
# Convert from [S, D] (interleaved) back to [S, D/2]
cos_half = cos[:, ::2] # [S, D/2]
sin_half = sin[:, ::2] # [S, D/2]
# xq/xk are [B, S, H, D], need to reshape for each batch
B, S, H, D = xq.shape
xq_out = _apply_rotary_emb(xq, cos_half, sin_half, is_neox_style=False)
xk_out = _apply_rotary_emb(xk, cos_half, sin_half, is_neox_style=False)
return xq_out, xk_out
class ActionModule(nn.Module):
"""
action module from https://arxiv.org/pdf/2501.08325
"""
def __init__(
self,
mouse_dim_in: int = 2,
keyboard_dim_in: int = 6,
hidden_size: int = 128,
img_hidden_size: int = 1536,
keyboard_hidden_dim: int = 1024,
mouse_hidden_dim: int = 1024,
vae_time_compression_ratio: int = 4,
windows_size: int = 3,
heads_num: int = 16,
patch_size: list | None = None,
qk_norm: bool = True,
qkv_bias: bool = False,
rope_dim_list: list | None = None,
rope_theta = 256,
mouse_qk_dim_list: list | None = None,
enable_mouse = True,
enable_keyboard = True,
local_attn_size = 6,
blocks: list | None = None,
):
super().__init__()
# Initialize mutable defaults
patch_size = patch_size if patch_size is not None else [1, 2, 2]
rope_dim_list = rope_dim_list if rope_dim_list is not None else [8, 28, 28]
mouse_qk_dim_list = mouse_qk_dim_list if mouse_qk_dim_list is not None else [8, 28, 28]
blocks = blocks if blocks is not None else []
self.local_attn_size = local_attn_size
self.enable_mouse = enable_mouse
self.enable_keyboard = enable_keyboard
self.rope_dim_list = rope_dim_list
self.rope_theta = rope_theta
if self.enable_keyboard:
self.keyboard_embed = nn.Sequential(
nn.Linear(keyboard_dim_in, hidden_size, bias=True),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=True)
)
self.mouse_qk_dim_list = mouse_qk_dim_list
self.heads_num = heads_num
if self.enable_mouse:
c = mouse_hidden_dim
self.mouse_mlp = nn.Sequential(
nn.Linear(mouse_dim_in * vae_time_compression_ratio * windows_size + img_hidden_size, c, bias=True),
nn.GELU(approximate="tanh"),
nn.Linear(c, c),
FP32LayerNorm(c, elementwise_affine=True),
)
head_dim = c // heads_num
self.t_qkv = ReplicatedLinear(c, c*3, bias=qkv_bias)
self.img_attn_q_norm = (
RMSNorm(head_dim, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.img_attn_k_norm = (
RMSNorm(head_dim, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.proj_mouse = ReplicatedLinear(c, img_hidden_size, bias=qkv_bias)
if self.enable_keyboard:
head_dim_key = keyboard_hidden_dim // heads_num
self.key_attn_q_norm = (
RMSNorm(head_dim_key, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.key_attn_k_norm = (
RMSNorm(head_dim_key, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.mouse_attn_q = ReplicatedLinear(img_hidden_size, keyboard_hidden_dim, bias=qkv_bias)
self.keyboard_attn_kv = ReplicatedLinear(hidden_size * windows_size * vae_time_compression_ratio, keyboard_hidden_dim * 2, bias=qkv_bias)
self.proj_keyboard = ReplicatedLinear(keyboard_hidden_dim, img_hidden_size, bias=qkv_bias)
self.mouse_attn_layer = LocalAttention(
num_heads=heads_num,
head_size=mouse_hidden_dim // heads_num,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) if self.enable_mouse else None
self.keyboard_attn_layer = LocalAttention(
num_heads=heads_num,
head_size=keyboard_hidden_dim // heads_num,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) if self.enable_keyboard else None
self.vae_time_compression_ratio = vae_time_compression_ratio
self.windows_size = windows_size
self.patch_size = patch_size
# Lazy initialization: freqs will be created on first forward pass
self._freqs_cos = None
self._freqs_sin = None
def patchify(self, x, patch_size):
"""
x : (N C T H W)
"""
pt, ph, pw = self.patch_size
t, h, w = x.shape[2] // pt, x.shape[3] // ph, x.shape[4] // pw
c = x.shape[1]
x = x.reshape(shape=(x.shape[0], c, t , pt, h , ph, w , pw))
x = torch.einsum("nctohpwq->nthwcopq", x)
x = x.reshape(shape=(x.shape[0], t*h*w, c*pt*ph*pw))
return x
def unpatchify(self, x, t, h, w, patch_size):
"""
x: (N, T, patch_size**2 * C)
imgs: (N, H, W, C)
"""
c = x.shape[2] // patch_size #self.unpatchify_channels
pt, ph, pw = self.patch_size
assert t * h * w == x.shape[1]
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
def get_rotary_pos_embed(self, video_length, height, width, head_dim, rope_dim_list = None, start_offset=0):
target_ndim = 3
ndim = 5 - 2
latents_size = [video_length+start_offset, height, width]
if isinstance(self.patch_size, int):
assert all(s % self.patch_size == 0 for s in latents_size), (
f"Latent size(last {ndim} dimensions) should be divisible by patch size({self.patch_size}), "
f"but got {latents_size}."
)
rope_sizes = [s // self.patch_size for s in latents_size]
elif isinstance(self.patch_size, list):
assert all(
s % self.patch_size[idx] == 0
for idx, s in enumerate(latents_size)
), (
f"Latent size(last {ndim} dimensions) should be divisible by patch size({self.patch_size}), "
f"but got {latents_size}."
)
rope_sizes = [
s // self.patch_size[idx] for idx, s in enumerate(latents_size)
]
if len(rope_sizes) != target_ndim:
rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes # time axis
if rope_dim_list is None:
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
# Use Matrix-Game wrapper for FastVideo's function
freqs_cos, freqs_sin = _get_nd_rotary_pos_embed_matrixgame(
rope_dim_list,
rope_sizes,
theta=self.rope_theta,
theta_rescale_factor=1,
)
return freqs_cos[-video_length*rope_sizes[1]*rope_sizes[2]//self.patch_size[0]:], freqs_sin[-video_length*rope_sizes[1]*rope_sizes[2]//self.patch_size[0]:]
def forward(self, x, tt, th, tw, mouse_condition=None, keyboard_condition=None, block_mask_mouse=None, block_mask_keyboard=None, is_causal=False, kv_cache_mouse=None, kv_cache_keyboard=None, start_frame=0, use_rope_keyboard=True, num_frame_per_block=3):
'''
hidden_states: B, tt*th*tw, C
mouse_condition: B, N_frames, C1
keyboard_condition: B, N_frames, C2
'''
assert use_rope_keyboard
B, N_frames, C = keyboard_condition.shape
assert tt*th*tw == x.shape[1]
assert ((N_frames - 1) + self.vae_time_compression_ratio) % self.vae_time_compression_ratio == 0
N_feats = int((N_frames - 1) / self.vae_time_compression_ratio) + 1
# Lazy initialization of freqs on first forward pass
if self._freqs_cos is None or self._freqs_sin is None:
self._freqs_cos, self._freqs_sin = self.get_rotary_pos_embed(
7500, self.patch_size[1], self.patch_size[2], 64,
self.mouse_qk_dim_list, start_offset=0
)
# Defined freqs_cis early so it's available for both mouse and keyboard
freqs_cis = (self._freqs_cos, self._freqs_sin)
assert (N_feats == tt and ((is_causal and kv_cache_mouse is None) or not is_causal)) or ((N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block and is_causal)
if self.enable_mouse and mouse_condition is not None:
hidden_states = rearrange(x, "B (T S) C -> (B S) T C", T=tt, S=th*tw) # 65*272*480 -> 17*(272//16)*(480//16) -> 8670
B, N_frames, C = mouse_condition.shape
else:
hidden_states = x
# padding
pad_t = self.vae_time_compression_ratio * self.windows_size
if self.enable_mouse and mouse_condition is not None:
pad = mouse_condition[:, 0:1, :].expand(-1, pad_t, -1)
mouse_condition = torch.cat([pad, mouse_condition], dim=1)
if is_causal and kv_cache_mouse is not None:
mouse_condition = mouse_condition[:, self.vae_time_compression_ratio*(N_feats - num_frame_per_block - self.windows_size) + pad_t:, :]
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
else:
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
group_mouse = torch.stack(group_mouse, dim = 1)
S = th * tw
group_mouse = group_mouse.unsqueeze(-1).expand(B, num_frame_per_block, pad_t, C, S)
group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(B * S, num_frame_per_block, pad_t * C)
group_mouse = torch.cat([hidden_states, group_mouse], dim = -1)
group_mouse = self.mouse_mlp(group_mouse)
# qkv
mouse_qkv, _ = self.t_qkv(group_mouse)
q, k, v = rearrange(mouse_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num) # BHW F H C
q = self.img_attn_q_norm(q).to(v)
k = self.img_attn_k_norm(k).to(v)
# rope embd
# freqs_cis = (self.freqs_cos, self.freqs_sin)
q, k = _apply_rotary_emb_qk(q, k, freqs_cis[0], freqs_cis[1], start_offset=start_frame)
## TODO: adding cache here
if is_causal:
if kv_cache_mouse is None:
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0 # == 880, f"{q.shape[0]},{k.shape[0]}"
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_k = torch.cat(
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
attn = flex_attention(
query=padded_q.transpose(2, 1), # after: B, HW, F, C
key=padded_k.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask_mouse
)[:, :, :-padded_length].transpose(2, 1)
else:
current_start = start_frame
current_end = current_start + q.shape[1]
assert q.shape[1] == num_frame_per_block
sink_size = 0
max_attention_size = self.local_attn_size
sink_tokens = sink_size * 1
kv_cache_size = kv_cache_mouse["k"].shape[1]
num_new_tokens = q.shape[1]
if (current_end > kv_cache_mouse["global_end_index"].item()) and (
num_new_tokens + kv_cache_mouse["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache_mouse["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_mouse["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache_mouse["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_mouse["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache_mouse["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_mouse["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache_mouse["local_end_index"].item() + current_end - \
kv_cache_mouse["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
else:
local_end_index = kv_cache_mouse["local_end_index"].item() + current_end - kv_cache_mouse["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache_mouse["k"][:, local_start_index:local_end_index] = k
kv_cache_mouse["v"][:, local_start_index:local_end_index] = v
attn = self.mouse_attn_layer(
q,
kv_cache_mouse["k"][:, max(0, local_end_index - max_attention_size):local_end_index],
kv_cache_mouse["v"][:, max(0, local_end_index - max_attention_size):local_end_index],
)
kv_cache_mouse["global_end_index"].fill_(current_end)
kv_cache_mouse["local_end_index"].fill_(local_end_index)
else:
attn = self.mouse_attn_layer(q, k, v)
# Compute cu_squlens and max_seqlen for flash attention
# qk norm
attn = rearrange(attn, '(b S) T h d -> b (T S) (h d)',b=B)
hidden_states = rearrange(x, "(B S) T C -> B (T S) C", B=B)
attn, _ = self.proj_mouse(attn)
hidden_states = hidden_states + attn
if self.enable_keyboard and keyboard_condition is not None:
pad = keyboard_condition[:, 0:1, :].expand(-1, pad_t, -1)
keyboard_condition = torch.cat([pad, keyboard_condition], dim=1)
if is_causal and kv_cache_keyboard is not None:
keyboard_condition = keyboard_condition[:, self.vae_time_compression_ratio*(N_feats - num_frame_per_block - self.windows_size) + pad_t:, :] # keyboard_condition[:, self.vae_time_compression_ratio*(start_frame - self.windows_size) + pad_t:start_frame * self.vae_time_compression_ratio + pad_t,:]
keyboard_condition = self.keyboard_embed(keyboard_condition)
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
else:
keyboard_condition = self.keyboard_embed(keyboard_condition)
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
group_keyboard = torch.stack(group_keyboard, dim = 1) # B F RW C
group_keyboard = group_keyboard.reshape(shape=(group_keyboard.shape[0],group_keyboard.shape[1],-1))
# apply cross attn
mouse_q, _ = self.mouse_attn_q(hidden_states)
keyboard_kv, _ = self.keyboard_attn_kv(group_keyboard)
B, L, HD = mouse_q.shape
D = HD // self.heads_num
q = mouse_q.view(B, L, self.heads_num, D)
B, L, KHD = keyboard_kv.shape
k, v = keyboard_kv.view(B, L, 2, self.heads_num, D).permute(2, 0, 1, 3, 4)
# Compute cu_squlens and max_seqlen for flash attention
# qk norm
q = self.key_attn_q_norm(q).to(v)
k = self.key_attn_k_norm(k).to(v)
S = th * tw
assert S == 880
# position embed
if use_rope_keyboard:
B, TS, H, D = q.shape
T_ = TS // S
q = q.view(B, T_, S, H, D).transpose(1, 2).reshape(B * S, T_, H, D)
q, k = _apply_rotary_emb_qk(q, k, freqs_cis[0], freqs_cis[1], start_offset=start_frame)
k1, k2, k3, k4 = k.shape
k = k.expand(S, k2, k3, k4)
v = v.expand(S, k2, k3, k4)
if is_causal:
if kv_cache_keyboard is None:
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_k = torch.cat(
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
attn = flex_attention(
query=padded_q.transpose(2, 1), # after: B, HW, F, C
key=padded_k.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask_keyboard
)[:, :, :-padded_length].transpose(2, 1)
else:
current_start = start_frame
current_end = current_start + k.shape[1]
assert k.shape[1] == num_frame_per_block
sink_size = 0
max_attention_size = self.local_attn_size
sink_tokens = sink_size * 1
kv_cache_size = kv_cache_keyboard["k"].shape[1]
num_new_tokens = k.shape[1]
if (current_end > kv_cache_keyboard["global_end_index"].item()) and (
num_new_tokens + kv_cache_keyboard["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache_keyboard["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_keyboard["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache_keyboard["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache_keyboard["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - \
kv_cache_keyboard["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
else:
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
assert k.shape[0] == 880 # BS == 1 or the cache should not be saved/ load method should be modified
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k[:1]
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v[:1]
attn = self.keyboard_attn_layer(
q,
kv_cache_keyboard["k"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat(S, 1, 1, 1),
kv_cache_keyboard["v"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat(S, 1, 1, 1),
)
kv_cache_keyboard["global_end_index"].fill_(current_end)
kv_cache_keyboard["local_end_index"].fill_(local_end_index)
else:
attn = self.keyboard_attn_layer(q, k, v)
attn = rearrange(attn, '(B S) T H D -> B (T S) (H D)', S=S)
else:
if is_causal:
if kv_cache_keyboard is None:
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_k = torch.cat(
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
attn = flex_attention(
query=padded_q.transpose(2, 1), # after: B, HW, F, C
key=padded_k.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask_keyboard
)[:, :, :-padded_length].transpose(2, 1)
else:
current_start = start_frame
current_end = current_start + k.shape[1]
assert k.shape[1] == num_frame_per_block
sink_size = 0
max_attention_size = self.local_attn_size
sink_tokens = sink_size * 1
kv_cache_size = kv_cache_keyboard["k"].shape[1]
num_new_tokens = k.shape[1]
if (current_end > kv_cache_keyboard["global_end_index"].item()) and (
num_new_tokens + kv_cache_keyboard["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache_keyboard["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_keyboard["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache_keyboard["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache_keyboard["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - \
kv_cache_keyboard["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
else:
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v
attn = self.keyboard_attn_layer(
q,
kv_cache_keyboard["k"][:, max(0, local_end_index - max_attention_size):local_end_index],
kv_cache_keyboard["v"][:, max(0, local_end_index - max_attention_size):local_end_index],
)
kv_cache_keyboard["global_end_index"].fill_(current_end)
kv_cache_keyboard["local_end_index"].fill_(local_end_index)
else:
attn = self.keyboard_attn_layer(q, k, v)
attn = rearrange(attn, 'B L H D -> B L (H D)')
attn, _ = self.proj_keyboard(attn)
hidden_states = hidden_states + attn
return hidden_states
File diff suppressed because it is too large Load Diff
+469
View File
@@ -0,0 +1,469 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.attention import DistributedAttention
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import (WanSelfAttention,
WanI2VCrossAttention,
WanT2VCrossAttention,
WanImageEmbedding)
from fastvideo.platforms import AttentionBackendEnum, current_platform
# Import ActionModule
from .action_module import ActionModule
logger = init_logger(__name__)
class MatrixGameTimeImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
image_embed_dim: int | None = None,
):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep, timestep_seq_len)
timestep_proj = self.time_modulation(temb)
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
# and return None for the text embedding part
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
return temb, timestep_proj, None, encoder_hidden_states_image
class MatrixGameCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens=None, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C] - typically 257 image tokens
context_lens(Tensor): Shape [B]
crossattn_cache(dict): Optional cache for k/v during inference
"""
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
# output
x = x.flatten(2)
x, _ = self.to_out(x)
return x
class MatrixGameTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
action_config: dict | None = None):
super().__init__()
action_config = action_config or {}
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
print("QK Norm type not supported")
raise Exception
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2.1. Action Module Integration
self.use_action_module = len(action_config) > 0
if self.use_action_module:
self.action_model = ActionModule(**action_config)
else:
self.action_model = None
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
# Action Module specific args
grid_sizes: torch.Tensor | None = None,
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
orig_dtype = hidden_states.dtype
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query, key, value)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# ================= Action Module =================
if self.action_model is not None:
if mouse_cond is not None or keyboard_cond is not None:
# grid_sizes is expected to be [F, H, W]
# ActionModule implementation takes hidden_states directly
hidden_states = self.action_model(
hidden_states,
grid_sizes[0], grid_sizes[1], grid_sizes[2],
mouse_cond, keyboard_cond,
num_frame_per_block=grid_sizes[0],
)
# =================================================
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
_DEFAULT_MATRIXGAME_CONFIG = MatrixGameWanVideoConfig()
class MatrixGameWanModel(BaseDiT):
# Marker for action input support (Matrix-Game)
supports_action_input = True
_fsdp_shard_conditions = _DEFAULT_MATRIXGAME_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_MATRIXGAME_CONFIG._compile_conditions
_supported_attention_backends = _DEFAULT_MATRIXGAME_CONFIG._supported_attention_backends
param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.param_names_mapping
reverse_param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.reverse_param_names_mapping
lora_param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.lora_param_names_mapping
def __init__(self,
config: MatrixGameWanVideoConfig,
hf_config: dict[str, Any],
**kwargs) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.patch_size = config.patch_size
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings
self.condition_embedder = MatrixGameTimeImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
image_embed_dim=config.image_dim,
)
# 2.1. Get action config
self.action_config = getattr(config, 'action_config', {})
# 3. Transformer blocks
self.blocks = nn.ModuleList([
MatrixGameTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{getattr(config, 'prefix', 'Wan')}.blocks.{i}",
action_config=self.action_config)
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor
| list[torch.Tensor] | None = None,
# Action inputs
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
timestep = timestep.flatten()
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states is not None:
if isinstance(encoder_hidden_states, list):
encoder_hidden_states = encoder_hidden_states[0]
elif encoder_hidden_states.ndim == 2:
encoder_hidden_states = encoder_hidden_states.unsqueeze(0)
else:
# encoder_hidden_states is None (e.g. no text encoder)
# MatrixGame uses image-action cross-attn.
pass
if encoder_hidden_states_image is not None:
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
else:
encoder_hidden_states = encoder_hidden_states_image
# This is [F, H, W] for the ActionModule
grid_sizes = torch.tensor([
post_patch_num_frames, post_patch_height, post_patch_width
],
device=hidden_states.device)
# Blocks
for block in self.blocks:
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj,
freqs_cis,
grid_sizes=grid_sizes,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond)
else:
hidden_states = block(
hidden_states,
encoder_hidden_states,
timestep_proj,
freqs_cis,
grid_sizes=grid_sizes,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond)
# Output
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
+307
View File
@@ -0,0 +1,307 @@
from __future__ import annotations
import os
import random
# import cv2
import numpy as np
import torch
from diffusers.utils import export_to_video
from PIL import Image
from fastvideo.utils import logger
CAM_VALUE = 0.1
CAMERA_MAP = {
"i": [CAM_VALUE, 0], "k": [-CAM_VALUE, 0],
"j": [0, -CAM_VALUE], "l": [0, CAM_VALUE], "u": [0, 0]
}
KEYBOARD_MAP_4 = { # base_distilled_model (universal): W/S/A/D
"w": [1, 0, 0, 0], "s": [0, 1, 0, 0],
"a": [0, 0, 1, 0], "d": [0, 0, 0, 1], "q": [0, 0, 0, 0]
}
KEYBOARD_MAP_2 = { # gta_distilled_model: W/S only (steering via mouse)
"w": [1, 0], "s": [0, 1], "q": [0, 0]
}
KEYBOARD_MAP_7 = { # templerun_distilled_model: still/w/s/left/right/a/d
"q": [1, 0, 0, 0, 0, 0, 0], # still
"w": [0, 1, 0, 0, 0, 0, 0], # forward
"s": [0, 0, 1, 0, 0, 0, 0], # back
"j": [0, 0, 0, 1, 0, 0, 0], # left (swipe)
"l": [0, 0, 0, 0, 1, 0, 0], # right (swipe)
"a": [0, 0, 0, 0, 0, 1, 0], # a
"d": [0, 0, 0, 0, 0, 0, 1], # d
}
KEYBOARD_MAP = KEYBOARD_MAP_4 # Default for backward compatibility
def load_initial_image(image_path: str = None) -> Image.Image:
if image_path and os.path.exists(image_path):
return Image.open(image_path).convert("RGB")
logger.warning("No image provided, creating placeholder...")
return Image.new("RGB", (640, 352), (128, 128, 128))
def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = None):
if keyboard_dim not in (2, 4, 7):
raise ValueError(f"keyboard_dim must be 2, 4, or 7, got {keyboard_dim}")
if num_frames % 4 != 1:
raise ValueError("Matrix-Game conditioning expects num_frames to be 4k+1.")
# Set seed for reproducibility if provided
if seed is not None:
random.seed(seed)
num_samples_per_action = 4
# Define actions based on keyboard_dim
if keyboard_dim == 4:
# Universal model: W, S, A, D
actions_single_action = ["forward", "left", "right"]
actions_double_action = ["forward_left", "forward_right"]
actions_single_camera = ["camera_l", "camera_r"]
keyboard_idx = {"forward": 0, "back": 1, "left": 2, "right": 3}
elif keyboard_dim == 2:
# GTA model: W, S only (steering via mouse)
actions_single_action = ["forward", "back"]
actions_double_action = []
actions_single_camera = ["camera_l", "camera_r"]
keyboard_idx = {"forward": 0, "back": 1}
else: # keyboard_dim == 7
# Temple Run model: still, w, s, left, right, a, d (no mouse)
actions_single_action = ["forward", "back", "left", "right"]
actions_double_action = []
actions_single_camera = [] # No mouse for Temple Run
keyboard_idx = {"still": 0, "forward": 1, "back": 2, "left": 3, "right": 4, "a": 5, "d": 6}
actions_to_test = (
actions_double_action * 5 + actions_single_camera * 5 + actions_single_action * 5
)
for action in (actions_single_action + actions_double_action):
for camera in actions_single_camera:
actions_to_test.append(f"{action}_{camera}")
# Ensure we have at least some actions
if not actions_to_test:
actions_to_test = actions_single_action * 5
base_action = actions_single_action + actions_single_camera
cam_value = 0.1
camera_value_map = {
"camera_up": [cam_value, 0],
"camera_down": [-cam_value, 0],
"camera_l": [0, -cam_value],
"camera_r": [0, cam_value],
"camera_ur": [cam_value, cam_value],
"camera_ul": [cam_value, -cam_value],
"camera_dr": [-cam_value, cam_value],
"camera_dl": [-cam_value, -cam_value],
}
data = []
for action_name in actions_to_test:
keyboard_condition = torch.zeros((num_samples_per_action, keyboard_dim))
mouse_condition = torch.zeros((num_samples_per_action, 2))
for sub_act in base_action:
if sub_act not in action_name:
continue
if sub_act in camera_value_map:
mouse_condition = torch.tensor(
[camera_value_map[sub_act] for _ in range(num_samples_per_action)],
dtype=mouse_condition.dtype,
)
elif sub_act in keyboard_idx:
keyboard_condition[:, keyboard_idx[sub_act]] = 1
data.append({
"keyboard_condition": keyboard_condition,
"mouse_condition": mouse_condition,
})
keyboard_condition = torch.zeros((num_frames, keyboard_dim))
mouse_condition = torch.zeros((num_frames, 2))
current_frame = 0
selections = [12]
while current_frame < num_frames:
rd_frame = selections[random.randint(0, len(selections) - 1)]
entry = data[random.randint(0, len(data) - 1)]
key_seq = entry["keyboard_condition"]
mouse_seq = entry["mouse_condition"]
if current_frame == 0:
keyboard_condition[:1] = key_seq[:1]
mouse_condition[:1] = mouse_seq[:1]
current_frame = 1
else:
rd_frame = min(rd_frame, num_frames - current_frame)
repeat_time = rd_frame // 4
keyboard_condition[current_frame:current_frame + rd_frame] = key_seq.repeat(repeat_time, 1)
mouse_condition[current_frame:current_frame + rd_frame] = mouse_seq.repeat(repeat_time, 1)
current_frame += rd_frame
return {"keyboard": keyboard_condition, "mouse": mouse_condition}
def parse_config(config, mode="universal"):
assert mode in ['universal', 'gta_drive', 'templerun']
key_data = {}
mouse_data = {}
if mode != 'templerun':
key, mouse = config
else:
key = config
for i in range(len(key)):
if mode == 'templerun':
still, w, s, left, right, a, d = key[i]
elif mode == 'universal':
w, s, a, d = key[i]
else:
w, s, a, d = key[i][0], key[i][1], mouse[i][1] < 0, mouse[i][1] > 0
if mode == 'universal':
mouse_y, mouse_x = mouse[i]
mouse_y = -1 * mouse_y
key_data[i] = {"W": bool(w), "A": bool(a), "S": bool(s), "D": bool(d)}
if mode == 'templerun':
key_data[i].update({"left": bool(left), "right": bool(right)})
if mode == 'universal':
if i == 0:
mouse_data[i] = (320, 352 // 2)
else:
global_scale_factor = 0.1
mouse_scale_x = 15 * global_scale_factor
mouse_scale_y = 15 * 4 * global_scale_factor
mouse_data[i] = (
mouse_data[i - 1][0] + mouse_x * mouse_scale_x,
mouse_data[i - 1][1] + mouse_y * mouse_scale_y,
)
return key_data, mouse_data
# NOTE: drawing functions are commented out to avoid cv2/libGL dependency.
#
# def draw_rounded_rectangle(image, top_left, bottom_right, color, radius=10, alpha=0.5):
# overlay = image.copy()
# x1, y1 = top_left
# x2, y2 = bottom_right
#
# cv2.rectangle(overlay, (x1 + radius, y1), (x2 - radius, y2), color, -1)
# cv2.rectangle(overlay, (x1, y1 + radius), (x2, y2 - radius), color, -1)
# cv2.ellipse(overlay, (x1 + radius, y1 + radius), (radius, radius), 180, 0, 90, color, -1)
# cv2.ellipse(overlay, (x2 - radius, y1 + radius), (radius, radius), 270, 0, 90, color, -1)
# cv2.ellipse(overlay, (x1 + radius, y2 - radius), (radius, radius), 90, 0, 90, color, -1)
# cv2.ellipse(overlay, (x2 - radius, y2 - radius), (radius, radius), 0, 0, 90, color, -1)
# cv2.addWeighted(overlay, alpha, image, 1 - alpha, 0, image)
#
#
# def draw_keys_on_frame(frame, keys, key_size=(80, 50), spacing=20, bottom_margin=30, mode='universal'):
# h, w, _ = frame.shape
# horison_shift = 90
# vertical_shift = -20
# horizon_shift_all = 50
# key_positions = {
# "W": (w // 2 - key_size[0] // 2 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] * 2 + vertical_shift - 20),
# "A": (w // 2 - key_size[0] * 2 + 5 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# "S": (w // 2 - key_size[0] // 2 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# "D": (w // 2 + key_size[0] - 5 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# }
# key_icon = {"W": "W", "A": "A", "S": "S", "D": "D", "left": "left", "right": "right"}
# if mode == 'templerun':
# key_positions.update({
# "left": (w // 2 + key_size[0] * 2 + spacing * 2 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# "right": (w // 2 + key_size[0] * 3 + spacing * 7 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift)
# })
#
# for key, (x, y) in key_positions.items():
# is_pressed = keys.get(key, False)
# top_left = (x, y)
# if key in ["left", "right"]:
# bottom_right = (x + key_size[0] + 40, y + key_size[1])
# else:
# bottom_right = (x + key_size[0], y + key_size[1])
#
# color = (0, 255, 0) if is_pressed else (200, 200, 200)
# alpha = 0.8 if is_pressed else 0.5
# draw_rounded_rectangle(frame, top_left, bottom_right, color, radius=10, alpha=alpha)
#
# text_size = cv2.getTextSize(key, cv2.FONT_HERSHEY_SIMPLEX, 0.8, 2)[0]
# if key in ["left", "right"]:
# text_x = x + (key_size[0] + 40 - text_size[0]) // 2
# else:
# text_x = x + (key_size[0] - text_size[0]) // 2
# text_y = y + (key_size[1] + text_size[1]) // 2
# cv2.putText(frame, key_icon[key], (text_x, text_y), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 0, 0), 2)
#
#
# def overlay_icon(frame, icon, position, scale=1.0, rotation=0):
# x, y = position
# h, w, _ = icon.shape
#
# scaled_width = int(w * scale)
# scaled_height = int(h * scale)
# icon_resized = cv2.resize(icon, (scaled_width, scaled_height), interpolation=cv2.INTER_AREA)
#
# center = (scaled_width // 2, scaled_height // 2)
# rotation_matrix = cv2.getRotationMatrix2D(center, rotation, 1.0)
# icon_rotated = cv2.warpAffine(
# icon_resized, rotation_matrix, (scaled_width, scaled_height),
# flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0, 0)
# )
#
# h, w, _ = icon_rotated.shape
# frame_h, frame_w, _ = frame.shape
#
# top_left_x = max(0, int(x - w // 2))
# top_left_y = max(0, int(y - h // 2))
# bottom_right_x = min(frame_w, int(x + w // 2))
# bottom_right_y = min(frame_h, int(y + h // 2))
#
# icon_x_start = max(0, int(-x + w // 2))
# icon_y_start = max(0, int(-y + h // 2))
# icon_x_end = icon_x_start + (bottom_right_x - top_left_x)
# icon_y_end = icon_y_start + (bottom_right_y - top_left_y)
#
# icon_region = icon_rotated[icon_y_start:icon_y_end, icon_x_start:icon_x_end]
# alpha = icon_region[:, :, 3] / 255.0
# icon_rgb = icon_region[:, :, :3]
#
# frame_region = frame[top_left_y:bottom_right_y, top_left_x:bottom_right_x]
# for c in range(3):
# frame_region[:, :, c] = (1 - alpha) * frame_region[:, :, c] + alpha * icon_rgb[:, :, c]
# frame[top_left_y:bottom_right_y, top_left_x:bottom_right_x] = frame_region
#
#
# def process_video(input_video, output_video, config, mouse_icon_path,
# mouse_scale=1.0, mouse_rotation=0, process_icon=True, mode='universal'):
# key_data, mouse_data = parse_config(config, mode=mode)
# fps = 12
#
# mouse_icon = cv2.imread(mouse_icon_path, cv2.IMREAD_UNCHANGED)
#
# out_video = []
# for frame_idx, frame in enumerate(input_video):
# frame = np.ascontiguousarray(frame)
# if process_icon:
# keys = key_data.get(frame_idx, {"W": False, "A": False, "S": False, "D": False, "left": False, "right": False})
# draw_keys_on_frame(frame, keys, key_size=(50, 50), spacing=10, bottom_margin=20, mode=mode)
# if mode == 'universal':
# frame_width = frame.shape[1]
# frame_height = frame.shape[0]
# mouse_position = mouse_data.get(frame_idx, (frame_width // 2, frame_height // 2))
# overlay_icon(frame, mouse_icon, mouse_position, scale=mouse_scale, rotation=mouse_rotation)
# out_video.append(frame / 255)
#
# export_to_video(out_video, output_video, fps=fps)
# logger.info(f"Video saved to {output_video}")
+53 -24
View File
@@ -12,7 +12,10 @@ from fastvideo.attention import (DistributedAttention, DistributedAttention_VSA,
LocalAttention)
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.sample.wan import WanTeaCacheParams
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather,
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
@@ -21,14 +24,16 @@ from fastvideo.layers.linear import ReplicatedLinear
# from torch.nn import RMSNorm
# TODO: RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
TimestepEmbedder)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
logger = init_logger(__name__)
@@ -314,6 +319,7 @@ class WanTransformerBlock(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
@@ -356,13 +362,7 @@ class WanTransformerBlock(nn.Module):
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query, key, value)
attn_output, _ = self.attn1(query, key, value, freqs_cis=freqs_cis, attention_mask=attention_mask)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
@@ -474,6 +474,7 @@ class WanTransformerBlock_VSA(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
@@ -504,16 +505,12 @@ class WanTransformerBlock_VSA(nn.Module):
gate_compress = gate_compress.squeeze(1).unflatten(
2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query,
key,
value,
gate_compress=gate_compress)
freqs_cis = freqs_cis,
gate_compress=gate_compress,
attention_mask=attention_mask)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
@@ -563,6 +560,8 @@ class WanTransformer3DModel(CachableDiT):
self.patch_size = config.patch_size
self.text_len = config.text_len
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
@@ -606,6 +605,7 @@ class WanTransformer3DModel(CachableDiT):
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
self._logged_attention_mask = False
# For type checking
self.previous_e0_even = None
@@ -650,21 +650,43 @@ class WanTransformer3DModel(CachableDiT):
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
(post_patch_num_frames, post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
freqs_sin.to(hidden_states.device).float())
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# Shard with padding support - returns (sharded_tensor, original_seq_len)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
# Create attention mask for padded tokens if padding was applied
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
if not self._logged_attention_mask:
logger.info(f"Padding applied, original seq len: {original_seq_len}, padded seq len: {padded_seq_len}")
self._logged_attention_mask = True
attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
if not self._logged_attention_mask:
logger.info(f"Padding not applied")
self._logged_attention_mask = True
attention_mask = None
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
@@ -698,6 +720,7 @@ class WanTransformer3DModel(CachableDiT):
timestep_proj=timestep_proj, temb=temb)
if should_skip_forward:
print("skipping forward, cached")
hidden_states = self.retrieve_cached_states(hidden_states)
else:
# if teacache is enabled, we need to cache the original hidden states
@@ -708,12 +731,13 @@ class WanTransformer3DModel(CachableDiT):
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
timestep_proj, freqs_cis, attention_mask)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
timestep_proj, freqs_cis, attention_mask)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
@@ -726,7 +750,12 @@ class WanTransformer3DModel(CachableDiT):
# batch_size, inner_dim
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
# Gather and unpad in one operation
hidden_states = sequence_model_parallel_all_gather_with_unpad(
hidden_states, original_seq_len, dim=1)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
+387
View File
@@ -0,0 +1,387 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from transformers: https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py
import math
from typing import Any, Optional, Tuple, Union, List, Callable, Iterable
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.encoders import BaseEncoderOutput, Qwen2_5_VLConfig
from fastvideo.distributed import get_tp_rank, get_tp_world_size
from fastvideo.layers.activation import get_act_fn, SiluAndMul
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import MergedColumnParallelLinear, QKVParallelLinear, RowParallelLinear
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.models.mask_utils import sdpa_mask
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
"""
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
"""
batch, num_key_value_heads, slen, head_dim = hidden_states.shape
if n_rep == 1:
return hidden_states
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
def sdpa_attention_forward(
module: torch.nn.Module,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attention_mask: Optional[torch.Tensor],
dropout: float = 0.0,
scaling: Optional[float] = None,
is_causal: Optional[bool] = None,
**kwargs,
) -> tuple[torch.Tensor, None]:
if kwargs.get("output_attentions", False) or kwargs.get("head_mask") is not None:
logger.warning_once(
"`sdpa` attention does not support `output_attentions=True` or `head_mask`."
" Please set your attention to `eager` if you want any of these features."
)
if hasattr(module, "num_key_value_groups"):
key = repeat_kv(key, module.num_key_value_groups)
value = repeat_kv(value, module.num_key_value_groups)
if attention_mask is not None and attention_mask.ndim == 4:
attention_mask = attention_mask[:, :, :, : key.shape[-2]]
# If attention_mask is not None, convert it to boolean type
if attention_mask is not None and attention_mask.dtype != torch.bool:
attention_mask = attention_mask.bool()
attn_output = torch.nn.functional.scaled_dot_product_attention(
query,
key,
value,
attn_mask=attention_mask,
dropout_p=dropout,
scale=scaling,
is_causal=is_causal,
)
attn_output = attn_output.transpose(1, 2).contiguous()
return attn_output, None
def rotate_half(x):
"""Rotates half the hidden dims of the input."""
x1 = x[..., : x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim=-1)
def apply_multimodal_rotary_pos_emb(q, k, cos, sin, mrope_section, unsqueeze_dim=1):
mrope_section = [s * 2 for s in mrope_section]
cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
unsqueeze_dim
)
sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
unsqueeze_dim
)
q_embed = (q * cos) + (rotate_half(q) * sin)
k_embed = (k * cos) + (rotate_half(k) * sin)
return q_embed, k_embed
class Qwen2_5_VLRotaryEmbedding(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, device=None):
super().__init__()
self.max_seq_len_cached = config.max_position_embeddings
self.original_max_seq_len = config.max_position_embeddings
self.config = config
self.rope_type = config.rope_scaling.get("rope_type", "default")
self.base = config.rope_theta
# Simplified initialization
head_dim = config.hidden_size // config.num_attention_heads
dim = head_dim
self.attention_scaling = 1.0
inv_freq = 1.0 / (
self.base ** (torch.arange(0, dim, 2, dtype=torch.int64).float().to(device) / dim)
)
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.original_inv_freq = inv_freq
def forward(self, x, position_ids):
# In contrast to other models, Qwen2_5_VL has different position ids for the grids
# So we expand the inv_freq to shape (3, ...)
inv_freq_expanded = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1)
position_ids_expanded = position_ids[:, :, None, :].float() # shape (3, bs, 1, positions)
device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
with torch.autocast(device_type=device_type, enabled=False): # Force float32
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3)
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos() * self.attention_scaling
sin = emb.sin() * self.attention_scaling
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
class Qwen2_5_VLMLP(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, quant_config: QuantizationConfig | None = None, prefix: str = ""):
super().__init__()
self.gate_up_proj = MergedColumnParallelLinear(
input_size=config.hidden_size,
output_sizes=[config.intermediate_size] * 2,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.gate_up_proj",
)
self.down_proj = RowParallelLinear(
input_size=config.intermediate_size,
output_size=config.hidden_size,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.down_proj",
)
self.act_fn = SiluAndMul()
def forward(self, x):
x, _ = self.gate_up_proj(x)
x = self.act_fn(x)
x, _ = self.down_proj(x)
return x
class Qwen2_5_VLAttention(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int, quant_config: QuantizationConfig | None = None, prefix: str = ""):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = self.hidden_size // self.num_heads
self.num_key_value_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
tp_size = get_tp_world_size()
self.total_num_heads = self.num_heads
assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size
self.total_num_kv_heads = self.num_key_value_heads
if self.total_num_kv_heads >= tp_size:
assert self.total_num_kv_heads % tp_size == 0
self.num_kv_heads = self.total_num_kv_heads // tp_size
else:
assert tp_size % self.total_num_kv_heads == 0
self.num_kv_heads = 1
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
self.q_size = self.num_heads * self.head_dim
self.kv_size = self.num_kv_heads * self.head_dim
self.scaling = self.head_dim**-0.5
self.qkv_proj = QKVParallelLinear(
hidden_size=self.hidden_size,
head_size=self.head_dim,
total_num_heads=self.total_num_heads,
total_num_kv_heads=self.total_num_kv_heads,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.o_proj = RowParallelLinear(
input_size=self.total_num_heads * self.head_dim,
output_size=self.hidden_size,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.o_proj",
)
self.layer_type = config.layer_types[layer_idx] if config.layer_types else "full_attention"
self.sliding_window = config.sliding_window if self.layer_type == "sliding_attention" else None
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
output_attentions: bool = False,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
bsz, q_len, _ = hidden_states.size()
qkv, _ = self.qkv_proj(hidden_states)
query_states, key_states, value_states = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
key_states = key_states.view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
value_states = value_states.view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
cos, sin = position_embeddings
query_states, key_states = apply_multimodal_rotary_pos_emb(
query_states, key_states, cos, sin, self.config.rope_scaling["mrope_section"]
)
attn_output = sdpa_attention_forward(self, query_states, key_states, value_states, attention_mask, dropout=self.config.attention_dropout, scaling=self.scaling, is_causal=False)[0].reshape(bsz, q_len, -1)
attn_output, _ = self.o_proj(attn_output)
return attn_output
class Qwen2_5_VLDecoderLayer(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int, quant_config: QuantizationConfig | None = None, prefix: str = ""):
super().__init__()
self.self_attn = Qwen2_5_VLAttention(config, layer_idx, quant_config=quant_config, prefix=f"{prefix}.self_attn")
self.mlp = Qwen2_5_VLMLP(config, quant_config=quant_config, prefix=f"{prefix}.mlp")
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
output_attentions: Optional[bool] = False,
) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.self_attn(
hidden_states=hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
position_embeddings=position_embeddings,
output_attentions=output_attentions,
)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
outputs = (hidden_states,)
return outputs
class Qwen2_5_VLTextModel(TextEncoder):
def __init__(self, config: Qwen2_5_VLConfig):
super().__init__(config)
quant_config = None
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size
)
self.layers = nn.ModuleList([
Qwen2_5_VLDecoderLayer(config, layer_idx, quant_config=quant_config, prefix=f"{config.prefix}.layers.{layer_idx}")
for layer_idx in range(config.num_hidden_layers)
])
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config)
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.embed_tokens(input_ids)
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
output_hidden_states: Optional[bool] = None,
**kwargs,
) -> BaseEncoderOutput:
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
if inputs_embeds is None:
inputs_embeds = self.get_input_embeddings(input_ids)
hidden_states = inputs_embeds
if position_ids is None:
seq_length = hidden_states.shape[1]
cache_position = torch.arange(seq_length, device=hidden_states.device)
position_ids = cache_position.view(1, 1, -1).expand(3, hidden_states.shape[0], -1)
mask_kwargs = {
"batch_size": hidden_states.shape[0],
"cache_position": cache_position,
"kv_length": attention_mask.shape[-1],
"kv_offset": 0,
"attention_mask": attention_mask
}
position_embeddings = self.rotary_emb(hidden_states, position_ids)
all_hidden_states = () if output_hidden_states else None
for decoder_layer in self.layers:
if output_hidden_states:
all_hidden_states += (hidden_states,)
layer_outputs = decoder_layer(
hidden_states,
attention_mask=sdpa_mask(**mask_kwargs),
position_ids=position_ids,
position_embeddings=position_embeddings,
output_attentions=False,
)
hidden_states = layer_outputs[0]
hidden_states = self.norm(hidden_states)
if output_hidden_states:
all_hidden_states += (hidden_states,)
return BaseEncoderOutput(
last_hidden_state=hidden_states,
hidden_states=all_hidden_states,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name:
continue
for param_name, weight_name, shard_id in self.config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
# Skip loading extra bias for GPTQ models.
# if name.endswith(".bias") and name not in params_dict:
# continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
# Skip loading extra bias for GPTQ models.
# if name.endswith(".bias") and name not in params_dict:
# continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
+2 -1
View File
@@ -19,7 +19,8 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from transformers.modeling_utils import PretrainedConfig, PreTrainedModel
from transformers import PretrainedConfig
from transformers.modeling_utils import PreTrainedModel
from fastvideo.models.dits.stepvideo import StepVideoRMSNorm
+5 -8
View File
@@ -16,6 +16,7 @@ import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file
from torch.distributed import init_device_mesh
from transformers import AutoImageProcessor, AutoTokenizer
from transformers import UMT5EncoderModel
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.configs.models import EncoderConfig
@@ -405,18 +406,16 @@ class VAELoader(ComponentLoader):
target_device = get_local_torch_device()
with set_default_torch_dtype(PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]):
fastvideo_args.pipeline_config.vae_precision] if fastvideo_args.pipeline_config.vae_precision else torch.bfloat16):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors"))
# TODO(PY)
assert len(
safetensors_list
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
vae.load_state_dict(
loaded, strict=False) # We might only load encoder or decoder
@@ -478,8 +477,6 @@ class TransformerLoader(ComponentLoader):
fastvideo_args.pipeline_config.dit_precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
default_dtype)
assert fastvideo_args.hsdp_shard_dim is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
+201
View File
@@ -0,0 +1,201 @@
import torch
from typing import Callable, Optional
def and_masks(*mask_functions: Callable) -> Callable:
"""Returns a mask function that is the intersection of provided mask functions"""
if not all(callable(arg) for arg in mask_functions):
raise RuntimeError(f"All inputs should be callable mask_functions: {mask_functions}")
def and_mask(batch_idx, head_idx, q_idx, kv_idx):
result = q_idx.new_ones((), dtype=torch.bool)
for mask in mask_functions:
result = result & mask(batch_idx, head_idx, q_idx, kv_idx).to(result.device)
return result
return and_mask
def causal_mask_function(batch_idx: int, head_idx: int, q_idx: int, kv_idx: int) -> bool:
"""
This creates a basic lower-diagonal causal mask.
"""
return kv_idx <= q_idx
def padding_mask_function(padding_mask: torch.Tensor) -> Callable:
"""
This return the mask_function function corresponding to a 2D padding mask.
"""
def inner_mask(batch_idx: int, head_idx: int, q_idx: int, kv_idx: int) -> bool:
# Note that here the mask should ALWAYS be at least of the max `kv_index` size in the dimension 1. This is because
# we cannot pad it here in the mask_function as we don't know the final size, and we cannot try/except, as it is not
# vectorizable on accelerator devices
return padding_mask[batch_idx, kv_idx]
return inner_mask
def prepare_padding_mask(
attention_mask: Optional[torch.Tensor], kv_length: int, kv_offset: int
) -> Optional[torch.Tensor]:
"""
From the 2D attention mask, prepare the correct padding mask to use by potentially padding it.
"""
local_padding_mask = attention_mask
if attention_mask is not None:
# Pad it if necessary
if (padding_length := kv_length + kv_offset - attention_mask.shape[-1]) > 0:
local_padding_mask = torch.nn.functional.pad(attention_mask, (0, padding_length))
return local_padding_mask
def _non_vmap_expansion_sdpa(
batch_indices: torch.Tensor, head_indices: torch.Tensor, q_indices: torch.Tensor, kv_indices: torch.Tensor
):
"""
Used to broadcast our mask_functions over the all 4 dimensions (b_idx, h_idx, q_idx, kv_idx) of the inputs.
Allows the usage of any index-based mask function without relying on vmap.
NOTE: This is limited to index based functions only and is not guaranteed to work otherwise.
Reference:
- https://github.com/huggingface/optimum-onnx/blob/c123e8f4fab61b54a8e0e31ce74462bcacca576e/optimum/exporters/onnx/model_patcher.py#L362-L365
"""
batch_indices = batch_indices[:, None, None, None]
head_indices = head_indices[None, :, None, None]
q_indices = q_indices[None, None, :, None]
kv_indices = kv_indices[None, None, None, :]
return batch_indices, head_indices, q_indices, kv_indices
def sdpa_mask(
batch_size: int,
cache_position: torch.Tensor,
kv_length: int,
kv_offset: int = 0,
mask_function: Callable = causal_mask_function,
attention_mask: Optional[torch.Tensor] = None,
local_size: Optional[int] = None,
allow_is_causal_skip: bool = True,
allow_is_bidirectional_skip: bool = False,
allow_torch_fix: bool = True,
use_vmap: bool = False,
**kwargs,
) -> Optional[torch.Tensor]:
"""
Create a 4D boolean mask of shape `(batch_size, 1, query_length, kv_length)` where a value of True indicates that
the element should take part in the attention computation, and False that it should not.
This function can only be used with torch>=2.5, as the context manager is otherwise not available.
Args:
batch_size (`int`):
The batch size of the input sequence.
cache_position (`torch.Tensor`):
A tensor of shape (query_length,) indicating the current indices of the input sequence elements.
kv_length (`int`):
The size that the key and value states will have during the attention computation.
kv_offset (`int`, optional):
An optional offset to indicate at which first position the key and values states will refer to.
mask_function (`Callable`):
The mask factory function describing the mask pattern.
attention_mask (`torch.Tensor`, optional):
The 2D attention mask corresponding to padded tokens of shape (batch_size, number_of_seen_tokens+q_length)
local_size (`int`, optional):
The size of the local attention, if we do not use full attention. This is used only if `allow_is_causal_skip=True`
to try to skip mask creation if possible.
allow_is_causal_skip (`bool`, optional):
Whether to allow to return `None` for the mask under conditions where we can use the `is_causal` argument in
`torch.sdpa` instead. Default to `True`.
allow_is_bidirectional_skip (`bool`, optional):
Whether to allow to return `None` for the mask under conditions where we do not have to add any bias,
i.e. full attention without any padding. Default to `False`.
allow_torch_fix (`bool`, optional):
Whether to update the mask in case a query is not attending to any tokens, to solve a bug in torch's older
versions. We need an arg to skip it when using eager. By default `True`.
use_vmap (`bool`, optional):
Whether to use `vmap` during the mask construction or not. Allows powerful custom patterns that may not be
index-based (for the cost of speed performance). By default `False`.
## Creating a simple causal mask:
To create the following causal mask:
0 ■ ⬚ ⬚ ⬚ ⬚
1 ■ ■ ⬚ ⬚ ⬚
2 ■ ■ ■ ⬚ ⬚
3 ■ ■ ■ ■ ⬚
4 ■ ■ ■ ■ ■
You can do
```python
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5)
>>> tensor([[[[ True, False, False, False, False],
[ True, True, False, False, False],
[ True, True, True, False, False],
[ True, True, True, True, False],
[ True, True, True, True, True]]]])
```
## Creating a sliding window mask:
To create the following sliding window mask (`sliding_window=3`):
0 ■ ⬚ ⬚ ⬚ ⬚
1 ■ ■ ⬚ ⬚ ⬚
2 ■ ■ ■ ⬚ ⬚
3 ⬚ ■ ■ ■ ⬚
4 ⬚ ⬚ ■ ■ ■
You can do
```python
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5, mask_function=sliding_window_causal_mask_function(3))
>>> tensor([[[[ True, False, False, False, False],
[ True, True, False, False, False],
[ True, True, True, False, False],
[False, True, True, True, False],
[False, False, True, True, True]]]])
```
## Creating a chunked attention mask
To create the following chunked attention mask (`chunk_size=3`):
0 ■ ⬚ ⬚ ⬚ ⬚
1 ■ ■ ⬚ ⬚ ⬚
2 ■ ■ ■ ⬚ ⬚
3 ⬚ ⬚ ⬚ ■ ⬚
4 ⬚ ⬚ ⬚ ■ ■
You can do
```python
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5, mask_function=chunked_causal_mask_function(3, torch.zeros(1, dtype=int)))
>>> tensor([[[[ True, False, False, False, False],
[ True, True, False, False, False],
[ True, True, True, False, False],
[False, False, False, True, False],
[False, False, False, True, True]]]])
```
"""
q_length = cache_position.shape[0]
# Potentially pad the 2D mask
padding_mask = prepare_padding_mask(attention_mask, kv_length, kv_offset)
# Potentially add the padding 2D mask
if padding_mask is not None:
mask_function = and_masks(mask_function, padding_mask_function(padding_mask))
batch_arange = torch.arange(batch_size, device=cache_position.device)
head_arange = torch.arange(1, device=cache_position.device)
# Similar to `kv_arange = torch.arange(start=kv_offset, end=kv_offset + kv_length, device=cache_position.device)`
# but without data-dependent slicing (i.e. torch.compile friendly)
kv_arange = torch.arange(kv_length, device=cache_position.device) + kv_offset
# Actual mask creation
# Apply mask function element-wise through broadcasting
attention_mask = mask_function(*_non_vmap_expansion_sdpa(batch_arange, head_arange, cache_position, kv_arange))
# Expand the mask to match batch size and query length if they weren't used in the mask function
attention_mask = attention_mask.expand(batch_size, -1, q_length, kv_length)
return attention_mask
+10 -1
View File
@@ -24,16 +24,22 @@ logger = init_logger(__name__)
_TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"HunyuanVideo15Transformer3DModel":
("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel")
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrix_game", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrix_game", "CausalMatrixGameWanModel"),
}
_TEXT_ENCODER_MODELS = {
@@ -43,16 +49,19 @@ _TEXT_ENCODER_MODELS = {
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
"CLIPVisionModel": ("encoders", "clip", "CLIPVisionModel"),
}
_VAE_MODELS = {
"AutoencoderKLHunyuanVideo":
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
}
+703
View File
@@ -0,0 +1,703 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from diffusers
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.layers.activation import get_act_fn
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from fastvideo.models.vaes.common import ParallelTiledVAE
class HunyuanVideo15CausalConv3d(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3,
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1,
bias: bool = True,
pad_mode: str = "replicate",
) -> None:
super().__init__()
kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
self.pad_mode = pad_mode
self.time_causal_padding = (
kernel_size[0] // 2,
kernel_size[0] // 2,
kernel_size[1] // 2,
kernel_size[1] // 2,
kernel_size[2] - 1,
0,
)
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode)
return self.conv(hidden_states)
class HunyuanVideo15RMS_norm(nn.Module):
r"""
A custom RMS normalization layer.
Args:
dim (int): The number of dimensions to normalize over.
channel_first (bool, optional): Whether the input tensor has channels as the first dimension.
Default is True.
images (bool, optional): Whether the input represents image data. Default is True.
bias (bool, optional): Whether to include a learnable bias term. Default is False.
"""
def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None:
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
def forward(self, x):
return F.normalize(x, dim=(1 if self.channel_first else -1)) * self.scale * self.gamma + self.bias
class HunyuanVideo15AttnBlock(nn.Module):
def __init__(self, in_channels: int):
super().__init__()
self.in_channels = in_channels
self.norm = HunyuanVideo15RMS_norm(in_channels, images=False)
self.to_q = nn.Conv3d(in_channels, in_channels, kernel_size=1)
self.to_k = nn.Conv3d(in_channels, in_channels, kernel_size=1)
self.to_v = nn.Conv3d(in_channels, in_channels, kernel_size=1)
self.proj_out = nn.Conv3d(in_channels, in_channels, kernel_size=1)
@staticmethod
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
"""Prepare a causal attention mask for 3D videos.
Args:
n_frame (int): Number of frames (temporal length).
n_hw (int): Product of height and width.
dtype: Desired mask dtype.
device: Device for the mask.
batch_size (int, optional): If set, expands for batch.
Returns:
torch.Tensor: Causal attention mask.
"""
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
def forward(self, x: torch.Tensor) -> torch.Tensor:
identity = x
x = self.norm(x)
query = self.to_q(x)
key = self.to_k(x)
value = self.to_v(x)
batch_size, channels, frames, height, width = query.shape
query = query.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
key = key.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
value = value.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
attention_mask = self.prepare_causal_attention_mask(
frames, height * width, query.dtype, query.device, batch_size=batch_size
)
x = nn.functional.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask)
# batch_size, 1, frames * height * width, channels
x = x.squeeze(1).reshape(batch_size, frames, height, width, channels).permute(0, 4, 1, 2, 3)
x = self.proj_out(x)
return x + identity
class HunyuanVideo15Upsample(nn.Module):
def __init__(self, in_channels: int, out_channels: int, add_temporal_upsample: bool = True):
super().__init__()
factor = 2 * 2 * 2 if add_temporal_upsample else 1 * 2 * 2
self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels * factor, kernel_size=3)
self.add_temporal_upsample = add_temporal_upsample
self.repeats = factor * out_channels // in_channels
@staticmethod
def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2):
"""
Convert (b, r1*r2*r3*c, f, h, w) -> (b, c, r1*f, r2*h, r3*w)
Args:
tensor: Input tensor of shape (b, r1*r2*r3*c, f, h, w)
r1: temporal upsampling factor
r2: height upsampling factor
r3: width upsampling factor
"""
b, packed_c, f, h, w = tensor.shape
factor = r1 * r2 * r3
c = packed_c // factor
tensor = tensor.view(b, r1, r2, r3, c, f, h, w)
tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3)
return tensor.reshape(b, c, f * r1, h * r2, w * r3)
def forward(self, x: torch.Tensor):
r1 = 2 if self.add_temporal_upsample else 1
h = self.conv(x)
if self.add_temporal_upsample:
h_first = h[:, :, :1, :, :]
h_first = self._dcae_upsample_rearrange(h_first, r1=1, r2=2, r3=2)
h_first = h_first[:, : h_first.shape[1] // 2]
h_next = h[:, :, 1:, :, :]
h_next = self._dcae_upsample_rearrange(h_next, r1=r1, r2=2, r3=2)
h = torch.cat([h_first, h_next], dim=2)
# shortcut computation
x_first = x[:, :, :1, :, :]
x_first = self._dcae_upsample_rearrange(x_first, r1=1, r2=2, r3=2)
x_first = x_first.repeat_interleave(repeats=self.repeats // 2, dim=1)
x_next = x[:, :, 1:, :, :]
x_next = self._dcae_upsample_rearrange(x_next, r1=r1, r2=2, r3=2)
x_next = x_next.repeat_interleave(repeats=self.repeats, dim=1)
shortcut = torch.cat([x_first, x_next], dim=2)
else:
h = self._dcae_upsample_rearrange(h, r1=r1, r2=2, r3=2)
shortcut = x.repeat_interleave(repeats=self.repeats, dim=1)
shortcut = self._dcae_upsample_rearrange(shortcut, r1=r1, r2=2, r3=2)
return h + shortcut
class HunyuanVideo15Downsample(nn.Module):
def __init__(self, in_channels: int, out_channels: int, add_temporal_downsample: bool = True):
super().__init__()
factor = 2 * 2 * 2 if add_temporal_downsample else 1 * 2 * 2
self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels // factor, kernel_size=3)
self.add_temporal_downsample = add_temporal_downsample
self.group_size = factor * in_channels // out_channels
@staticmethod
def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2):
"""
Convert (b, c, r1*f, r2*h, r3*w) -> (b, r1*r2*r3*c, f, h, w)
This packs spatial/temporal dimensions into channels (opposite of upsample)
"""
b, c, packed_f, packed_h, packed_w = tensor.shape
f, h, w = packed_f // r1, packed_h // r2, packed_w // r3
tensor = tensor.view(b, c, f, r1, h, r2, w, r3)
tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6)
return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w)
def forward(self, x: torch.Tensor):
r1 = 2 if self.add_temporal_downsample else 1
h = self.conv(x)
if self.add_temporal_downsample:
h_first = h[:, :, :1, :, :]
h_first = self._dcae_downsample_rearrange(h_first, r1=1, r2=2, r3=2)
h_first = torch.cat([h_first, h_first], dim=1)
h_next = h[:, :, 1:, :, :]
h_next = self._dcae_downsample_rearrange(h_next, r1=r1, r2=2, r3=2)
h = torch.cat([h_first, h_next], dim=2)
# shortcut computation
x_first = x[:, :, :1, :, :]
x_first = self._dcae_downsample_rearrange(x_first, r1=1, r2=2, r3=2)
B, C, T, H, W = x_first.shape
x_first = x_first.view(B, h.shape[1], self.group_size // 2, T, H, W).mean(dim=2)
x_next = x[:, :, 1:, :, :]
x_next = self._dcae_downsample_rearrange(x_next, r1=r1, r2=2, r3=2)
B, C, T, H, W = x_next.shape
x_next = x_next.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
shortcut = torch.cat([x_first, x_next], dim=2)
else:
h = self._dcae_downsample_rearrange(h, r1=r1, r2=2, r3=2)
shortcut = self._dcae_downsample_rearrange(x, r1=r1, r2=2, r3=2)
B, C, T, H, W = shortcut.shape
shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
return h + shortcut
class HunyuanVideo15ResnetBlock(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
non_linearity: str = "swish",
) -> None:
super().__init__()
out_channels = out_channels or in_channels
self.nonlinearity = get_act_fn(non_linearity)
self.norm1 = HunyuanVideo15RMS_norm(in_channels, images=False)
self.conv1 = HunyuanVideo15CausalConv3d(in_channels, out_channels, kernel_size=3)
self.norm2 = HunyuanVideo15RMS_norm(out_channels, images=False)
self.conv2 = HunyuanVideo15CausalConv3d(out_channels, out_channels, kernel_size=3)
self.conv_shortcut = None
if in_channels != out_channels:
self.conv_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
residual = hidden_states
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv2(hidden_states)
if self.conv_shortcut is not None:
residual = self.conv_shortcut(residual)
return hidden_states + residual
class HunyuanVideo15MidBlock(nn.Module):
def __init__(
self,
in_channels: int,
num_layers: int = 1,
add_attention: bool = True,
) -> None:
super().__init__()
self.add_attention = add_attention
# There is always at least one resnet
resnets = [
HunyuanVideo15ResnetBlock(
in_channels=in_channels,
out_channels=in_channels,
)
]
attentions = []
for _ in range(num_layers):
if self.add_attention:
attentions.append(HunyuanVideo15AttnBlock(in_channels))
else:
attentions.append(None)
resnets.append(
HunyuanVideo15ResnetBlock(
in_channels=in_channels,
out_channels=in_channels,
)
)
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.resnets[0](hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
hidden_states = attn(hidden_states)
hidden_states = resnet(hidden_states)
return hidden_states
class HunyuanVideo15DownBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 1,
downsample_out_channels: Optional[int] = None,
add_temporal_downsample: int = True,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
HunyuanVideo15ResnetBlock(
in_channels=in_channels,
out_channels=out_channels,
)
)
self.resnets = nn.ModuleList(resnets)
if downsample_out_channels is not None:
self.downsamplers = nn.ModuleList(
[
HunyuanVideo15Downsample(
out_channels,
out_channels=downsample_out_channels,
add_temporal_downsample=add_temporal_downsample,
)
]
)
else:
self.downsamplers = None
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states)
return hidden_states
class HunyuanVideo15UpBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 1,
upsample_out_channels: Optional[int] = None,
add_temporal_upsample: bool = True,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
input_channels = in_channels if i == 0 else out_channels
resnets.append(
HunyuanVideo15ResnetBlock(
in_channels=input_channels,
out_channels=out_channels,
)
)
self.resnets = nn.ModuleList(resnets)
if upsample_out_channels is not None:
self.upsamplers = nn.ModuleList(
[
HunyuanVideo15Upsample(
out_channels,
out_channels=upsample_out_channels,
add_temporal_upsample=add_temporal_upsample,
)
]
)
else:
self.upsamplers = None
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if torch.is_grad_enabled() and self.gradient_checkpointing:
for resnet in self.resnets:
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
else:
for resnet in self.resnets:
hidden_states = resnet(hidden_states)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states)
return hidden_states
class HunyuanVideo15Encoder3D(nn.Module):
r"""
3D vae encoder for HunyuanImageRefiner.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 64,
block_out_channels: Tuple[int, ...] = (128, 256, 512, 1024, 1024),
layers_per_block: int = 2,
temporal_compression_ratio: int = 4,
spatial_compression_ratio: int = 16,
downsample_match_channel: bool = True,
) -> None:
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.group_size = block_out_channels[-1] // self.out_channels
self.conv_in = HunyuanVideo15CausalConv3d(in_channels, block_out_channels[0], kernel_size=3)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
input_channel = block_out_channels[0]
for i in range(len(block_out_channels)):
add_spatial_downsample = i < np.log2(spatial_compression_ratio)
output_channel = block_out_channels[i]
if not add_spatial_downsample:
down_block = HunyuanVideo15DownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
downsample_out_channels=None,
add_temporal_downsample=False,
)
input_channel = output_channel
else:
add_temporal_downsample = i >= np.log2(spatial_compression_ratio // temporal_compression_ratio)
downsample_out_channels = block_out_channels[i + 1] if downsample_match_channel else output_channel
down_block = HunyuanVideo15DownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
downsample_out_channels=downsample_out_channels,
add_temporal_downsample=add_temporal_downsample,
)
input_channel = downsample_out_channels
self.down_blocks.append(down_block)
self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[-1])
self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False)
self.conv_act = nn.SiLU()
self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv_in(hidden_states)
if torch.is_grad_enabled() and self.gradient_checkpointing:
for down_block in self.down_blocks:
hidden_states = self._gradient_checkpointing_func(down_block, hidden_states)
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
else:
for down_block in self.down_blocks:
hidden_states = down_block(hidden_states)
hidden_states = self.mid_block(hidden_states)
batch_size, _, frame, height, width = hidden_states.shape
short_cut = hidden_states.view(batch_size, -1, self.group_size, frame, height, width).mean(dim=2)
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
hidden_states += short_cut
return hidden_states
class HunyuanVideo15Decoder3D(nn.Module):
r"""
Causal decoder for 3D video-like data used for HunyuanImage-1.5 Refiner.
"""
def __init__(
self,
in_channels: int = 32,
out_channels: int = 3,
block_out_channels: Tuple[int, ...] = (1024, 1024, 512, 256, 128),
layers_per_block: int = 2,
spatial_compression_ratio: int = 16,
temporal_compression_ratio: int = 4,
upsample_match_channel: bool = True,
):
super().__init__()
self.layers_per_block = layers_per_block
self.in_channels = in_channels
self.out_channels = out_channels
self.repeat = block_out_channels[0] // self.in_channels
self.conv_in = HunyuanVideo15CausalConv3d(self.in_channels, block_out_channels[0], kernel_size=3)
self.up_blocks = nn.ModuleList([])
# mid
self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[0])
# up
input_channel = block_out_channels[0]
for i in range(len(block_out_channels)):
output_channel = block_out_channels[i]
add_spatial_upsample = i < np.log2(spatial_compression_ratio)
add_temporal_upsample = i < np.log2(temporal_compression_ratio)
if add_spatial_upsample or add_temporal_upsample:
upsample_out_channels = block_out_channels[i + 1] if upsample_match_channel else output_channel
up_block = HunyuanVideo15UpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=input_channel,
out_channels=output_channel,
upsample_out_channels=upsample_out_channels,
add_temporal_upsample=add_temporal_upsample,
)
input_channel = upsample_out_channels
else:
up_block = HunyuanVideo15UpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=input_channel,
out_channels=output_channel,
upsample_out_channels=None,
add_temporal_upsample=False,
)
input_channel = output_channel
self.up_blocks.append(up_block)
# out
self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False)
self.conv_act = nn.SiLU()
self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv_in(hidden_states) + hidden_states.repeat_interleave(repeats=self.repeat, dim=1)
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
for up_block in self.up_blocks:
hidden_states = self._gradient_checkpointing_func(up_block, hidden_states)
else:
hidden_states = self.mid_block(hidden_states)
for up_block in self.up_blocks:
hidden_states = up_block(hidden_states)
# post-process
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
r"""
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used for
HunyuanVideo-1.5.
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
"""
_supports_gradient_checkpointing = True
def __init__(
self,
config: Hunyuan15VAEConfig,
) -> None:
nn.Module.__init__(self)
ParallelTiledVAE.__init__(self, config)
if config.load_encoder:
self.encoder = HunyuanVideo15Encoder3D(
in_channels=config.in_channels,
out_channels=config.latent_channels * 2,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
downsample_match_channel=config.downsample_match_channel,
)
if config.load_decoder:
self.decoder = HunyuanVideo15Decoder3D(
in_channels=config.latent_channels,
out_channels=config.out_channels,
block_out_channels=list(reversed(config.block_out_channels)),
layers_per_block=config.layers_per_block,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
upsample_match_channel=config.upsample_match_channel,
)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = False
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 2000 # Fill in a random large number, as hy1.5 vae does not use temporal tiling
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
return x
def _decode(self, z: torch.Tensor) -> torch.Tensor:
dec = self.decoder(z)
return dec
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
return_dict: bool = True,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z)
return dec
+6
View File
@@ -44,6 +44,12 @@ def build_pipeline(
config = verify_model_config_and_directory(model_path)
pipeline_name = config.get("_class_name")
if fastvideo_args.override_pipeline_cls_name:
logger.info("Overriding pipeline class name from %s to %s",
pipeline_name, fastvideo_args.override_pipeline_cls_name)
pipeline_name = fastvideo_args.override_pipeline_cls_name
if pipeline_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "

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