Compare commits

..
80 changed files with 102 additions and 10417 deletions
@@ -1,236 +0,0 @@
name: Publish FastVideo Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/fastvideo_kernel/pyproject.toml"
workflow_dispatch:
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
cd csrc/fastvideo_kernel
# Get current commit's version from pyproject.toml
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build_wheels:
name: Build Wheel
needs: check-version-change
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12', '3.13']
torch-cuda:
- torch-version: '2.5.1'
cuda-version: '12.4.1'
torch-cuda-short: 'cu124'
- torch-version: '2.6.0'
cuda-version: '12.6.3'
torch-cuda-short: 'cu126'
- torch-version: '2.7.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
steps:
- name: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.torch-cuda.cuda-version }}
linux-local-args: '["--toolkit"]'
method: 'network'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
run: |
pip install --upgrade pip
pip install typing-extensions==4.12.2
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
pip install setuptools ninja packaging wheel triton
cd csrc/fastvideo_kernel
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/fastvideo_kernel
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
# Get the correct version format
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
# Rename with version information
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
- name: Upload wheel artifact
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}-py${{ matrix.python-version }}
path: csrc/fastvideo_kernel/dist/*.whl
retention-days: 90
publish_package:
name: Publish package
needs: [build_wheels, check-version-change]
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Install CUDA 12.4.1
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: 12.4.1
linux-local-args: '["--toolkit"]'
method: 'network'
sub-packages: '["nvcc"]'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-12.4.1
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch 2.5.1+cu12.4.1
run: |
pip install --upgrade pip
pip install typing-extensions==4.12.2
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
pip install setuptools ninja packaging wheel triton
cd csrc/fastvideo_kernel
git submodule update --init --recursive
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/fastvideo_kernel/dist/
-2
View File
@@ -30,8 +30,6 @@ env
**/build/
**.pyc
**.txt
*.log
weights/
# Distribution / packaging
build/
+7 -14
View File
@@ -2,12 +2,12 @@
# Attention Kernel Used in FastVideo
## Sliding Tile Attention (STA)
We support H100 (via TK) and any other GPU (via triton) for STA.
We only support H100 for STA.
### Installation
```bash
pip install st_attn
```
```
Install from source:
@@ -16,14 +16,6 @@ git submodule update --init --recursive
python setup.py install
```
If you want to skip the compilation of the TK kernel and only use the Triton version, try below:
```bash
SKIP_SM90_EXT=1 python setup.py install
or
SKIP_SM90_EXT=1 pip install --no-build-isolation .
```
If you encounter error during installation, try below:
Install C++20 for ThunderKittens:
```bash
@@ -38,7 +30,7 @@ sudo apt install clang-11
(If you use CUDA12.8)
```bash
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
@@ -51,7 +43,7 @@ bash scripts/inference/v1_inference_wan_STA.sh
If you want to use sliding tile attention in your custom model:
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
# a tile is a cube of size (6, 8, 8)
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
@@ -66,6 +58,7 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
### Test
```bash
python ../tests/test_sta.py # test STA
python ../tests/test_vsa.py # test VSA
```
### Benchmark
```bash
@@ -74,7 +67,7 @@ python ../benchmarks/bench_sta.py
### How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
@@ -89,7 +82,7 @@ Here is a diagram of how the window is configured and passed through the FastVid
## Why is STA Fast?
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
STA removes mixed blocks.
+9 -16
View File
@@ -51,28 +51,21 @@ for k in kernels:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
ext_modules = []
if os.environ.get("SKIP_SM90_EXT", "0") != "1":
ext_modules.append(
CUDAExtension('st_attn_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
)
else:
print("ENV SKIP_SM90_EXT=1, skip st_attn_cuda compile")
setup(name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
ext_modules=ext_modules,
ext_modules=[
CUDAExtension('st_attn_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
],
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
@@ -7,17 +7,12 @@ try:
except ImportError:
sta_fwd = None
try:
from st_attn.st_attn_triton import sliding_tile_attention_triton
except ImportError:
sliding_tile_attention_triton = None
def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
seq_length = q_all.shape[2]
dit_seq_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
@@ -51,13 +46,4 @@ def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, h
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
return hidden_states[:, :, :seq_length]
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
major, minor = torch.cuda.get_device_capability(q_all.device)
if major == 9 and minor == 0 and sta_fwd is not None:
return sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
elif sliding_tile_attention_triton is not None:
return sliding_tile_attention_triton(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
else:
raise ImportError("No suitable sliding tile attention implementation found.")
return hidden_states[:, :, :seq_length]
@@ -1,327 +0,0 @@
import math
import torch
import triton
import triton.language as tl
def is_cuda():
return triton.runtime.driver.active.get_current_target().backend == "cuda"
def is_hip():
target = triton.runtime.driver.active.get_current_target()
return target.backend == 'hip'
def get_common_autotune_config():
configs = [
triton.Config({'BLOCK_Q': BLOCK_Q, 'BLOCK_KV': BLOCK_KV}, num_stages=s, num_warps=w) \
for BLOCK_Q in [32, 64, 128]\
for BLOCK_KV in [32, 64, 128]\
for s in [1, 2, 3, 4]\
for w in [4, 8]\
]
return configs
def get_cuda_autotune_config():
# cuda and hip can use differnt autotune configs
return get_common_autotune_config()
def get_hip_autotune_config():
# cuda and hip can use differnt autotune configs
return get_common_autotune_config()
def get_autotune_config():
if is_cuda():
return get_cuda_autotune_config()
else:
return get_hip_autotune_config()
@triton.jit
def clamp_int(value, min_val, max_val):
ret = tl.where(value > max_val, max_val, value)
ret = tl.where(ret < min_val, min_val, ret)
return ret
@triton.jit
def _attn_fwd_loop(
q, k, v, kv_mask, m, l, acc, sm_scale,
MASK_KV: tl.constexpr,
):
scores = tl.dot(q, k.T) #[BLOCK_Q, BLOCK_KV]
scores = scores * sm_scale
if MASK_KV:
scores = tl.where(kv_mask[None, :], scores, -float('inf'))
current_m = tl.max(scores, axis=1)
new_m = tl.maximum(m, current_m)
exp_scores = tl.math.exp2(scores - new_m[:, None])
current_l = tl.sum(exp_scores, axis=1)
# Update L <- L * exp(M - M') + L1, M <- M'
alpha = tl.math.exp2(m - new_m)
l = l * alpha + current_l
m = new_m
# Update O <- O * exp(M - M') + P @ V
acc = (acc * alpha[:, None] + tl.dot(exp_scores.to(v.type.element_ty), v))
return m, l, acc
@triton.autotune(
configs=get_autotune_config(),
key=['head_dim'],
)
@triton.jit
def triton_sta_kernel(
Q, K, V, output,
batch_size: int, num_heads: int, seq_len: int, head_dim: int,
img_seq_len: int,
text_length: int,
canvas_t: int, canvas_h: int, canvas_w: int,
kernel_t: int, kernel_h: int, kernel_w: int,
tile_t: int, tile_h: int, tile_w: int,
scale: float,
has_text: tl.constexpr,
text_q: tl.constexpr,
BLOCK_Q: tl.constexpr,
BLOCK_KV: tl.constexpr,
BLOCK_DIM: tl.constexpr,
):
total_tile_size = tile_t * tile_h * tile_w
q_block_per_tile = (total_tile_size + BLOCK_Q - 1) // BLOCK_Q
batch_idx = tl.program_id(0)
head_idx = tl.program_id(1)
if text_q:
q_block_idx = tl.program_id(2)
else:
q_tile_flat = tl.program_id(2) // q_block_per_tile
q_block_idx = tl.program_id(2) % q_block_per_tile
m = tl.full((BLOCK_Q,), -float('inf'), dtype=tl.float32)
l = tl.zeros((BLOCK_Q,), dtype=tl.float32)
acc = tl.zeros((BLOCK_Q, BLOCK_DIM), dtype=tl.float32)
q_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
if text_q:
q_base_idx = img_seq_len + q_block_idx * BLOCK_Q
else:
q_base_idx = q_tile_flat * total_tile_size + q_block_idx * BLOCK_Q
q_offset_in_tile = tl.arange(0, BLOCK_Q)
q_idx = q_base_idx + q_offset_in_tile
q_mask = (q_block_idx * BLOCK_Q + tl.arange(0, BLOCK_Q)) < total_tile_size
q = tl.load(
Q + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
mask=q_mask[:, None],
other=0.0
) # [BLOCK_Q, BLOCK_DIM]
# Scale sm_scale by log_2(e) and use 2^x instead of exp
sm_scale = scale * 1.4426950408889634
num_tiles_t = canvas_t // tile_t
num_tiles_h = canvas_h // tile_h
num_tiles_w = canvas_w // tile_w
tiles_per_hw = num_tiles_h * num_tiles_w
if text_q:
kv_tile_start_t = 0
kv_tile_end_t = num_tiles_t
kv_tile_start_h = 0
kv_tile_end_h = num_tiles_h
kv_tile_start_w = 0
kv_tile_end_w = num_tiles_w
else:
q_tile_t = q_tile_flat // tiles_per_hw
remaining = q_tile_flat % tiles_per_hw
q_tile_h = remaining // num_tiles_w
q_tile_w = remaining % num_tiles_w
kernel_center_t = clamp_int(q_tile_t, kernel_t // 2, (num_tiles_t - 1) - kernel_t // 2)
kernel_center_h = clamp_int(q_tile_h, kernel_h // 2, (num_tiles_h - 1) - kernel_h // 2)
kernel_center_w = clamp_int(q_tile_w, kernel_w // 2, (num_tiles_w - 1) - kernel_w // 2)
kv_tile_start_t = kernel_center_t - kernel_t // 2
kv_tile_end_t = kernel_center_t + kernel_t // 2 + 1
kv_tile_end_t = tl.where(kv_tile_end_t > num_tiles_t, num_tiles_t, kv_tile_end_t)
kv_tile_start_h = kernel_center_h - kernel_h // 2
kv_tile_end_h = kernel_center_h + kernel_h // 2 + 1
kv_tile_end_h = tl.where(kv_tile_end_h > num_tiles_h, num_tiles_h, kv_tile_end_h)
kv_tile_start_w = kernel_center_w - kernel_w // 2
kv_tile_end_w = kernel_center_w + kernel_w // 2 + 1
kv_tile_end_w = tl.where(kv_tile_end_w > num_tiles_w, num_tiles_w, kv_tile_end_w)
# for kv_img
for kv_tile_t in tl.range(kv_tile_start_t, kv_tile_end_t):
for kv_tile_h in tl.range(kv_tile_start_h, kv_tile_end_h):
for kv_tile_w in tl.range(kv_tile_start_w, kv_tile_end_w):
kv_base_idx = (kv_tile_t * num_tiles_h * num_tiles_w + kv_tile_h * num_tiles_w + kv_tile_w) * total_tile_size
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
kv_offset_in_block = tl.arange(0, BLOCK_KV)
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < total_tile_size
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
k = tl.load(
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
mask=kv_mask[:, None],
other=0.0
) # [BLOCK_KV, BLOCK_DIM]
v = tl.load(
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
mask=kv_mask[:, None],
other=0.0
) # [BLOCK_KV, BLOCK_DIM]
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, False)
# for kv_text
if has_text:
kv_base_idx = img_seq_len
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
kv_offset_in_block = tl.arange(0, BLOCK_KV)
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < text_length
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
k = tl.load(
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
mask=kv_mask[:, None],
other=0.0
) # [BLOCK_KV, BLOCK_DIM]
v = tl.load(
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
mask=kv_mask[:, None],
other=0.0
) # [BLOCK_KV, BLOCK_DIM]
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, True)
output_acc = acc / l[:, None]
tl.store(
output + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
output_acc,
mask=q_mask[:, None]
) # [BLOCK_Q, BLOCK_DIM]
def sliding_tile_attention_triton(
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
window_size, text_length: int,
has_text=True, dit_seq_shape='30x48x80') -> torch.Tensor:
seq_length = q.shape[2]
if has_text:
assert q.shape[2] >= 115200 and q.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '30x48x80' for HunyuanVideo"
target_size = math.ceil(seq_length / 384) * 384
pad_size = target_size - seq_length
if pad_size > 0:
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
else:
if dit_seq_shape == '36x48x48': # Stepvideo
assert q.shape[2] == 82944
elif dit_seq_shape == '18x48x80': # Wan
assert q.shape[2] == 69120
else:
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
assert q.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
batch_size, num_heads, seq_len, head_dim = q.shape
if dit_seq_shape == '30x48x80': # Hunyuan
canvas_t, canvas_h, canvas_w = 30, 48, 80
tile_t, tile_h, tile_w = 6, 8, 8
elif dit_seq_shape == '36x48x48': # Stepvideo
canvas_t, canvas_h, canvas_w = 36, 48, 48
tile_t, tile_h, tile_w = 6, 8, 8
elif dit_seq_shape == '18x48x80': # Wan
canvas_t, canvas_h, canvas_w = 18, 48, 80
tile_t, tile_h, tile_w = 6, 8, 8
img_seq_len = canvas_t * canvas_h * canvas_w
num_tiles_t = canvas_t // tile_t
num_tiles_h = canvas_h // tile_h
num_tiles_w = canvas_w // tile_w
num_tiles = num_tiles_t * num_tiles_h * num_tiles_w
total_tile_size = tile_t * tile_h * tile_w
# BLOCK_Q=128
# BLOCK_KV=128
BLOCK_DIM = head_dim
output = torch.empty_like(q)
# for q_img
# kernel_size maybe different for different head
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (kernel_t, kernel_h, kernel_w) in enumerate(window_size):
for batch in range(batch_size):
q_head, k_head, v_head, o_head = (q[batch:batch + 1, head_index:head_index + 1],
k[batch:batch + 1, head_index:head_index + 1],
v[batch:batch + 1, head_index:head_index + 1],
output[batch:batch + 1, head_index:head_index + 1])
# triton_sta_kernel[(1, 1, num_tiles * triton.cdiv(total_tile_size, BLOCK_Q))](
grid = lambda META: (1, 1, num_tiles * triton.cdiv(total_tile_size, META['BLOCK_Q']))
triton_sta_kernel[grid](
q_head, k_head, v_head, o_head,
1, 1, seq_len, head_dim,
img_seq_len,
text_length,
canvas_t, canvas_h, canvas_w,
kernel_t, kernel_h, kernel_w,
tile_t, tile_h, tile_w,
scale=1.0 / (head_dim ** 0.5),
has_text=has_text,
text_q=False,
# BLOCK_Q=BLOCK_Q,
# BLOCK_KV=BLOCK_KV,
BLOCK_DIM=BLOCK_DIM,
)
# for q_text
# kernel_t, kernel_h, kernel_w is not used, set to (3, 3, 3)
if has_text:
# triton_sta_kernel[(batch_size, num_heads, triton.cdiv(total_tile_size, BLOCK_Q))](
grid = lambda META: (batch_size, num_heads, triton.cdiv(total_tile_size, META['BLOCK_Q']))
triton_sta_kernel[grid](
q, k, v, output,
batch_size, num_heads, seq_len, head_dim,
img_seq_len,
text_length,
canvas_t, canvas_h, canvas_w,
3, 3, 3,
#kernel_t, kernel_h, kernel_w,
tile_t, tile_h, tile_w,
scale=1.0 / (head_dim ** 0.5),
has_text=has_text,
text_q=True,
# BLOCK_Q=BLOCK_Q,
# BLOCK_KV=BLOCK_KV,
BLOCK_DIM=BLOCK_DIM,
)
if has_text:
if pad_size > 0:
output = output[:, :, :seq_length]
return output
-7
View File
@@ -1,7 +0,0 @@
build/
dist/
*.egg-info/
__pycache__/
*.so
*.pyc
.ipynb_checkpoints/
-187
View File
@@ -1,187 +0,0 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
-6
View File
@@ -1,6 +0,0 @@
include LICENSE
include README.md
include pyproject.toml
recursive-include src/fastvideo_kernel *.cu *.cuh *.cpp *.h
recursive-include csrc *.cu *.cuh *.cpp *.h
recursive-include tk *.cu *.cuh *.cpp *.h
-31
View File
@@ -1,31 +0,0 @@
# FastVideo Kernel
CUDA kernels for FastVideo video generation.
## Installation
```bash
git submodule update --init --recursive
cd csrc/fastvideo_kernel
pip install .
```
## Usage
```python
from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_attn_varlen
# Example: Sliding Tile Attention
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
# Example: Video Sparse Attention (with Triton fallback)
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
# Example: VMoBA
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
```
## Requirements
- H100 GPU (sm_90a) for CUDA kernels
- Triton for non-H100 fallback
File diff suppressed because it is too large Load Diff
-23
View File
@@ -1,23 +0,0 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ST_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
-573
View File
@@ -1,573 +0,0 @@
// # Define TORCH_COMPILE macro
#include "kittens.cuh"
#include <cooperative_groups.h>
#include <iostream>
#include <stdio.h>
#include <c10/cuda/CUDAGuard.h>
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
return (value < min) ? min : ((value > max) ? max : value);
}
// #define ABS(x) ((x) < 0 ? -(x) : (x))
__device__ __forceinline__ int abs_int(int value) {
return (value < 0) ? -value : value;
}
constexpr int CONSUMER_WARPGROUPS = (3);
constexpr int PRODUCER_WARPGROUPS = (1);
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
using namespace kittens;
namespace cg = cooperative_groups;
template<int D> struct fwd_attend_ker_tile_dims {};
template<> struct fwd_attend_ker_tile_dims<64> {
constexpr static int tile_width = (64);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (4);
};
template<> struct fwd_attend_ker_tile_dims<128> {
constexpr static int tile_width = (128);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (2);
};
template<int D> struct fwd_globals {
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
q_gl q;
k_gl k;
v_gl v;
l_gl l;
o_gl o;
const int N;
const int text_L;
const int hr;
};
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
extern __shared__ int __shm[];
tma_swizzle_allocator al((int*)&__shm[0]);
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
using K = fwd_attend_ker_tile_dims<D>;
using q_tile = st_bf<K::qo_height, K::tile_width>;
using k_tile = st_bf<K::kv_height, K::tile_width>;
using v_tile = st_bf<K::kv_height, K::tile_width>;
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
using o_tile = st_bf<K::qo_height, K::tile_width>;
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
int img_kv_blocks;
int kv_blocks = g.N / (K::kv_height);
if constexpr (text_kv) {
img_kv_blocks = kv_blocks - 3;
} else {
img_kv_blocks = kv_blocks;
}
int kv_head_idx = blockIdx.y / g.hr;
int seq_idx;
if constexpr (text_q) {
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
} else {
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
}
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
if (threadIdx.x == 0) {
init_semaphore(qsmem_semaphore, 0, 1);
for(int j = 0; j < K::stages; j++) {
init_semaphore(k_smem_arrived[j], 0, 1);
init_semaphore(v_smem_arrived[j], 0, 1);
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
}
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
}
if constexpr (text_q){
for (int j = 0; j < K::stages - 1; j++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int count = 0;
int j = 0;
while (count < K::stages - 1) {
int kt = j / 3 / (CH * CW);
int kh = (j / 3) % (CH * CW) / CW;
int kw = (j / 3) % CW;
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
if (mask){
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
count += 1;
}
j += 1;
}
}
}
__syncthreads();
int pipe_idx = K::stages - 1;
if(warpgroupid == NUM_WARPGROUPS-1) {
warpgroup::decrease_registers<32>();
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
}
else { kv_iters = kv_blocks-2;}
if(warpid == NUM_WORKERS-4) {
if constexpr (text_q){
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int k_t_min = clamp_int(qt-DT, 0, CT-1);
int k_t_max = clamp_int(qt+DT, 0, CT-1);
int k_h_min = clamp_int(qh-DH, 0, CH-1);
int k_h_max = clamp_int(qh+DH, 0, CH-1);
int k_w_min = clamp_int(qw-DW, 0, CW-1);
int k_w_max = clamp_int(qw+DW, 0, CW-1);
int count = 0;
for (int kt = k_t_min; kt <= k_t_max; kt++) {
for (int kh = k_h_min; kh <= k_h_max; kh++) {
for (int kw = k_w_min; kw <= k_w_max; kw++) {
for (int j = 0; j <= 2; j++){
if (count >= K::stages - 1) {
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
} else {
count += 1;
}
}
}
}
}
// for text
for (int index = img_kv_blocks; index < kv_blocks; index++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
}
}
}
}
else {
warpgroup::increase_registers<160>();
rt_fl<16, K::kv_height> att_block;
rt_bf<16, K::kv_height> att_block_mma;
rt_fl<16, K::tile_width> o_reg;
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
neg_infty(max_vec);
zero(norm_vec);
zero(o_reg);
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
kv_iters = (kv_iters/8);
}
else if constexpr (text_q){
// the last three kv blocks are for text, we process them separately
kv_iters = img_kv_blocks - 1;
} else {
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
// the last three kv blocks are for text, we process them separately
if constexpr(text_kv) {
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
// apply non-pad mask
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
}
div_row(o_reg, o_reg, norm_vec);
warpgroup::store(o_smem[warpgroupid], o_reg);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
}
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
log(norm_vec, norm_vec);
add(norm_vec, norm_vec, max_vec_scaled);
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
else { mul(norm_vec, norm_vec, -11.313708499f); }
warpgroup::store(l_smem[warpgroupid], norm_vec);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
}
tma::store_async_wait();
}
}
#include "pyutils/torch_helpers.cuh"
#include <ATen/cuda/CUDAContext.h>
#include <iostream>
torch::Tensor
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
CHECK_INPUT(v);
auto batch = q.size(0);
auto seq_len = q.size(2);
auto head_dim = q.size(3);
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
auto hr = qo_heads / kv_heads;
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<128>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
// Shared memory size for the kernel.
// We use the maximum available shared memory (kittens::MAX_SHARED_MEMORY)
// which is approximately 227KB on H100, necessary for the high-performance
// TMA-based attention tiles with multiple stages.
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
dim3 grid_text(2, qo_heads, batch);
if (!process_text) {
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
cudaFuncSetAttribute( \
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10>, \
cudaFuncAttributeMaxDynamicSharedMemorySize, \
mem_size \
); \
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(2, 1, 1); }
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 3, 0); }
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 1, 2); }
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 2, 2); }
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(2, 2, 3); }
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 3, 5); }
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 0, 0); }
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 0, 5); }
else {
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
}
#undef LAUNCH_IMAGE_KER
} else {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
}
} else {
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (kernel_aspect_ratio_flag == 2){
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
cudaFuncSetAttribute( \
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6>, \
cudaFuncAttributeMaxDynamicSharedMemorySize, \
mem_size \
); \
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 1, 3); }
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(3, 1, 1); }
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 3, 3); }
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 1, 3); }
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 3, 3); }
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 0, 0); }
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 0, 3); }
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 3, 0); }
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 3, 3); }
else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 0, 3); }
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(0, 3, 0); }
else {
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
}
#undef LAUNCH_IMAGE_KER
}
else if (kernel_aspect_ratio_flag == 3) {
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
cudaFuncSetAttribute( \
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10>, \
cudaFuncAttributeMaxDynamicSharedMemorySize, \
mem_size \
); \
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 3, 0); }
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(1, 2, 3); }
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(1, 2, 4); }
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 0, 0); }
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 2, 5); }
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 3, 3); }
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 2, 3); }
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(0, 2, 4); }
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 0, 5); }
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 1, 5); }
else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 1, 5); }
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(0, 3, 2); }
else {
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
}
#undef LAUNCH_IMAGE_KER
}
else {
TORCH_CHECK(false, "Unsupported kernel_aspect_ratio_flag: ", kernel_aspect_ratio_flag);
}
}
CHECK_CUDA_ERROR(cudaGetLastError());
// cudaStreamSynchronize(stream);
}
return o;
//cudadevicesynchronize();
}
-27
View File
@@ -1,27 +0,0 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_BLOCK_SPARSE
extern std::vector<torch::Tensor> block_sparse_attention_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
);
extern std::vector<torch::Tensor> block_sparse_attention_backward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_BLOCK_SPARSE
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
#endif
}
-27
View File
@@ -1,27 +0,0 @@
[build-system]
requires = ["setuptools>=61.0", "torch>=2.5.0", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "fastvideo-kernel"
version = "0.1.0"
description = "CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
license = {text = "Apache-2.0"}
authors = [{name = "Hao AI Lab"}]
classifiers = [
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
]
dependencies = [
"torch>=2.5.0",
"triton>=2.0.0"
]
[project.urls]
Repository = "https://github.com/hao-ai-lab/FastVideo"
[tool.setuptools.packages.find]
where = ["src"]
-132
View File
@@ -1,132 +0,0 @@
import os
import subprocess
import sys
from pathlib import Path
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
ROOT = Path(__file__).parent.absolute()
CSRC_DIR = ROOT / "csrc"
# Path to ThunderKittens (TK)
def get_tk_dir():
tk_env = os.getenv("THUNDERKITTENS_ROOT")
if tk_env:
return tk_env
# Check common locations
possible_paths = [
ROOT / "tk",
ROOT / "csrc" / "tk",
ROOT.parent / "attn" / "sliding_tile_attn" / "tk",
ROOT.parent / "attn" / "video_sparse_attn" / "tk",
]
for p in possible_paths:
if (p / "include" / "kittens.cuh").exists():
return str(p)
# Default fallback
return str(ROOT.parent / "attn" / "sliding_tile_attn" / "tk")
TK_DIR = get_tk_dir()
def get_cuda_flags(tk_root: str) -> list:
python_include = subprocess.check_output(
["python", "-c", "import sysconfig; print(sysconfig.get_path('include'))"]
).decode().strip()
torch_includes = subprocess.check_output([
"python", "-c",
"import torch; from torch.utils.cpp_extension import include_paths; "
"print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip().split()
return [
"-DNDEBUG",
"-Xcompiler=-Wno-psabi",
"-Xcompiler=-fno-strict-aliasing",
"--expt-extended-lambda",
"--expt-relaxed-constexpr",
"-forward-unknown-to-host-compiler",
"--use_fast_math",
"-std=c++20",
"-O3",
"-Xnvlink=--verbose",
"-Xptxas=--verbose",
"-Xptxas=--warn-on-spills",
f"-I{tk_root}/include",
f"-I{tk_root}/prototype",
f"-I{python_include}",
"-DTORCH_COMPILE",
"-DKITTENS_HOPPER",
"-arch=sm_90a",
] + torch_includes
def get_extensions():
if not torch.cuda.is_available():
return []
extensions = []
cpp_flags = ["-std=c++20", "-O3"]
# Check if TK is available
if not os.path.exists(os.path.join(TK_DIR, "include", "kittens.cuh")):
print(f"Warning: ThunderKittens not found at {TK_DIR}. CUDA kernels will not be built.")
return []
cuda_flags = get_cuda_flags(TK_DIR)
# STA Extension
extensions.append(CUDAExtension(
"fastvideo_kernel._C.st_attn",
sources=[
"csrc/st_attn.cpp",
"csrc/st_attn_h100.cu",
],
extra_compile_args={
"cxx": cpp_flags + ["-DTK_COMPILE_ST_ATTN"],
"nvcc": cuda_flags + ["-DTK_COMPILE_ST_ATTN"]
},
libraries=["cuda"],
))
# VSA Extension
extensions.append(CUDAExtension(
"fastvideo_kernel._C.vsa",
sources=[
"csrc/vsa.cpp",
"csrc/block_sparse_h100.cu",
],
extra_compile_args={
"cxx": cpp_flags + ["-DTK_COMPILE_BLOCK_SPARSE"],
"nvcc": cuda_flags + ["-DTK_COMPILE_BLOCK_SPARSE"]
},
libraries=["cuda"],
))
return extensions
ext_modules = []
if not any(arg in sys.argv for arg in ["clean", "egg_info", "--version"]):
try:
import torch
ext_modules = get_extensions()
except Exception as e:
print(f"Warning: Failed to configure CUDA extensions: {e}")
setup(
name="fastvideo-kernel",
version="0.1.0",
description="Unified CUDA kernels for FastVideo",
long_description=open("README.md").read(),
long_description_content_type="text/markdown",
license="Apache-2.0",
author="Hao AI Lab",
url="https://github.com/hao-ai-lab/FastVideo",
package_dir={"": "src"},
packages=find_packages(where="src"),
ext_modules=ext_modules,
cmdclass={"build_ext": BuildExtension} if ext_modules else {},
python_requires=">=3.10",
install_requires=["torch>=2.5.0", "triton>=2.0.0"],
)
@@ -1,21 +0,0 @@
__version__ = "0.1.0"
from fastvideo_kernel.ops import (
sliding_tile_attention,
video_sparse_attn,
)
from fastvideo_kernel.vmoba import (
moba_attn_varlen,
process_moba_input,
process_moba_output,
)
__all__ = [
"sliding_tile_attention",
"video_sparse_attn",
"moba_attn_varlen",
"process_moba_input",
"process_moba_output",
"__version__",
]
@@ -1,103 +0,0 @@
import math
import torch
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
from .triton_kernels.index import map_to_index
try:
from fastvideo_kernel._C.st_attn import sta_fwd
except ImportError:
sta_fwd = None
try:
from fastvideo_kernel._C.vsa import block_sparse_fwd, block_sparse_bwd
except ImportError:
block_sparse_fwd = None
block_sparse_bwd = None
def sliding_tile_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
window_size: list,
text_length: int,
has_text: bool = True,
seq_shape: str = "30x48x80",
) -> torch.Tensor:
if sta_fwd is None:
raise RuntimeError("STA kernel not compiled. Requires H100 and ThunderKittens at build time.")
seq_length = q.shape[2]
shape_map = {"30x48x80": 1, "36x48x48": 2, "18x48x80": 3}
if has_text:
target_size = math.ceil(seq_length / 384) * 384
pad_size = target_size - seq_length
if pad_size > 0:
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
output = torch.empty_like(q)
flag = shape_map[seq_shape]
for head_idx, (t, h, w) in enumerate(window_size):
sta_fwd(
q[:, head_idx:head_idx+1],
k[:, head_idx:head_idx+1],
v[:, head_idx:head_idx+1],
output[:, head_idx:head_idx+1],
t, h, w, text_length, False, has_text, flag
)
if has_text:
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
return output[:, :, :seq_length]
def video_sparse_attn(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
variable_block_sizes: torch.Tensor,
topk: int,
block_size: int | tuple = 64,
compress_attn_weight: torch.Tensor = None,
) -> torch.Tensor:
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
batch, heads, seq_len, dim = q.shape
# Compression branch
q_c = q.view(batch, heads, seq_len // block_elements, block_elements, dim)
k_c = k.view(batch, heads, seq_len // block_elements, block_elements, dim)
v_c = v.view(batch, heads, seq_len // block_elements, block_elements, dim)
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
v_c = (v_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
scores = torch.matmul(q_c, k_c.transpose(-2, -1)) / (dim ** 0.5)
attn = torch.softmax(scores, dim=-1)
out_c = torch.matmul(attn, v_c)
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
out_c = out_c.repeat(1, 1, 1, block_elements, 1).view(batch, heads, seq_len, dim)
# Sparse branch
topk_idx = torch.topk(scores, topk, dim=-1).indices
mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
if block_sparse_fwd is not None:
idx, num = map_to_index(mask)
out_s, _ = block_sparse_fwd(q, k, v, idx, num, variable_block_sizes.int())
else:
idx, num = map_to_index(mask)
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
return out_c + out_s
@@ -1,449 +0,0 @@
"""
Fused Attention
===============
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
(https://tridao.me/publications/flash2/flash2.pdf)
Credits: OpenAI kernel team
"""
import torch
import triton
import triton.language as tl
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
import math # small utility needed by the sparse wrapper
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
# the code below and commenting out the equivalent parameters is convenient for
# re-tuning.
configs = [
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BM in [64]\
for BN in [64]\
for s in [3, 4, 7]\
for w in [4, 8]\
]
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(Q, K, V, sm_scale, #
q2k_index, q2k_num, max_kv_blks, #
variable_block_sizes,
M, Out, #
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vk, stride_vn,
stride_oz, stride_oh, stride_om, stride_on,
Z, H, N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
STAGE: tl.constexpr):
"""
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
(32×64 and 64×32) – memory footprint unchanged.
"""
# ----- program-id mapping -----
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
qvk_off = (b.to(tl.int64) * stride_qz +
h.to(tl.int64) * stride_qh)
Q_ptr = tl.make_block_ptr(
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
K_base = tl.make_block_ptr(
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
V_base = tl.make_block_ptr(
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
O_ptr = tl.make_block_ptr(
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
# ----- accumulators -----
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
qk_scale = sm_scale * 1.44269504 # 1/ln2
q = tl.load(Q_ptr)
# ----- sparse loop over valid K/V tiles -----
for i in range(0, kv_blocks):
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
block_size = tl.load(variable_block_sizes + kv_idx)
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
k = tl.load(K_ptr)
qk = tl.dot(q, k)
# mask out invalid columns
mask = tl.arange(0, BLOCK_N) < block_size
qk = tl.where(mask[None, :], qk, -float("inf"))
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
l_ij = tl.sum(p, 1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
v = tl.load(V_ptr)
acc = tl.dot(p.to(tl.bfloat16), v, acc)
m_i = m_ij
# ----- epilogue -----
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
@triton.jit
def _attn_bwd_preprocess(O, DO, #
Delta, #
Z, H, N_CTX, #
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
):
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
off_hz = tl.program_id(1)
off_n = tl.arange(0, HEAD_DIM)
# load
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
delta = tl.sum(o * do, axis=1)
# write-back
tl.store(Delta + off_hz * N_CTX + off_m, delta)
# The main inner-loop logic for computing dK and dV.
@triton.jit
def _attn_bwd_dkdv(dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_tok, stride_d, #
H, N_CTX, BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
# Filled in by the wrapper.
start_n, start_m, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M1)
offs_n = start_n + tl.arange(0, BLOCK_N1)
offs_k = tl.arange(0, HEAD_DIM)
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
step_m = BLOCK_M1
kv_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_N1
meta_base = ((b * H + h) * q_tiles + kv_blk)
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
block_size = tl.load(variable_block_sizes + kv_blk)
for blk_idx in range(q_blocks*2):
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
qkT = tl.dot(k, qT)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
# Compute dV.
ppT = pT
ppT = ppT.to(tl.bfloat16)
dv += tl.dot(ppT, do)
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# Compute dP and dS.
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
dsT = pT * (dpT - Di[None, :])
dsT = dsT.to(tl.bfloat16)
dk += tl.dot(dsT, tl.trans(qT))
# Increment pointers.
return dk, dv
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq(dq, q, K, V, #
do, m, D,
# shared by Q/K/V/DO.
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
# Filled in by the wrapper.
start_m, start_n, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M2)
offs_n = start_n + tl.arange(0, BLOCK_N2)
offs_k = tl.arange(0, HEAD_DIM)
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
step_n = BLOCK_N2
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M2
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
block_size = tl.load(variable_block_sizes + q_blk)
for blk_idx in range(kv_blocks*2):
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
p = tl.where(mask[None, :], p , 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@triton.jit
def _attn_bwd(Q, K, V, sm_scale, #
DO, #
DQ, DK, DV, #
M, D,
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_z, stride_h, stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
off_chz = (bhid * N_CTX).to(tl.int64)
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
pid = tl.program_id(0)
# offset pointers for batch/head
Q += adj
K += adj
V += adj
DO += adj
DQ += adj
DK += adj
DV += adj
M += off_chz
D += off_chz
# load scales
offs_k = tl.arange(0, HEAD_DIM)
start_n = pid * BLOCK_N1
start_m = 0
offs_n = start_n + tl.arange(0, BLOCK_N1)
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
num_steps = N_CTX // BLOCK_M1
dk, dv = _attn_bwd_dkdv( #
dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
start_n, start_m, num_steps #
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv)
# Write back dK.
dk *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk)
# THIS BLOCK DOES DQ:
start_m = pid * BLOCK_M2
end_n = 0
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)
m = m[:, None]
num_steps = N_CTX // BLOCK_N2
dq = _attn_bwd_dq(dq, q, K, V, #
do, m, D, #
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
start_m, end_n, num_steps #
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= LN2
tl.store(dq_ptrs, dq)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
o = torch.empty_like(q)
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
_attn_fwd_sparse[grid](
q, k, v, sm_scale,
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
M, o,
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
B, H, T,
HEAD_DIM=D, STAGE=3
)
return o, M
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
assert do.is_contiguous()
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
BATCH, N_HEAD, N_CTX = q.shape[:3]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert N_CTX % PRE_BLOCK == 0
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o, do, #
delta, #
BATCH, N_HEAD, N_CTX, #
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
)
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
q, arg_k, v, sm_scale, do, dq, dk, dv, #
M, delta, #
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
N_HEAD, N_CTX, #
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
HEAD_DIM=D #
)
return dq, dk, dv
@@ -1,152 +0,0 @@
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
@@ -1,868 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
import random
import time
import os
import torch
from typing import Tuple
try:
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
from functools import lru_cache
from einops import rearrange
@lru_cache(maxsize=16)
def calc_chunks(cu_seqlen, moba_chunk_size):
"""
Calculate chunk boundaries.
For vision tasks we include all chunks (even the last one which might be shorter)
so that every chunk can be selected.
"""
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
cu_num_chunk = torch.ones(
batch_num_chunk.numel() + 1,
device=cu_seqlen.device,
dtype=batch_num_chunk.dtype,
)
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
num_chunk = cu_num_chunk[-1]
chunk_sizes = torch.full(
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
)
chunk_sizes[0] = 0
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
chunk_to_batch = torch.zeros(
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
)
chunk_to_batch[cu_num_chunk[1:-1]] = 1
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
# Do not filter out any chunk
filtered_chunk_indices = torch.arange(
num_chunk, device=cu_seqlen.device, dtype=torch.int32
)
num_filtered_chunk = num_chunk
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
# --- Threshold Selection Helper Functions ---
def _select_threshold_query_head(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects chunks for each <query, head> pair based on threshold.
Normalization and sorting happen along the chunk dimension (dim=0).
"""
C, H, S = gate.shape
eps = 1e-6
# LSE‐style normalization per <head, query> (across chunks)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
row_min = gate_min_val.amin(dim=0) # (H, S)
row_max = gate_masked.amax(dim=0) # (H, S)
denom = row_max - row_min
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
# 2) compute how much more normalized weight we need beyond self
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
# 3) zero out the self‐chunk in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0
# 4) sort the other chunks by descending norm, per <head,seq>
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
# 5) cumulative‑sum the sorted norms per <head,seq>
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
any_cond = cond.any(dim=0) # (H, S)
# Find the index of the first True value along dim 0. If none, use C-1.
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
# 7) build a mask in sorted order up to that cutoff
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
# 8) scatter it back to original chunk order
others_mask = torch.zeros_like(gate, dtype=torch.bool)
others_mask.scatter_(0, sorted_idx, sorted_mask)
# 9) finally, include every self‐chunk plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_block(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <query, head> pairs for each block based on threshold.
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
"""
C, H, S = gate.shape
HS = H * S
eps = 1e-6
# LSE‐style normalization per block (across heads and queries)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_denom = block_max - block_min
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
gate_norm = (gate - block_min) / block_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights *per block*
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort the other <head, seq> pairs by descending norm, per block
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
# 5) cumulative‑sum the sorted norms per block
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
any_cond = cond_flat.any(dim=1) # (C,)
# Find the index of the first True value along dim 1. If none, use HS-1.
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
# 7) build a mask in sorted order up to that cutoff per block
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
# 8) scatter it back to original <head, seq> order per block
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_overall(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query, head> triplets globally based on threshold.
Normalization and sorting happen across all valid entries.
"""
C, H, S = gate.shape
CHS = C * H * S
eps = 1e-6
# LSE‐style normalization globally across all valid entries
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
overall_max = gate_masked.max() # scalar
overall_min = gate_min_val.min() # scalar
overall_denom = overall_max - overall_min
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights globally
self_norm_sum_overall = self_norm_entries.sum() # scalar
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
total_norm_sum_overall = gate_norm.sum() # scalar
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort all other entries by descending norm, globally
others_flat = others_norm.flatten() # (C*H*S,)
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
# Only sort the valid 'other' entries
valid_others_indices = torch.where(valid_others_mask_flat)[0]
valid_others_values = others_flat[valid_others_indices]
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
# 5) cumulative‑sum the sorted valid 'other' norms globally
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
any_cond = cond_values.any() # scalar
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
cutoff_idx_in_sorted = torch.where(
any_cond,
cond_values.float().argmax(dim=0),
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
)
# 7) build a mask selecting the top-k others based on the cutoff
# Select the original indices corresponding to the top entries in the sorted list
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
# 8) create the mask in the original flat shape
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
others_mask_flat[selected_other_indices] = True
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_head_global(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query> globally for each head based on threshold.
"""
C, H, S = gate.shape
eps = 1e-6
# 1) LSE‐style normalization per head (across chunks and sequence dims)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
denom = max_per_head - min_per_head
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
gate_norm = (gate - min_per_head) / denom
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 2) sum normalized self‐chunk contributions per head
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
# 3) total normalized sum per head
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
# 4) how much more normalized weight needed per head
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
remain_ratio = torch.clamp(remain_ratio, min=0.0)
# 5) zero out self‐chunk entries to focus on "others"
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
# 6) flatten chunk and sequence dims, per head
CS = C * S
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
# 7) vectorized selection of “others” per head
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
has_cutoff = cond.any(dim=1) # (H,)
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
# 8) reshape selection mask back to (C, H, S)
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
# 9) include self‐chunks plus selected others, and obey valid mask
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
return final_gate_mask
class MixedAttention(torch.autograd.Function):
@staticmethod
def forward(
ctx,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
):
ctx.max_seqlen = max_seqlen
ctx.moba_chunk_size = moba_chunk_size
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
# Non-causal self-attention branch
# return out, softmax_lse, S_dmask, rng_state
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=q,
k=k,
v=v,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
# MOBA attention branch (non-causal)
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
output_2d = output.view(-1, q.shape[2])
max_lse_1d = self_attn_lse_sh.view(-1)
max_lse_1d = max_lse_1d.index_reduce(
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
)
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
moba_attn_lse = (
moba_attn_lse.view(-1)
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
.reshape_as(moba_attn_lse)
)
mixed_attn_se_sh = self_attn_lse_sh.exp()
moba_attn_se = moba_attn_lse.exp()
mixed_attn_se_sh.view(-1).index_add_(
0, moba_q_sh_indices, moba_attn_se.view(-1)
)
mixed_attn_lse_sh = mixed_attn_se_sh.log()
# Combine self-attention output
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
output_2d += self_attn_out_sh.reshape_as(output_2d)
# Combine MOBA attention output
mixed_attn_lse = (
mixed_attn_lse_sh.view(-1)
.index_select(0, moba_q_sh_indices)
.view_as(moba_attn_lse)
)
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
output = output.to(q.dtype)
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
ctx.save_for_backward(
output,
mixed_attn_lse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
)
return output
@staticmethod
def backward(ctx, d_output):
max_seqlen = ctx.max_seqlen
moba_chunk_size = ctx.moba_chunk_size
softmax_scale = ctx.softmax_scale
(
output,
mixed_attn_vlse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
) = ctx.saved_tensors
d_output = d_output.contiguous()
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
_ = _flash_attn_varlen_backward(
dout=d_output,
q=q,
k=k,
v=v,
out=output,
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
dq=dq,
dk=dk,
dv=dv,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
headdim = q.shape[-1]
d_moba_output = (
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
moba_output = (
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
mixed_attn_vlse = (
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
)
dmq = torch.empty_like(moba_q)
dmkv = torch.empty_like(moba_kv)
_ = _flash_attn_varlen_backward(
dout=d_moba_output,
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
out=moba_output,
softmax_lse=mixed_attn_vlse,
dq=dmq,
dk=dmkv[:,0],
dv=dmkv[:,1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
def moba_attn_varlen(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens: torch.Tensor,
max_seqlen: int,
moba_chunk_size: int,
moba_topk: int,
select_mode: str = 'threshold', # "topk" or "threshold"
simsum_threshold: float = 0.25,
threshold_type: str = 'query_head',
) -> torch.Tensor:
"""
Accelerated MOBA attention for vision tasks with proper LSE normalization.
This version:
- Splits KV into chunks.
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
by amplifying the diagonal (self-chunk) logits.
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
reduction so that attending to each query over the selected chunks is equivalent
to the original algorithm.
"""
# Stack keys and values.
kv = torch.stack((k, v), dim=1)
seqlen, num_head, head_dim = q.shape
# Compute chunk boundaries.
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
cu_seqlens, moba_chunk_size
)
self_attn_cu_seqlen = cu_chunk
# Update top-k selection to include the self chunk.
moba_topk = min(moba_topk, num_filtered_chunk)
# --- Build filtered KV from chunks ---
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
max_chunk_len = int(chunk_lengths.max().item())
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
indices = chunk_starts.unsqueeze(1) + range_tensor
indices = torch.clamp(indices, max=kv.shape[0] - 1)
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
# Compute key_gate_weight over valid tokens.
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
key_sum = (key_values * valid_mask_exp).sum(dim=1)
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
# Compute gate logits between key_gate_weight and queries.
q_float = q.float()
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
# Amplify the diagonal (self chunk) contributions.
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
amplification_factor = 1e9 # Example factor; adjust as needed.
origin_gate = gate.clone()
gate = gate.clone()
if select_mode == "topk":
gate[gate_self_chunk_mask] += amplification_factor
# Exclude positions that are outside the valid batch boundaries.
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
if select_mode == 'topk':
# We amplify self‐chunk in gate already, so self entries will rank highest.
valid_gate_mask = gate != -float("inf")
if threshold_type == 'query_head':
# === per‐<head,seq> top-k across chunks (original behavior) ===
# gate: (C, H, S)
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
gate_idx_mask.scatter_(0, gate_topk_idx, True)
gate_mask = valid_gate_mask & gate_idx_mask
elif threshold_type == 'overall':
# === global top-k across all (chunk, head, seq) entries ===
C, H, S = gate.shape
flat_gate = gate.flatten()
flat_mask = valid_gate_mask.flatten()
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
# pick topk global entries
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
others_mask_flat[idx] = True
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
elif threshold_type == 'head_global':
# per-head top-k across all chunks and sequence positions
C, H, S = gate.shape
CS = C * S
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
# pick top-k indices per head
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
gate_idx_flat.scatter_(1, topk_idx, True)
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
else:
raise ValueError(
f"Invalid threshold_type for topk: {threshold_type}. "
"Choose 'query_head', 'block', or 'overall'."
)
elif select_mode == 'threshold':
# Delegate to the specific thresholding function
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
if threshold_type == 'query_head':
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'block':
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'overall':
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'head_global':
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
else:
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
else:
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
# eliminate self_chunk in MoBA branch
gate_mask = gate_mask & ~gate_self_chunk_mask
# if gate_mask is all false, perform flash_attn instead
if gate_mask.sum() == 0:
return flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
)
# Determine which query positions are selected.
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
# Build cumulative sequence lengths for the selected queries.
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
q_zero_mask = moba_seqlen_q == 0
valid_expert_mask = ~q_zero_mask
if q_zero_mask.sum() > 0:
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
moba_cu_seqlen_q = torch.cat(
(
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
moba_seqlen_q.cumsum(dim=0),
),
dim=0,
).to(torch.int32)
# Rearrange gathered KV for the MOBA branch.
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
if q_zero_mask.sum() > 0:
experts_tensor = experts_tensor[valid_expert_mask]
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
mask = seq_range < valid_expert_lengths.unsqueeze(1)
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
moba_cu_seqlen_kv = torch.cat(
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
valid_expert_lengths.cumsum(dim=0)],
dim=0,
).to(torch.int32)
assert (
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
return MixedAttention.apply(
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
)
def process_moba_input(
x,
patch_resolution,
chunk_size,
):
"""
Process inputs for the attention function.
Args:
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
Returns:
torch.Tensor: Processed input tensor.
"""
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
else:
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
if len(chunk_size) == 2:
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
elif len(chunk_size) == 3:
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
else:
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
return x, moba_chunk_size
def process_moba_output(
x,
patch_resolution,
chunk_size,
):
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
pass
elif len(chunk_size) == 2:
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
elif len(chunk_size) == 3:
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
return x
# TEST
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
random.seed(0)
torch.manual_seed(0)
torch.cuda.manual_seed(0)
device = torch.cuda.current_device()
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
max_seqlen = q.shape[1]
q = rearrange(q, "b s ... -> (b s) ...")
k = rearrange(k, "b s ... -> (b s) ...")
v = rearrange(v, "b s ... -> (b s) ...")
return q, k, v, cu_seqlens, max_seqlen
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
"""Speed test comparing flash_attn vs moba_attention"""
# Get data
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
vo_grad = torch.randn_like(q)
# Warmup
warmup_iters = 3
perf_test_iters = 10
# Warmup
for _ in range(warmup_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
start_flash = time.perf_counter()
for _ in range(perf_test_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
# Warmup
for _ in range(warmup_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
start_moba = time.perf_counter()
for _ in range(perf_test_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
print(f"Speedup: {time_flash / time_moba:.2f}x")
if __name__ == "__main__":
"""
CUDA_VISIBLE_DEVICES=1 \
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
"""
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
@@ -1,71 +0,0 @@
from typing import Tuple
import torch
from torch import BoolTensor, IntTensor
from torch.nn.attention.flex_attention import create_block_mask
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
torch._inductor.config.realize_opcount_threshold = 100
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
"""Generates a 3D NATTEN attention mask with a given kernel size.
Args:
canvas_t: The time dimension of the canvas.
canvas_h: The height of the canvas.
canvas_w: The width of the canvas.
kernel_t: The time dimension of the kernel.
kernel_h: The height of the kernel.
kernel_w: The width of the kernel.
"""
canvas_t, canvas_h, canvas_w = canvas_twh
kernel_t, kernel_h, kernel_w = kernel_twh
tile_t_size, tile_h_size, tile_w_size = tile_twh
total_tile_size = tile_t_size * tile_h_size * tile_w_size
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
img_seq_len = canvas_t * canvas_h * canvas_w
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
tile_id = idx // total_tile_size
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
tile_w = tile_id % canvas_tile_w
return tile_t, tile_h, tile_w
def sta_mask_3d(
b: IntTensor,
h: IntTensor,
q_idx: IntTensor,
kv_idx: IntTensor,
) -> BoolTensor:
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
# kernel nominally attempts to center itself on the query, but kernel center
# is clamped to a fixed distance (kernel half-length) from the canvas edge
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
return sta_mask_3d
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
img_seq_len = img_size[0] * img_size[1] * img_size[2]
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
mask = create_block_mask(image_mask,
B=None,
H=None,
Q_LEN=img_seq_len + text_max_len,
KV_LEN=img_seq_len + text_max_len,
device=device,
_compile=True)
return mask
@@ -1,63 +0,0 @@
import torch
import sys
import os
from tqdm import tqdm
# Local support import
from .support_flex_sta import get_sliding_tile_attention_mask
# USE OUR NEW PACKAGE!
from fastvideo_kernel import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
# Using the same parameters as the original test
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
for kernel_size in kernel_size_ls:
print(f"Testing kernel_size: {kernel_size}")
for xi in tqdm(range(num_iterations)):
torch.manual_seed(xi)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
pt_o = flex_test(Q, K, V, kernel_size)
diff = pt_o - tk_o
abs_diff = torch.abs(diff)
max_d = torch.max(abs_diff).item()
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
if max_d > 0.1:
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
if __name__ == "__main__":
b, h, d = 2, 24, 128
n = 69120
causal = False
mean = 1e-1
std = 10
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
@@ -1,97 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import pytest
import random
from fastvideo_kernel.vmoba import moba_attn_varlen
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
"""
Generates random data for testing the variable-length attention function.
"""
torch.manual_seed(42)
random.seed(42)
torch.cuda.manual_seed_all(42)
# Generate sequence lengths for each item in the batch
if batch_size > 1:
# Ensure sequence lengths are reasonably distributed
avg_seqlen = total_seqlen // batch_size
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
remaining_len = total_seqlen - sum(seqlens)
if remaining_len > 0:
seqlens.append(remaining_len)
else: # Adjust if sum exceeds total_seqlen
seqlens.append(avg_seqlen)
current_sum = sum(seqlens)
seqlens[-1] -= (current_sum - total_seqlen)
# Ensure all lengths are positive
seqlens = [max(1, s) for s in seqlens]
# Final adjustment to match total_seqlen
seqlens[-1] += total_seqlen - sum(seqlens)
else:
seqlens = [total_seqlen]
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
max_seqlen = max(seqlens) if seqlens else 0
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
return q, k, v, cu_seqlens, max_seqlen
@pytest.mark.parametrize("batch_size", [1, 2])
@pytest.mark.parametrize("total_seqlen", [512, 1024])
@pytest.mark.parametrize("num_heads", [8])
@pytest.mark.parametrize("head_dim", [64])
@pytest.mark.parametrize("moba_chunk_size", [64])
@pytest.mark.parametrize("moba_topk", [2, 4])
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_moba_attn_varlen_forward(
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
):
"""
Tests the forward pass of moba_attn_varlen for basic correctness.
It checks output shape, dtype, and for the presence of NaNs/Infs.
"""
if dtype == torch.float32:
pytest.skip("float32 is not supported in flash attention")
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
batch_size, total_seqlen, num_heads, head_dim, dtype
)
# Ensure chunk size is not larger than the smallest sequence length
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
if moba_chunk_size > min_seqlen:
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
try:
output = moba_attn_varlen(
q=q,
k=k,
v=v,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=moba_chunk_size,
moba_topk=moba_topk,
select_mode=select_mode,
threshold_type=threshold_type,
simsum_threshold=0.5, # A reasonable default for threshold mode
)
except Exception as e:
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
# 1. Check output shape
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
# 2. Check output dtype
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
# 3. Check for NaNs or Infs in the output
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
+9 -2
View File
@@ -55,10 +55,17 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Kernels
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/fastvideo_kernel && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
+9 -2
View File
@@ -55,10 +55,17 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Kernels
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/fastvideo_kernel && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
+9 -2
View File
@@ -55,10 +55,17 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Kernels
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/fastvideo_kernel && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
-58
View File
@@ -1,58 +0,0 @@
FROM rocm/pytorch:rocm7.1_ubuntu22.04_py3.10_pytorch_release_2.9.1
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject_other.toml ./pyproject.toml
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.10 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip
COPY . .
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[rocm] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo 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 requirements-mkdocs.txt
pip install -r docs/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 supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
+2 -2
View File
@@ -1,7 +1,7 @@
# Profiling FastVideo
!!! warning
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down inference.
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
## Profiling with PyTorch
@@ -49,5 +49,5 @@ Traces can be visualized using <https://ui.perfetto.dev/>.
### Best Practices
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
- After profiling, clean up trace directories to avoid filling disk storage.
- After profiling, clean up trace directories to avoid filling disks.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
+3 -5
View File
@@ -74,11 +74,9 @@ To add a new SSIM test, follow these steps:
generator.generate_video(prompt, ...)
# Compare with Reference
ssim_values = compute_video_ssim_torchvision(
reference_path, generated_path, use_ms_ssim=True
)
assert ssim_values[0] >= 0.98 # Threshold
```
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
assert ssim_values[0] >= 0.98 # Threshold
```
4. **Reference Videos**:
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
+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 computation, enabling much faster video generation.
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computations, enabling much faster video generation.
## 📊 Model Overview
+18 -40
View File
@@ -7,11 +7,6 @@ Get up and running with FastVideo in minutes!
First, install FastVideo:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Install FastVideo
pip install fastvideo
```
@@ -20,53 +15,36 @@ pip install fastvideo
### Text-to-Video Generation
```python
from fastvideo import VideoGenerator
from fastvideo import FastVideoPipeline
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Initialize the pipeline
pipe = FastVideoPipeline.from_pretrained("wan2.1-t2v-1.3B")
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate a video
prompt = "A cat playing with a ball of yarn"
video = pipe(prompt, num_frames=16, height=512, width=512)
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
# Save the video
video.save("output.mp4")
```
### Image-to-Video Generation
```python
from fastvideo import VideoGenerator, SamplingParam
from fastvideo import FastVideoPipeline
from PIL import Image
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Load an image
image = Image.open("input.jpg")
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
# Initialize the pipeline
pipe = FastVideoPipeline.from_pretrained("wan2.1-i2v-14B-480p")
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
# Generate a video from the image
video = pipe(image, num_frames=16, height=480, width=480)
if __name__ == '__main__':
main()
# Save the video
video.save("output.mp4")
```
## Next Steps
+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 an implementation for H100s.
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
+1 -1
View File
@@ -30,7 +30,7 @@ path_to_your_dataset_folder/
└── prompt.txt
```
To generate the `videos2caption.json` and `merge.txt`, run
To geranate the `videos2caption.json` and `merge.txt`, run
``` python
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
+2 -2
View File
@@ -7,9 +7,9 @@ pip install vsa
```
# Building from Source
We support H100s (via ThunderKittens) and any other GPU (via Triton) for VSA.
We support H100 (via ThunderKittens) and any other GPU (via Triton) for VSA.
First, install C++20 for ThunderKittens (if using an H100):
First, install C++20 for ThunderKittens (if using H100):
```bash
sudo apt update
@@ -5,7 +5,7 @@ from typing import Any
import torch
from einops import rearrange
from fastvideo_kernel import sliding_tile_attention
from st_attn import sliding_tile_attention
import fastvideo.envs as envs
from fastvideo.attention.backends.abstract import (AttentionBackend,
@@ -6,7 +6,7 @@ from dataclasses import dataclass
import torch
try:
from fastvideo_kernel import video_sparse_attn
from vsa import video_sparse_attn
except ImportError:
video_sparse_attn = None
+2 -2
View File
@@ -6,8 +6,8 @@ from dataclasses import dataclass
import torch
from einops import rearrange
from fastvideo_kernel import (moba_attn_varlen, process_moba_input,
process_moba_output)
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
+1 -3
View File
@@ -2,12 +2,10 @@ from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig"
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig"
]
-149
View File
@@ -1,149 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat Video DiT configuration for native FastVideo implementation.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def is_longcat_blocks(n: str, m) -> bool:
"""FSDP shard condition for LongCat transformer blocks."""
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LongCatVideoArchConfig(DiTArchConfig):
"""Architecture configuration for native LongCat Video DiT."""
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_longcat_blocks])
# Enable torch.compile for transformer blocks (major speedup!)
_compile_conditions: list = field(
default_factory=lambda: [is_longcat_blocks])
# Parameter name mapping for weight conversion
# Maps original LongCat third_party names -> native FastVideo names
param_names_mapping: dict = field(
default_factory=lambda: {
# Embedders
r"^x_embedder\.(.*)$": r"patch_embed.\1",
r"^t_embedder\.mlp\.0\.(.*)$": r"time_embedder.linear_1.\1",
r"^t_embedder\.mlp\.2\.(.*)$": r"time_embedder.linear_2.\1",
r"^y_embedder\.y_proj\.0\.(.*)$": r"caption_embedder.linear_1.\1",
r"^y_embedder\.y_proj\.2\.(.*)$": r"caption_embedder.linear_2.\1",
# Transformer blocks - AdaLN modulation
r"^blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
r"blocks.\1.adaln_linear_1.\2",
# Transformer blocks - Normalization
r"^blocks\.(\d+)\.mod_norm_attn\.(.*)$": r"blocks.\1.norm_attn.\2",
r"^blocks\.(\d+)\.mod_norm_ffn\.(.*)$": r"blocks.\1.norm_ffn.\2",
r"^blocks\.(\d+)\.pre_crs_attn_norm\.(.*)$":
r"blocks.\1.norm_cross.\2",
# Self-attention: QKV fused -> separate (will need splitting in converter)
# Original has attn.qkv.weight -> need to split into to_q, to_k, to_v
r"^blocks\.(\d+)\.attn\.qkv\.(.*)$":
r"blocks.\1.self_attn.qkv_fused.\2", # Marker for splitting
r"^blocks\.(\d+)\.attn\.proj\.(.*)$":
r"blocks.\1.self_attn.to_out.\2",
r"^blocks\.(\d+)\.attn\.q_norm\.(.*)$":
r"blocks.\1.self_attn.q_norm.\2",
r"^blocks\.(\d+)\.attn\.k_norm\.(.*)$":
r"blocks.\1.self_attn.k_norm.\2",
# Cross-attention
r"^blocks\.(\d+)\.cross_attn\.q_linear\.(.*)$":
r"blocks.\1.cross_attn.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.kv_linear\.(.*)$":
r"blocks.\1.cross_attn.kv_fused.\2", # Marker for splitting
r"^blocks\.(\d+)\.cross_attn\.proj\.(.*)$":
r"blocks.\1.cross_attn.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.q_norm\.(.*)$":
r"blocks.\1.cross_attn.q_norm.\2",
r"^blocks\.(\d+)\.cross_attn\.k_norm\.(.*)$":
r"blocks.\1.cross_attn.k_norm.\2",
# FFN (SwiGLU)
r"^blocks\.(\d+)\.ffn\.w1\.(.*)$": r"blocks.\1.ffn.w1.\2", # gate
r"^blocks\.(\d+)\.ffn\.w2\.(.*)$": r"blocks.\1.ffn.w2.\2", # down
r"^blocks\.(\d+)\.ffn\.w3\.(.*)$": r"blocks.\1.ffn.w3.\2", # up
# Final layer
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
r"final_layer.adaln_linear.\1",
r"^final_layer\.norm_final\.(.*)$": r"final_layer.norm.\1",
r"^final_layer\.linear\.(.*)$": r"final_layer.proj.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# LoRA parameter name mapping
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Model architecture parameters
hidden_size: int = 4096
depth: int = 48 # Number of transformer blocks
num_attention_heads: int = 32
attention_head_dim: int = 128 # hidden_size / num_attention_heads
in_channels: int = 16 # Latent space channels
out_channels: int = 16
num_channels_latents: int = 16
# Patch embedding
patch_size: tuple[int, int,
int] = (1, 2, 2) # [T, H, W] - no temporal compression
# Text/caption embedding
caption_channels: int = 4096 # UMT5 d_model
# Timestep embedding
adaln_tembed_dim: int = 512
frequency_embedding_size: int = 256
# FFN
mlp_ratio: int = 4
# Attention backend support
_supported_attention_backends: tuple = field(default_factory=lambda: (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
))
# Text padding behavior
text_tokens_zero_pad: bool = True
# Block Sparse Attention (BSA)
enable_bsa: bool = False
bsa_params: dict | None = field(
default_factory=lambda: {
"sparsity": 0.9375,
"cdf_threshold": None,
"chunk_3d_shape_q": [4, 4, 4],
"chunk_3d_shape_k": [4, 4, 4],
})
# LoRA exclusions
exclude_lora_layers: list[str] = field(default_factory=lambda: [])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
# Ensure attention_head_dim matches
self.attention_head_dim = self.hidden_size // self.num_attention_heads
@dataclass
class LongCatVideoConfig(DiTConfig):
"""Main configuration for LongCat Video DiT."""
arch_config: DiTArchConfig = field(default_factory=LongCatVideoArchConfig)
prefix: str = "longcat"
-355
View File
@@ -1,355 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
import html
import ftfy
import regex as re
import torch
from fastvideo.configs.models import DiTConfig, VAEConfig
from fastvideo.configs.models.dits.base import DiTArchConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class LongCatDiTArchConfig(DiTArchConfig):
"""Extended DiTArchConfig with LongCat-specific fields.
NOTE: This is for Phase 1 wrapper compatibility. For native model (Phase 2),
use LongCatVideoConfig from fastvideo.configs.models.dits.longcat instead.
"""
# LongCat-specific architecture parameters
adaln_tembed_dim: int = 512
caption_channels: int = 4096
depth: int = 48
enable_bsa: bool = False
enable_flashattn3: bool = False
enable_flashattn2: bool = True
enable_xformers: bool = False
frequency_embedding_size: int = 256
in_channels: int = 16
mlp_ratio: int = 4
num_heads: int = 32
out_channels: int = 16
text_tokens_zero_pad: bool = True
patch_size: list[int] = field(default_factory=lambda: [1, 2, 2])
cp_split_hw: list[int] | None = None
bsa_params: dict | None = None
def longcat_preprocess_text(prompt: str) -> str:
"""Clean and preprocess text like original LongCat implementation.
This function applies the same text cleaning pipeline as the original
LongCat-Video implementation to ensure identical tokenization results.
Steps:
1. basic_clean: Fix unicode issues and unescape HTML entities
2. whitespace_clean: Normalize whitespace to single spaces
Args:
prompt: Raw input text prompt
Returns:
Cleaned and normalized text prompt
"""
# basic_clean: fix unicode and HTML entities
text = ftfy.fix_text(prompt)
text = html.unescape(html.unescape(text))
text = text.strip()
# whitespace_clean: normalize whitespace
text = re.sub(r"\s+", " ", text)
text = text.strip()
return text
def umt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
"""
Postprocess UMT5/T5 encoder outputs to fixed length 512 embeddings.
"""
mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.Tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
prompt_embeds_tensor: torch.Tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in prompt_embeds
],
dim=0)
return prompt_embeds_tensor
@dataclass
class LongCatT2V480PConfig(PipelineConfig):
"""Configuration for LongCat pipeline (480p) aligned to LongCat-Video modules.
Components expected by loaders:
- tokenizer: AutoTokenizer
- text_encoder: UMT5EncoderModel
- transformer: LongCatVideoTransformer3DModel (Phase 1 wrapper)
OR LongCatTransformer3DModel (Phase 2 native)
- vae: AutoencoderKLWan (Wan VAE, 4x8 compression)
- scheduler: FlowMatchEulerDiscreteScheduler
"""
# DiT config with LongCat-specific arch_config
# NOTE: For Phase 1 wrapper, uses LongCatDiTArchConfig
# For Phase 2 native model, can use LongCatVideoConfig directly
dit_config: DiTConfig = field(
default_factory=lambda: DiTConfig(arch_config=LongCatDiTArchConfig()))
# VAE config: Wan VAE with encoder+decoder enabled
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# Precision defaults
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", ))
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
text_encoder_configs: tuple[T5Config, ...] = field(
default_factory=lambda: (T5Config(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (longcat_preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(umt5_postprocess_text, ))
# LongCat-specific runtime toggles (consumed by pipeline/stages)
enable_kv_cache: bool = True
offload_kv_cache: bool = False
enable_bsa: bool = False
use_distill: bool = False
enhance_hf: bool = False
# Optional BSA parameter dict (kept for backward/phase-1 compatibility).
# `LongCatPipeline.initialize_pipeline()` uses this as a base and then applies
# CLI overrides (bsa_sparsity / bsa_chunk_{q,k} / bsa_cdf_threshold).
bsa_params: dict | None = None
# BSA runtime overrides (preferred over bsa_params if provided via CLI)
bsa_sparsity: float | None = None
bsa_cdf_threshold: float | None = None
bsa_chunk_q: list[int] | None = None
bsa_chunk_k: list[int] | None = None
t_thresh: float | None = None # refine stage default controlled by sampling args
# LongCat does not need flow_shift
flow_shift: float | None = None
dmd_denoising_steps: list[int] | None = None
def __post_init__(self):
# LongCat inference requires vae encoder and decoder
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class LongCatT2V704PConfig(LongCatT2V480PConfig):
"""Configuration for LongCat pipeline (704p) with BSA enabled by default.
Uses the same resolution and BSA parameters as original LongCat refinement stage.
BSA parameters configured in transformer config.json with chunk_3d_shape=[4,4,4]:
- Input: 704×1280×96
- VAE (8x): 88×160×96
- Patch [1,2,2]: 44×80×96
- chunk [4,4,4]: 96%4=0, 44%4=0, 80%4=0 ✅
This configuration matches the original LongCat refinement stage parameters.
"""
# Enable BSA by default for 704p
enable_bsa: bool = True
ASPECT_RATIO_627 = {
'0.26': ([320, 1216], 1),
'0.31': ([352, 1120], 1),
'0.38': ([384, 1024], 1),
'0.43': ([416, 960], 1),
'0.52': ([448, 864], 1),
'0.58': ([480, 832], 1),
'0.67': ([512, 768], 1),
'0.74': ([544, 736], 1),
'0.86': ([576, 672], 1),
'0.95': ([608, 640], 1),
'1.05': ([640, 608], 1),
'1.17': ([672, 576], 1),
'1.29': ([704, 544], 1),
'1.35': ([736, 544], 1),
'1.50': ([768, 512], 1),
'1.67': ([800, 480], 1),
'1.73': ([832, 480], 1),
'2.00': ([896, 448], 1),
'2.31': ([960, 416], 1),
'2.58': ([992, 384], 1),
'2.75': ([1056, 384], 1),
'3.09': ([1088, 352], 1),
'3.70': ([1184, 320], 1),
'3.80': ([1216, 320], 1),
'3.90': ([1248, 320], 1),
'4.00': ([1280, 320], 1)
}
ASPECT_RATIO_627_F64 = {
'0.26': ([320, 1216], 1),
'0.38': ([384, 1024], 1),
'0.50': ([448, 896], 1),
'0.67': ([512, 768], 1),
'0.82': ([576, 704], 1),
'1.00': ([640, 640], 1),
'1.22': ([704, 576], 1),
'1.50': ([768, 512], 1),
'1.86': ([832, 448], 1),
'2.00': ([896, 448], 1),
'2.50': ([960, 384], 1),
'2.83': ([1088, 384], 1),
'3.60': ([1152, 320], 1),
'3.80': ([1216, 320], 1),
'4.00': ([1280, 320], 1)
}
ASPECT_RATIO_627_F128 = {
'0.25': ([256, 1024], 1),
'0.38': ([384, 1024], 1),
'0.43': ([384, 896], 1),
'0.57': ([512, 896], 1),
'0.67': ([512, 768], 1),
'1.00': ([640, 640], 1),
'1.50': ([768, 512], 1),
'1.75': ([896, 512], 1),
'2.33': ([896, 384], 1),
'2.67': ([1024, 384], 1),
'4.00': ([1024, 256], 1),
}
ASPECT_RATIO_627_F256 = {
'0.25': ([256, 1024], 1),
'0.33': ([256, 768], 1),
'0.50': ([256, 512], 1),
'0.67': ([512, 768], 1),
'1.00': ([512, 512], 1),
'1.50': ([768, 512], 1),
'2.00': ([512, 256], 1),
'3.00': ([768, 256], 1),
'4.00': ([1024, 256], 1),
}
ASPECT_RATIO_960 = {
'0.25': ([480, 1920], 1),
'0.29': ([512, 1792], 1),
'0.32': ([544, 1696], 1),
'0.36': ([576, 1600], 1),
'0.40': ([608, 1504], 1),
'0.49': ([672, 1376], 1),
'0.54': ([704, 1312], 1),
'0.59': ([736, 1248], 1),
'0.69': ([800, 1152], 1),
'0.74': ([832, 1120], 1),
'0.82': ([864, 1056], 1),
'0.88': ([896, 1024], 1),
'0.94': ([928, 992], 1),
'1.00': ([960, 960], 1),
'1.07': ([992, 928], 1),
'1.14': ([1024, 896], 1),
'1.22': ([1056, 864], 1),
'1.31': ([1088, 832], 1),
'1.35': ([1120, 832], 1),
'1.44': ([1152, 800], 1),
'1.70': ([1248, 736], 1),
'2.00': ([1344, 672], 1),
'2.05': ([1376, 672], 1),
'2.47': ([1504, 608], 1),
'2.53': ([1536, 608], 1),
'2.83': ([1632, 576], 1),
'3.06': ([1664, 544], 1),
'3.12': ([1696, 544], 1),
'3.62': ([1856, 512], 1),
'3.93': ([1888, 480], 1),
'4.00': ([1920, 480], 1)
}
ASPECT_RATIO_960_F64 = {
'0.22': ([448, 2048], 1),
'0.29': ([512, 1792], 1),
'0.36': ([576, 1600], 1),
'0.45': ([640, 1408], 1),
'0.55': ([704, 1280], 1),
'0.63': ([768, 1216], 1),
'0.76': ([832, 1088], 1),
'0.88': ([896, 1024], 1),
'1.00': ([960, 960], 1),
'1.14': ([1024, 896], 1),
'1.31': ([1088, 832], 1),
'1.50': ([1152, 768], 1),
'1.58': ([1216, 768], 1),
'1.82': ([1280, 704], 1),
'1.91': ([1344, 704], 1),
'2.20': ([1408, 640], 1),
'2.30': ([1472, 640], 1),
'2.67': ([1536, 576], 1),
'2.89': ([1664, 576], 1),
'3.62': ([1856, 512], 1),
'3.75': ([1920, 512], 1)
}
ASPECT_RATIO_960_F128 = {
'0.20': ([384, 1920], 1),
'0.27': ([512, 1920], 1),
'0.33': ([512, 1536], 1),
'0.42': ([640, 1536], 1),
'0.50': ([640, 1280], 1),
'0.60': ([768, 1280], 1),
'0.67': ([768, 1152], 1),
'0.78': ([896, 1152], 1),
'1.00': ([1024, 1024], 1),
'1.29': ([1152, 896], 1),
'1.50': ([1152, 768], 1),
'1.67': ([1280, 768], 1),
'2.00': ([1280, 640], 1),
'2.40': ([1536, 640], 1),
'3.00': ([1536, 512], 1),
'3.75': ([1920, 512], 1),
'5.00': ([1920, 384], 1),
}
ASPECT_RATIO_960_F256 = {
'0.33': ([512, 1536], 1),
'0.60': ([768, 1280], 1),
'1.00': ([1024, 1024], 1),
'1.67': ([1280, 768], 1),
'3.00': ([1536, 512], 1),
}
def get_bucket_config(resolution, scale_factor_spatial):
if resolution == '480p':
if scale_factor_spatial == 16 or scale_factor_spatial == 32:
return ASPECT_RATIO_627
elif scale_factor_spatial == 64:
return ASPECT_RATIO_627_F64
elif scale_factor_spatial == 128:
return ASPECT_RATIO_627_F128
elif scale_factor_spatial == 256:
return ASPECT_RATIO_627_F256
elif resolution == '720p':
if scale_factor_spatial == 16 or scale_factor_spatial == 32:
return ASPECT_RATIO_960
elif scale_factor_spatial == 64:
return ASPECT_RATIO_960_F64
elif scale_factor_spatial == 128:
return ASPECT_RATIO_960_F128
elif scale_factor_spatial == 256:
return ASPECT_RATIO_960_F256
raise ValueError(
f"Unsupported resolution '{resolution}' or scale_factor_spatial '{scale_factor_spatial}'"
)
-4
View File
@@ -9,7 +9,6 @@ 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 (
@@ -78,14 +77,11 @@ 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,
-37
View File
@@ -3,7 +3,6 @@ from dataclasses import dataclass
from typing import Any
from fastvideo.logger import init_logger
from fastvideo.utils import StoreBoolean
logger = init_logger(__name__)
@@ -28,17 +27,6 @@ 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"
@@ -228,31 +216,6 @@ class SamplingParam:
default=SamplingParam.video_path,
help="Path to input video for video-to-video generation",
)
parser.add_argument(
"--refine-from",
type=str,
default=SamplingParam.refine_from,
help="Path to stage1 video for refinement (LongCat 480p->720p)",
)
parser.add_argument(
"--t-thresh",
type=float,
default=SamplingParam.t_thresh,
help=
"Threshold for timestep scheduling in refinement (default: 0.5)",
)
parser.add_argument(
"--spatial-refine-only",
action=StoreBoolean,
default=SamplingParam.spatial_refine_only,
help="Only perform spatial super-resolution (no temporal doubling)",
)
parser.add_argument(
"--num-cond-frames",
type=int,
default=SamplingParam.num_cond_frames,
help="Number of conditioning frames for refinement",
)
parser.add_argument(
"--moba-config-path",
type=str,
-56
View File
@@ -319,62 +319,6 @@ 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",
-209
View File
@@ -1,209 +0,0 @@
# 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)
-866
View File
@@ -1,866 +0,0 @@
# 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
+1 -2
View File
@@ -16,7 +16,6 @@ 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
@@ -406,7 +405,7 @@ class VAELoader(ComponentLoader):
target_device = get_local_torch_device()
with set_default_torch_dtype(PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision] if fastvideo_args.pipeline_config.vae_precision else torch.bfloat16):
fastvideo_args.pipeline_config.vae_precision]):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
+1 -3
View File
@@ -29,9 +29,7 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel")
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -1,6 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""LongCat pipeline module."""
from fastvideo.pipelines.basic.longcat.longcat_pipeline import LongCatPipeline
__all__ = ["LongCatPipeline"]
@@ -1,145 +0,0 @@
# 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,7 +287,6 @@ 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,14 +91,6 @@ 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
-1
View File
@@ -29,7 +29,6 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"Cosmos2VideoToWorldPipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
"MatrixGameCausalDMDPipeline": "matrixgame",
"LongCatPipeline": "longcat",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
-15
View File
@@ -187,21 +187,6 @@ 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
@@ -136,26 +136,16 @@ class LatentPreparationStage(PipelineStage):
)
# Generate or use provided latents
if latents is None:
latents = randn_tensor(
shape,
generator=generator,
device=device,
dtype=dtype,
)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
latents = randn_tensor(shape,
generator=generator,
device=device,
dtype=dtype)
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
@@ -1,179 +0,0 @@
# 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
@@ -1,310 +0,0 @@
# 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
@@ -1,104 +0,0 @@
# 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
@@ -7,8 +7,6 @@ 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
@@ -73,12 +71,7 @@ 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."
)
# 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,
scheduler.set_timesteps(timesteps=timesteps,
device=device,
**extra_set_timesteps_kwargs)
timesteps = scheduler.timesteps
+4 -3
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 fastvideo_kernel import sliding_tile_attention # noqa: F401
from st_attn 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 fastvideo_kernel import video_sparse_attn # noqa: F401
from vsa import block_sparse_attn # noqa: F401
from fastvideo.attention.backends.video_sparse_attn import ( # noqa: F401
VideoSparseAttentionBackend)
@@ -188,7 +188,8 @@ class CudaPlatformBase(Platform):
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
try:
from fastvideo_kernel import moba_attn_varlen # noqa: F401
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
moba_attn_varlen)
from fastvideo.attention.backends.vmoba import ( # noqa: F401
VMOBAAttentionBackend)
logger.info("Using Video MOBA Attention backend.")
+2 -24
View File
@@ -58,13 +58,6 @@ 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:
@@ -78,23 +71,8 @@ class RocmPlatform(Platform):
elif selected_backend in (AttentionBackendEnum.FLASH_ATTN, None):
pass
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):
elif selected_backend in (AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN):
raise ValueError(
f"{selected_backend.name} is not supported on {cls.device_name}."
)
+2 -9
View File
@@ -4,7 +4,6 @@ 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}")
@@ -75,15 +74,9 @@ 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:4",
image=image,
timeout=2700,
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
volumes={"/root/data": model_vol}
)
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
def run_ssim_tests():
run_test("export MODEL_PATH='/root/data/weights' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
run_test("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():
@@ -1,3 +0,0 @@
"""Block-sparse attention kernels for LongCat."""
@@ -1,656 +0,0 @@
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
@@ -1,111 +0,0 @@
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)
@@ -1,43 +0,0 @@
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
@@ -1,946 +0,0 @@
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))
@@ -901,58 +901,6 @@ class DistillationPipeline(TrainingPipeline):
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
with self.tracker.timed("timing/get_next_batch"):
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
device = get_local_torch_device()
dtype = torch.bfloat16
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
if self.training_args.simulate_generator_forward:
batch_size = encoder_hidden_states.shape[0]
vae_config = self.training_args.pipeline_config.vae_config.arch_config
num_channels = vae_config.z_dim
spatial_compression_ratio = vae_config.spatial_compression_ratio
latent_height = self.training_args.num_height // spatial_compression_ratio
latent_width = self.training_args.num_width // spatial_compression_ratio
latents = torch.zeros(
batch_size,
num_channels,
self.training_args.num_latent_t,
latent_height,
latent_width,
device=device,
dtype=dtype,
)
else:
if 'vae_latent' not in batch:
raise ValueError(
"vae_latent not found in batch and simulate_generator_forward is False"
)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
latents = latents.to(device, dtype=dtype)
training_batch.latents = latents
training_batch.encoder_hidden_states = encoder_hidden_states.to(
device, dtype=dtype)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
device, dtype=dtype)
training_batch.infos = infos
return training_batch
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
gradient_accumulation_steps = getattr(self.training_args,
'gradient_accumulation_steps', 1)
-9
View File
@@ -543,18 +543,9 @@ def maybe_download_lora(model_name_or_path: str,
Local path to the model
"""
# If it's already a file path, return it directly
if os.path.isfile(model_name_or_path):
return model_name_or_path
local_path = maybe_download_model(model_name_or_path, local_dir, download)
weight_name = _best_guess_weight_name(model_name_or_path,
file_extension=".safetensors")
# If weight_name is None, assume local_path is already the full path
if weight_name is None:
return local_path
return os.path.join(local_path, weight_name)
+3 -8
View File
@@ -27,7 +27,7 @@ dependencies = [
"timm==1.0.11",
"peft>=0.15.0",
"diffusers>=0.33.1",
"torch>=2.9.0",
"torch==2.9.0",
"torchvision",
# Acceleration & Optimization
@@ -74,7 +74,6 @@ dependencies = [
# Preprocessing Dependencies
"torchcodec==0.5.0",
"ray>=2.49.1",
"ftfy==6.3.1",
]
[tool.uv]
@@ -102,14 +101,14 @@ explicit = true
[project.optional-dependencies]
# flash-attn: pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
# flash-attn: pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
lint = [
"pre-commit==4.0.1",
]
test = [
test = [
"av==14.3.0",
"pytorch-msssim==1.0.0",
"pytest",
@@ -117,10 +116,6 @@ test = [
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
rocm = [
"amdsmi",
]
[project.scripts]
fastvideo = "fastvideo.entrypoints.cli.main:main"
-182
View File
@@ -1,182 +0,0 @@
[build-system]
requires = ["setuptools>=61.0"]
build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.1.6"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.10"
classifiers = [
"Programming Language :: Python :: 3",
"License :: OSI Approved :: Apache Software License",
]
dependencies = [
# Core Libraries
"scipy==1.14.1",
"six==1.16.0",
"h5py==3.12.1",
"requests>=2.32.2",
# Machine Learning & Transformers
"transformers==4.57.3",
"tokenizers>=0.20.1",
"sentencepiece==0.2.0",
"timm==1.0.11",
"peft>=0.15.0",
"diffusers>=0.33.1",
"torch>=2.9.1",
"torchvision",
# Acceleration & Optimization
"accelerate==1.0.1",
# Computer Vision & Image Processing
"opencv-python==4.10.0.84",
"pillow>=10.3.0",
"imageio==2.36.0",
"imageio-ffmpeg==0.5.1",
"einops",
# Experiment Tracking & Logging
"wandb>=0.21.0",
"loguru",
"test-tube==0.7.5",
# Miscellaneous Utilities
"tqdm",
"pytest",
"PyYAML==6.0.1",
"protobuf>=5.28.3",
"gradio==5.32.0",
"moviepy>=2.0.0",
"flask",
"flask_restful",
"aiohttp",
"huggingface_hub",
"cloudpickle",
# System & Monitoring Tools
"gpustat",
"watch",
"remote-pdb",
# Kernel & Packaging
"wheel",
# Training Dependencies
"torchdata",
"pyarrow",
"datasets==4.0.0",
"av",
# Preprocessing Dependencies
"torchcodec==0.5.0",
"ray>=2.49.1",
]
[tool.uv]
prerelease = "allow"
[project.optional-dependencies]
# flash-attn: pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
lint = [
"pre-commit==4.0.1",
]
test = [
"av==14.3.0",
"pytorch-msssim==1.0.0",
"pytest",
]
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
rocm = [
"amdsmi",
]
[project.scripts]
fastvideo = "fastvideo.entrypoints.cli.main:main"
[tool.setuptools.packages.find]
exclude = ["assets*", "docker*", "docs", "scripts*"]
[tool.wheel]
exclude = ["assets*", "docker*", "docs", "scripts*"]
[tool.mypy]
warn_unused_configs = true
ignore_missing_imports = true
disallow_untyped_calls = true
check_untyped_defs = true
follow_imports = "silent"
[tool.codespell]
skip ="./data,./wandb,./csrc/sliding_tile_attention/tk"
[tool.ruff]
# Allow lines to be as long as 80.
line-length = 80
[tool.ruff.lint]
select = [
# pycodestyle
"E",
# Pyflakes
"F",
# pyupgrade
"UP",
# flake8-bugbear
"B",
# flake8-simplify
"SIM",
# isort
# "I",
"G",
]
ignore = [
# star imports
"F405", "F403",
# lambda expression assignment
"E731",
# Loop control variable not used within loop body
"B007",
# f-string format
"UP032",
# line too long
"E501",
]
[tool.ruff.lint.per-file-ignores]
"fastvideo/models/stepvideo/diffusion/video_pipeline.py" = ["F821"]
"fastvideo/sample/call_remote_server_stepvideo.py" = ["E722"]
"csrc/sliding_tile_attention/test/bench.py" = ["F841"]
"fastvideo/models/stepvideo/__init__.py" = ["F403"]
"fastvideo/models/stepvideo/utils/__init__.py" = ["F403"]
# Ignore all files that end in `_test.py`.
"fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py" = ["E741"]
[tool.yapf]
column_limit = 80
[tool.isort]
line_length = 80
use_parentheses = true
skip_gitignore = true
[project.urls]
Repository = "https://github.com/hao-ai-lab/FastVideo"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "fastvideo"
DisplayName = "FastVideo"
Icon = "https://raw.githubusercontent.com/hao-ai-lab/FastVideo/main/comfyui/assets/icon_simple.svg"
@@ -1,708 +0,0 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""
Convert LongCat weights to FastVideo native format.
This script performs a complete conversion from original LongCat weights
to FastVideo native implementation in a single step:
1. Converts transformer weights (with QKV/KV splitting)
2. Copies other components (VAE, text encoder, tokenizer, scheduler)
3. Converts LoRA weights (cfg_step_lora, refinement_lora)
4. Updates config files to point to native model
Usage:
python scripts/checkpoint_conversion/longcat_to_fastvideo.py \
--source /path/to/LongCat-Video/weights/LongCat-Video \
--output weights/longcat-native \
--validate
"""
import argparse
import glob
import json
import re
import shutil
from pathlib import Path
from collections import OrderedDict
import torch
from safetensors.torch import load_file, save_file
from tqdm import tqdm
def split_qkv(qkv_weight: torch.Tensor, qkv_bias: torch.Tensor | None = None):
"""Split fused QKV projection into separate Q, K, V."""
dim = qkv_weight.shape[0] // 3
q, k, v = torch.chunk(qkv_weight, 3, dim=0)
if qkv_bias is not None:
q_bias, k_bias, v_bias = torch.chunk(qkv_bias, 3, dim=0)
else:
q_bias = k_bias = v_bias = None
return (q, k, v), (q_bias, k_bias, v_bias)
def split_kv(kv_weight: torch.Tensor, kv_bias: torch.Tensor | None = None):
"""Split fused KV projection into separate K, V."""
dim = kv_weight.shape[0] // 2
k, v = torch.chunk(kv_weight, 2, dim=0)
if kv_bias is not None:
k_bias, v_bias = torch.chunk(kv_bias, 2, dim=0)
else:
k_bias = v_bias = None
return (k, v), (k_bias, v_bias)
def convert_transformer_weights(source_weights: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
"""
Convert LongCat transformer weights to native FastVideo format.
Main transformations:
1. Split fused QKV projections (self-attention)
2. Split fused KV projections (cross-attention)
3. Rename parameters according to mapping
"""
converted = OrderedDict()
processed_keys = set()
print(" Converting transformer weights...")
for key, value in tqdm(source_weights.items(), desc=" Processing parameters"):
if key in processed_keys:
continue
# === Embedders ===
if key.startswith("x_embedder."):
new_key = key.replace("x_embedder.", "patch_embed.")
converted[new_key] = value
elif key.startswith("t_embedder.mlp.0."):
new_key = key.replace("t_embedder.mlp.0.", "time_embedder.linear_1.")
converted[new_key] = value
elif key.startswith("t_embedder.mlp.2."):
new_key = key.replace("t_embedder.mlp.2.", "time_embedder.linear_2.")
converted[new_key] = value
elif key.startswith("y_embedder.y_proj.0."):
new_key = key.replace("y_embedder.y_proj.0.", "caption_embedder.linear_1.")
converted[new_key] = value
elif key.startswith("y_embedder.y_proj.2."):
new_key = key.replace("y_embedder.y_proj.2.", "caption_embedder.linear_2.")
converted[new_key] = value
# === Self-Attention QKV Splitting ===
elif ".attn.qkv.weight" in key:
block_idx = key.split(".")[1]
qkv_weight = value
qkv_bias_key = key.replace(".weight", ".bias")
qkv_bias = source_weights.get(qkv_bias_key)
(q, k, v), (q_bias, k_bias, v_bias) = split_qkv(qkv_weight, qkv_bias)
converted[f"blocks.{block_idx}.self_attn.to_q.weight"] = q
converted[f"blocks.{block_idx}.self_attn.to_k.weight"] = k
converted[f"blocks.{block_idx}.self_attn.to_v.weight"] = v
if q_bias is not None:
converted[f"blocks.{block_idx}.self_attn.to_q.bias"] = q_bias
converted[f"blocks.{block_idx}.self_attn.to_k.bias"] = k_bias
converted[f"blocks.{block_idx}.self_attn.to_v.bias"] = v_bias
processed_keys.add(key)
if qkv_bias is not None:
processed_keys.add(qkv_bias_key)
elif ".attn.qkv.bias" in key:
continue
elif ".attn.proj." in key:
new_key = key.replace(".attn.proj.", ".self_attn.to_out.")
converted[new_key] = value
elif ".attn.q_norm." in key or ".attn.k_norm." in key:
new_key = key.replace(".attn.", ".self_attn.")
converted[new_key] = value
# === Cross-Attention ===
elif ".cross_attn.q_linear." in key:
new_key = key.replace(".cross_attn.q_linear.", ".cross_attn.to_q.")
converted[new_key] = value
elif ".cross_attn.kv_linear.weight" in key:
block_idx = key.split(".")[1]
kv_weight = value
kv_bias_key = key.replace(".weight", ".bias")
kv_bias = source_weights.get(kv_bias_key)
(k, v), (k_bias, v_bias) = split_kv(kv_weight, kv_bias)
converted[f"blocks.{block_idx}.cross_attn.to_k.weight"] = k
converted[f"blocks.{block_idx}.cross_attn.to_v.weight"] = v
if k_bias is not None:
converted[f"blocks.{block_idx}.cross_attn.to_k.bias"] = k_bias
converted[f"blocks.{block_idx}.cross_attn.to_v.bias"] = v_bias
processed_keys.add(key)
if kv_bias is not None:
processed_keys.add(kv_bias_key)
elif ".cross_attn.kv_linear.bias" in key:
continue
elif ".cross_attn.proj." in key:
new_key = key.replace(".cross_attn.proj.", ".cross_attn.to_out.")
converted[new_key] = value
elif ".cross_attn.q_norm." in key or ".cross_attn.k_norm." in key:
converted[key] = value
# === Final Layer (must come BEFORE general transformer block patterns) ===
elif key.startswith("final_layer.adaLN_modulation.1."):
new_key = key.replace("final_layer.adaLN_modulation.1.", "final_layer.adaln_linear.")
converted[new_key] = value
# === Transformer Block AdaLN ===
elif ".adaLN_modulation.1." in key:
new_key = key.replace(".adaLN_modulation.1.", ".adaln_linear_1.")
converted[new_key] = value
# === Transformer Block Normalization ===
elif ".mod_norm_attn." in key or ".mod_norm_ffn." in key:
continue
elif ".pre_crs_attn_norm.weight" in key:
new_key = key.replace(".pre_crs_attn_norm.", ".norm_cross.")
converted[new_key] = value
elif ".pre_crs_attn_norm.bias" in key:
new_key = key.replace(".pre_crs_attn_norm.", ".norm_cross.")
converted[new_key] = value
# === FFN (SwiGLU) ===
elif ".ffn.w1." in key or ".ffn.w2." in key or ".ffn.w3." in key:
converted[key] = value
elif key.startswith("final_layer.norm_final."):
continue
elif key.startswith("final_layer.linear."):
new_key = key.replace("final_layer.linear.", "final_layer.proj.")
converted[new_key] = value
else:
print(f" ⚠️ Unknown key: {key}")
converted[key] = value
return converted
def validate_conversion(original: dict, converted: dict) -> bool:
"""Validate that conversion preserved all parameters correctly."""
print("\n Validating conversion...")
orig_count = sum(p.numel() for p in original.values())
conv_count = sum(p.numel() for p in converted.values())
dropped_count = 0
for key, value in original.items():
if ".mod_norm_attn." in key or ".mod_norm_ffn." in key:
dropped_count += value.numel()
elif "final_layer.norm_final." in key:
dropped_count += value.numel()
expected_conv_count = orig_count - dropped_count
print(f" Original parameters: {orig_count:,}")
print(f" Converted parameters: {conv_count:,}")
print(f" Dropped parameters (norms without params): {dropped_count:,}")
if conv_count != expected_conv_count:
print(f" ⚠️ Parameter count mismatch!")
return False
print(f" ✓ Parameter count matches")
# Verify QKV/KV splits
print("\n Verifying QKV/KV splits...")
num_blocks = 48
for i in range(num_blocks):
orig_qkv_weight = original.get(f"blocks.{i}.attn.qkv.weight")
if orig_qkv_weight is not None:
conv_q = converted[f"blocks.{i}.self_attn.to_q.weight"]
conv_k = converted[f"blocks.{i}.self_attn.to_k.weight"]
conv_v = converted[f"blocks.{i}.self_attn.to_v.weight"]
reconstructed = torch.cat([conv_q, conv_k, conv_v], dim=0)
if not torch.allclose(orig_qkv_weight, reconstructed):
print(f" ❌ QKV weight mismatch in block {i}")
return False
orig_kv_weight = original.get(f"blocks.{i}.cross_attn.kv_linear.weight")
if orig_kv_weight is not None:
conv_k = converted[f"blocks.{i}.cross_attn.to_k.weight"]
conv_v = converted[f"blocks.{i}.cross_attn.to_v.weight"]
reconstructed = torch.cat([conv_k, conv_v], dim=0)
if not torch.allclose(orig_kv_weight, reconstructed):
print(f" ❌ KV weight mismatch in block {i}")
return False
print(f" ✓ All splits verified successfully")
return True
def copy_component(source_dir: Path, output_dir: Path, component: str, mapping: dict = None) -> bool:
"""Copy a component directory, optionally with name mapping."""
source_name = mapping.get(component, component) if mapping else component
source_path = source_dir / source_name
if source_path.exists():
output_path = output_dir / component
if output_path.exists():
shutil.rmtree(output_path)
shutil.copytree(source_path, output_path)
print(f" ✓ {component} copied")
return True
else:
print(f" ⚠️ {component} not found, skipping")
return False
def create_model_index():
"""Create model_index.json for FastVideo native model."""
return {
"_class_name": "LongCatPipeline",
"_diffusers_version": "0.32.0",
"workload_type": "video-generation",
"tokenizer": ["transformers", "AutoTokenizer"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"transformer": ["diffusers", "LongCatTransformer3DModel"] # Native model
}
def update_transformer_config(transformer_dir: Path):
"""Update transformer config.json to point to native model."""
config_path = transformer_dir / "config.json"
if not config_path.exists():
print(" ⚠️ Transformer config not found, skipping")
return
with open(config_path, 'r') as f:
config = json.load(f)
if '_class_name' in config:
old_class = config['_class_name']
config['_class_name'] = 'LongCatTransformer3DModel'
print(f" Updated _class_name: {old_class} → LongCatTransformer3DModel")
else:
config['_class_name'] = 'LongCatTransformer3DModel'
print(f" Added _class_name: LongCatTransformer3DModel")
# Fix num_heads -> num_attention_heads for FastVideo compatibility
if 'num_heads' in config and 'num_attention_heads' not in config:
config['num_attention_heads'] = config.pop('num_heads')
print(f" Updated num_heads → num_attention_heads")
with open(config_path, 'w') as f:
json.dump(config, f, indent=2)
print(" ✓ Transformer config updated")
# ============================================================================
# LoRA Conversion Functions
# ============================================================================
def parse_lora_key(key: str) -> tuple[str, str]:
"""Parse LongCat LoRA key into module path and weight type."""
if key.startswith("lora___lorahyphen___"):
key = key[len("lora___lorahyphen___"):]
key = key.replace("___lorahyphen___", ".")
if ".lora_down.weight" in key:
return key.replace(".lora_down.weight", ""), "lora_down.weight"
elif ".lora_up.weight" in key:
return key.replace(".lora_up.weight", ""), "lora_up.weight"
elif ".lora_up.blocks." in key:
match = re.match(r"(.+)\.lora_up\.blocks\.(\d+)\.weight", key)
if match:
return match.group(1), f"lora_up.blocks.{match.group(2)}.weight"
elif ".alpha_scale" in key:
return key.replace(".alpha_scale", ""), "alpha_scale"
elif ".lora_alpha" in key:
return key.replace(".lora_alpha", ""), "lora_alpha"
raise ValueError(f"Unknown LoRA key format: {key}")
def map_lora_module(module_path: str) -> list[tuple[str, str]]:
"""Map LongCat module path to FastVideo paths. Returns [(path, component)]."""
# Self-attention QKV → Q, K, V
match = re.match(r"blocks\.(\d+)\.attn\.qkv", module_path)
if match:
b = match.group(1)
return [(f"blocks.{b}.self_attn.to_q", "q"),
(f"blocks.{b}.self_attn.to_k", "k"),
(f"blocks.{b}.self_attn.to_v", "v")]
# Self-attention output
match = re.match(r"blocks\.(\d+)\.attn\.proj", module_path)
if match:
return [(f"blocks.{match.group(1)}.self_attn.to_out", "single")]
# Cross-attention Q
match = re.match(r"blocks\.(\d+)\.cross_attn\.q_linear", module_path)
if match:
return [(f"blocks.{match.group(1)}.cross_attn.to_q", "single")]
# Cross-attention KV → K, V
match = re.match(r"blocks\.(\d+)\.cross_attn\.kv_linear", module_path)
if match:
b = match.group(1)
return [(f"blocks.{b}.cross_attn.to_k", "k"),
(f"blocks.{b}.cross_attn.to_v", "v")]
# FFN
match = re.match(r"blocks\.(\d+)\.ffn\.(w[123])", module_path)
if match:
return [(f"blocks.{match.group(1)}.ffn.{match.group(2)}", "single")]
# AdaLN modulation
match = re.match(r"blocks\.(\d+)\.adaLN_modulation\.1", module_path)
if match:
return [(f"blocks.{match.group(1)}.adaln_linear_1", "single")]
# Final layer
if module_path == "final_layer.adaLN_modulation.1":
return [("final_layer.adaln_linear", "single")]
if module_path == "final_layer.linear":
return [("final_layer.proj", "single")]
raise ValueError(f"Unknown LoRA module: {module_path}")
def convert_lora_weights(source_weights: dict[str, torch.Tensor], lora_name: str) -> dict[str, torch.Tensor]:
"""Convert LongCat LoRA to FastVideo format."""
print(f" Converting {lora_name}...")
print(f" Source keys: {len(source_weights)}")
converted = OrderedDict()
# Group by module
modules = {}
for key in source_weights.keys():
try:
module_path, weight_type = parse_lora_key(key)
if module_path not in modules:
modules[module_path] = {}
modules[module_path][weight_type] = key
except ValueError:
continue
# Process each module
for module_path, weight_keys in modules.items():
try:
targets = map_lora_module(module_path)
except ValueError:
continue
# Get alpha_scale if present (defaults to 1.0 if missing)
alpha_scale = 1.0
if "alpha_scale" in weight_keys:
alpha_scale_tensor = source_weights[weight_keys["alpha_scale"]]
alpha_scale = alpha_scale_tensor.item() if alpha_scale_tensor.numel() == 1 else float(alpha_scale_tensor.mean())
# Handle lora_down (lora_A)
if "lora_down.weight" in weight_keys:
lora_down = source_weights[weight_keys["lora_down.weight"]]
if len(targets) == 1:
converted[f"{targets[0][0]}.lora_A"] = lora_down
# Compute alpha from alpha_scale and rank
rank = lora_down.shape[0]
alpha = alpha_scale * rank
converted[f"{targets[0][0]}.lora_alpha"] = torch.tensor(alpha, dtype=torch.float32)
else:
# Split for fused projections
n = len(targets)
rank = lora_down.shape[0] // n
for i, (path, _) in enumerate(targets):
converted[f"{path}.lora_A"] = lora_down[i*rank:(i+1)*rank, :]
# Compute alpha from alpha_scale and rank for each split
alpha = alpha_scale * rank
converted[f"{path}.lora_alpha"] = torch.tensor(alpha, dtype=torch.float32)
# Handle lora_up (lora_B) - may have multiple blocks
lora_up_blocks = []
i = 0
while f"lora_up.blocks.{i}.weight" in weight_keys:
lora_up_blocks.append(source_weights[weight_keys[f"lora_up.blocks.{i}.weight"]])
i += 1
if lora_up_blocks:
# Multi-block LoRA: construct block-diagonal lora_B
# This is equivalent to the multi-block computation without modifying fastvideo
n_blocks = len(lora_up_blocks)
out_per_block, rank_per_block = lora_up_blocks[0].shape # e.g., [4096, 128]
if len(targets) == 1:
# Single layer with multi-block: create block-diagonal matrix
total_out = out_per_block * n_blocks
total_rank = rank_per_block * n_blocks
lora_B_blockdiag = torch.zeros(total_out, total_rank, dtype=lora_up_blocks[0].dtype)
for i in range(n_blocks):
lora_B_blockdiag[i*out_per_block:(i+1)*out_per_block,
i*rank_per_block:(i+1)*rank_per_block] = lora_up_blocks[i]
converted[f"{targets[0][0]}.lora_B"] = lora_B_blockdiag
# Note: rank for alpha calculation should be total_rank (will be computed from lora_A.shape[0])
else:
# Multi-block with split targets (e.g., QKV split)
# Each target gets one block
for i, (path, _) in enumerate(targets):
if i < n_blocks:
converted[f"{path}.lora_B"] = lora_up_blocks[i]
elif "lora_up.weight" in weight_keys:
lora_up = source_weights[weight_keys["lora_up.weight"]]
# Split if needed
if len(targets) == 1:
converted[f"{targets[0][0]}.lora_B"] = lora_up
else:
n = len(targets)
out_dim = lora_up.shape[0] // n
for i, (path, _) in enumerate(targets):
converted[f"{path}.lora_B"] = lora_up[i*out_dim:(i+1)*out_dim, :]
else:
continue
print(f" Output keys: {len(converted)} (including lora_alpha)")
# Count how many lora_alpha values were added
alpha_count = sum(1 for k in converted.keys() if "lora_alpha" in k)
print(f" Alpha values saved: {alpha_count}")
return converted
def convert_loras(source_dir: Path, output_dir: Path) -> bool:
"""Convert all LoRA files in source directory."""
lora_source = source_dir / "lora"
if not lora_source.exists():
print(" No LoRA directory found, skipping")
return False
lora_files = list(lora_source.glob("*.safetensors"))
if not lora_files:
print(" No LoRA files found, skipping")
return False
print(f" Found {len(lora_files)} LoRA file(s)")
# Map LoRA filenames to subdirectory names for FastVideo compatibility
# Each LoRA gets its own directory under lora/
lora_subdir_mapping = {
"cfg_step_lora.safetensors": "distilled",
"refinement_lora.safetensors": "refinement",
}
for lora_file in lora_files:
try:
# Determine output subdirectory - use mapping if available, otherwise generic name
if lora_file.name in lora_subdir_mapping:
lora_subdir_name = lora_subdir_mapping[lora_file.name]
else:
# For unknown LoRAs, create subdirectory based on filename
lora_subdir_name = lora_file.stem
lora_output = output_dir / "lora" / lora_subdir_name
lora_output.mkdir(parents=True, exist_ok=True)
# Load
source_weights = load_file(str(lora_file))
# Convert
converted = convert_lora_weights(source_weights, lora_file.stem)
# Save
output_file = lora_output / lora_file.name
save_file(converted, str(output_file))
size_mb = output_file.stat().st_size / (1024**2)
print(f" ✓ {lora_file.name} → lora/{lora_subdir_name}/ ({size_mb:.1f} MB)")
except Exception as e:
print(f" ❌ Failed to convert {lora_file.name}: {e}")
return False
return True
def main():
parser = argparse.ArgumentParser(
description="Convert LongCat weights to FastVideo native format"
)
parser.add_argument(
"--source",
type=str,
required=True,
help="Path to original LongCat weights (LongCat-Video/weights/LongCat-Video/)",
)
parser.add_argument(
"--output",
type=str,
required=True,
help="Path to output directory for native weights",
)
parser.add_argument(
"--validate",
action="store_true",
help="Run validation after conversion",
)
args = parser.parse_args()
source_dir = Path(args.source)
output_dir = Path(args.output)
# Check source directory
if not source_dir.exists():
print(f"❌ Error: Source directory not found: {source_dir}")
return 1
# Check for dit/transformer directory (original uses 'dit', we output to 'transformer')
transformer_source = source_dir / "dit"
if not transformer_source.exists():
print(f"❌ Error: DiT directory not found in source")
return 1
print("=" * 60)
print("LongCat → FastVideo Native Conversion")
print("=" * 60)
print(f"Source: {source_dir}")
print(f"Output: {output_dir}")
print()
# Step 1: Convert transformer weights
print("[Step 1/4] Converting transformer weights...")
# Load source weights
shard_files = sorted(glob.glob(str(transformer_source / "*.safetensors")))
if not shard_files:
print(f"❌ Error: No safetensors files found in {transformer_source}")
return 1
print(f" Found {len(shard_files)} shard(s)")
source_weights = {}
for shard_file in shard_files:
print(f" Loading {Path(shard_file).name}...")
source_weights.update(load_file(shard_file))
print(f" Loaded {len(source_weights)} parameters")
# Convert
converted_weights = convert_transformer_weights(source_weights)
print(f" Converted to {len(converted_weights)} parameters")
# Validate if requested
if args.validate:
if not validate_conversion(source_weights, converted_weights):
print("\n❌ Validation failed!")
return 1
print("\n✓ Validation passed!")
# Save
transformer_output = output_dir / "transformer"
transformer_output.mkdir(parents=True, exist_ok=True)
output_file = transformer_output / "model.safetensors"
print(f"\n Saving to {output_file}...")
save_file(converted_weights, str(output_file))
size_gb = output_file.stat().st_size / (1024**3)
print(f" ✓ Saved ({size_gb:.2f} GB)")
print()
# Step 2: Copy other components
print("[Step 2/5] Copying other components...")
output_dir.mkdir(parents=True, exist_ok=True)
components = ["vae", "text_encoder", "tokenizer", "scheduler"]
for component in components:
copy_component(source_dir, output_dir, component)
print()
# Step 3: Convert LoRA weights
print("[Step 3/5] Converting LoRA weights...")
convert_loras(source_dir, output_dir)
print()
# Step 4: Update transformer config
print("[Step 4/5] Updating transformer config...")
# Copy config.json from source if exists
source_config = transformer_source / "config.json"
output_config = transformer_output / "config.json"
if source_config.exists():
shutil.copy(source_config, output_config)
print(f" Copied config.json")
update_transformer_config(transformer_output)
print()
# Step 5: Create model_index.json
print("[Step 5/5] Creating model_index.json...")
model_index_path = output_dir / "model_index.json"
with open(model_index_path, 'w') as f:
json.dump(create_model_index(), f, indent=2)
print(f" ✓ Created {model_index_path}")
print()
print("=" * 60)
print("✓ Conversion Complete!")
print("=" * 60)
print(f"Native weights ready at: {output_dir}")
print()
print("Converted components:")
print(" ✓ Transformer (native FastVideo implementation)")
print(" ✓ VAE, text encoder, tokenizer, scheduler")
if (output_dir / "lora").exists():
lora_dirs = [d for d in (output_dir / "lora").iterdir() if d.is_dir()]
if lora_dirs:
print(f" ✓ LoRA weights ({len(lora_dirs)} adapters)")
for lora_dir in sorted(lora_dirs):
print(f" - lora/{lora_dir.name}/")
print()
print("Next steps:")
print()
print(" 1. Test basic generation:")
print(" from fastvideo import VideoGenerator")
print(f" generator = VideoGenerator.from_pretrained('{output_dir}')")
print(" video = generator.generate_video(")
print(" prompt='A cat playing piano',")
print(" num_inference_steps=50")
print(" )")
print()
if (output_dir / "lora" / "distilled").exists():
print(" 2. Test distilled generation (16 steps with LoRA):")
print(f" generator = VideoGenerator.from_pretrained('{output_dir}',")
print(f" lora_path='{output_dir}/lora/distilled',")
print(" lora_nickname='distilled')")
print(" video = generator.generate_video(")
print(" prompt='A cat playing piano',")
print(" num_inference_steps=16,")
print(" guidance_scale=1.0)")
print()
print()
return 0
if __name__ == "__main__":
exit(main())
@@ -1,241 +0,0 @@
import os
import json
import argparse
from pathlib import Path
from safetensors import safe_open
def validate_components(model_path):
"""Validate all model components exist."""
print("=" * 60)
print("VALIDATING COMPONENTS")
print("=" * 60)
components = {
"tokenizer": ["special_tokens_map.json", "tokenizer_config.json"],
"text_encoder": ["config.json", "model.safetensors.index.json"],
"vae": ["config.json", "diffusion_pytorch_model.safetensors"],
"scheduler": ["scheduler_config.json"],
"transformer": ["config.json", "diffusion_pytorch_model.safetensors.index.json"]
}
all_valid = True
for component, required_files in components.items():
component_path = os.path.join(model_path, component)
print(f"\n{component}:")
if not os.path.exists(component_path):
print(f" ✗ Directory not found")
all_valid = False
continue
for req_file in required_files:
file_path = os.path.join(component_path, req_file)
exists = os.path.exists(file_path)
symbol = "✓" if exists else "✗"
print(f" {symbol} {req_file}")
if not exists:
all_valid = False
return all_valid
def validate_dit_weights(dit_path):
"""Validate DiT weights structure."""
print("\n" + "=" * 60)
print("VALIDATING DiT WEIGHTS")
print("=" * 60)
# Load config
config_path = os.path.join(dit_path, "config.json")
with open(config_path) as f:
config = json.load(f)
hidden_size = config["hidden_size"]
depth = config["depth"]
num_heads = config["num_heads"]
print(f"\nArchitecture:")
print(f" - hidden_size: {hidden_size}")
print(f" - depth: {depth}")
print(f" - num_heads: {num_heads}")
# Load weight index
index_path = os.path.join(dit_path, "diffusion_pytorch_model.safetensors.index.json")
with open(index_path) as f:
index = json.load(f)
weight_map = index["weight_map"]
all_keys = list(weight_map.keys())
print(f"\nWeight statistics:")
print(f" - Total keys: {len(all_keys)}")
print(f" - Total size: {index['metadata']['total_size'] / 1e9:.2f} GB")
# Check structure
embedder_keys = [k for k in all_keys if 'embedder' in k]
block_keys = [k for k in all_keys if k.startswith('blocks.')]
final_keys = [k for k in all_keys if k.startswith('final_layer.')]
print(f"\nKey distribution:")
print(f" - Embedder layers: {len(embedder_keys)}")
print(f" - Transformer blocks: {len(block_keys)}")
print(f" - Final layer: {len(final_keys)}")
# Verify all blocks present
block_nums = set()
for key in block_keys:
if key.startswith('blocks.'):
block_num = int(key.split('.')[1])
block_nums.add(block_num)
expected_blocks = set(range(depth))
missing_blocks = expected_blocks - block_nums
if missing_blocks:
print(f"\n✗ Missing blocks: {sorted(missing_blocks)}")
return False
else:
print(f"\n✓ All {depth} blocks present (0-{depth-1})")
# Sample weights
first_shard = os.path.join(dit_path, "diffusion_pytorch_model-00001-of-00006.safetensors")
print(f"\nSampling weights from first shard:")
with safe_open(first_shard, framework="pt", device="cpu") as f:
sample_keys = [k for k in f.keys() if k in all_keys][:5]
for key in sample_keys:
tensor = f.get_tensor(key)
print(f" - {key}")
print(f" Shape: {tuple(tensor.shape)}, Dtype: {tensor.dtype}")
return True
def validate_shapes(dit_path):
"""Validate weight shapes match expected architecture."""
print("\n" + "=" * 60)
print("VALIDATING WEIGHT SHAPES")
print("=" * 60)
# Load config
config_path = os.path.join(dit_path, "config.json")
with open(config_path) as f:
config = json.load(f)
hidden_size = config["hidden_size"]
num_heads = config["num_heads"]
head_dim = hidden_size // num_heads
adaln_dim = config.get("adaln_tembed_dim", 512)
mlp_ratio = config.get("mlp_ratio", 4)
# Calculate FFN hidden_dim using SwiGLU formula from blocks.py
# hidden_dim = int(2 * (hidden_size * mlp_ratio) / 3)
# rounded to multiple_of=256
multiple_of = 256
ffn_hidden = int(2 * hidden_size * mlp_ratio / 3)
ffn_hidden = multiple_of * ((ffn_hidden + multiple_of - 1) // multiple_of)
# Expected shapes
expected = {
"x_embedder.proj.weight": (hidden_size, 16, 1, 2, 2),
"x_embedder.proj.bias": (hidden_size,),
"t_embedder.mlp.0.weight": (adaln_dim, 256),
"t_embedder.mlp.2.weight": (adaln_dim, adaln_dim),
"y_embedder.y_proj.0.weight": (hidden_size, 4096),
"blocks.0.attn.qkv.weight": (3 * hidden_size, hidden_size),
"blocks.0.attn.q_norm.weight": (head_dim,),
"blocks.0.attn.proj.weight": (hidden_size, hidden_size),
"blocks.0.cross_attn.q_linear.weight": (hidden_size, hidden_size),
"blocks.0.cross_attn.kv_linear.weight": (2 * hidden_size, hidden_size),
"blocks.0.ffn.w1.weight": (ffn_hidden, hidden_size),
"blocks.0.adaLN_modulation.1.weight": (6 * hidden_size, adaln_dim),
"final_layer.linear.weight": (64, hidden_size),
}
# Load and check
first_shard = os.path.join(dit_path, "diffusion_pytorch_model-00001-of-00006.safetensors")
all_valid = True
with safe_open(first_shard, framework="pt", device="cpu") as f:
for key, expected_shape in expected.items():
if key in f.keys():
tensor = f.get_tensor(key)
actual_shape = tuple(tensor.shape)
if actual_shape == expected_shape:
print(f"✓ {key}: {actual_shape}")
else:
print(f"✗ {key}: expected {expected_shape}, got {actual_shape}")
all_valid = False
return all_valid
def validate_model_index(model_path):
"""Validate model_index.json exists and is correct."""
print("\n" + "=" * 60)
print("VALIDATING MODEL INDEX")
print("=" * 60)
model_index_path = os.path.join(model_path, "model_index.json")
if not os.path.exists(model_index_path):
print("✗ model_index.json not found")
return False
with open(model_index_path) as f:
index = json.load(f)
required_keys = ["_class_name", "workload_type", "tokenizer", "text_encoder",
"vae", "scheduler", "transformer"]
all_valid = True
for key in required_keys:
if key in index:
print(f"✓ {key}: {index[key]}")
else:
print(f"✗ {key}: missing")
all_valid = False
return all_valid
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Validate LongCat weights for FastVideo")
parser.add_argument("--model-path", type=str, required=True,
help="Path to LongCat model directory")
parser.add_argument("--check-shapes", action="store_true",
help="Also validate weight shapes (slower)")
args = parser.parse_args()
print(f"\nValidating: {args.model_path}\n")
# Run validations
components_valid = validate_components(args.model_path)
model_index_valid = validate_model_index(args.model_path)
dit_valid = validate_dit_weights(os.path.join(args.model_path, "transformer"))
if args.check_shapes:
shapes_valid = validate_shapes(os.path.join(args.model_path, "transformer"))
else:
shapes_valid = True
print("\nSkipping shape validation (use --check-shapes to enable)")
# Summary
print("\n" + "=" * 60)
print("VALIDATION SUMMARY")
print("=" * 60)
all_valid = components_valid and model_index_valid and dit_valid and shapes_valid
if all_valid:
print("✓ All validations passed!")
print("✓ Model ready for FastVideo")
else:
print("✗ Some validations failed")
print("✗ Please check errors above")
print("=" * 60)
@@ -9,7 +9,6 @@ fastvideo generate \
--model-path $MODEL_BASE \
--sp-size ${num_gpus} \
--tp-size 1 \
--num-gpus ${num_gpus} \
--height 768 \
--width 1280 \
--num-frames 117 \
-32
View File
@@ -1,32 +0,0 @@
#!/bin/bash
num_gpus=2
export FASTVIDEO_ATTENTION_BACKEND=
# For longcat, we must first convert the official weights to FastVideo native format
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
# --source /path/to/LongCat-Video/weights/LongCat-Video
# --output weights/longcat-native
export MODEL_BASE=weights/longcat-native
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--dit-cpu-offload False \
--vae-cpu-offload False \
--text-encoder-cpu-offload False \
--pin-cpu-memory False \
--enable-bsa False \
--height 480 \
--width 832 \
--num-frames 93 \
--num-inference-steps 50 \
--fps 15 \
--guidance-scale 4.0 \
--prompt-txt assets/prompt.txt \
--negative-prompt "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" \
--seed 42 \
--output-path outputs_video/longcat_480p
@@ -1,31 +0,0 @@
#!/bin/bash
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=
# For longcat, we must first convert the official weights to FastVideo native format
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
# --source /path/to/LongCat-Video/weights/LongCat-Video
# --output weights/longcat-native
export MODEL_BASE=weights/longcat-native
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--dit-cpu-offload False \
--vae-cpu-offload False \
--text-encoder-cpu-offload False \
--pin-cpu-memory False \
--enable-bsa False \
--lora-path "$MODEL_BASE/lora/distilled" \
--lora-nickname "distilled" \
--height 480 \
--width 832 \
--num-frames 93 \
--num-inference-steps 16 \
--fps 15 \
--guidance-scale 1.0 \
--prompt "In a realistic photography style, an asian boy around seven or eight years old sits on a park bench, wearing a light yellow T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene." \
--seed 42 \
--output-path outputs_video/longcat_distill
@@ -1,73 +0,0 @@
#!/bin/bash
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=
# For longcat, we must first convert the official weights to FastVideo native format
# conversion method: python scripts/checkpoint_conversion/longcat_to_fastvideo.py
# --source /path/to/LongCat-Video/weights/LongCat-Video
# --output weights/longcat-native
export MODEL_BASE=weights/longcat-native
INPUT_VIDEO="outputs_video/longcat_distill/In a realistic photography style, an asian boy around seven or eight years old sits on a park bench,.mp4"
REFINE_OUTPUT="outputs_video/longcat_refine_720p"
# Prompt used for base generation
PROMPT="In a realistic photography style, an asian boy around seven or eight years old sits on a park bench, wearing a light yellow T-shirt, denim shorts, and white sneakers. He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, who eagerly licks it with its tongue. The sun is shining brightly, and the background features a green lawn and several tall trees, creating a warm and loving scene."
echo "=========================================="
echo "LongCat 480p -> 720p Refinement"
echo "=========================================="
echo ""
echo "Input: $INPUT_VIDEO"
echo "Output: $REFINE_OUTPUT"
echo ""
# Check if input video exists
if [ ! -f "$INPUT_VIDEO" ]; then
echo "Error: Input video not found: $INPUT_VIDEO"
echo "Please set INPUT_VIDEO to your 480p video path"
exit 1
fi
echo "🔧 Configuring refinement (BSA enabled, refinement LoRA)..."
echo "✅ Input video: $INPUT_VIDEO"
echo "✅ BSA enabled with sparsity=0.875"
echo "✅ Refinement LoRA loaded"
echo ""
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--dit-cpu-offload True \
--vae-cpu-offload False \
--text-encoder-cpu-offload True \
--pin-cpu-memory False \
--enable-bsa True \
--bsa-sparsity 0.875 \
--bsa-chunk-q 4 4 8 \
--bsa-chunk-k 4 4 8 \
--lora-path "$MODEL_BASE/lora/refinement" \
--lora-nickname "refinement" \
--refine-from "$INPUT_VIDEO" \
--t-thresh 0.5 \
--spatial-refine-only False \
--num-cond-frames 0 \
--height 720 \
--width 1280 \
--num-inference-steps 50 \
--fps 30 \
--guidance-scale 1.0 \
--prompt "$PROMPT" \
--seed 42 \
--output-path "$REFINE_OUTPUT"
echo ""
echo "=========================================="
echo "✓ Refinement Complete!"
echo "=========================================="
echo ""
echo "Output directory: $REFINE_OUTPUT"
echo ""