Compare commits

...
110 changed files with 13810 additions and 140 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/
+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
+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 -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
+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
+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`.
+1 -1
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
+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
+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()
+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,
+4 -2
View File
@@ -108,6 +108,7 @@ class DistributedAttention(nn.Module):
# 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, :, :]
@@ -140,8 +141,9 @@ 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)
@@ -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"
@@ -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}'"
)
+14
View File
@@ -7,7 +7,9 @@ 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 (
@@ -27,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,
@@ -56,6 +62,8 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"matrixgame":
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
"wanpipeline":
@@ -70,14 +78,19 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
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,
@@ -123,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():
+38
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__)
@@ -27,6 +28,17 @@ class SamplingParam:
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"
@@ -50,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
@@ -215,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
+9
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
@@ -35,6 +36,10 @@ logger = init_logger(__name__)
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Wan2.1
@@ -86,6 +91,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
@@ -105,6 +112,8 @@ 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,
+56
View File
@@ -319,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",
+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,
}
+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)
+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
+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
+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
+7 -1
View File
@@ -24,10 +24,14 @@ 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 = {
@@ -45,6 +49,7 @@ _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] = {
@@ -56,6 +61,7 @@ _IMAGE_ENCODER_MODELS: dict[str, tuple] = {
_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
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
Hy15ImageEncodingStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15Pipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HunyuanVideo15Pipeline
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""LongCat pipeline module."""
from fastvideo.pipelines.basic.longcat.longcat_pipeline import LongCatPipeline
__all__ = ["LongCatPipeline"]
@@ -0,0 +1,145 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat video diffusion pipeline implementation (Phase 1: Wrapper).
This module contains a wrapper implementation of the LongCat video diffusion pipeline
using FastVideo's modular pipeline architecture with the original LongCat modules.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
logger = init_logger(__name__)
class LongCatPipeline(LoRAPipeline, ComposedPipelineBase):
"""
LongCat video diffusion pipeline with LoRA support.
Phase 1 implementation using wrapper modules from third_party/longcat_video.
This validates the pipeline infrastructure before full FastVideo integration.
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize LongCat-specific components."""
# Enable BSA (Block Sparse Attention) if configured
pipeline_config = fastvideo_args.pipeline_config
transformer = self.get_module("transformer", None)
if transformer is None:
raise RuntimeError(
"Transformer module not found during initializing LongCat pipeline."
)
# If user toggles BSA via CLI/config
if pipeline_config.enable_bsa:
# Build effective BSA params:
# 1) from explicit CLI overrides if provided
# 2) else from pipeline_config.bsa_params
# 3) else fall back to reasonable defaults
bsa_params_cfg = pipeline_config.bsa_params
sparsity = pipeline_config.bsa_sparsity
cdf_threshold = pipeline_config.bsa_cdf_threshold
chunk_q = pipeline_config.bsa_chunk_q
chunk_k = pipeline_config.bsa_chunk_k
effective_bsa_params = dict(bsa_params_cfg) if isinstance(
bsa_params_cfg, dict) else {}
if sparsity is not None:
effective_bsa_params['sparsity'] = sparsity
if cdf_threshold is not None:
effective_bsa_params['cdf_threshold'] = cdf_threshold
if chunk_q is not None:
effective_bsa_params['chunk_3d_shape_q'] = chunk_q
if chunk_k is not None:
effective_bsa_params['chunk_3d_shape_k'] = chunk_k
# Provide defaults if still missing
effective_bsa_params.setdefault('sparsity', 0.9375)
effective_bsa_params.setdefault('chunk_3d_shape_q', [4, 4, 4])
effective_bsa_params.setdefault('chunk_3d_shape_k', [4, 4, 4])
if hasattr(transformer, 'enable_bsa'):
logger.info(
"Enabling Block Sparse Attention (BSA) for LongCat transformer"
)
transformer.enable_bsa()
# Propagate params to all attention modules
if hasattr(transformer, 'blocks'):
try:
for blk in transformer.blocks:
if hasattr(blk, 'self_attn'):
blk.self_attn.bsa_params = effective_bsa_params
except Exception as e:
logger.warning(
"Failed to set BSA params on all blocks: %s", e)
logger.info("BSA parameters in effect: %s",
effective_bsa_params)
else:
logger.warning(
"BSA is enabled in config but transformer does not support it"
)
else:
# Explicitly disable if present
if hasattr(transformer, 'disable_bsa'):
transformer.disable_bsa()
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
# Add refine initialization stage (will be skipped if not refining)
self.add_stage(stage_name="longcat_refine_init_stage",
stage=LongCatRefineInitStage(vae=self.get_module("vae")))
# First prepare generic timesteps (for non-refine paths)
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
# Then override timesteps for refinement (will be a no-op if not refining),
# matching LongCat's generate_refine schedule.
self.add_stage(stage_name="longcat_refine_timestep_stage",
stage=LongCatRefineTimestepStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=LongCatDenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae"),
pipeline=self))
EntryClass = LongCatPipeline
@@ -287,6 +287,7 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
model_index.pop("workload_type", None)
if "boundary_ratio" in model_index and model_index[
"boundary_ratio"] is not None:
logger.info(
@@ -91,6 +91,14 @@ class ForwardBatch:
video_path: str | None = None
video_latent: torch.Tensor | None = None
# Refine inputs (LongCat)
refine_from: str | None = None
t_thresh: float = 0.5
spatial_refine_only: bool = False
num_cond_frames: int = 0
stage1_video: list[
PIL.Image.Image] | None = None # Loaded frames from refine_from
# Primary encoder embeddings
prompt_embeds: list[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: list[torch.Tensor] | None = None
+2
View File
@@ -25,9 +25,11 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
"HunyuanVideo15Pipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
"MatrixGameCausalDMDPipeline": "matrixgame",
"LongCatPipeline": "longcat",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
+2 -1
View File
@@ -16,7 +16,7 @@ from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
ImageVAEEncodingStage, VideoVAEEncodingStage)
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import (
CosmosLatentPreparationStage, LatentPreparationStage)
@@ -44,6 +44,7 @@ __all__ = [
"DecodingStage",
"ImageEncodingStage",
"MatrixGameImageEncodingStage",
"Hy15ImageEncodingStage",
"RefImageEncodingStage",
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
+15
View File
@@ -187,6 +187,21 @@ class DecodingStage(PipelineStage):
# Convert to CPU float32 for compatibility
frames = frames.cpu().float()
# Crop padding if this is a LongCat refinement
if hasattr(batch, 'num_cond_frames_added') and hasattr(
batch, 'new_frame_size_before_padding'):
num_cond_frames_added = batch.num_cond_frames_added
new_frame_size = batch.new_frame_size_before_padding
if num_cond_frames_added > 0 or frames.shape[2] != new_frame_size:
# frames is [B, C, T, H, W], crop temporal dimension
frames = frames[:, :,
num_cond_frames_added:num_cond_frames_added +
new_frame_size, :, :]
logger.info(
"Cropped LongCat refinement padding: %s:%s, final shape: %s",
num_cond_frames_added,
num_cond_frames_added + new_frame_size, frames.shape)
# Update batch with decoded image
batch.output = frames
@@ -100,6 +100,33 @@ class ImageEncodingStage(PipelineStage):
return result
class Hy15ImageEncodingStage(ImageEncodingStage):
"""
Stage for encoding image prompts into embeddings for HunyuanVideo1.5 models.
"""
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify image encoding stage inputs."""
return VerificationResult()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states.
"""
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
raw_latent_shape = list(batch.raw_latent_shape)
raw_latent_shape[1] = 1
batch.video_latent = torch.zeros(tuple(raw_latent_shape),
device=get_local_torch_device())
return batch
class MatrixGameImageEncodingStage(ImageEncodingStage):
CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073]
CLIP_STD = [0.26862954, 0.26130258, 0.27577711]
@@ -136,16 +136,26 @@ class LatentPreparationStage(PipelineStage):
)
# Generate or use provided latents
if latents is None:
latents = randn_tensor(shape,
generator=generator,
device=device,
dtype=dtype)
latents = randn_tensor(
shape,
generator=generator,
device=device,
dtype=dtype,
)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
else:
# Pre-initialized latents:
# - For LongCat refine (refine_from or stage1_video present), we should not re-scale by init_noise_sigma.
# - For other models, keep the original behavior.
latents = latents.to(device)
is_longcat_refine = (batch.refine_from
is not None) or (batch.stage1_video
is not None)
if (not is_longcat_refine) and hasattr(self.scheduler,
"init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
# Scale the initial noise if needed
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
# Update batch with prepared latents
batch.latents = latents
batch.raw_latent_shape = bcthw_shape
@@ -0,0 +1,179 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat-specific denoising stage implementing CFG-zero optimized guidance.
"""
import torch
from tqdm import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.forward_context import set_forward_context
logger = init_logger(__name__)
class LongCatDenoisingStage(DenoisingStage):
"""
LongCat denoising stage with CFG-zero optimized guidance scale.
Implements:
1. Optimized CFG scale from CFG-zero paper
2. Negation of noise prediction before scheduler step (flow matching convention)
3. Batched CFG computation (unlike standard FastVideo separate passes)
"""
def optimized_scale(self, positive_flat, negative_flat) -> torch.Tensor:
"""
Calculate optimized scale from CFG-zero paper.
st_star = (v_cond^T * v_uncond) / ||v_uncond||^2
Args:
positive_flat: Conditional prediction, flattened [B, -1]
negative_flat: Unconditional prediction, flattened [B, -1]
Returns:
st_star: Optimized scale [B, 1]
"""
# Calculate dot product
dot_product = torch.sum(positive_flat * negative_flat,
dim=1,
keepdim=True)
# Squared norm of uncondition
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
st_star = dot_product / squared_norm
return st_star
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run LongCat denoising loop with optimized CFG.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
if not fastvideo_args.model_loaded["transformer"]:
from fastvideo.models.model_loader import TransformerLoader
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
pipeline = self.pipeline() if self.pipeline else None
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
# Get transformer dtype
if hasattr(self.transformer, 'module'):
transformer_dtype = next(self.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Extract batch parameters
latents = batch.latents
timesteps = batch.timesteps
prompt_embeds = batch.prompt_embeds[0] # LongCat uses single encoder
prompt_attention_mask = batch.prompt_attention_mask[
0] if batch.prompt_attention_mask else None
guidance_scale = batch.guidance_scale
do_classifier_free_guidance = batch.do_classifier_free_guidance
# Get negative prompts if doing CFG
if do_classifier_free_guidance:
negative_prompt_embeds = batch.negative_prompt_embeds[0]
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
if batch.negative_attention_mask
else None)
# Concatenate for batched processing
prompt_embeds_combined = torch.cat(
[negative_prompt_embeds, prompt_embeds], dim=0)
if prompt_attention_mask is not None:
prompt_attention_mask_combined = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask],
dim=0)
else:
prompt_attention_mask_combined = None
else:
prompt_embeds_combined = prompt_embeds
prompt_attention_mask_combined = prompt_attention_mask
# Denoising loop
num_inference_steps = len(timesteps)
with tqdm(total=num_inference_steps,
desc="LongCat Denoising") as progress_bar:
for i, t in enumerate(timesteps):
# Expand latents for CFG
if do_classifier_free_guidance:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
# Expand timestep to match batch size
timestep = t.expand(
latent_model_input.shape[0]).to(target_dtype)
# Run transformer with context
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type='cuda',
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask_combined,
)
# Apply CFG with optimized scale
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
# Calculate optimized scale (CFG-zero)
st_star = self.optimized_scale(positive, negative)
# Reshape for broadcasting
st_star = st_star.view(B, 1, 1, 1, 1)
# Apply optimized CFG formula
noise_pred = (
noise_pred_uncond * st_star + guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
# CRITICAL: Negate noise prediction for flow matching scheduler
noise_pred = -noise_pred
# Compute previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
progress_bar.update()
# Update batch with denoised latents
batch.latents = latents
return batch
@@ -0,0 +1,310 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat refinement initialization stage.
This stage prepares the latent variables for LongCat's 480p->720p refinement by:
1. Loading the stage1 (480p) video
2. Upsampling it to 720p resolution
3. Encoding it with VAE
4. Mixing with noise according to t_thresh
"""
import math
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vision_utils import load_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.configs.pipelines.longcat import get_bucket_config
logger = init_logger(__name__)
class LongCatRefineInitStage(PipelineStage):
"""
Stage for initializing LongCat refinement from a stage1 (480p) video.
This replicates the logic from LongCatVideoPipeline.generate_refine():
- Load stage1_video frames
- Upsample spatially and temporally
- VAE encode and normalize
- Mix with noise according to t_thresh
"""
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Initialize latents for refinement.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with initialized latents for refinement.
"""
refine_from = batch.refine_from
in_memory_stage1 = batch.stage1_video
# Only run for refinement tasks: either a path (refine_from) or in-memory video is provided
if refine_from is None and in_memory_stage1 is None:
# Not a refinement task, skip
return batch
# ------------------------------------------------------------------
# 1. Obtain stage1 frames (either from disk or from in-memory input)
# ------------------------------------------------------------------
if in_memory_stage1 is not None:
# User provided stage1 frames directly (e.g., from distilled stage output)
if len(in_memory_stage1) == 0:
raise ValueError(
"stage1_video is empty; expected a non-empty list of frames"
)
if isinstance(in_memory_stage1[0], Image.Image):
pil_images = in_memory_stage1
else:
# Assume numpy arrays or torch tensors with shape [H, W, C]
pil_images = [
Image.fromarray(np.array(frame))
for frame in in_memory_stage1
]
logger.info(
"Initializing LongCat refinement from in-memory stage1_video (%s frames)",
len(pil_images))
else:
# Path-based refine: load video from disk (original design)
logger.info("Initializing LongCat refinement from file: %s",
refine_from)
stage1_video_path = Path(refine_from)
if not stage1_video_path.exists():
raise FileNotFoundError(
f"Stage1 video not found: {refine_from}")
# Load video frames as PIL Images
pil_images, original_fps = load_video(str(stage1_video_path),
return_fps=True)
logger.info("Loaded stage1 video: %s frames @ %s fps",
len(pil_images), original_fps)
# Store in batch for reference (use PIL images, same as official demo)
batch.stage1_video = pil_images
# Get parameters from batch
num_frames = len(pil_images)
spatial_refine_only = batch.spatial_refine_only
t_thresh = batch.t_thresh
num_cond_frames = batch.num_cond_frames if hasattr(
batch, 'num_cond_frames') else 0
# Calculate new frame count (temporal upsampling if not spatial_refine_only)
new_num_frames = num_frames if spatial_refine_only else 2 * num_frames
logger.info(
"Refine mode: %s",
'spatial only' if spatial_refine_only else 'spatial + temporal')
# Update batch.num_frames to reflect the upsampled count
batch.num_frames = new_num_frames
# Use bucket system to select resolution (exactly like LongCat)
# Calculate scale_factor_spatial considering SP split
sp_size = fastvideo_args.sp_size if fastvideo_args.sp_size > 0 else 1
vae_scale_factor_spatial = 8 # VAE spatial downsampling
patch_size_spatial = 2 # LongCat patch size
bsa_latent_granularity = 4
scale_factor_spatial = vae_scale_factor_spatial * patch_size_spatial * bsa_latent_granularity # 64
# Calculate optimal split like LongCat (cp_split_hw logic)
# For sp_size=1: [1,1], max=1
# For sp_size=2: [1,2], max=2
# For sp_size=4: [2,2], max=2
# For sp_size=8: [2,4], max=4
if sp_size > 1:
# Get optimal 2D split factors (mimic context_parallel_util.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]))
scale_factor_spatial *= max(cp_split_hw)
logger.info("SP split: sp_size=%s, cp_split_hw=%s, max_split=%s",
sp_size, cp_split_hw, max(cp_split_hw))
else:
cp_split_hw = [1, 1]
# Get bucket config and find closest bucket for the input aspect ratio
bucket_config = get_bucket_config('720p', scale_factor_spatial)
# Get input aspect ratio from stage1 video
input_height, input_width = pil_images[0].height, pil_images[0].width
input_ratio = input_height / input_width
# Find closest bucket
closest_ratio = min(bucket_config.keys(),
key=lambda x: abs(float(x) - input_ratio))
height, width = bucket_config[closest_ratio][0]
logger.info("Input aspect ratio: %.2f (%sx%s)", input_ratio,
input_width, input_height)
logger.info("Matched bucket ratio: %s -> resolution: %sx%s",
closest_ratio, width, height)
logger.info("Target: %sx%s @ %s frames (sp_size=%s, scale_factor=%s)",
width, height, new_num_frames, sp_size,
scale_factor_spatial)
# Override batch height/width with bucket-selected resolution
batch.height = height
batch.width = width
# Convert PIL images to tensor [T, C, H, W]
stage1_video_tensor = torch.stack([
torch.from_numpy(np.array(img)).permute(2, 0, 1) # HWC -> CHW
for img in pil_images
]).float() # [T, C, H, W]
device = batch.prompt_embeds[0].device
dtype = batch.prompt_embeds[0].dtype
stage1_video_tensor = stage1_video_tensor.to(device=device, dtype=dtype)
# Replicate LongCat's exact preprocessing (lines 1227-1235 in pipeline_longcat_video.py)
# First: spatial interpolation to target (height, width) on [T, C, H, W]
video_down = F.interpolate(stage1_video_tensor,
size=(height, width),
mode='bilinear',
align_corners=True)
# Rearrange to [C, T, H, W] and add batch dimension -> [1, C, T, H, W]
video_down = video_down.permute(1, 0, 2,
3).unsqueeze(0) # [1, C, T, H, W]
video_down = video_down / 255.0 # Normalize to [0, 1]
# Then: temporal+spatial interpolation to (new_num_frames, height, width)
video_up = F.interpolate(video_down,
size=(new_num_frames, height, width),
mode='trilinear',
align_corners=True)
# Rescale to [-1, 1] for VAE
video_up = video_up * 2.0 - 1.0
logger.info("Upsampled video shape: %s", video_up.shape)
# Padding logic (exactly like LongCat lines 1237-1255)
# Only pad temporal dimension to ensure BSA compatibility
vae_scale_factor_temporal = 4
num_noise_frames = video_up.shape[2] - num_cond_frames
num_cond_latents = 0
num_cond_frames_added = 0
if num_cond_frames > 0:
num_cond_latents = 1 + math.ceil(
(num_cond_frames - 1) / vae_scale_factor_temporal)
num_cond_latents = math.ceil(
num_cond_latents /
bsa_latent_granularity) * bsa_latent_granularity
num_cond_frames_added = 1 + (
num_cond_latents -
1) * vae_scale_factor_temporal - num_cond_frames
num_cond_frames = num_cond_frames + num_cond_frames_added
num_noise_latents = math.ceil(num_noise_frames /
vae_scale_factor_temporal)
num_noise_latents = math.ceil(
num_noise_latents / bsa_latent_granularity) * bsa_latent_granularity
num_noise_frames_added = num_noise_latents * vae_scale_factor_temporal - num_noise_frames
if num_cond_frames_added > 0 or num_noise_frames_added > 0:
logger.info(
"Padding temporal dimension for BSA: cond_frames+=%s, noise_frames+=%s",
num_cond_frames_added, num_noise_frames_added)
pad_front = video_up[:, :, 0:1].repeat(1, 1, num_cond_frames_added,
1, 1)
pad_back = video_up[:, :, -1:].repeat(1, 1, num_noise_frames_added,
1, 1)
video_up = torch.cat([pad_front, video_up, pad_back], dim=2)
logger.info("Padded video shape: %s", video_up.shape)
# Update batch with actual frame count after padding
batch.num_frames = video_up.shape[2]
# Store padding info for later cropping (CRITICAL for correct output!)
batch.num_cond_frames_added = num_cond_frames_added
batch.num_noise_frames_added = num_noise_frames_added
batch.new_frame_size_before_padding = new_num_frames
# Store num_cond_latents for denoising stage
if num_cond_latents > 0:
batch.num_cond_latents = num_cond_latents
logger.info("Will use num_cond_latents=%s during denoising",
num_cond_latents)
logger.info("Padding info: cond+=%s, noise+=%s, original=%s",
num_cond_frames_added, num_noise_frames_added,
new_num_frames)
# VAE encode
logger.info("Encoding stage1 video with VAE...")
vae_dtype = next(self.vae.parameters()).dtype
vae_device = next(self.vae.parameters()).device
video_up = video_up.to(dtype=vae_dtype, device=vae_device)
with torch.no_grad():
latent_dist = self.vae.encode(video_up)
# Extract tensor from latent distribution
if hasattr(latent_dist, 'latent_dist'):
# Nested distribution wrapper
latent_up = latent_dist.latent_dist.sample()
elif hasattr(latent_dist, 'sample'):
# DiagonalGaussianDistribution or similar
latent_up = latent_dist.sample()
elif hasattr(latent_dist, 'latents'):
# Direct latents tensor
latent_up = latent_dist.latents
else:
# Assume it's already a tensor
latent_up = latent_dist
# Normalize latents using VAE config (exactly like LongCat)
if hasattr(self.vae.config, 'latents_mean') and hasattr(
self.vae.config, 'latents_std'):
latents_mean = torch.tensor(self.vae.config.latents_mean).view(
1, self.vae.config.z_dim, 1, 1, 1).to(latent_up.device,
latent_up.dtype)
# LongCat uses: 1.0 / latents_std (equivalent to dividing by latents_std)
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(
1, self.vae.config.z_dim, 1, 1, 1).to(latent_up.device,
latent_up.dtype)
# LongCat: (latents - mean) * (1/std)
latent_up = (latent_up - latents_mean) * latents_std
logger.info("Encoded latent shape: %s", latent_up.shape)
# Mix with noise according to t_thresh
# latent_up = (1 - t_thresh) * latent_up + t_thresh * noise
noise = torch.randn_like(latent_up).contiguous()
latent_up = (1 - t_thresh) * latent_up + t_thresh * noise
logger.info("Applied t_thresh=%s noise mixing", t_thresh)
# Store in batch
batch.latents = latent_up.to(dtype)
batch.raw_latent_shape = latent_up.shape
logger.info("LongCat refinement initialization complete")
return batch
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat refinement timestep preparation stage.
This stage prepares special timesteps for LongCat refinement that start from t_thresh.
"""
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
logger = init_logger(__name__)
class LongCatRefineTimestepStage(PipelineStage):
"""
Stage for preparing timesteps specific to LongCat refinement.
For refinement, we need to start from t_thresh instead of t=1.0, so we:
1. Generate normal timesteps for num_inference_steps
2. Filter to only keep timesteps < t_thresh * 1000
3. Prepend t_thresh * 1000 as the first timestep
"""
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Prepare refinement-specific timesteps.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with refinement timesteps.
"""
# Only apply if this is a refinement task
# Trigger when either a refine_from path or in-memory stage1_video is provided
if batch.refine_from is None and batch.stage1_video is None:
return batch
device = get_local_torch_device()
num_inference_steps = batch.num_inference_steps
t_thresh = batch.t_thresh
logger.info("Preparing LongCat refinement timesteps (t_thresh=%s)",
t_thresh)
# ------------------------------------------------------------------
# 1) Match LongCatVideoPipeline.get_timesteps_sigmas (non-distill):
# sigmas = linspace(1, 0.001, num_inference_steps) on CPU
# ------------------------------------------------------------------
base_sigmas = torch.linspace(
1.0,
0.001,
num_inference_steps,
dtype=torch.float32,
device=
"cpu", # scheduler.set_timesteps expects CPU-convertible sigmas
)
# Let the scheduler build its internal timestep schedule from sigmas
self.scheduler.set_timesteps(num_inference_steps,
sigmas=base_sigmas,
device=device)
base_timesteps = self.scheduler.timesteps
# ------------------------------------------------------------------
# 2) Apply t_thresh cropping exactly like generate_refine:
# timesteps = [t_thresh*1000] + [t for t in base_timesteps if t < t_thresh*1000]
# sigmas = timesteps / 1000 (with trailing zero)
# ------------------------------------------------------------------
t_thresh_value = t_thresh * 1000.0
t_thresh_tensor = torch.tensor(t_thresh_value,
dtype=base_timesteps.dtype,
device=device)
filtered_timesteps = base_timesteps[base_timesteps < t_thresh_tensor]
timesteps = torch.cat(
[t_thresh_tensor.unsqueeze(0), filtered_timesteps])
# Update scheduler with these custom timesteps and corresponding sigmas
self.scheduler.timesteps = timesteps
sigmas = torch.cat([timesteps / 1000.0, torch.zeros(1, device=device)])
self.scheduler.sigmas = sigmas
logger.info("Refinement timesteps: %s steps starting from t=%s",
len(timesteps), t_thresh)
logger.info("First few timesteps: %s", timesteps[:5].tolist())
# Store in batch so downstream stages (denoising) use the same schedule
batch.timesteps = timesteps
return batch
+52 -10
View File
@@ -6,6 +6,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
"""
import torch
from typing import Any
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -100,9 +101,9 @@ class TextEncodingStage(PipelineStage):
"""Verify text encoding stage inputs."""
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
result.add_check(
"negative_prompt", batch.negative_prompt, lambda x: not batch.
do_classifier_free_guidance or V.string_not_empty(x))
# result.add_check(
# "negative_prompt", batch.negative_prompt, lambda x: not batch.
# do_classifier_free_guidance or V.string_not_empty(x))
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check("prompt_embeds", batch.prompt_embeds, V.is_list)
@@ -203,20 +204,45 @@ class TextEncodingStage(PipelineStage):
preprocess_func = preprocess_funcs[i]
postprocess_func = postprocess_funcs[i]
processed_texts: list[str] = []
for prompt_str in texts:
processed_texts.append(preprocess_func(prompt_str))
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
if max_length is not None:
tok_kwargs["max_length"] = max_length
elif hasattr(fastvideo_args.pipeline_config,
"text_encoder_max_lengths"):
tok_kwargs[
"max_length"] = fastvideo_args.pipeline_config.text_encoder_max_lengths[
i]
if truncation is not None:
tok_kwargs["truncation"] = truncation
if padding is not None:
tok_kwargs["padding"] = padding
text_inputs = tokenizer(processed_texts,
**tok_kwargs).to(target_device)
processed_texts: list[str] = []
for prompt_str in texts:
processed_text = preprocess_func(prompt_str)
if processed_text is not None:
processed_texts.append(processed_text)
else:
# Assuming batch_size = 1
prompt_embeds = torch.zeros((1, tok_kwargs["max_length"],
encoder_config.hidden_size),
device=target_device)
attention_mask = torch.zeros((1, tok_kwargs["max_length"]),
device=target_device,
dtype=torch.int64)
embeds_list.append(prompt_embeds)
attn_masks_list.append(attention_mask)
return self.return_embeds(embeds_list, attn_masks_list,
return_type,
return_attention_mask, indices)
if encoder_config.is_chat_model:
text_inputs = tokenizer.apply_chat_template(
processed_texts, **tok_kwargs).to(target_device)
else:
text_inputs = tokenizer(processed_texts,
**tok_kwargs).to(target_device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
@@ -228,13 +254,29 @@ class TextEncodingStage(PipelineStage):
output_hidden_states=True,
)
prompt_embeds = postprocess_func(outputs)
try:
prompt_embeds = postprocess_func(outputs)
except Exception:
prompt_embeds, attention_mask = postprocess_func(
outputs, attention_mask)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
embeds_list.append(prompt_embeds)
if return_attention_mask:
attn_masks_list.append(attention_mask)
return self.return_embeds(embeds_list, attn_masks_list, return_type,
return_attention_mask, indices)
def return_embeds(
self,
embeds_list: list[torch.Tensor],
attn_masks_list: list[torch.Tensor],
return_type: str = "list",
return_attention_mask: bool = False,
indices: list[int] | None = None,
) -> Any:
# Shape results according to return_type
if return_type == "list":
if return_attention_mask:
@@ -7,6 +7,8 @@ This module contains implementations of timestep preparation stages for diffusio
import inspect
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
@@ -71,7 +73,12 @@ class TimestepPreparationStage(PipelineStage):
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps,
# Convert timesteps to CPU if it's a tensor (for numpy conversion in scheduler)
if isinstance(timesteps, torch.Tensor):
timesteps_for_scheduler = timesteps.cpu()
else:
timesteps_for_scheduler = timesteps
scheduler.set_timesteps(timesteps=timesteps_for_scheduler,
device=device,
**extra_set_timesteps_kwargs)
timesteps = scheduler.timesteps
+3 -4
View File
@@ -126,7 +126,7 @@ class CudaPlatformBase(Platform):
logger.info("Selected backend: %s", selected_backend)
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
try:
from st_attn import sliding_tile_attention # noqa: F401
from fastvideo_kernel import sliding_tile_attention # noqa: F401
from fastvideo.attention.backends.sliding_tile_attn import ( # noqa: F401
SlidingTileAttentionBackend)
@@ -169,7 +169,7 @@ class CudaPlatformBase(Platform):
)
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
try:
from vsa import block_sparse_attn # noqa: F401
from fastvideo_kernel import video_sparse_attn # noqa: F401
from fastvideo.attention.backends.video_sparse_attn import ( # noqa: F401
VideoSparseAttentionBackend)
@@ -188,8 +188,7 @@ class CudaPlatformBase(Platform):
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
try:
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
moba_attn_varlen)
from fastvideo_kernel import moba_attn_varlen # noqa: F401
from fastvideo.attention.backends.vmoba import ( # noqa: F401
VMOBAAttentionBackend)
logger.info("Using Video MOBA Attention backend.")
+24 -2
View File
@@ -58,6 +58,13 @@ class RocmPlatform(Platform):
torch.cuda.reset_peak_memory_stats(device)
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_torch_device(cls):
"""
Return torch.cuda
"""
return torch.cuda
@classmethod
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
head_size: int, dtype: torch.dtype) -> str:
@@ -71,8 +78,23 @@ class RocmPlatform(Platform):
elif selected_backend in (AttentionBackendEnum.FLASH_ATTN, None):
pass
elif selected_backend in (AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN):
elif selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
try:
from st_attn import sliding_tile_attention # noqa: F401
from fastvideo.attention.backends.sliding_tile_attn import ( # noqa: F401
SlidingTileAttentionBackend)
logger.info("Using Sliding Tile Attention backend.")
return "fastvideo.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
except ImportError as e:
logger.error(
"Failed to import Sliding Tile Attention backend: %s",
str(e))
raise ImportError(
"Sliding Tile Attention backend is not installed. ") from e
elif selected_backend in (AttentionBackendEnum.SAGE_ATTN):
raise ValueError(
f"{selected_backend.name} is not supported on {cls.device_name}."
)
@@ -0,0 +1,140 @@
# SPDX-License-Identifier: Apache-2.0
import os
import numpy as np
import pytest
import torch
from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, T5EncoderModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.models.encoders import T5Config
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
@pytest.fixture
def t5_model_paths():
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
model_path = maybe_download_model(base_model_path)
text_encoder_path = os.path.join(model_path, "text_encoder_2")
tokenizer_path = os.path.join(model_path, "tokenizer_2")
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_t5_encoder(t5_model_paths):
# Initialize the two model implementations
text_encoder_path, tokenizer_path = t5_model_paths
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision_str = "fp32"
precision = PRECISION_TO_TYPE[precision_str]
model1 = T5EncoderModel.from_pretrained(text_encoder_path).to(
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
model2 = model2.to(precision)
model2.eval()
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Check number of parameters
logger.info("Model1 has %s parameters", len(params1))
logger.info("Model2 has %s parameters", len(params2))
# check if embed_tokens are the same
weights = ["encoder.block.{}.layer.0.layer_norm.weight", \
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight"]
for idx in range(hf_config.num_hidden_layers):
for w in weights:
name1 = w.format(idx)
name2 = w.format(idx)
p1 = params1[name1]
p2 = params2[name2]
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
# Test with some sample prompts
prompts = [
"Once upon a time", "The quick brown fox jumps over",
"In a galaxy far, far away"
]
logger.info("Testing T5 encoder with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info("Testing prompt: %s", prompt)
# Tokenize the prompt
tokens = tokenizer(prompt,
padding="max_length",
max_length=512,
truncation=True,
add_special_tokens=True,
return_tensors="pt").to(device)
# Get outputs from HuggingFace implementation
# filter out padding input_ids
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
outputs1 = model1(input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask.float())[0]
print("--------------------------------")
logger.info("Testing model2")
# Get outputs from our implementation
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs2 = model2(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
).last_hidden_state
# Compare last hidden states
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info("Maximum difference in last hidden states: %s",
max_diff_hidden.item())
logger.info("Mean difference in last hidden states: %s",
mean_diff_hidden.item())
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
# Check if outputs are similar (allowing for small numerical differences)
assert mean_diff_hidden < 1e-4, \
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
assert max_diff_hidden < 1e-4, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
@@ -0,0 +1,150 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pytest
import torch
from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, Qwen2_5_VLTextModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.models.encoders import Qwen2_5_VLConfig
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29505"
@pytest.fixture
def qwen_model_path():
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
model_path = maybe_download_model(base_model_path)
text_encoder_path = os.path.join(model_path, "text_encoder")
tokenizer_path = os.path.join(model_path, "tokenizer")
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_qwen2_5_encoder(qwen_model_path):
text_encoder_path, tokenizer_path = qwen_model_path
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
# Qwen2.5-VL default dtype is usually bf16
precision_str = "fp32"
precision = PRECISION_TO_TYPE[precision_str]
logger.info(f"Using precision: {precision_str}")
# Load HF model (Base model)
model1 = Qwen2_5_VLTextModel.from_pretrained(text_encoder_path).to(
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
# Load FastVideo model
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=PipelineConfig(text_encoder_configs=(Qwen2_5_VLConfig(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
model2 = model2.to(precision)
model2.eval()
# Sanity check weights
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
logger.info("Model1 has %s parameters", len(params1))
logger.info("Model2 has %s parameters", len(params2))
# Check common layers like Norms which are likely not merged/sharded in a way that changes name significantly
# or simple linear layers if names match.
# Note: FastVideo uses QKVParallelLinear, so q_proj, k_proj, v_proj are merged.
# HF Qwen2_5_VL uses separate projections? No, usually they are separate nn.Linear in HF.
weights_to_check = [
"norm.weight",
"layers.{}.self_attn.o_proj.weight",
"layers.{}.input_layernorm.weight",
"layers.{}.post_attention_layernorm.weight",
"layers.{}.mlp.down_proj.weight"
]
for idx in range(hf_config.num_hidden_layers):
for w in weights_to_check:
name1 = w.format(idx)
name2 = w.format(idx)
p1 = params1[name1]
p2 = params2[name2]
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
# Check shape
assert p1.shape == p2.shape, f"Shape mismatch for {w}: {p1.shape} vs {p2.shape}"
# Check values
assert_close(p1, p2, atol=1e-7, rtol=1e-7, msg=f"Weight mismatch for {w}")
# Test with sample prompts
prompts = [
"Hello world",
"The quick brown fox jumps over the lazy dog."
]
logger.info("Testing with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info(f"Prompt: {prompt}")
tokens = tokenizer(prompt, return_tensors="pt", padding="max_length", max_length=1000, truncation=True).to(device)
# HF Forward
# AutoModel for Qwen2.5-VL usually returns BaseModelOutputWithPast
outputs1 = model1(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
).hidden_states[-3]
# FastVideo Forward
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs2 = model2(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
).hidden_states[-3]
# Compare
# Filter padding for comparison if needed, but here we just check raw output matching
# Check shapes
assert outputs1.shape == outputs2.shape, f"Output shape mismatch: {outputs1.shape} vs {outputs2.shape}"
diff = torch.abs(outputs1 - outputs2)
max_diff = diff.max().item()
mean_diff = diff.mean().item()
logger.info(f"Max diff: {max_diff}")
logger.info(f"Mean diff: {mean_diff}")
# Thresholds
# Qwen2.5-VL RoPE is complex, if our implementation is slightly off (e.g. float32 conversion logic in RoPE),
# differences might appear. But should be small.
if precision_str == "bf16":
atol = 5e-2 # relaxed for bf16
else:
atol = 1e-3
if max_diff > atol:
logger.warning(f"Max diff {max_diff} > {atol}. Checking if it's acceptable...")
# If mean diff is small, maybe just outliers
assert mean_diff < atol, f"Mean diff {mean_diff} too high"
else:
logger.info("Outputs match within tolerance.")
+9 -2
View File
@@ -4,6 +4,7 @@ app = modal.App()
import os
model_vol = modal.Volume.from_name("hf-model-weights")
image_version = os.getenv("IMAGE_VERSION")
image_tag = f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{image_version}"
print(f"Using image: {image_tag}")
@@ -74,9 +75,15 @@ def run_vae_tests():
def run_transformer_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
@app.function(
gpu="L40S:4",
image=image,
timeout=2700,
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol}
)
def run_ssim_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
run_test("export MODEL_PATH='/root/data/weights' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
@@ -0,0 +1,91 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import pytest
import torch
from diffusers import AutoencoderKLHunyuanVideo15
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.logger import init_logger
# from fastvideo.models.vaes.hunyuanvae import (
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.loader.component_loader import VAELoader
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
"data", BASE_MODEL_PATH))
VAE_PATH = os.path.join(MODEL_PATH, "vae")
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
@pytest.mark.usefixtures("distributed_setup")
def test_hunyuan_vae():
device = torch.device("cuda:0")
precision = torch.float32
precision_str = "fp32"
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=Hunyuan15VAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_cpu_offload = False
model1 = AutoencoderKLHunyuanVideo15.from_pretrained(
VAE_PATH, torch_dtype=precision).to(device).eval()
model1.enable_tiling()
loader = VAELoader()
model2 = loader.load(VAE_PATH, args)
model2.enable_tiling()
batch_size = 1
# Video input [B, C, T, H, W]
input_tensor = torch.randn(batch_size,
3,
81,
512,
512,
device=device,
dtype=precision)
# Disable gradients for inference
with torch.no_grad():
latent1 = model1.encode(input_tensor, return_dict=False)[0].mode()
latent2 = model2.encode(input_tensor).mode()
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
max_diff_encode = torch.max(torch.abs(latent1.float() - latent2.float()))
mean_diff_encode = torch.mean(torch.abs(latent1.float() - latent2.float()))
logger.info("Maximum difference between encoded latents: %s",
max_diff_encode.item())
logger.info("Mean difference between encoded latents: %s",
mean_diff_encode.item())
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {max_diff_encode.item()}, mean diff = {mean_diff_encode.item()}"
# Test decoding
latent1 = latent1 / model1.config.scaling_factor
latent2 = latent2 / model2.config.scaling_factor
with torch.no_grad():
video1 = model1.decode(latent1, return_dict=False)[0]
video2 = model2.decode(latent2)
assert video1.shape == video2.shape, f"Video shapes don't match: {video1.shape} vs {video2.shape}"
max_diff_decode = torch.max(torch.abs(video1.float() - video2.float()))
mean_diff_decode = torch.mean(torch.abs(video1.float() - video2.float()))
logger.info("Maximum difference between decoded videos: %s",
max_diff_decode.item())
logger.info("Mean difference between decoded videos: %s",
mean_diff_decode.item())
assert max_diff_decode < 1e-5, f"Decoded videos differ significantly: max diff = {max_diff_decode.item()}, mean diff = {mean_diff_decode.item()}"
@@ -0,0 +1,3 @@
"""Block-sparse attention kernels for LongCat."""
@@ -0,0 +1,656 @@
import os
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
import math
from .common import _attn_fwd_gating, _attn_bwd_preprocess, configs_gating_preset
from .flash_attn_bsa_varlen_mask import (
_attn_fwd_bsa_varlen, _attn_fwd_bsa_varlen_align, _attn_bwd_dkdv_bsa_varlen_wrapper, _attn_bwd_dq_bsa_varlen_wrapper, _attn_bwd_dq_bsa_varlen_align_wrapper,
configs_fwd_bsa_varlen_preset, configs_fwd_bsa_varlen_align_preset, configs_bwd_dkdv_bsa_varlen_preset, configs_bwd_dq_bsa_varlen_preset, configs_bwd_dq_bsa_varlen_align_preset
)
torch._dynamo.config.cache_size_limit = 32
def is_cuda():
return triton.runtime.driver.active.get_current_target().backend == "cuda"
def supports_tma():
return is_cuda() and torch.cuda.get_device_capability()[0] >= 9
HAS_TMA_DESC = "nv_tma_desc_type" in dir(tl)
if HAS_TMA_DESC:
print("TMA benchmarks will be running with experimental grid constant TMA descriptor.", )
else:
print("TMA benchmarks will be running without grid constant TMA descriptor.", )
# TmaAutoTuneHelper used in htyu's PR #5622
class TmaAutoTuneHelper:
# duck typing wrapper to implement the same interface as TmaDescKernelParam in Triton PR #4498
class KernelParamWrapper:
def __init__(self, desc):
self.desc = desc
def tma_desc_cpu_ptr(self):
return self.desc.data_ptr()
TMA_SIZE = 128
def __init__(self):
self.fill_1d_tma_descriptor_inner = (triton.runtime.driver.active.utils.fill_1d_tma_descriptor)
self.fill_2d_tma_descriptor_inner = (triton.runtime.driver.active.utils.fill_2d_tma_descriptor)
if HAS_TMA_DESC:
self.descriptors = {}
else:
self.cuda_descriptors = {}
# Call this method outside of the lambda function for grid size
def init_tma_descriptor(self, name):
if HAS_TMA_DESC:
self.descriptors[name] = torch.empty(TmaAutoTuneHelper.TMA_SIZE, device="cpu", dtype=torch.int8)
else:
self.cuda_descriptors[name] = torch.empty(TmaAutoTuneHelper.TMA_SIZE, device="cuda", dtype=torch.int8)
# Call this method inside the lambda function for grid size
def fill_1d_tma_descriptor(self, name, ptr, dim, block_dim, element_size):
if HAS_TMA_DESC:
desc_x = self.descriptors[name]
assert desc_x.data_ptr() % 64 == 0
self.fill_1d_tma_descriptor_inner(ptr, dim, block_dim, element_size, desc_x.data_ptr())
else:
desc_x = self.cuda_descriptors[name]
buf_x = torch.empty_like(desc_x, device="cpu", pin_memory=True)
self.fill_1d_tma_descriptor_inner(ptr, dim, block_dim, element_size, buf_x.data_ptr())
desc_x.copy_(buf_x, non_blocking=True)
# Call this method inside the lambda function for grid size
def fill_2d_tma_descriptor(self, name, ptr, dim1, dim0, block_dim1, block_dim0, element_size):
if HAS_TMA_DESC:
desc_x = self.descriptors[name]
assert desc_x.data_ptr() % 64 == 0
self.fill_2d_tma_descriptor_inner(ptr, dim1, dim0, block_dim1, block_dim0, element_size, desc_x.data_ptr())
else:
desc_x = self.cuda_descriptors[name]
buf_x = torch.empty_like(desc_x, device="cpu", pin_memory=True)
self.fill_2d_tma_descriptor_inner(ptr, dim1, dim0, block_dim1, block_dim0, element_size, buf_x.data_ptr())
desc_x.copy_(buf_x, non_blocking=True)
def get_tma_descriptor_kernel_param(self, name):
if HAS_TMA_DESC:
assert self.descriptors[name] is not None
return self.KernelParamWrapper(self.descriptors[name])
else:
assert self.cuda_descriptors[name] is not None
return self.cuda_descriptors[name]
@triton.jit
def create_mask_from_indices_kernel(
block_indices,
block_mask,
stride_bz, stride_bh, stride_bm, stride_bs,
stride_mz, stride_mh, stride_mm, stride_mn,
H,
):
i_zh, i_m, i_s = tl.program_id(0), tl.program_id(1), tl.program_id(2)
i_z, i_h = i_zh // H, i_zh % H
off_b = i_z.to(tl.int64) * stride_bz + i_h.to(tl.int64) * stride_bh + i_m.to(tl.int64) * stride_bm + i_s.to(tl.int64) * stride_bs
b_i = tl.load(block_indices + off_b)
off_m = i_z.to(tl.int64) * stride_mz + i_h.to(tl.int64) * stride_mh + i_m.to(tl.int64) * stride_mm + b_i.to(tl.int64) * stride_mn
b_m = 1
tl.store(block_mask + off_m, b_m.to(block_mask.dtype.element_ty))
def create_mask_from_indices_triton(
block_indices,
N_cols
):
B, H, N_rows, S = block_indices.shape
block_mask = torch.zeros((B, H, N_rows, N_cols), dtype=torch.bool, device=block_indices.device)
create_mask_from_indices_kernel[(B * H, N_rows, S)](
block_indices,
block_mask,
block_indices.stride(0), block_indices.stride(1), block_indices.stride(2), block_indices.stride(3),
block_mask.stride(0), block_mask.stride(1), block_mask.stride(2), block_mask.stride(3),
H,
)
return block_mask
@torch.compile
def create_mask_from_indices_varlen(block_indices, N_cols_mask):
B, H, M, _ = block_indices.shape
device = block_indices.device
mask = torch.zeros((B, H, M, N_cols_mask), dtype=torch.bool, device=device)
valid = block_indices < N_cols_mask
b_idx = torch.arange(B, device=device)[:, None, None, None].expand_as(block_indices)
h_idx = torch.arange(H, device=device)[None, :, None, None].expand_as(block_indices)
m_idx = torch.arange(M, device=device)[None, None, :, None].expand_as(block_indices)
valid_coords = (b_idx[valid], h_idx[valid], m_idx[valid], block_indices[valid])
mask[valid_coords] = True
return mask
@torch.compile
def create_indices_k_from_indices_q_varlen(
block_indices,
N_cols_mask # indicate the number of the last dimension of the bool mask, since this information cannot be determined by block_indices, which may contain invalid elements
):
block_mask_qk = create_mask_from_indices_varlen(block_indices, N_cols_mask)
B, H, M, N = block_mask_qk.shape
block_mask_kq = block_mask_qk.permute(0, 1, 3, 2)
indices = torch.arange(M, device=block_indices.device).view(1, 1, 1, -1).expand_as(block_mask_kq)
block_indices_k = torch.where(block_mask_kq, indices, M)
block_indices_k, _ = torch.sort(block_indices_k, dim=-1)
block_indices_k_lens = (block_indices_k < M).sum(dim=-1)
return block_indices_k, block_indices_k_lens
@torch.compile
def mean_pooling_compression(
x: torch.Tensor,
block_size: int
) -> torch.Tensor:
B, H, S = x.shape[:3]
num_block = math.ceil(S / block_size)
if S % block_size != 0:
x = F.pad(x, (0, 0, 0, num_block * block_size - S))
x_cmp = x.view(B, H, num_block, block_size, -1).mean(dim=3)
return x_cmp
@torch.compile
def cal_score(q, k):
k_transposed = k.transpose(-1, -2) # [b, h, d, s_k]
score = torch.matmul(q, k_transposed) # [b, h, s_q, s_k]
return score
def cal_score_triton(q, k):
B, H, s_q, D = q.shape
s_k = k.shape[2]
score = torch.empty(B, H, s_q, s_k, device=q.device, dtype=q.dtype)
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_gating_preset['default']
grid = lambda args: (triton.cdiv(s_q, args["BLOCK_M"]), B * H, 1)
_attn_fwd_gating[grid](
q, k, score,
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
score.stride(0), score.stride(1), score.stride(2), score.stride(3),
H, s_q, s_k,
HEAD_DIM=D,
**kernel_config
)
return score
@torch.compile
def get_select_indices_topk(q, k, sparsity):
score = cal_score(q, k)
block_indices, block_indices_lens = get_select_indices_topk_from_score(score, sparsity)
return block_indices, block_indices_lens
@torch.compile
def get_select_indices_topk_from_score(score, sparsity):
num_selected = int((1 - sparsity) * score.shape[-1])
block_indices = torch.topk(score, num_selected)[1]
block_indices_lens = torch.full(
(block_indices.shape[0], block_indices.shape[1], block_indices.shape[2]),
num_selected,
dtype=torch.int32,
device=block_indices.device
)
return block_indices, block_indices_lens
@torch.compile
def get_select_indices_cdf(q, k, cdf_threshold):
score = cal_score(q, k)
head_dim = q.shape[-1]
block_indices, block_indices_lens = get_select_indices_cdf_from_score(score, cdf_threshold, 1 / head_dim**0.5)
return block_indices, block_indices_lens
@torch.compile
def get_select_indices_cdf_from_score(score, cdf_threshold, sm_scale):
weights = torch.softmax(score * sm_scale, dim=-1)
B, H, Sq, Sk = weights.shape
cdf_threshold = torch.full((H,), cdf_threshold, device=weights.device).view(1, H, 1, 1).expand(B, -1, Sq, -1)
weights_sorted = torch.sort(weights, dim=-1, descending=True)
cdf = torch.cumsum(weights_sorted.values, dim=-1)
num_selected = torch.searchsorted(cdf, cdf_threshold, right=True)
return weights_sorted.indices, num_selected.squeeze(-1)
@torch.compile
def get_select_indices_cdf_topk(q, k, sparsity, cdf_threshold):
score = cal_score(q, k)
head_dim = q.shape[-1]
block_indices, block_indices_lens = get_select_indices_cdf_topk_from_score(score, sparsity, cdf_threshold, 1 / head_dim**0.5)
return block_indices, block_indices_lens
@torch.compile
def get_select_indices_cdf_topk_from_score(score, sparsity, cdf_threshold, sm_scale):
weights = torch.softmax(score * sm_scale, dim=-1)
B, H, Sq, Sk = weights.shape
cdf_threshold = torch.full((H,), cdf_threshold, device=weights.device).view(1, H, 1, 1).expand(B, -1, Sq, -1)
weights_sorted = torch.sort(weights, dim=-1, descending=True)
cdf = torch.cumsum(weights_sorted.values, dim=-1)
num_selected = torch.searchsorted(cdf, cdf_threshold, right=True)
# max(cdf, topk)
num_selected_topk = int((1 - sparsity) * score.shape[-1])
num_selected[num_selected < num_selected_topk] = num_selected_topk
return weights_sorted.indices, num_selected.squeeze(-1)
def get_select_indices(q, k, sparsity, cdf_threshold):
if sparsity is not None and cdf_threshold is None:
block_indices, block_indices_lens = get_select_indices_topk(q, k, sparsity)
elif sparsity is None and cdf_threshold is not None:
block_indices, block_indices_lens = get_select_indices_cdf(q, k, cdf_threshold)
elif sparsity is not None and cdf_threshold is not None:
block_indices, block_indices_lens = get_select_indices_cdf_topk(q, k, sparsity, cdf_threshold)
else:
raise ValueError
return block_indices, block_indices_lens
def get_select_indices_from_score(score, sparsity, cdf_threshold):
if sparsity is not None and cdf_threshold is None:
block_indices, block_indices_lens = get_select_indices_topk_from_score(score, sparsity)
elif sparsity is None and cdf_threshold is not None:
block_indices, block_indices_lens = get_select_indices_cdf_from_score(score, cdf_threshold)
elif sparsity is not None and cdf_threshold is not None:
block_indices, block_indices_lens = get_select_indices_cdf_topk_from_score(score, sparsity, cdf_threshold)
else:
raise ValueError
return block_indices, block_indices_lens
def attn_fwd_bsa_varlen_triton(
q,
k,
v,
sm_scale,
block_indices,
block_indices_lens,
chunk_size_q,
chunk_size_k,
sparsity
):
B, H, Seq, D = q.shape
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
grid = lambda args: (triton.cdiv(q.shape[2], args["BLOCK_M"]), q.shape[0] * q.shape[1], 1)
config_key = 'BLOCK_N_LG=64' if chunk_size_k == 64 else 'default'
if chunk_size_k > 128:
fwd_func = _attn_fwd_bsa_varlen
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_fwd_bsa_varlen_preset[config_key]
else:
fwd_func = _attn_fwd_bsa_varlen_align
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_fwd_bsa_varlen_align_preset[config_key]
block_indices = block_indices.contiguous()
block_indices_lens = block_indices_lens.contiguous()
fwd_func[grid](
q, k, v, sm_scale, M, o,
block_indices, # [B, H, M_COMPRESS, S]
block_indices_lens, # [B, H, M_COMPRESS, S_MAX]
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),
block_indices.stride(0), block_indices.stride(1), block_indices.stride(2), block_indices.stride(3),
block_indices_lens.stride(0), block_indices_lens.stride(1), block_indices_lens.stride(2),
H, Seq,
D,
BLOCK_M=chunk_size_q,
BLOCK_N_LG=chunk_size_k,
SPARSITY=sparsity,
**kernel_config
)
LN2 = 0.6931471824645996
lse = M * LN2 # convert back to natural units (M is of base 2)
return o, lse
def attn_bwd_bsa_varlen_triton(
do,
q,
k,
v,
o,
dq,
dk,
dv,
sm_scale,
M,
block_indices,
block_indices_lens,
chunk_size_q,
chunk_size_k,
sparsity
):
RCP_LN2 = 1.4426950408889634
M = M * RCP_LN2 # ln -> log2
do = do.contiguous()
# assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
BATCH, N_HEAD, N_CTX, HEAD_DIM = q.shape
N_CTX_KV = k.shape[-2]
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2) # reciprocal
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
if min(chunk_size_q, chunk_size_k) >= 128:
PRE_BLOCK = 128
else:
PRE_BLOCK = min(chunk_size_q, chunk_size_k)
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,
N_CTX,
BLOCK_M=PRE_BLOCK,
HEAD_DIM=HEAD_DIM
)
block_indices_k, block_indices_k_lens = create_indices_k_from_indices_q_varlen(
block_indices=block_indices,
N_cols_mask=N_CTX_KV // chunk_size_k
)
block_indices = block_indices.contiguous()
block_indices_lens = block_indices_lens.contiguous()
block_indices_k = block_indices_k.contiguous()
block_indices_k_lens = block_indices_k_lens.contiguous()
config_key = 'BLOCK_N_DQ_LG=64' if chunk_size_k == 64 else 'default'
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_bwd_dkdv_bsa_varlen_preset[config_key]
grid_dkdv = lambda args: (triton.cdiv(arg_k.shape[2], args["BLOCK_N"]), 1, arg_k.shape[0] * arg_k.shape[1])
_attn_bwd_dkdv_bsa_varlen_wrapper[grid_dkdv](
q, arg_k, v, sm_scale, # softmax scale
do,
dk, dv,
M, # lse (log2)
delta,
block_indices_k,
block_indices_k_lens,
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),
dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3),
dv.stride(0), dv.stride(1), dv.stride(2), dv.stride(3),
do.stride(0), do.stride(1), do.stride(2), do.stride(3),
M.stride(0), M.stride(1), M.stride(2),
delta.stride(0), delta.stride(1), delta.stride(2),
block_indices_k.stride(0), block_indices_k.stride(1), block_indices_k.stride(2), block_indices_k.stride(3),
block_indices_k_lens.stride(0), block_indices_k_lens.stride(1), block_indices_k_lens.stride(2),
N_HEAD, N_CTX,
BLOCK_M=chunk_size_q,
BLOCK_N_DQ_LG=chunk_size_k,
HEAD_DIM=HEAD_DIM,
SPARSITY=sparsity,
**kernel_config
)
config_key = 'BLOCK_N_DQ_LG=64' if chunk_size_k == 64 else 'default'
if chunk_size_k > 128:
bwd_dq_func = _attn_bwd_dq_bsa_varlen_wrapper
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_bwd_dq_bsa_varlen_preset[config_key]
else:
bwd_dq_func = _attn_bwd_dq_bsa_varlen_align_wrapper
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_bwd_dq_bsa_varlen_align_preset[config_key]
grid_dq = lambda args: (triton.cdiv(q.shape[2], args["BLOCK_M"]), 1, q.shape[0] * q.shape[1])
bwd_dq_func[grid_dq](
q, arg_k, v,
do,
dq,
M, # lse (log2)
delta,
block_indices,
block_indices_lens,
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),
dq.stride(0), dq.stride(1), dq.stride(2), dq.stride(3),
do.stride(0), do.stride(1), do.stride(2), do.stride(3),
M.stride(0), M.stride(1), M.stride(2),
delta.stride(0), delta.stride(1), delta.stride(2),
block_indices.stride(0), block_indices.stride(1), block_indices.stride(2), block_indices.stride(3),
block_indices_lens.stride(0), block_indices_lens.stride(1), block_indices_lens.stride(2),
N_HEAD, N_CTX,
BLOCK_M=chunk_size_q,
BLOCK_N_DQ_LG=chunk_size_k,
HEAD_DIM=HEAD_DIM,
SPARSITY=sparsity,
**kernel_config
)
@torch.compile
def make_block_indices_varlen_cp_list(block_indices, cp_size, num_blocks_k_full):
"""
Args:
block_indices: [B, H, num_blocks_q_per_cp_rank, num_blocks_k_full]
Return:
a list of [block_indices, block_indices_lens] for k from each cp_rank
- each block_indices starts from zero
- block_indices_lens indicates the valid number of elements in the last dimension of block_indices
"""
res = []
num_blocks_per_rank = num_blocks_k_full // cp_size
for i in range(cp_size):
block_indices_tmp = block_indices.clone()
min_block_idx = i * num_blocks_per_rank
block_indices_tmp -= min_block_idx
block_indices_tmp[block_indices_tmp < 0] = num_blocks_per_rank # block_indices_tmp < 0 indicate invalid indices, set them to num_blocks_per_rank in order to sort them to the tail, so that the first N elements of the block_indices indicated by block_indices_lens are valid
block_indices_tmp, _ = torch.sort(block_indices_tmp, dim=-1)
block_indices_tmp_lens = (block_indices_tmp < num_blocks_per_rank).sum(dim=-1)
res.append([block_indices_tmp, block_indices_tmp_lens])
return res
@torch.compile
def flash_attn_fwd_softmax_lse_correction(
softmax_lse: torch.Tensor,
softmax_lse_per_step: torch.Tensor,
):
"""Merge softmax stats of each step in Attention with context parallelism"""
max_scale = torch.max(softmax_lse, softmax_lse_per_step)
min_scale = torch.min(softmax_lse, softmax_lse_per_step)
lse_diff = min_scale - max_scale
lse_diff = lse_diff.nan_to_num(nan=0.) # handle cases: tensor(-inf) - tensor(-inf) = tensor(nan); In the current cp implementation, it is possible that lses of 2 cp ranks are both -inf, if no block is selected from both cp ranks. In such cases, the finally corrected lse should remain -inf.
new_scale = max_scale + torch.log1p(torch.exp(lse_diff)) # a + ln(1 + e^(b - a)) = ln(e^a) + ln(1 + e^(b - a)) = ln(e^a + e^b)
softmax_lse.copy_(new_scale)
@torch.compile
def flash_attn_fwd_out_correction_init(
out_init_step: torch.Tensor, # b h s d
softmax_lse: torch.Tensor, # b h s
softmax_lse_init_step: torch.Tensor,
):
"""Merge partial outputs of the first step in Attention with context parallelism"""
softmax_lse_corrected_exp = torch.exp(softmax_lse_init_step - softmax_lse)
softmax_lse_corrected_exp = softmax_lse_corrected_exp.unsqueeze(-1)
out_corrected = out_init_step * softmax_lse_corrected_exp
return out_corrected.to(out_init_step.dtype)
@torch.compile
def flash_attn_fwd_out_correction(
out: torch.Tensor,
out_per_step: torch.Tensor,
softmax_lse: torch.Tensor,
softmax_lse_per_step: torch.Tensor,
):
"""Merge partial outputs of each step in Attention with context parallelism"""
softmax_lse_corrected_exp = torch.exp(softmax_lse_per_step - softmax_lse)
softmax_lse_corrected_exp = softmax_lse_corrected_exp.unsqueeze(-1)
out_corrected = out_per_step * softmax_lse_corrected_exp
out.add_(out_corrected)
@torch.compile
def topk_sort(score, num_chunks_selected):
block_indices = torch.topk(score, num_chunks_selected)[1]
block_indices, _ = torch.sort(block_indices, dim=-1)
return block_indices
class _attention_bsa(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, chunk_size_q, chunk_size_k, sparsity, cdf_threshold, sm_scale, use_tma=False):
# shape constraints
HEAD_DIM_Q, HEAD_DIM_K = q.shape[-1], k.shape[-1]
# when v is in float8_e5m2 it is transposed.
HEAD_DIM_V = v.shape[-1]
assert HEAD_DIM_Q == HEAD_DIM_K and HEAD_DIM_K == HEAD_DIM_V
assert HEAD_DIM_K in {16, 32, 64, 128, 256}
# ---------------------- gating ----------------------
q_cmp = mean_pooling_compression(q, chunk_size_q)
k_cmp = mean_pooling_compression(k, chunk_size_k)
block_indices, block_indices_lens = get_select_indices(q_cmp, k_cmp, sparsity, cdf_threshold)
# ---------------------- bsa ----------------------
o, lse = attn_fwd_bsa_varlen_triton(
q, k, v,
sm_scale, block_indices, block_indices_lens,
chunk_size_q, chunk_size_k,
sparsity
)
ctx.save_for_backward(q, k, v, o, lse, block_indices, block_indices_lens)
ctx.sm_scale = sm_scale
ctx.HEAD_DIM = HEAD_DIM_K
ctx.chunk_size_q = chunk_size_q
ctx.chunk_size_k = chunk_size_k
ctx.use_tma = use_tma
ctx.sparsity = sparsity
return o
@staticmethod
def backward(ctx, do):
q, k, v, o, lse, block_indices, block_indices_lens = ctx.saved_tensors
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
attn_bwd_bsa_varlen_triton(
do,
q,
k,
v,
o,
dq,
dk,
dv,
ctx.sm_scale,
lse,
block_indices,
block_indices_lens,
ctx.chunk_size_q,
ctx.chunk_size_k,
ctx.sparsity
)
return dq, dk, dv, None, None, None, None, None, None
flash_attn_bsa = _attention_bsa.apply
def rearrange_THW_to_3d_block(x, Nt, Nh, Nw, t, h, w, D):
B, H, _, D = x.shape
x = x.view(B, H, Nt, t, Nh, h, Nw, w, D)
x = x.permute(0, 1, 2, 4, 6, 3, 5, 7, 8) # B H Nt Nh Nw t h w D
return x.contiguous().view(B, H, Nt * Nh * Nw * t * h * w, D)
def rearrange_3d_block_to_THW(x, Nt, Nh, Nw, t, h, w, D):
B, H, _, D = x.shape
x = x.view(B, H, Nt, Nh, Nw, t, h, w, D)
x = x.permute(0, 1, 2, 5, 3, 6, 4, 7, 8) # B H Nt t Nh h Nw w D
return x.contiguous().view(B, H, Nt * t * Nh * h * Nw * w, D)
def flash_attn_bsa_3d(
q: torch.Tensor, # [B, H, Sq, D]
k: torch.Tensor, # [B, H, Skv, D]
v: torch.Tensor, # [B, H, Skv, D]
latent_shape_q,
latent_shape_k,
# bsa_params
sparsity=0.875,
cdf_threshold=None,
chunk_3d_shape_q=[4, 4, 8],
chunk_3d_shape_k=[4, 4, 8],
) -> torch.Tensor:
_, _, Sq, head_dim_q = q.shape
_, _, Sk, head_dim_k = k.shape
assert head_dim_q == head_dim_k
head_dim = head_dim_q
Tq, Hq, Wq = latent_shape_q
Tk, Hk, Wk = latent_shape_k
assert Tq * Hq * Wq == Sq
assert Tk * Hk * Wk == Sk
tq, hq, wq = chunk_3d_shape_q
tk, hk, wk = chunk_3d_shape_k
assert Tq % tq == 0 and Hq % hq == 0 and Wq % wq == 0
assert Tk % tk == 0 and Hk % hk == 0 and Wk % wk == 0
Ntq = Tq // tq
Nhq = Hq // hq
Nwq = Wq // wq
Ntk = Tk // tk
Nhk = Hk // hk
Nwk = Wk // wk
q = rearrange_THW_to_3d_block(q, Ntq, Nhq, Nwq, tq, hq, wq, q.shape[-1])
k = rearrange_THW_to_3d_block(k, Ntk, Nhk, Nwk, tk, hk, wk, k.shape[-1])
v = rearrange_THW_to_3d_block(v, Ntk, Nhk, Nwk, tk, hk, wk, v.shape[-1])
chunk_size_q = tq * hq * wq
chunk_size_k = tk * hk * wk
output = flash_attn_bsa(q, k, v, chunk_size_q, chunk_size_k, sparsity, cdf_threshold, 1 / head_dim**0.5)
output = rearrange_3d_block_to_THW(output, Ntq, Nhq, Nwq, tq, hq, wq, output.shape[-1])
return output
@@ -0,0 +1,111 @@
import triton
import triton.language as tl
import os
if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1':
autotune = triton.autotune
else:
def autotune(*args, **kwargs):
def decorator(func):
return func
return decorator
configs_gating_preset = {
'default': {
'BLOCK_M': 64,
'BLOCK_N': 64,
'num_stages': 3,
'num_warps': 8,
}
}
configs_gating = [
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BM in [64, 128] \
for BN in [32, 64] \
for s in [2, 3, 4, 5] \
for w in [4, 8] \
]
gating_reevaluate_keys = ["M", "N"] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
@autotune(configs_gating, key=gating_reevaluate_keys)
@triton.jit
def _attn_fwd_gating(
Q, K, Out,
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_oz, stride_oh, stride_om, stride_on,
H, M, N,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
tl.static_assert(BLOCK_N <= HEAD_DIM)
start_m = tl.program_id(0)
off_hz = tl.program_id(1)
off_z = off_hz // H
off_h = off_hz % H
q_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
k_offset = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
o_offset = off_z.to(tl.int64) * stride_oz + off_h.to(tl.int64) * stride_oh
# block pointers
Q_block_ptr = tl.make_block_ptr(
base=Q + q_offset,
shape=(M, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
K_block_ptr = tl.make_block_ptr(
base=K + k_offset,
shape=(HEAD_DIM, N),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N),
order=(0, 1),
)
O_block_ptr = tl.make_block_ptr(
base=Out + o_offset,
shape=(M, N),
strides=(stride_om, stride_on),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, BLOCK_N),
order=(1, 0),
)
# load q: it will stay in SRAM throughout
q = tl.load(Q_block_ptr, boundary_check=(0,))
for start_n in range(0, N, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
# -- compute qk ----
k = tl.load(K_block_ptr, boundary_check=(1,))
qk = tl.dot(q, k)
tl.store(O_block_ptr, qk.to(Out.type.element_ty), boundary_check=(0, 1))
K_block_ptr = tl.advance(K_block_ptr, (0, BLOCK_N))
O_block_ptr = tl.advance(O_block_ptr, (0, BLOCK_N))
@triton.jit
def _attn_bwd_preprocess(
O, DO,
Delta, # output
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)
@@ -0,0 +1,43 @@
import torch
def p2p_communicate(
rank, send_tensor, send_dst, recv_tensor, recv_src, cp_group, batch_p2p_comm
):
"""Point-to-point communications of KV and dKV in Attention with context parallelism"""
send_recv_ops = []
if batch_p2p_comm: # int(os.getenv("NVTE_BATCH_MHA_P2P_COMM", "0")) or (cp_size == 2) 为啥呢
if rank % 2 == 0:
send_op = torch.distributed.P2POp(
torch.distributed.isend, send_tensor, send_dst, cp_group
)
recv_op = torch.distributed.P2POp(
torch.distributed.irecv, recv_tensor, recv_src, cp_group
)
send_recv_ops.append(send_op)
send_recv_ops.append(recv_op)
else:
recv_op = torch.distributed.P2POp(
torch.distributed.irecv, recv_tensor, recv_src, cp_group
)
send_op = torch.distributed.P2POp(
torch.distributed.isend, send_tensor, send_dst, cp_group
)
send_recv_ops.append(recv_op)
send_recv_ops.append(send_op)
send_recv_reqs = torch.distributed.batch_isend_irecv(send_recv_ops)
else:
if rank % 2 == 0:
send_op = torch.distributed.isend(send_tensor, send_dst, cp_group)
recv_op = torch.distributed.irecv(recv_tensor, recv_src, cp_group)
send_recv_ops.append(send_op)
send_recv_ops.append(recv_op)
else:
recv_op = torch.distributed.irecv(recv_tensor, recv_src, cp_group)
send_op = torch.distributed.isend(send_tensor, send_dst, cp_group)
send_recv_ops.append(recv_op)
send_recv_ops.append(send_op)
send_recv_reqs = send_recv_ops
return send_recv_reqs
@@ -0,0 +1,946 @@
import triton
import triton.language as tl
import os
from .common import autotune
"""
TRITON_REEVALUATE_KEY=1
- autotune whenever params in reevaluate keys change
- use in benchmark script to fine the best config
TRITON_AUTOTUNE_ENBALE=1
- if set to 0, autotune will not work, and the related params must be passed to the function call.
"""
configs_fwd_bsa_varlen_preset = {
'default': {
'BLOCK_N': 64,
'num_stages': 3,
'num_warps': 8,
},
'BLOCK_N_LG=64': {
'BLOCK_N': 64,
'num_stages': 3,
'num_warps': 4,
},
}
configs_fwd_bsa_varlen = [
triton.Config({'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BN in [32, 64, 128] \
for s in [2, 3, 4, 5] \
for w in [4, 8] \
]
fwd_bsa_reevaluate_varlen_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
@autotune(list(configs_fwd_bsa_varlen), key=fwd_bsa_reevaluate_varlen_keys)
@triton.jit
def _attn_fwd_bsa_varlen(
Q, K, V, sm_scale, M, Out,
block_indices, # [B, H, M_COMPRESS, S_MAX]
block_indices_lens, # [B, H, M_COMPRESS]
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vn, stride_vk,
stride_oz, stride_oh, stride_om, stride_ok,
stride_bz, stride_bh, stride_bm, stride_bs,
stride_lz, stride_lh, stride_lm,
H, N_CTX,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N_LG: tl.constexpr,
BLOCK_N: tl.constexpr,
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
):
start_m = tl.program_id(0)
off_hz = tl.program_id(1)
off_z = off_hz // H
off_h = off_hz % H
q_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
k_offset = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
v_offset = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
o_offset = off_z.to(tl.int64) * stride_oz + off_h.to(tl.int64) * stride_oh
b_offset = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
l_offset = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
# block pointers
Q_block_ptr = tl.make_block_ptr(
base=Q + q_offset,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
V_block_ptr = tl.make_block_ptr(
base=V + v_offset,
shape=(N_CTX, HEAD_DIM),
strides=(stride_vn, stride_vk),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=(1, 0),
)
KT_block_ptr = tl.make_block_ptr(
base=K + k_offset,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N),
order=(0, 1),
)
O_block_ptr = tl.make_block_ptr(
base=Out + o_offset,
shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_ok),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
block_indices += b_offset + start_m * stride_bm
block_indices_lens += l_offset + start_m * stride_lm
# initialize offsets
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
# initialize pointer to m and l
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
# load scales
qk_scale = sm_scale
qk_scale *= 1.44269504 # 1/ln2; exp2(x/ln2) == exp2(ln(e^x) / ln2) == exp2(log2(e^x)) == exp(x)
# load q: it will stay in SRAM throughout
q = tl.load(Q_block_ptr)
S = tl.load(block_indices_lens)
for i in range(S):
block_id = tl.load(block_indices + i * stride_bs).to(tl.int32)
lo, hi = block_id * BLOCK_N_LG, (block_id + 1) * BLOCK_N_LG
lo = tl.multiple_of(lo, BLOCK_N)
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
V_block_ptr_i = tl.advance(V_block_ptr, (lo, 0))
# loop over k, v and update accumulator
for start_n in range(lo, hi, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
# -- compute qk ----
kT = tl.load(KT_block_ptr_i)
qkT = tl.dot(q, kT)
m_ij = tl.maximum(m_i, tl.max(qkT, 1) * qk_scale)
qkT = qkT * qk_scale - m_ij[:, None]
p = tl.math.exp2(qkT)
# -- update m_i and l_i
alpha = tl.math.exp2(m_i - m_ij)
l_ij = tl.sum(p, 1)
# -- update output accumulator --
acc = acc * alpha[:, None]
# update acc
v = tl.load(V_block_ptr_i)
acc = tl.dot(p.to(v.dtype), v, acc)
# update m_i and l_i
# place this at the end of the loop to reduce register pressure: https://github.com/triton-lang/triton/commit/ee6abd9
l_i = l_i * alpha + l_ij
m_i = m_ij
V_block_ptr_i = tl.advance(V_block_ptr_i, (BLOCK_N, 0))
KT_block_ptr_i = tl.advance(KT_block_ptr_i, (0, BLOCK_N))
# epilogue
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
m_ptrs = M + off_hz * N_CTX + offs_m
tl.store(m_ptrs, m_i)
tl.store(O_block_ptr, acc.to(Out.type.element_ty))
configs_fwd_bsa_varlen_align_preset = {
'default': {
'num_stages': 3,
'num_warps': 8,
},
'BLOCK_N_LG=64': {
'num_stages': 3,
'num_warps': 4,
},
}
configs_fwd_bsa_varlen_align = [
triton.Config({}, num_stages=s, num_warps=w) \
for s in [2, 3, 4, 5] \
for w in [4, 8] \
]
fwd_bsa_reevaluate_varlen_align_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
@autotune(list(configs_fwd_bsa_varlen_align), key=fwd_bsa_reevaluate_varlen_align_keys)
@triton.jit
def _attn_fwd_bsa_varlen_align(
Q, K, V, sm_scale, M, Out,
block_indices, # [B, H, M_COMPRESS, S_MAX]
block_indices_lens, # [B, H, M_COMPRESS]
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vn, stride_vk,
stride_oz, stride_oh, stride_om, stride_on,
stride_bz, stride_bh, stride_bm, stride_bs,
stride_lz, stride_lh, stride_lm,
H, N_CTX,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N_LG: tl.constexpr,
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
):
start_m = tl.program_id(0)
off_hz = tl.program_id(1)
off_z = off_hz // H
off_h = off_hz % H
q_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
k_offset = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
v_offset = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
o_offset = off_z.to(tl.int64) * stride_oz + off_h.to(tl.int64) * stride_oh
b_offset = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
l_offset = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
# block pointers
Q_block_ptr = tl.make_block_ptr(
base=Q + q_offset,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
V_block_ptr = tl.make_block_ptr(
base=V + v_offset,
shape=(N_CTX, HEAD_DIM),
strides=(stride_vn, stride_vk),
offsets=(0, 0),
block_shape=(BLOCK_N_LG, HEAD_DIM),
order=(1, 0),
)
KT_block_ptr = tl.make_block_ptr(
base=K + k_offset,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N_LG),
order=(0, 1),
)
O_block_ptr = tl.make_block_ptr(
base=Out + o_offset,
shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
block_indices += b_offset + start_m * stride_bm
block_indices_lens += l_offset + start_m * stride_lm
# initialize offsets
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
# initialize pointer to m and l
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
# load scales
qk_scale = sm_scale
qk_scale *= 1.44269504 # 1/ln2; exp2(x/ln2) == exp2(ln(e^x) / ln2) == exp2(log2(e^x)) == exp(x);乘1/ln2后,exp2(x/ln2) == exp(x),exp2速度更快
# load q: it will stay in SRAM throughout
q = tl.load(Q_block_ptr)
S = tl.load(block_indices_lens)
for i in range(S):
block_id = tl.load(block_indices + i * stride_bs).to(tl.int32)
lo = block_id * BLOCK_N_LG
lo = tl.multiple_of(lo, BLOCK_N_LG)
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
V_block_ptr_i = tl.advance(V_block_ptr, (lo, 0))
# -- compute qk ----
kT = tl.load(KT_block_ptr_i)
qkT = tl.dot(q, kT)
m_ij = tl.maximum(m_i, tl.max(qkT, 1) * qk_scale)
qkT = qkT * qk_scale - m_ij[:, None]
p = tl.math.exp2(qkT)
# -- update m_i and l_i
alpha = tl.math.exp2(m_i - m_ij)
l_ij = tl.sum(p, 1)
# -- update output accumulator --
acc = acc * alpha[:, None]
# update acc
v = tl.load(V_block_ptr_i)
acc = tl.dot(p.to(v.dtype), v, acc) # 没除se,fa2引入的优化
# update m_i and l_i
# place this at the end of the loop to reduce register pressure: https://github.com/triton-lang/triton/commit/ee6abd9
l_i = l_i * alpha + l_ij # 当前总se
m_i = m_ij
# epilogue
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
m_ptrs = M + off_hz * N_CTX + offs_m
tl.store(m_ptrs, m_i)
tl.store(O_block_ptr, acc.to(Out.type.element_ty))
# The main inner-loop logic for computing dK and dV.
@triton.jit
def _attn_bwd_dkdv_bsa_varlen(
dk, dv,
k, v,
Q, DO,
M, D,
block_indices,
block_indices_lens,
# shared by Q/K/V/DO.
# stride_tok, stride_d,
stride_qm, stride_qk,
stride_dom, stride_dok,
stride_mm,
stride_dm,
stride_bm,
N_CTX,
BLOCK_M1: tl.constexpr,
HEAD_DIM: tl.constexpr,
):
QT_block_ptr = tl.make_block_ptr(
base=Q,
shape=(HEAD_DIM, N_CTX),
strides=(stride_qk, stride_qm),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_M1),
order=(0, 1),
)
DO_block_ptr = tl.make_block_ptr(
base=DO,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dom, stride_dok),
offsets=(0, 0),
block_shape=(BLOCK_M1, HEAD_DIM),
order=(1, 0),
)
S = tl.load(block_indices_lens)
for i in range(S):
block_id = tl.load(block_indices + i * stride_bm).to(tl.int32)
start_m = block_id * BLOCK_M1
start_m = tl.multiple_of(start_m, BLOCK_M1)
QT_block_ptr_i = tl.advance(QT_block_ptr, (0, start_m))
DO_block_ptr_i = tl.advance(DO_block_ptr, (start_m, 0))
qT = tl.load(QT_block_ptr_i)
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + tl.arange(0, BLOCK_M1) * stride_mm
m = tl.load(M + offs_m)
kqT = tl.dot(k, qT)
pT = tl.math.exp2(kqT - m[None, :])
do = tl.load(DO_block_ptr_i)
# Compute dV.
ppT = pT
ppT = ppT.to(v.dtype)
dv += tl.dot(ppT, do)
# D (= delta) is pre-divided by ds_scale.
offs_d = start_m + tl.arange(0, BLOCK_M1) * stride_dm
Di = tl.load(D + offs_d)
# Compute dP and dS.
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
dsT = pT * (dpT - Di[None, :])
dsT = dsT.to(v.dtype)
dk += tl.dot(dsT, tl.trans(qT))
return dk, dv
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq_bsa_varlen(
dq,
q, do,
m, d,
K, V,
N_CTX,
BLOCK_N2: tl.constexpr,
BLOCK_N_LG: tl.constexpr,
HEAD_DIM: tl.constexpr,
block_indices,
block_indices_lens,
stride_bn,
# stride_tok, stride_d,
stride_kn, stride_kk,
stride_vn, stride_vk,
):
VT_block_ptr = tl.make_block_ptr(
base=V,
shape=(HEAD_DIM, N_CTX),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N2),
order=(0, 1),
)
KT_block_ptr = tl.make_block_ptr(
base=K,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N2),
order=(0, 1),
)
S = tl.load(block_indices_lens)
for i in range(S):
block_id = tl.load(block_indices + i * stride_bn).to(tl.int32)
lo, hi = block_id * BLOCK_N_LG, (block_id + 1) * BLOCK_N_LG
lo = tl.multiple_of(lo, BLOCK_N2)
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
VT_block_ptr_i = tl.advance(VT_block_ptr, (0, lo))
for start_n in range(lo, hi, BLOCK_N2):
start_n = tl.multiple_of(start_n, BLOCK_N2)
kT = tl.load(KT_block_ptr_i)
vT = tl.load(VT_block_ptr_i)
qkT = tl.dot(q, kT)
p = tl.math.exp2(qkT - m)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - d)
ds = ds.to(kT.dtype) # https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/flash_attn_triton.py: Converting ds to q.dtype here reduces register pressure and makes it much faster for BLOCK_HEADDIM=128
# 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.
KT_block_ptr_i = tl.advance(KT_block_ptr_i, (0, BLOCK_N2))
VT_block_ptr_i = tl.advance(VT_block_ptr_i, (0, BLOCK_N2))
return dq
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq_bsa_varlen_align(
dq,
q, do,
m, d,
K, V,
N_CTX,
BLOCK_N_LG: tl.constexpr,
HEAD_DIM: tl.constexpr,
block_indices,
block_indices_lens,
stride_bn,
stride_kn, stride_kk,
stride_vn, stride_vk,
):
VT_block_ptr = tl.make_block_ptr(
base=V,
shape=(HEAD_DIM, N_CTX),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N_LG),
order=(0, 1),
)
KT_block_ptr = tl.make_block_ptr(
base=K,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N_LG),
order=(0, 1),
)
S = tl.load(block_indices_lens)
for i in range(S):
block_id = tl.load(block_indices + i * stride_bn).to(tl.int32)
lo = block_id * BLOCK_N_LG
lo = tl.multiple_of(lo, BLOCK_N_LG)
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
VT_block_ptr_i = tl.advance(VT_block_ptr, (0, lo))
kT = tl.load(KT_block_ptr_i)
vT = tl.load(VT_block_ptr_i)
qkT = tl.dot(q, kT)
p = tl.math.exp2(qkT - m)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - d)
ds = ds.to(kT.dtype) # https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/flash_attn_triton.py: Converting ds to q.dtype here reduces register pressure and makes it much faster for BLOCK_HEADDIM=128
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
dq += tl.dot(ds, tl.trans(kT))
return dq
configs_bwd_dkdv_bsa_varlen_preset = {
'default': {
'BLOCK_N': 128,
'num_stages': 2,
'num_warps': 8,
},
'BLOCK_N_DQ_LG=64': {
'BLOCK_N': 64,
'num_stages': 2,
'num_warps': 4,
}
}
configs_bwd_dkdv_bsa_varlen = [
triton.Config({'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BN in [32, 64, 128] \
for s in [2, 3, 4, 5] \
for w in [4, 8] \
]
bwd_dkdv_bsa_varlen_reevaluate_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_DQ_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
@autotune(list(configs_bwd_dkdv_bsa_varlen), key=bwd_dkdv_bsa_varlen_reevaluate_keys)
@triton.jit
def _attn_bwd_dkdv_bsa_varlen_wrapper(
Q, K, V, sm_scale, # softmax scale
DO,
DK, DV,
M, # lse (log2)
D,
block_indices,
block_indices_lens,
# stride_z, stride_h, stride_tok, stride_d, # shared by Q/K/V/DO.
# qkv
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vn, stride_vk,
# dk dv do
stride_dkz, stride_dkh, stride_dkn, stride_dkk,
stride_dvz, stride_dvh, stride_dvn, stride_dvk,
stride_doz, stride_doh, stride_dom, stride_dok,
# m, d
stride_mz, stride_mh, stride_mm,
stride_dz, stride_dh, stride_dm,
#
stride_bz, stride_bh, stride_bn, stride_bm, # block_indices
stride_lz, stride_lh, stride_ln, # block_indices_lens
#
H, N_CTX,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_N_DQ_LG: tl.constexpr, # logical block size
HEAD_DIM: tl.constexpr,
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
):
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_N == 0)
start_n = tl.program_id(0)
off_hz = tl.program_id(2)
off_z = off_hz // H
off_h = off_hz % H
off_q = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
off_k = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
off_v = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
off_dk = off_z.to(tl.int64) * stride_dkz + off_h.to(tl.int64) * stride_dkh
off_dv = off_z.to(tl.int64) * stride_dvz + off_h.to(tl.int64) * stride_dvh
off_do = off_z.to(tl.int64) * stride_doz + off_h.to(tl.int64) * stride_doh
off_m = off_z.to(tl.int64) * stride_mz + off_h.to(tl.int64) * stride_mh
off_d = off_z.to(tl.int64) * stride_dz + off_h.to(tl.int64) * stride_dh
off_block_incides = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
off_block_incides_lens = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
# offset pointers for batch/head
Q += off_q
K += off_k
V += off_v
DO += off_do
DK += off_dk
DV += off_dv
M += off_m
D += off_d
block_indices += off_block_incides
block_indices_lens += off_block_incides_lens
# ---------------------------------------- [DKDV] ----------------------------------------
dv = tl.zeros([BLOCK_N, HEAD_DIM], dtype=tl.float32)
dk = tl.zeros([BLOCK_N, HEAD_DIM], dtype=tl.float32)
# load K and V: they stay in SRAM throughout the inner loop.
K_block_ptr = tl.make_block_ptr(
base=K,
shape=(N_CTX, HEAD_DIM),
strides=(stride_kn, stride_kk),
offsets=(start_n * BLOCK_N, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=(1, 0),
)
V_block_ptr = tl.make_block_ptr(
base=V,
shape=(N_CTX, HEAD_DIM),
strides=(stride_vn, stride_vk),
offsets=(start_n * BLOCK_N, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=(1, 0),
)
DK_block_ptr = tl.make_block_ptr(
base=DK,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dkn, stride_dkk),
offsets=(start_n * BLOCK_N, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=(1, 0),
)
DV_block_ptr = tl.make_block_ptr(
base=DV,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dvn, stride_dvk),
offsets=(start_n * BLOCK_N, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=(1, 0),
)
k = tl.load(K_block_ptr)
v = tl.load(V_block_ptr)
k_compress_idx = start_n * BLOCK_N // BLOCK_N_DQ_LG
block_indices_i = block_indices + k_compress_idx * stride_bn
block_indices_lens_i = block_indices_lens + k_compress_idx * stride_ln
dk, dv = _attn_bwd_dkdv_bsa_varlen(
dk, dv,
k, v,
Q, DO,
M, D,
block_indices_i,
block_indices_lens_i,
# shared by Q/K/V/DO.
stride_qm, stride_qk,
stride_dom, stride_dok,
stride_mm,
stride_dm,
#
stride_bm,
N_CTX,
BLOCK_M,
HEAD_DIM,
)
# Write back dk
dk *= sm_scale # S = scale * QKT; dK = scale * QdST
tl.store(DK_block_ptr, dk.to(k.dtype))
# Write back dv
tl.store(DV_block_ptr, dv.to(v.dtype))
configs_bwd_dq_bsa_varlen_preset = {
'default': {
'BLOCK_N_DQ': 64,
'num_stages': 2,
'num_warps': 8,
},
'BLOCK_N_DQ_LG=64': {
'BLOCK_N_DQ': 64,
'num_stages': 2,
'num_warps': 4,
},
}
configs_bwd_dq_bsa_varlen = [
triton.Config({'BLOCK_N_DQ': BN}, num_stages=s, num_warps=w) \
for BN in [32, 64, 128] \
for s in [2, 3, 4, 5] \
for w in [4, 8] \
]
bwd_dq_bsa_varlen_reevaluate_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_DQ_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
@autotune(list(configs_bwd_dq_bsa_varlen), key=bwd_dq_bsa_varlen_reevaluate_keys)
@triton.jit
def _attn_bwd_dq_bsa_varlen_wrapper(
Q, K, V, # softmax scale
DO,
DQ,
M, # lse (log2)
D,
block_indices,
block_indices_lens,
# stride_z, stride_h, stride_tok, stride_d, # shared by Q/K/V/DO.
# qkv
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vn, stride_vk,
# dq do
stride_dqz, stride_dqh, stride_dqm, stride_dqk,
stride_doz, stride_doh, stride_dom, stride_dok,
# m, d
stride_mz, stride_mh, stride_mm,
stride_dz, stride_dh, stride_dm,
#
stride_bz, stride_bh, stride_bm, stride_bn, # block_indices
stride_lz, stride_lh, stride_lm, # block_indices_lens
#
H, N_CTX,
BLOCK_M: tl.constexpr,
BLOCK_N_DQ_LG: tl.constexpr, # logical block size
BLOCK_N_DQ: tl.constexpr,
HEAD_DIM: tl.constexpr,
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
):
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_N_DQ == 0)
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_M == 0)
LN2: tl.constexpr = 0.6931471824645996 # = ln(2)
start_m = tl.program_id(0)
off_hz = tl.program_id(2)
off_z = off_hz // H
off_h = off_hz % H
off_q = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
off_k = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
off_v = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
off_dq = off_z.to(tl.int64) * stride_dqz + off_h.to(tl.int64) * stride_dqh
off_do = off_z.to(tl.int64) * stride_doz + off_h.to(tl.int64) * stride_doh
off_m = off_z.to(tl.int64) * stride_mz + off_h.to(tl.int64) * stride_mh
off_d = off_z.to(tl.int64) * stride_dz + off_h.to(tl.int64) * stride_dh
off_block_incides = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
off_block_incides_lens = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
# offset pointers for batch/head
Q += off_q
K += off_k
V += off_v
DO += off_do
DQ += off_dq
M += off_m
D += off_d
block_indices += off_block_incides
block_indices_lens += off_block_incides_lens
# ---------------------------------------- [DQ] ----------------------------------------
Q_block_ptr = tl.make_block_ptr(
base=Q,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
DO_block_ptr = tl.make_block_ptr(
base=DO,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dom, stride_dok),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
DQ_block_ptr = tl.make_block_ptr(
base=DQ,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dqm, stride_dqk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
q = tl.load(Q_block_ptr)
do = tl.load(DO_block_ptr)
dq = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
start_m = start_m * BLOCK_M
offs_m = start_m + tl.arange(0, BLOCK_M) * stride_mm
offs_d = start_m + tl.arange(0, BLOCK_M) * stride_dm
m = tl.load(M + offs_m)
m = m[:, None]
d = tl.load(D + offs_d)
d = d[:, None]
block_indices_m = block_indices + (start_m // BLOCK_M) * stride_bm
block_indices_lens_m = block_indices_lens + (start_m // BLOCK_M) * stride_lm
dq = _attn_bwd_dq_bsa_varlen(
dq,
q, do,
m, d,
K, V,
N_CTX,
BLOCK_N_DQ,
BLOCK_N_DQ_LG,
HEAD_DIM,
block_indices_m,
block_indices_lens_m,
stride_bn,
stride_kn, stride_kk,
stride_vn, stride_vk,
)
# Write back dQ.
dq *= LN2
tl.store(DQ_block_ptr, dq.to(q.dtype))
configs_bwd_dq_bsa_varlen_align_preset = {
'default': {
'num_stages': 2,
'num_warps': 8,
},
'BLOCK_N_DQ_LG=64': {
'num_stages': 2,
'num_warps': 4,
},
}
configs_bwd_dq_bsa_varlen_align = [
triton.Config({}, num_stages=s, num_warps=w) \
for s in [2, 3, 4, 5] \
for w in [4, 8] \
]
bwd_dq_bsa_varlen_align_reevaluate_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_DQ_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
@autotune(list(configs_bwd_dq_bsa_varlen_align), key=bwd_dq_bsa_varlen_align_reevaluate_keys)
@triton.jit
def _attn_bwd_dq_bsa_varlen_align_wrapper(
Q, K, V, # softmax scale
DO,
DQ,
M, # lse (log2)
D,
block_indices,
block_indices_lens,
# qkv
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vn, stride_vk,
# dq do
stride_dqz, stride_dqh, stride_dqm, stride_dqk,
stride_doz, stride_doh, stride_dom, stride_dok,
# m, d
stride_mz, stride_mh, stride_mm,
stride_dz, stride_dh, stride_dm,
#
stride_bz, stride_bh, stride_bm, stride_bn, # block_indices
stride_lz, stride_lh, stride_lm, # block_indices_lens
#
H, N_CTX,
BLOCK_M: tl.constexpr,
BLOCK_N_DQ_LG: tl.constexpr, # logical block size
HEAD_DIM: tl.constexpr,
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
):
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_N_DQ_LG == 0)
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_M == 0)
LN2: tl.constexpr = 0.6931471824645996 # = ln(2)
start_m = tl.program_id(0)
off_hz = tl.program_id(2)
off_z = off_hz // H
off_h = off_hz % H
off_q = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
off_k = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
off_v = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
off_dq = off_z.to(tl.int64) * stride_dqz + off_h.to(tl.int64) * stride_dqh
off_do = off_z.to(tl.int64) * stride_doz + off_h.to(tl.int64) * stride_doh
off_m = off_z.to(tl.int64) * stride_mz + off_h.to(tl.int64) * stride_mh
off_d = off_z.to(tl.int64) * stride_dz + off_h.to(tl.int64) * stride_dh
off_block_incides = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
off_block_incides_lens = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
# offset pointers for batch/head
Q += off_q
K += off_k
V += off_v
DO += off_do
DQ += off_dq
M += off_m
D += off_d
block_indices += off_block_incides
block_indices_lens += off_block_incides_lens
# ---------------------------------------- [DQ] ----------------------------------------
Q_block_ptr = tl.make_block_ptr(
base=Q,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
DO_block_ptr = tl.make_block_ptr(
base=DO,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dom, stride_dok),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
DQ_block_ptr = tl.make_block_ptr(
base=DQ,
shape=(N_CTX, HEAD_DIM),
strides=(stride_dqm, stride_dqk),
offsets=(start_m * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0),
)
q = tl.load(Q_block_ptr)
do = tl.load(DO_block_ptr)
dq = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
start_m = start_m * BLOCK_M
offs_m = start_m + tl.arange(0, BLOCK_M) * stride_mm
offs_d = start_m + tl.arange(0, BLOCK_M) * stride_dm
m = tl.load(M + offs_m)
m = m[:, None]
# D (= delta) is pre-divided by ds_scale.
d = tl.load(D + offs_d)
d = d[:, None]
block_indices_m = block_indices + (start_m // BLOCK_M) * stride_bm
block_indices_lens_m = block_indices_lens + (start_m // BLOCK_M) * stride_lm
dq = _attn_bwd_dq_bsa_varlen_align(
dq,
q, do,
m, d,
K, V,
N_CTX,
BLOCK_N_DQ_LG,
HEAD_DIM,
block_indices_m,
block_indices_lens_m,
stride_bn,
stride_kn, stride_kk,
stride_vn, stride_vk,
)
# Write back dQ.
dq *= LN2
tl.store(DQ_block_ptr, dq.to(q.dtype))

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