Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9d822d464b | ||
|
|
9f86b28ecb | ||
|
|
95b9cde729 | ||
|
|
54b85d8931 | ||
|
|
922b082cb2 | ||
|
|
fb3bfecd18 | ||
|
|
f01fb88c6a | ||
|
|
e8f298c1ab | ||
|
|
d7c7d23375 | ||
|
|
44808ce145 | ||
|
|
0d5124e092 |
@@ -1,236 +0,0 @@
|
||||
name: Publish FastVideo Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/fastvideo_kernel/pyproject.toml"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/fastvideo_kernel
|
||||
# Get current commit's version from pyproject.toml
|
||||
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
torch-cuda:
|
||||
- torch-version: '2.5.1'
|
||||
cuda-version: '12.4.1'
|
||||
torch-cuda-short: 'cu124'
|
||||
- torch-version: '2.6.0'
|
||||
cuda-version: '12.6.3'
|
||||
torch-cuda-short: 'cu126'
|
||||
- torch-version: '2.7.1'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton
|
||||
|
||||
cd csrc/fastvideo_kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/fastvideo_kernel
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}-py${{ matrix.python-version }}
|
||||
path: csrc/fastvideo_kernel/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton
|
||||
|
||||
cd csrc/fastvideo_kernel
|
||||
git submodule update --init --recursive
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/fastvideo_kernel/dist/
|
||||
@@ -30,8 +30,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
@@ -2,12 +2,12 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We support H100 (via TK) and any other GPU (via triton) for STA.
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
@@ -16,14 +16,6 @@ git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you want to skip the compilation of the TK kernel and only use the Triton version, try below:
|
||||
|
||||
```bash
|
||||
SKIP_SM90_EXT=1 python setup.py install
|
||||
or
|
||||
SKIP_SM90_EXT=1 pip install --no-build-isolation .
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
@@ -38,7 +30,7 @@ sudo apt install clang-11
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
@@ -51,7 +43,7 @@ bash scripts/inference/v1_inference_wan_STA.sh
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
@@ -66,6 +58,7 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
@@ -74,7 +67,7 @@ python ../benchmarks/bench_sta.py
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
@@ -89,7 +82,7 @@ Here is a diagram of how the window is configured and passed through the FastVid
|
||||
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
@@ -51,28 +51,21 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
ext_modules = []
|
||||
|
||||
if os.environ.get("SKIP_SM90_EXT", "0") != "1":
|
||||
ext_modules.append(
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
)
|
||||
else:
|
||||
print("ENV SKIP_SM90_EXT=1, skip st_attn_cuda compile")
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=ext_modules,
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
|
||||
@@ -7,17 +7,12 @@ try:
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
try:
|
||||
from st_attn.st_attn_triton import sliding_tile_attention_triton
|
||||
except ImportError:
|
||||
sliding_tile_attention_triton = None
|
||||
|
||||
def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
seq_length = q_all.shape[2]
|
||||
dit_seq_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
@@ -51,13 +46,4 @@ def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, h
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
major, minor = torch.cuda.get_device_capability(q_all.device)
|
||||
if major == 9 and minor == 0 and sta_fwd is not None:
|
||||
return sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
|
||||
elif sliding_tile_attention_triton is not None:
|
||||
return sliding_tile_attention_triton(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
|
||||
else:
|
||||
raise ImportError("No suitable sliding tile attention implementation found.")
|
||||
return hidden_states[:, :, :seq_length]
|
||||
@@ -1,327 +0,0 @@
|
||||
import math
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
def is_cuda():
|
||||
return triton.runtime.driver.active.get_current_target().backend == "cuda"
|
||||
|
||||
|
||||
def is_hip():
|
||||
target = triton.runtime.driver.active.get_current_target()
|
||||
return target.backend == 'hip'
|
||||
|
||||
def get_common_autotune_config():
|
||||
configs = [
|
||||
triton.Config({'BLOCK_Q': BLOCK_Q, 'BLOCK_KV': BLOCK_KV}, num_stages=s, num_warps=w) \
|
||||
for BLOCK_Q in [32, 64, 128]\
|
||||
for BLOCK_KV in [32, 64, 128]\
|
||||
for s in [1, 2, 3, 4]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
return configs
|
||||
|
||||
|
||||
def get_cuda_autotune_config():
|
||||
# cuda and hip can use differnt autotune configs
|
||||
return get_common_autotune_config()
|
||||
|
||||
|
||||
def get_hip_autotune_config():
|
||||
# cuda and hip can use differnt autotune configs
|
||||
return get_common_autotune_config()
|
||||
|
||||
|
||||
def get_autotune_config():
|
||||
if is_cuda():
|
||||
return get_cuda_autotune_config()
|
||||
else:
|
||||
return get_hip_autotune_config()
|
||||
|
||||
|
||||
@triton.jit
|
||||
def clamp_int(value, min_val, max_val):
|
||||
ret = tl.where(value > max_val, max_val, value)
|
||||
ret = tl.where(ret < min_val, min_val, ret)
|
||||
return ret
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_loop(
|
||||
q, k, v, kv_mask, m, l, acc, sm_scale,
|
||||
MASK_KV: tl.constexpr,
|
||||
):
|
||||
scores = tl.dot(q, k.T) #[BLOCK_Q, BLOCK_KV]
|
||||
scores = scores * sm_scale
|
||||
if MASK_KV:
|
||||
scores = tl.where(kv_mask[None, :], scores, -float('inf'))
|
||||
|
||||
current_m = tl.max(scores, axis=1)
|
||||
new_m = tl.maximum(m, current_m)
|
||||
exp_scores = tl.math.exp2(scores - new_m[:, None])
|
||||
current_l = tl.sum(exp_scores, axis=1)
|
||||
|
||||
# Update L <- L * exp(M - M') + L1, M <- M'
|
||||
alpha = tl.math.exp2(m - new_m)
|
||||
l = l * alpha + current_l
|
||||
m = new_m
|
||||
|
||||
# Update O <- O * exp(M - M') + P @ V
|
||||
acc = (acc * alpha[:, None] + tl.dot(exp_scores.to(v.type.element_ty), v))
|
||||
|
||||
return m, l, acc
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=get_autotune_config(),
|
||||
key=['head_dim'],
|
||||
)
|
||||
@triton.jit
|
||||
def triton_sta_kernel(
|
||||
Q, K, V, output,
|
||||
batch_size: int, num_heads: int, seq_len: int, head_dim: int,
|
||||
img_seq_len: int,
|
||||
text_length: int,
|
||||
canvas_t: int, canvas_h: int, canvas_w: int,
|
||||
kernel_t: int, kernel_h: int, kernel_w: int,
|
||||
tile_t: int, tile_h: int, tile_w: int,
|
||||
scale: float,
|
||||
has_text: tl.constexpr,
|
||||
text_q: tl.constexpr,
|
||||
BLOCK_Q: tl.constexpr,
|
||||
BLOCK_KV: tl.constexpr,
|
||||
BLOCK_DIM: tl.constexpr,
|
||||
):
|
||||
total_tile_size = tile_t * tile_h * tile_w
|
||||
q_block_per_tile = (total_tile_size + BLOCK_Q - 1) // BLOCK_Q
|
||||
|
||||
batch_idx = tl.program_id(0)
|
||||
head_idx = tl.program_id(1)
|
||||
if text_q:
|
||||
q_block_idx = tl.program_id(2)
|
||||
else:
|
||||
q_tile_flat = tl.program_id(2) // q_block_per_tile
|
||||
q_block_idx = tl.program_id(2) % q_block_per_tile
|
||||
|
||||
m = tl.full((BLOCK_Q,), -float('inf'), dtype=tl.float32)
|
||||
l = tl.zeros((BLOCK_Q,), dtype=tl.float32)
|
||||
acc = tl.zeros((BLOCK_Q, BLOCK_DIM), dtype=tl.float32)
|
||||
|
||||
q_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
if text_q:
|
||||
q_base_idx = img_seq_len + q_block_idx * BLOCK_Q
|
||||
else:
|
||||
q_base_idx = q_tile_flat * total_tile_size + q_block_idx * BLOCK_Q
|
||||
|
||||
q_offset_in_tile = tl.arange(0, BLOCK_Q)
|
||||
q_idx = q_base_idx + q_offset_in_tile
|
||||
q_mask = (q_block_idx * BLOCK_Q + tl.arange(0, BLOCK_Q)) < total_tile_size
|
||||
|
||||
q = tl.load(
|
||||
Q + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=q_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_Q, BLOCK_DIM]
|
||||
|
||||
# Scale sm_scale by log_2(e) and use 2^x instead of exp
|
||||
sm_scale = scale * 1.4426950408889634
|
||||
|
||||
num_tiles_t = canvas_t // tile_t
|
||||
num_tiles_h = canvas_h // tile_h
|
||||
num_tiles_w = canvas_w // tile_w
|
||||
tiles_per_hw = num_tiles_h * num_tiles_w
|
||||
|
||||
if text_q:
|
||||
kv_tile_start_t = 0
|
||||
kv_tile_end_t = num_tiles_t
|
||||
|
||||
kv_tile_start_h = 0
|
||||
kv_tile_end_h = num_tiles_h
|
||||
|
||||
kv_tile_start_w = 0
|
||||
kv_tile_end_w = num_tiles_w
|
||||
|
||||
else:
|
||||
q_tile_t = q_tile_flat // tiles_per_hw
|
||||
remaining = q_tile_flat % tiles_per_hw
|
||||
q_tile_h = remaining // num_tiles_w
|
||||
q_tile_w = remaining % num_tiles_w
|
||||
|
||||
kernel_center_t = clamp_int(q_tile_t, kernel_t // 2, (num_tiles_t - 1) - kernel_t // 2)
|
||||
kernel_center_h = clamp_int(q_tile_h, kernel_h // 2, (num_tiles_h - 1) - kernel_h // 2)
|
||||
kernel_center_w = clamp_int(q_tile_w, kernel_w // 2, (num_tiles_w - 1) - kernel_w // 2)
|
||||
|
||||
kv_tile_start_t = kernel_center_t - kernel_t // 2
|
||||
kv_tile_end_t = kernel_center_t + kernel_t // 2 + 1
|
||||
kv_tile_end_t = tl.where(kv_tile_end_t > num_tiles_t, num_tiles_t, kv_tile_end_t)
|
||||
|
||||
kv_tile_start_h = kernel_center_h - kernel_h // 2
|
||||
kv_tile_end_h = kernel_center_h + kernel_h // 2 + 1
|
||||
kv_tile_end_h = tl.where(kv_tile_end_h > num_tiles_h, num_tiles_h, kv_tile_end_h)
|
||||
|
||||
kv_tile_start_w = kernel_center_w - kernel_w // 2
|
||||
kv_tile_end_w = kernel_center_w + kernel_w // 2 + 1
|
||||
kv_tile_end_w = tl.where(kv_tile_end_w > num_tiles_w, num_tiles_w, kv_tile_end_w)
|
||||
|
||||
# for kv_img
|
||||
for kv_tile_t in tl.range(kv_tile_start_t, kv_tile_end_t):
|
||||
for kv_tile_h in tl.range(kv_tile_start_h, kv_tile_end_h):
|
||||
for kv_tile_w in tl.range(kv_tile_start_w, kv_tile_end_w):
|
||||
kv_base_idx = (kv_tile_t * num_tiles_h * num_tiles_w + kv_tile_h * num_tiles_w + kv_tile_w) * total_tile_size
|
||||
|
||||
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
|
||||
kv_offset_in_block = tl.arange(0, BLOCK_KV)
|
||||
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
|
||||
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < total_tile_size
|
||||
|
||||
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
|
||||
k = tl.load(
|
||||
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
v = tl.load(
|
||||
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
|
||||
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, False)
|
||||
|
||||
|
||||
# for kv_text
|
||||
if has_text:
|
||||
kv_base_idx = img_seq_len
|
||||
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
|
||||
kv_offset_in_block = tl.arange(0, BLOCK_KV)
|
||||
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
|
||||
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < text_length
|
||||
|
||||
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
|
||||
k = tl.load(
|
||||
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
v = tl.load(
|
||||
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
|
||||
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, True)
|
||||
|
||||
|
||||
output_acc = acc / l[:, None]
|
||||
tl.store(
|
||||
output + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
output_acc,
|
||||
mask=q_mask[:, None]
|
||||
) # [BLOCK_Q, BLOCK_DIM]
|
||||
|
||||
|
||||
def sliding_tile_attention_triton(
|
||||
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
window_size, text_length: int,
|
||||
has_text=True, dit_seq_shape='30x48x80') -> torch.Tensor:
|
||||
seq_length = q.shape[2]
|
||||
if has_text:
|
||||
assert q.shape[2] >= 115200 and q.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
|
||||
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
|
||||
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo
|
||||
assert q.shape[2] == 82944
|
||||
elif dit_seq_shape == '18x48x80': # Wan
|
||||
assert q.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
assert q.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
if dit_seq_shape == '30x48x80': # Hunyuan
|
||||
canvas_t, canvas_h, canvas_w = 30, 48, 80
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
elif dit_seq_shape == '36x48x48': # Stepvideo
|
||||
canvas_t, canvas_h, canvas_w = 36, 48, 48
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
elif dit_seq_shape == '18x48x80': # Wan
|
||||
canvas_t, canvas_h, canvas_w = 18, 48, 80
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
num_tiles_t = canvas_t // tile_t
|
||||
num_tiles_h = canvas_h // tile_h
|
||||
num_tiles_w = canvas_w // tile_w
|
||||
num_tiles = num_tiles_t * num_tiles_h * num_tiles_w
|
||||
|
||||
total_tile_size = tile_t * tile_h * tile_w
|
||||
|
||||
# BLOCK_Q=128
|
||||
# BLOCK_KV=128
|
||||
BLOCK_DIM = head_dim
|
||||
|
||||
output = torch.empty_like(q)
|
||||
|
||||
# for q_img
|
||||
# kernel_size maybe different for different head
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (kernel_t, kernel_h, kernel_w) in enumerate(window_size):
|
||||
for batch in range(batch_size):
|
||||
q_head, k_head, v_head, o_head = (q[batch:batch + 1, head_index:head_index + 1],
|
||||
k[batch:batch + 1, head_index:head_index + 1],
|
||||
v[batch:batch + 1, head_index:head_index + 1],
|
||||
output[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
# triton_sta_kernel[(1, 1, num_tiles * triton.cdiv(total_tile_size, BLOCK_Q))](
|
||||
grid = lambda META: (1, 1, num_tiles * triton.cdiv(total_tile_size, META['BLOCK_Q']))
|
||||
triton_sta_kernel[grid](
|
||||
q_head, k_head, v_head, o_head,
|
||||
1, 1, seq_len, head_dim,
|
||||
img_seq_len,
|
||||
text_length,
|
||||
canvas_t, canvas_h, canvas_w,
|
||||
kernel_t, kernel_h, kernel_w,
|
||||
tile_t, tile_h, tile_w,
|
||||
scale=1.0 / (head_dim ** 0.5),
|
||||
has_text=has_text,
|
||||
text_q=False,
|
||||
# BLOCK_Q=BLOCK_Q,
|
||||
# BLOCK_KV=BLOCK_KV,
|
||||
BLOCK_DIM=BLOCK_DIM,
|
||||
)
|
||||
|
||||
# for q_text
|
||||
# kernel_t, kernel_h, kernel_w is not used, set to (3, 3, 3)
|
||||
if has_text:
|
||||
# triton_sta_kernel[(batch_size, num_heads, triton.cdiv(total_tile_size, BLOCK_Q))](
|
||||
grid = lambda META: (batch_size, num_heads, triton.cdiv(total_tile_size, META['BLOCK_Q']))
|
||||
triton_sta_kernel[grid](
|
||||
q, k, v, output,
|
||||
batch_size, num_heads, seq_len, head_dim,
|
||||
img_seq_len,
|
||||
text_length,
|
||||
canvas_t, canvas_h, canvas_w,
|
||||
3, 3, 3,
|
||||
#kernel_t, kernel_h, kernel_w,
|
||||
tile_t, tile_h, tile_w,
|
||||
scale=1.0 / (head_dim ** 0.5),
|
||||
has_text=has_text,
|
||||
text_q=True,
|
||||
# BLOCK_Q=BLOCK_Q,
|
||||
# BLOCK_KV=BLOCK_KV,
|
||||
BLOCK_DIM=BLOCK_DIM,
|
||||
)
|
||||
|
||||
if has_text:
|
||||
if pad_size > 0:
|
||||
output = output[:, :, :seq_length]
|
||||
|
||||
return output
|
||||
@@ -1,7 +0,0 @@
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
__pycache__/
|
||||
*.so
|
||||
*.pyc
|
||||
.ipynb_checkpoints/
|
||||
@@ -1,187 +0,0 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
@@ -1,6 +0,0 @@
|
||||
include LICENSE
|
||||
include README.md
|
||||
include pyproject.toml
|
||||
recursive-include src/fastvideo_kernel *.cu *.cuh *.cpp *.h
|
||||
recursive-include csrc *.cu *.cuh *.cpp *.h
|
||||
recursive-include tk *.cu *.cuh *.cpp *.h
|
||||
@@ -1,31 +0,0 @@
|
||||
# FastVideo Kernel
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
cd csrc/fastvideo_kernel
|
||||
pip install .
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_attn_varlen
|
||||
|
||||
# Example: Sliding Tile Attention
|
||||
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
|
||||
|
||||
# Example: Video Sparse Attention (with Triton fallback)
|
||||
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
|
||||
|
||||
# Example: VMoBA
|
||||
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
|
||||
```
|
||||
|
||||
## Requirements
|
||||
|
||||
- H100 GPU (sm_90a) for CUDA kernels
|
||||
- Triton for non-H100 fallback
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,23 +0,0 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
@@ -1,573 +0,0 @@
|
||||
// # Define TORCH_COMPILE macro
|
||||
|
||||
#include "kittens.cuh"
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <stdio.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
|
||||
return (value < min) ? min : ((value > max) ? max : value);
|
||||
}
|
||||
// #define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
__device__ __forceinline__ int abs_int(int value) {
|
||||
return (value < 0) ? -value : value;
|
||||
}
|
||||
|
||||
|
||||
constexpr int CONSUMER_WARPGROUPS = (3);
|
||||
constexpr int PRODUCER_WARPGROUPS = (1);
|
||||
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
|
||||
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
|
||||
|
||||
using namespace kittens;
|
||||
namespace cg = cooperative_groups;
|
||||
|
||||
template<int D> struct fwd_attend_ker_tile_dims {};
|
||||
template<> struct fwd_attend_ker_tile_dims<64> {
|
||||
constexpr static int tile_width = (64);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (4);
|
||||
};
|
||||
template<> struct fwd_attend_ker_tile_dims<128> {
|
||||
constexpr static int tile_width = (128);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (2);
|
||||
};
|
||||
|
||||
template<int D> struct fwd_globals {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
|
||||
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
q_gl q;
|
||||
k_gl k;
|
||||
v_gl v;
|
||||
l_gl l;
|
||||
o_gl o;
|
||||
|
||||
const int N;
|
||||
const int text_L;
|
||||
const int hr;
|
||||
};
|
||||
|
||||
|
||||
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
|
||||
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
|
||||
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
extern __shared__ int __shm[];
|
||||
tma_swizzle_allocator al((int*)&__shm[0]);
|
||||
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
|
||||
|
||||
using K = fwd_attend_ker_tile_dims<D>;
|
||||
|
||||
using q_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
using k_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using v_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
|
||||
using o_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
|
||||
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
|
||||
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
|
||||
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
|
||||
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
|
||||
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
|
||||
int img_kv_blocks;
|
||||
int kv_blocks = g.N / (K::kv_height);
|
||||
if constexpr (text_kv) {
|
||||
img_kv_blocks = kv_blocks - 3;
|
||||
} else {
|
||||
img_kv_blocks = kv_blocks;
|
||||
}
|
||||
int kv_head_idx = blockIdx.y / g.hr;
|
||||
int seq_idx;
|
||||
if constexpr (text_q) {
|
||||
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
} else {
|
||||
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
}
|
||||
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
|
||||
if (threadIdx.x == 0) {
|
||||
init_semaphore(qsmem_semaphore, 0, 1);
|
||||
for(int j = 0; j < K::stages; j++) {
|
||||
init_semaphore(k_smem_arrived[j], 0, 1);
|
||||
init_semaphore(v_smem_arrived[j], 0, 1);
|
||||
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
|
||||
}
|
||||
|
||||
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
|
||||
|
||||
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
|
||||
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
|
||||
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
|
||||
}
|
||||
|
||||
if constexpr (text_q){
|
||||
for (int j = 0; j < K::stages - 1; j++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
|
||||
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
|
||||
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
|
||||
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int count = 0;
|
||||
int j = 0;
|
||||
while (count < K::stages - 1) {
|
||||
int kt = j / 3 / (CH * CW);
|
||||
int kh = (j / 3) % (CH * CW) / CW;
|
||||
int kw = (j / 3) % CW;
|
||||
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
|
||||
if (mask){
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
|
||||
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
|
||||
count += 1;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int pipe_idx = K::stages - 1;
|
||||
|
||||
if(warpgroupid == NUM_WARPGROUPS-1) {
|
||||
warpgroup::decrease_registers<32>();
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
|
||||
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
|
||||
}
|
||||
else { kv_iters = kv_blocks-2;}
|
||||
|
||||
if(warpid == NUM_WORKERS-4) {
|
||||
if constexpr (text_q){
|
||||
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
|
||||
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int k_t_min = clamp_int(qt-DT, 0, CT-1);
|
||||
int k_t_max = clamp_int(qt+DT, 0, CT-1);
|
||||
int k_h_min = clamp_int(qh-DH, 0, CH-1);
|
||||
int k_h_max = clamp_int(qh+DH, 0, CH-1);
|
||||
int k_w_min = clamp_int(qw-DW, 0, CW-1);
|
||||
int k_w_max = clamp_int(qw+DW, 0, CW-1);
|
||||
int count = 0;
|
||||
for (int kt = k_t_min; kt <= k_t_max; kt++) {
|
||||
for (int kh = k_h_min; kh <= k_h_max; kh++) {
|
||||
for (int kw = k_w_min; kw <= k_w_max; kw++) {
|
||||
for (int j = 0; j <= 2; j++){
|
||||
if (count >= K::stages - 1) {
|
||||
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
} else {
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// for text
|
||||
for (int index = img_kv_blocks; index < kv_blocks; index++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
warpgroup::increase_registers<160>();
|
||||
|
||||
rt_fl<16, K::kv_height> att_block;
|
||||
rt_bf<16, K::kv_height> att_block_mma;
|
||||
rt_fl<16, K::tile_width> o_reg;
|
||||
|
||||
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
|
||||
|
||||
neg_infty(max_vec);
|
||||
zero(norm_vec);
|
||||
zero(o_reg);
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
|
||||
kv_iters = (kv_iters/8);
|
||||
}
|
||||
else if constexpr (text_q){
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
kv_iters = img_kv_blocks - 1;
|
||||
} else {
|
||||
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
}
|
||||
|
||||
kittens::wait(qsmem_semaphore, 0);
|
||||
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
if constexpr(text_kv) {
|
||||
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
// apply non-pad mask
|
||||
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
|
||||
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
|
||||
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
|
||||
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
}
|
||||
|
||||
div_row(o_reg, o_reg, norm_vec);
|
||||
warpgroup::store(o_smem[warpgroupid], o_reg);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
|
||||
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
|
||||
}
|
||||
|
||||
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
|
||||
log(norm_vec, norm_vec);
|
||||
add(norm_vec, norm_vec, max_vec_scaled);
|
||||
|
||||
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
|
||||
else { mul(norm_vec, norm_vec, -11.313708499f); }
|
||||
|
||||
warpgroup::store(l_smem[warpgroupid], norm_vec);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
|
||||
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
|
||||
}
|
||||
tma::store_async_wait();
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
#include "pyutils/torch_helpers.cuh"
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <iostream>
|
||||
|
||||
torch::Tensor
|
||||
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
|
||||
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
|
||||
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
|
||||
auto hr = qo_heads / kv_heads;
|
||||
|
||||
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
|
||||
|
||||
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
|
||||
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
|
||||
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
|
||||
|
||||
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
|
||||
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
|
||||
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
|
||||
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
|
||||
if (head_dim == 128) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<128>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
// Shared memory size for the kernel.
|
||||
// We use the maximum available shared memory (kittens::MAX_SHARED_MEMORY)
|
||||
// which is approximately 227KB on H100, necessary for the high-performance
|
||||
// TMA-based attention tiles with multiple stages.
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
dim3 grid_text(2, qo_heads, batch);
|
||||
if (!process_text) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(2, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 3, 0); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(2, 2, 3); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 0, 5); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
} else {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
}
|
||||
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
if (kernel_aspect_ratio_flag == 2){
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(3, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 3, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 3, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 0, 0); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 0, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 3, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 0, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(0, 3, 0); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
else if (kernel_aspect_ratio_flag == 3) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 3, 0); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(1, 2, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(1, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 2, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 2, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(0, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 0, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(0, 3, 2); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
|
||||
else {
|
||||
TORCH_CHECK(false, "Unsupported kernel_aspect_ratio_flag: ", kernel_aspect_ratio_flag);
|
||||
}
|
||||
|
||||
}
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return o;
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
|
||||
#endif
|
||||
}
|
||||
@@ -1,27 +0,0 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=61.0", "torch>=2.5.0", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.1.0"
|
||||
description = "CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
license = {text = "Apache-2.0"}
|
||||
authors = [{name = "Hao AI Lab"}]
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
]
|
||||
dependencies = [
|
||||
"torch>=2.5.0",
|
||||
"triton>=2.0.0"
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -1,132 +0,0 @@
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
ROOT = Path(__file__).parent.absolute()
|
||||
CSRC_DIR = ROOT / "csrc"
|
||||
|
||||
# Path to ThunderKittens (TK)
|
||||
def get_tk_dir():
|
||||
tk_env = os.getenv("THUNDERKITTENS_ROOT")
|
||||
if tk_env:
|
||||
return tk_env
|
||||
|
||||
# Check common locations
|
||||
possible_paths = [
|
||||
ROOT / "tk",
|
||||
ROOT / "csrc" / "tk",
|
||||
ROOT.parent / "attn" / "sliding_tile_attn" / "tk",
|
||||
ROOT.parent / "attn" / "video_sparse_attn" / "tk",
|
||||
]
|
||||
for p in possible_paths:
|
||||
if (p / "include" / "kittens.cuh").exists():
|
||||
return str(p)
|
||||
|
||||
# Default fallback
|
||||
return str(ROOT.parent / "attn" / "sliding_tile_attn" / "tk")
|
||||
|
||||
TK_DIR = get_tk_dir()
|
||||
|
||||
def get_cuda_flags(tk_root: str) -> list:
|
||||
python_include = subprocess.check_output(
|
||||
["python", "-c", "import sysconfig; print(sysconfig.get_path('include'))"]
|
||||
).decode().strip()
|
||||
|
||||
torch_includes = subprocess.check_output([
|
||||
"python", "-c",
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; "
|
||||
"print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip().split()
|
||||
|
||||
return [
|
||||
"-DNDEBUG",
|
||||
"-Xcompiler=-Wno-psabi",
|
||||
"-Xcompiler=-fno-strict-aliasing",
|
||||
"--expt-extended-lambda",
|
||||
"--expt-relaxed-constexpr",
|
||||
"-forward-unknown-to-host-compiler",
|
||||
"--use_fast_math",
|
||||
"-std=c++20",
|
||||
"-O3",
|
||||
"-Xnvlink=--verbose",
|
||||
"-Xptxas=--verbose",
|
||||
"-Xptxas=--warn-on-spills",
|
||||
f"-I{tk_root}/include",
|
||||
f"-I{tk_root}/prototype",
|
||||
f"-I{python_include}",
|
||||
"-DTORCH_COMPILE",
|
||||
"-DKITTENS_HOPPER",
|
||||
"-arch=sm_90a",
|
||||
] + torch_includes
|
||||
|
||||
def get_extensions():
|
||||
if not torch.cuda.is_available():
|
||||
return []
|
||||
|
||||
extensions = []
|
||||
cpp_flags = ["-std=c++20", "-O3"]
|
||||
|
||||
# Check if TK is available
|
||||
if not os.path.exists(os.path.join(TK_DIR, "include", "kittens.cuh")):
|
||||
print(f"Warning: ThunderKittens not found at {TK_DIR}. CUDA kernels will not be built.")
|
||||
return []
|
||||
|
||||
cuda_flags = get_cuda_flags(TK_DIR)
|
||||
|
||||
# STA Extension
|
||||
extensions.append(CUDAExtension(
|
||||
"fastvideo_kernel._C.st_attn",
|
||||
sources=[
|
||||
"csrc/st_attn.cpp",
|
||||
"csrc/st_attn_h100.cu",
|
||||
],
|
||||
extra_compile_args={
|
||||
"cxx": cpp_flags + ["-DTK_COMPILE_ST_ATTN"],
|
||||
"nvcc": cuda_flags + ["-DTK_COMPILE_ST_ATTN"]
|
||||
},
|
||||
libraries=["cuda"],
|
||||
))
|
||||
|
||||
# VSA Extension
|
||||
extensions.append(CUDAExtension(
|
||||
"fastvideo_kernel._C.vsa",
|
||||
sources=[
|
||||
"csrc/vsa.cpp",
|
||||
"csrc/block_sparse_h100.cu",
|
||||
],
|
||||
extra_compile_args={
|
||||
"cxx": cpp_flags + ["-DTK_COMPILE_BLOCK_SPARSE"],
|
||||
"nvcc": cuda_flags + ["-DTK_COMPILE_BLOCK_SPARSE"]
|
||||
},
|
||||
libraries=["cuda"],
|
||||
))
|
||||
|
||||
return extensions
|
||||
|
||||
ext_modules = []
|
||||
if not any(arg in sys.argv for arg in ["clean", "egg_info", "--version"]):
|
||||
try:
|
||||
import torch
|
||||
ext_modules = get_extensions()
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to configure CUDA extensions: {e}")
|
||||
|
||||
setup(
|
||||
name="fastvideo-kernel",
|
||||
version="0.1.0",
|
||||
description="Unified CUDA kernels for FastVideo",
|
||||
long_description=open("README.md").read(),
|
||||
long_description_content_type="text/markdown",
|
||||
license="Apache-2.0",
|
||||
author="Hao AI Lab",
|
||||
url="https://github.com/hao-ai-lab/FastVideo",
|
||||
package_dir={"": "src"},
|
||||
packages=find_packages(where="src"),
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={"build_ext": BuildExtension} if ext_modules else {},
|
||||
python_requires=">=3.10",
|
||||
install_requires=["torch>=2.5.0", "triton>=2.0.0"],
|
||||
)
|
||||
@@ -1,21 +0,0 @@
|
||||
__version__ = "0.1.0"
|
||||
|
||||
from fastvideo_kernel.ops import (
|
||||
sliding_tile_attention,
|
||||
video_sparse_attn,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.vmoba import (
|
||||
moba_attn_varlen,
|
||||
process_moba_input,
|
||||
process_moba_output,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
"__version__",
|
||||
]
|
||||
@@ -1,103 +0,0 @@
|
||||
import math
|
||||
import torch
|
||||
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
|
||||
from .triton_kernels.index import map_to_index
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C.st_attn import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C.vsa import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
|
||||
def sliding_tile_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
window_size: list,
|
||||
text_length: int,
|
||||
has_text: bool = True,
|
||||
seq_shape: str = "30x48x80",
|
||||
) -> torch.Tensor:
|
||||
if sta_fwd is None:
|
||||
raise RuntimeError("STA kernel not compiled. Requires H100 and ThunderKittens at build time.")
|
||||
|
||||
seq_length = q.shape[2]
|
||||
shape_map = {"30x48x80": 1, "36x48x48": 2, "18x48x80": 3}
|
||||
|
||||
if has_text:
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
|
||||
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
|
||||
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
|
||||
|
||||
output = torch.empty_like(q)
|
||||
flag = shape_map[seq_shape]
|
||||
|
||||
for head_idx, (t, h, w) in enumerate(window_size):
|
||||
sta_fwd(
|
||||
q[:, head_idx:head_idx+1],
|
||||
k[:, head_idx:head_idx+1],
|
||||
v[:, head_idx:head_idx+1],
|
||||
output[:, head_idx:head_idx+1],
|
||||
t, h, w, text_length, False, has_text, flag
|
||||
)
|
||||
|
||||
if has_text:
|
||||
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
|
||||
|
||||
return output[:, :, :seq_length]
|
||||
|
||||
|
||||
def video_sparse_attn(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
topk: int,
|
||||
block_size: int | tuple = 64,
|
||||
compress_attn_weight: torch.Tensor = None,
|
||||
) -> torch.Tensor:
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
batch, heads, seq_len, dim = q.shape
|
||||
|
||||
# Compression branch
|
||||
q_c = q.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
k_c = k.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
v_c = v.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
|
||||
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
|
||||
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
|
||||
v_c = (v_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
|
||||
|
||||
scores = torch.matmul(q_c, k_c.transpose(-2, -1)) / (dim ** 0.5)
|
||||
attn = torch.softmax(scores, dim=-1)
|
||||
out_c = torch.matmul(attn, v_c)
|
||||
|
||||
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
|
||||
out_c = out_c.repeat(1, 1, 1, block_elements, 1).view(batch, heads, seq_len, dim)
|
||||
|
||||
# Sparse branch
|
||||
topk_idx = torch.topk(scores, topk, dim=-1).indices
|
||||
mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
|
||||
|
||||
if block_sparse_fwd is not None:
|
||||
idx, num = map_to_index(mask)
|
||||
out_s, _ = block_sparse_fwd(q, k, v, idx, num, variable_block_sizes.int())
|
||||
else:
|
||||
idx, num = map_to_index(mask)
|
||||
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
return out_c + out_s
|
||||
-449
@@ -1,449 +0,0 @@
|
||||
"""
|
||||
Fused Attention
|
||||
===============
|
||||
|
||||
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
(https://tridao.me/publications/flash2/flash2.pdf)
|
||||
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
import math # small utility needed by the sparse wrapper
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
|
||||
# the code below and commenting out the equivalent parameters is convenient for
|
||||
# re-tuning.
|
||||
configs = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64]\
|
||||
for BN in [64]\
|
||||
for s in [3, 4, 7]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(Q, K, V, sm_scale, #
|
||||
q2k_index, q2k_num, max_kv_blks, #
|
||||
variable_block_sizes,
|
||||
M, Out, #
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vk, stride_vn,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
Z, H, N_CTX, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr):
|
||||
"""
|
||||
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
|
||||
(32×64 and 64×32) – memory footprint unchanged.
|
||||
"""
|
||||
|
||||
# ----- program-id mapping -----
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(1) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
# ----- base pointers -----
|
||||
qvk_off = (b.to(tl.int64) * stride_qz +
|
||||
h.to(tl.int64) * stride_qh)
|
||||
|
||||
Q_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
K_base = tl.make_block_ptr(
|
||||
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
|
||||
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_base = tl.make_block_ptr(
|
||||
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
|
||||
|
||||
O_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
# ----- accumulators -----
|
||||
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
q = tl.load(Q_ptr)
|
||||
|
||||
# ----- sparse loop over valid K/V tiles -----
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
|
||||
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
|
||||
|
||||
k = tl.load(K_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
# mask out invalid columns
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptr)
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
# ----- epilogue -----
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(M + off_hz * N_CTX + offs_m, m_i)
|
||||
tl.store(O_ptr, acc.to(Out.type.element_ty))
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO, #
|
||||
Delta, #
|
||||
Z, H, N_CTX, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
|
||||
):
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
# Filled in by the wrapper.
|
||||
start_n, start_m, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
|
||||
step_m = BLOCK_M1
|
||||
kv_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_N1
|
||||
meta_base = ((b * H + h) * q_tiles + kv_blk)
|
||||
|
||||
q_blocks = tl.load(k2q_num + meta_base) # int32
|
||||
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + kv_blk)
|
||||
|
||||
|
||||
|
||||
for blk_idx in range(q_blocks*2):
|
||||
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
|
||||
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
|
||||
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(tl.bfloat16)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(tl.bfloat16)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
# Increment pointers.
|
||||
return dk, dv
|
||||
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr,
|
||||
# Filled in by the wrapper.
|
||||
start_m, start_n, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N2)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
|
||||
step_n = BLOCK_N2
|
||||
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M2
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + q_blk)
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
|
||||
p = tl.where(mask[None, :], p , 0.0)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd(Q, K, V, sm_scale, #
|
||||
DO, #
|
||||
DQ, DK, DV, #
|
||||
M, D,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_z, stride_h, stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr):
|
||||
LN2 = 0.6931471824645996 # = ln(2)
|
||||
|
||||
bhid = tl.program_id(2)
|
||||
off_chz = (bhid * N_CTX).to(tl.int64)
|
||||
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
|
||||
pid = tl.program_id(0)
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += adj
|
||||
K += adj
|
||||
V += adj
|
||||
DO += adj
|
||||
DQ += adj
|
||||
DK += adj
|
||||
DV += adj
|
||||
M += off_chz
|
||||
D += off_chz
|
||||
|
||||
# load scales
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
start_n = pid * BLOCK_N1
|
||||
start_m = 0
|
||||
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
|
||||
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
|
||||
num_steps = N_CTX // BLOCK_M1
|
||||
|
||||
dk, dv = _attn_bwd_dkdv( #
|
||||
dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
|
||||
start_n, start_m, num_steps #
|
||||
)
|
||||
|
||||
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dv_ptrs, dv)
|
||||
|
||||
# Write back dK.
|
||||
dk *= sm_scale
|
||||
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dk_ptrs, dk)
|
||||
|
||||
# THIS BLOCK DOES DQ:
|
||||
start_m = pid * BLOCK_M2
|
||||
end_n = 0
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
|
||||
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
num_steps = N_CTX // BLOCK_N2
|
||||
dq = _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
|
||||
start_m, end_n, num_steps #
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
|
||||
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
|
||||
|
||||
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
q, k, v, sm_scale,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
M, o,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
B, H, T,
|
||||
HEAD_DIM=D, STAGE=3
|
||||
)
|
||||
|
||||
return o, M
|
||||
|
||||
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do, #
|
||||
delta, #
|
||||
BATCH, N_HEAD, N_CTX, #
|
||||
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
|
||||
)
|
||||
|
||||
|
||||
max_q_blks = k2q_index.shape[-1]
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
|
||||
N_HEAD, N_CTX, #
|
||||
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
|
||||
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
|
||||
HEAD_DIM=D #
|
||||
)
|
||||
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@@ -1,152 +0,0 @@
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
@@ -1,868 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
|
||||
|
||||
import random
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
except ImportError:
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
_flash_attn_varlen_backward = _unsupported
|
||||
flash_attn_varlen_func = _unsupported
|
||||
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def calc_chunks(cu_seqlen, moba_chunk_size):
|
||||
"""
|
||||
Calculate chunk boundaries.
|
||||
|
||||
For vision tasks we include all chunks (even the last one which might be shorter)
|
||||
so that every chunk can be selected.
|
||||
"""
|
||||
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
|
||||
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
|
||||
cu_num_chunk = torch.ones(
|
||||
batch_num_chunk.numel() + 1,
|
||||
device=cu_seqlen.device,
|
||||
dtype=batch_num_chunk.dtype,
|
||||
)
|
||||
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
|
||||
num_chunk = cu_num_chunk[-1]
|
||||
chunk_sizes = torch.full(
|
||||
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_sizes[0] = 0
|
||||
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
|
||||
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
|
||||
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
|
||||
chunk_to_batch = torch.zeros(
|
||||
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_to_batch[cu_num_chunk[1:-1]] = 1
|
||||
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
|
||||
|
||||
# Do not filter out any chunk
|
||||
filtered_chunk_indices = torch.arange(
|
||||
num_chunk, device=cu_seqlen.device, dtype=torch.int32
|
||||
)
|
||||
num_filtered_chunk = num_chunk
|
||||
|
||||
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
|
||||
|
||||
|
||||
# --- Threshold Selection Helper Functions ---
|
||||
|
||||
def _select_threshold_query_head(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects chunks for each <query, head> pair based on threshold.
|
||||
Normalization and sorting happen along the chunk dimension (dim=0).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per <head, query> (across chunks)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
row_min = gate_min_val.amin(dim=0) # (H, S)
|
||||
row_max = gate_masked.amax(dim=0) # (H, S)
|
||||
denom = row_max - row_min
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
|
||||
|
||||
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
|
||||
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
|
||||
|
||||
# 2) compute how much more normalized weight we need beyond self
|
||||
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
|
||||
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
|
||||
|
||||
# 3) zero out the self‐chunk in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0
|
||||
|
||||
# 4) sort the other chunks by descending norm, per <head,seq>
|
||||
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per <head,seq>
|
||||
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
|
||||
|
||||
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
|
||||
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
|
||||
any_cond = cond.any(dim=0) # (H, S)
|
||||
# Find the index of the first True value along dim 0. If none, use C-1.
|
||||
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff
|
||||
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
|
||||
|
||||
# 8) scatter it back to original chunk order
|
||||
others_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
others_mask.scatter_(0, sorted_idx, sorted_mask)
|
||||
|
||||
# 9) finally, include every self‐chunk plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_block(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <query, head> pairs for each block based on threshold.
|
||||
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
HS = H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per block (across heads and queries)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_denom = block_max - block_min
|
||||
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
|
||||
|
||||
gate_norm = (gate - block_min) / block_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights *per block*
|
||||
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
|
||||
|
||||
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
|
||||
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
|
||||
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort the other <head, seq> pairs by descending norm, per block
|
||||
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
|
||||
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per block
|
||||
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
|
||||
|
||||
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
|
||||
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
|
||||
any_cond = cond_flat.any(dim=1) # (C,)
|
||||
# Find the index of the first True value along dim 1. If none, use HS-1.
|
||||
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff per block
|
||||
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
|
||||
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
|
||||
|
||||
# 8) scatter it back to original <head, seq> order per block
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
|
||||
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_overall(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query, head> triplets globally based on threshold.
|
||||
Normalization and sorting happen across all valid entries.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
CHS = C * H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization globally across all valid entries
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
overall_max = gate_masked.max() # scalar
|
||||
overall_min = gate_min_val.min() # scalar
|
||||
overall_denom = overall_max - overall_min
|
||||
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
|
||||
|
||||
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights globally
|
||||
self_norm_sum_overall = self_norm_entries.sum() # scalar
|
||||
|
||||
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
|
||||
total_norm_sum_overall = gate_norm.sum() # scalar
|
||||
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort all other entries by descending norm, globally
|
||||
others_flat = others_norm.flatten() # (C*H*S,)
|
||||
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
|
||||
|
||||
# Only sort the valid 'other' entries
|
||||
valid_others_indices = torch.where(valid_others_mask_flat)[0]
|
||||
valid_others_values = others_flat[valid_others_indices]
|
||||
|
||||
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
|
||||
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
|
||||
|
||||
# 5) cumulative‑sum the sorted valid 'other' norms globally
|
||||
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
|
||||
|
||||
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
|
||||
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
|
||||
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
|
||||
any_cond = cond_values.any() # scalar
|
||||
|
||||
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
|
||||
cutoff_idx_in_sorted = torch.where(
|
||||
any_cond,
|
||||
cond_values.float().argmax(dim=0),
|
||||
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
|
||||
)
|
||||
|
||||
# 7) build a mask selecting the top-k others based on the cutoff
|
||||
# Select the original indices corresponding to the top entries in the sorted list
|
||||
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
|
||||
|
||||
# 8) create the mask in the original flat shape
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
|
||||
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
|
||||
others_mask_flat[selected_other_indices] = True
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_head_global(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query> globally for each head based on threshold.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# 1) LSE‐style normalization per head (across chunks and sequence dims)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
|
||||
|
||||
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
denom = max_per_head - min_per_head
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
|
||||
|
||||
gate_norm = (gate - min_per_head) / denom
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 2) sum normalized self‐chunk contributions per head
|
||||
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 3) total normalized sum per head
|
||||
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 4) how much more normalized weight needed per head
|
||||
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0)
|
||||
|
||||
# 5) zero out self‐chunk entries to focus on "others"
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
|
||||
|
||||
# 6) flatten chunk and sequence dims, per head
|
||||
CS = C * S
|
||||
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
|
||||
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
|
||||
# 7) vectorized selection of “others” per head
|
||||
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
|
||||
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
|
||||
|
||||
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
|
||||
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
|
||||
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
|
||||
|
||||
has_cutoff = cond.any(dim=1) # (H,)
|
||||
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
|
||||
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
|
||||
|
||||
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
|
||||
|
||||
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
|
||||
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
|
||||
|
||||
# 8) reshape selection mask back to (C, H, S)
|
||||
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
|
||||
|
||||
# 9) include self‐chunks plus selected others, and obey valid mask
|
||||
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
class MixedAttention(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
):
|
||||
ctx.max_seqlen = max_seqlen
|
||||
ctx.moba_chunk_size = moba_chunk_size
|
||||
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
|
||||
|
||||
# Non-causal self-attention branch
|
||||
# return out, softmax_lse, S_dmask, rng_state
|
||||
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
# MOBA attention branch (non-causal)
|
||||
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
|
||||
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
|
||||
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
|
||||
|
||||
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
output_2d = output.view(-1, q.shape[2])
|
||||
|
||||
max_lse_1d = self_attn_lse_sh.view(-1)
|
||||
max_lse_1d = max_lse_1d.index_reduce(
|
||||
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
|
||||
)
|
||||
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
|
||||
moba_attn_lse = (
|
||||
moba_attn_lse.view(-1)
|
||||
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
|
||||
.reshape_as(moba_attn_lse)
|
||||
)
|
||||
|
||||
mixed_attn_se_sh = self_attn_lse_sh.exp()
|
||||
moba_attn_se = moba_attn_lse.exp()
|
||||
|
||||
mixed_attn_se_sh.view(-1).index_add_(
|
||||
0, moba_q_sh_indices, moba_attn_se.view(-1)
|
||||
)
|
||||
mixed_attn_lse_sh = mixed_attn_se_sh.log()
|
||||
|
||||
# Combine self-attention output
|
||||
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
|
||||
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
|
||||
output_2d += self_attn_out_sh.reshape_as(output_2d)
|
||||
|
||||
# Combine MOBA attention output
|
||||
mixed_attn_lse = (
|
||||
mixed_attn_lse_sh.view(-1)
|
||||
.index_select(0, moba_q_sh_indices)
|
||||
.view_as(moba_attn_lse)
|
||||
)
|
||||
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
|
||||
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
|
||||
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
|
||||
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
|
||||
output = output.to(q.dtype)
|
||||
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
|
||||
ctx.save_for_backward(
|
||||
output,
|
||||
mixed_attn_lse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, d_output):
|
||||
|
||||
max_seqlen = ctx.max_seqlen
|
||||
moba_chunk_size = ctx.moba_chunk_size
|
||||
softmax_scale = ctx.softmax_scale
|
||||
|
||||
(
|
||||
output,
|
||||
mixed_attn_vlse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
) = ctx.saved_tensors
|
||||
|
||||
d_output = d_output.contiguous()
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_output,
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
out=output,
|
||||
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
|
||||
dq=dq,
|
||||
dk=dk,
|
||||
dv=dv,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
headdim = q.shape[-1]
|
||||
d_moba_output = (
|
||||
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
moba_output = (
|
||||
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
|
||||
mixed_attn_vlse = (
|
||||
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
|
||||
)
|
||||
|
||||
dmq = torch.empty_like(moba_q)
|
||||
dmkv = torch.empty_like(moba_kv)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_moba_output,
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
out=moba_output,
|
||||
softmax_lse=mixed_attn_vlse,
|
||||
dq=dmq,
|
||||
dk=dmkv[:,0],
|
||||
dv=dmkv[:,1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
|
||||
|
||||
|
||||
def moba_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int,
|
||||
moba_chunk_size: int,
|
||||
moba_topk: int,
|
||||
select_mode: str = 'threshold', # "topk" or "threshold"
|
||||
simsum_threshold: float = 0.25,
|
||||
threshold_type: str = 'query_head',
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Accelerated MOBA attention for vision tasks with proper LSE normalization.
|
||||
|
||||
This version:
|
||||
- Splits KV into chunks.
|
||||
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
|
||||
by amplifying the diagonal (self-chunk) logits.
|
||||
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
|
||||
reduction so that attending to each query over the selected chunks is equivalent
|
||||
to the original algorithm.
|
||||
"""
|
||||
# Stack keys and values.
|
||||
kv = torch.stack((k, v), dim=1)
|
||||
seqlen, num_head, head_dim = q.shape
|
||||
|
||||
# Compute chunk boundaries.
|
||||
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
|
||||
cu_seqlens, moba_chunk_size
|
||||
)
|
||||
|
||||
self_attn_cu_seqlen = cu_chunk
|
||||
|
||||
# Update top-k selection to include the self chunk.
|
||||
moba_topk = min(moba_topk, num_filtered_chunk)
|
||||
|
||||
# --- Build filtered KV from chunks ---
|
||||
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
|
||||
max_chunk_len = int(chunk_lengths.max().item())
|
||||
|
||||
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
|
||||
indices = chunk_starts.unsqueeze(1) + range_tensor
|
||||
indices = torch.clamp(indices, max=kv.shape[0] - 1)
|
||||
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
|
||||
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
|
||||
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
|
||||
|
||||
# Compute key_gate_weight over valid tokens.
|
||||
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
|
||||
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
|
||||
key_sum = (key_values * valid_mask_exp).sum(dim=1)
|
||||
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
|
||||
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
|
||||
|
||||
# Compute gate logits between key_gate_weight and queries.
|
||||
q_float = q.float()
|
||||
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
|
||||
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
|
||||
|
||||
# Amplify the diagonal (self chunk) contributions.
|
||||
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
|
||||
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
|
||||
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
|
||||
amplification_factor = 1e9 # Example factor; adjust as needed.
|
||||
origin_gate = gate.clone()
|
||||
gate = gate.clone()
|
||||
if select_mode == "topk":
|
||||
gate[gate_self_chunk_mask] += amplification_factor
|
||||
|
||||
# Exclude positions that are outside the valid batch boundaries.
|
||||
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
|
||||
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
|
||||
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
|
||||
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
|
||||
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
|
||||
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
|
||||
|
||||
if select_mode == 'topk':
|
||||
# We amplify self‐chunk in gate already, so self entries will rank highest.
|
||||
valid_gate_mask = gate != -float("inf")
|
||||
if threshold_type == 'query_head':
|
||||
# === per‐<head,seq> top-k across chunks (original behavior) ===
|
||||
# gate: (C, H, S)
|
||||
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
|
||||
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
gate_idx_mask.scatter_(0, gate_topk_idx, True)
|
||||
gate_mask = valid_gate_mask & gate_idx_mask
|
||||
elif threshold_type == 'overall':
|
||||
# === global top-k across all (chunk, head, seq) entries ===
|
||||
C, H, S = gate.shape
|
||||
flat_gate = gate.flatten()
|
||||
flat_mask = valid_gate_mask.flatten()
|
||||
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
|
||||
# pick topk global entries
|
||||
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
|
||||
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
|
||||
others_mask_flat[idx] = True
|
||||
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
|
||||
elif threshold_type == 'head_global':
|
||||
# per-head top-k across all chunks and sequence positions
|
||||
C, H, S = gate.shape
|
||||
CS = C * S
|
||||
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
|
||||
# pick top-k indices per head
|
||||
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
|
||||
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
|
||||
gate_idx_flat.scatter_(1, topk_idx, True)
|
||||
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid threshold_type for topk: {threshold_type}. "
|
||||
"Choose 'query_head', 'block', or 'overall'."
|
||||
)
|
||||
elif select_mode == 'threshold':
|
||||
# Delegate to the specific thresholding function
|
||||
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
|
||||
if threshold_type == 'query_head':
|
||||
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'block':
|
||||
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'overall':
|
||||
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'head_global':
|
||||
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
else:
|
||||
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
|
||||
else:
|
||||
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
|
||||
|
||||
# eliminate self_chunk in MoBA branch
|
||||
gate_mask = gate_mask & ~gate_self_chunk_mask
|
||||
# if gate_mask is all false, perform flash_attn instead
|
||||
if gate_mask.sum() == 0:
|
||||
return flash_attn_varlen_func(
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
|
||||
)
|
||||
|
||||
# Determine which query positions are selected.
|
||||
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
|
||||
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
|
||||
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
|
||||
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
|
||||
|
||||
# Build cumulative sequence lengths for the selected queries.
|
||||
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
|
||||
q_zero_mask = moba_seqlen_q == 0
|
||||
valid_expert_mask = ~q_zero_mask
|
||||
if q_zero_mask.sum() > 0:
|
||||
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
|
||||
moba_cu_seqlen_q = torch.cat(
|
||||
(
|
||||
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
|
||||
moba_seqlen_q.cumsum(dim=0),
|
||||
),
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
# Rearrange gathered KV for the MOBA branch.
|
||||
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
|
||||
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
|
||||
if q_zero_mask.sum() > 0:
|
||||
experts_tensor = experts_tensor[valid_expert_mask]
|
||||
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
|
||||
|
||||
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
|
||||
mask = seq_range < valid_expert_lengths.unsqueeze(1)
|
||||
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
|
||||
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
|
||||
|
||||
moba_cu_seqlen_kv = torch.cat(
|
||||
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
|
||||
valid_expert_lengths.cumsum(dim=0)],
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
assert (
|
||||
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
|
||||
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
|
||||
|
||||
return MixedAttention.apply(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
|
||||
def process_moba_input(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
"""
|
||||
Process inputs for the attention function.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
|
||||
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
|
||||
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed input tensor.
|
||||
"""
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
|
||||
else:
|
||||
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
|
||||
if len(chunk_size) == 2:
|
||||
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
|
||||
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
|
||||
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
|
||||
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
|
||||
elif len(chunk_size) == 3:
|
||||
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
|
||||
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
|
||||
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
|
||||
else:
|
||||
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
|
||||
|
||||
return x, moba_chunk_size
|
||||
|
||||
|
||||
def process_moba_output(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
pass
|
||||
elif len(chunk_size) == 2:
|
||||
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
|
||||
elif len(chunk_size) == 3:
|
||||
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# TEST
|
||||
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
|
||||
random.seed(0)
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
|
||||
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
|
||||
max_seqlen = q.shape[1]
|
||||
q = rearrange(q, "b s ... -> (b s) ...")
|
||||
k = rearrange(k, "b s ... -> (b s) ...")
|
||||
v = rearrange(v, "b s ... -> (b s) ...")
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
|
||||
"""Speed test comparing flash_attn vs moba_attention"""
|
||||
# Get data
|
||||
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
|
||||
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
|
||||
vo_grad = torch.randn_like(q)
|
||||
|
||||
# Warmup
|
||||
warmup_iters = 3
|
||||
perf_test_iters = 10
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_flash = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_moba = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
|
||||
|
||||
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
|
||||
print(f"Speedup: {time_flash / time_moba:.2f}x")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
CUDA_VISIBLE_DEVICES=1 \
|
||||
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
"""
|
||||
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
|
||||
@@ -1,71 +0,0 @@
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import BoolTensor, IntTensor
|
||||
from torch.nn.attention.flex_attention import create_block_mask
|
||||
|
||||
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
|
||||
torch._inductor.config.realize_opcount_threshold = 100
|
||||
|
||||
|
||||
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
|
||||
"""Generates a 3D NATTEN attention mask with a given kernel size.
|
||||
|
||||
Args:
|
||||
canvas_t: The time dimension of the canvas.
|
||||
canvas_h: The height of the canvas.
|
||||
canvas_w: The width of the canvas.
|
||||
kernel_t: The time dimension of the kernel.
|
||||
kernel_h: The height of the kernel.
|
||||
kernel_w: The width of the kernel.
|
||||
"""
|
||||
canvas_t, canvas_h, canvas_w = canvas_twh
|
||||
kernel_t, kernel_h, kernel_w = kernel_twh
|
||||
tile_t_size, tile_h_size, tile_w_size = tile_twh
|
||||
total_tile_size = tile_t_size * tile_h_size * tile_w_size
|
||||
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
|
||||
tile_id = idx // total_tile_size
|
||||
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
|
||||
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
|
||||
tile_w = tile_id % canvas_tile_w
|
||||
return tile_t, tile_h, tile_w
|
||||
|
||||
def sta_mask_3d(
|
||||
b: IntTensor,
|
||||
h: IntTensor,
|
||||
q_idx: IntTensor,
|
||||
kv_idx: IntTensor,
|
||||
) -> BoolTensor:
|
||||
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
|
||||
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
|
||||
# kernel nominally attempts to center itself on the query, but kernel center
|
||||
# is clamped to a fixed distance (kernel half-length) from the canvas edge
|
||||
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
|
||||
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
|
||||
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
|
||||
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
|
||||
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
|
||||
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
|
||||
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
|
||||
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
|
||||
|
||||
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return sta_mask_3d
|
||||
|
||||
|
||||
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
|
||||
img_seq_len = img_size[0] * img_size[1] * img_size[2]
|
||||
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
|
||||
mask = create_block_mask(image_mask,
|
||||
B=None,
|
||||
H=None,
|
||||
Q_LEN=img_seq_len + text_max_len,
|
||||
KV_LEN=img_seq_len + text_max_len,
|
||||
device=device,
|
||||
_compile=True)
|
||||
return mask
|
||||
@@ -1,63 +0,0 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
|
||||
# Local support import
|
||||
from .support_flex_sta import get_sliding_tile_attention_mask
|
||||
|
||||
# USE OUR NEW PACKAGE!
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
return output
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
# Using the same parameters as the original test
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
|
||||
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
|
||||
for kernel_size in kernel_size_ls:
|
||||
print(f"Testing kernel_size: {kernel_size}")
|
||||
for xi in tqdm(range(num_iterations)):
|
||||
torch.manual_seed(xi)
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
|
||||
pt_o = flex_test(Q, K, V, kernel_size)
|
||||
|
||||
diff = pt_o - tk_o
|
||||
abs_diff = torch.abs(diff)
|
||||
max_d = torch.max(abs_diff).item()
|
||||
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
|
||||
|
||||
if max_d > 0.1:
|
||||
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
|
||||
|
||||
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
|
||||
@@ -1,97 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from fastvideo_kernel.vmoba import moba_attn_varlen
|
||||
|
||||
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
|
||||
"""
|
||||
Generates random data for testing the variable-length attention function.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
random.seed(42)
|
||||
torch.cuda.manual_seed_all(42)
|
||||
|
||||
# Generate sequence lengths for each item in the batch
|
||||
if batch_size > 1:
|
||||
# Ensure sequence lengths are reasonably distributed
|
||||
avg_seqlen = total_seqlen // batch_size
|
||||
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
|
||||
remaining_len = total_seqlen - sum(seqlens)
|
||||
if remaining_len > 0:
|
||||
seqlens.append(remaining_len)
|
||||
else: # Adjust if sum exceeds total_seqlen
|
||||
seqlens.append(avg_seqlen)
|
||||
current_sum = sum(seqlens)
|
||||
seqlens[-1] -= (current_sum - total_seqlen)
|
||||
# Ensure all lengths are positive
|
||||
seqlens = [max(1, s) for s in seqlens]
|
||||
# Final adjustment to match total_seqlen
|
||||
seqlens[-1] += total_seqlen - sum(seqlens)
|
||||
|
||||
else:
|
||||
seqlens = [total_seqlen]
|
||||
|
||||
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
|
||||
max_seqlen = max(seqlens) if seqlens else 0
|
||||
|
||||
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [1, 2])
|
||||
@pytest.mark.parametrize("total_seqlen", [512, 1024])
|
||||
@pytest.mark.parametrize("num_heads", [8])
|
||||
@pytest.mark.parametrize("head_dim", [64])
|
||||
@pytest.mark.parametrize("moba_chunk_size", [64])
|
||||
@pytest.mark.parametrize("moba_topk", [2, 4])
|
||||
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
|
||||
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_moba_attn_varlen_forward(
|
||||
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
|
||||
):
|
||||
"""
|
||||
Tests the forward pass of moba_attn_varlen for basic correctness.
|
||||
It checks output shape, dtype, and for the presence of NaNs/Infs.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
pytest.skip("float32 is not supported in flash attention")
|
||||
|
||||
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
|
||||
batch_size, total_seqlen, num_heads, head_dim, dtype
|
||||
)
|
||||
|
||||
# Ensure chunk size is not larger than the smallest sequence length
|
||||
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
|
||||
if moba_chunk_size > min_seqlen:
|
||||
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
|
||||
|
||||
try:
|
||||
output = moba_attn_varlen(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
moba_chunk_size=moba_chunk_size,
|
||||
moba_topk=moba_topk,
|
||||
select_mode=select_mode,
|
||||
threshold_type=threshold_type,
|
||||
simsum_threshold=0.5, # A reasonable default for threshold mode
|
||||
)
|
||||
except Exception as e:
|
||||
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
|
||||
|
||||
# 1. Check output shape
|
||||
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
|
||||
|
||||
# 2. Check output dtype
|
||||
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
|
||||
|
||||
# 3. Check for NaNs or Infs in the output
|
||||
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
|
||||
@@ -55,10 +55,17 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Kernels
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -55,10 +55,17 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Kernels
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -55,10 +55,17 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Kernels
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
FROM rocm/pytorch:rocm7.1_ubuntu22.04_py3.10_pytorch_release_2.9.1
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject_other.toml ./pyproject.toml
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[rocm] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Kernels
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r requirements-mkdocs.txt
|
||||
pip install -r docs/requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -6,7 +6,7 @@ Thank you for your interest in contributing to FastVideo. We want to make the pr
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -30,7 +30,7 @@ path_to_your_dataset_folder/
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
To generate the `videos2caption.json` and `merge.txt`, run
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
|
||||
@@ -7,9 +7,9 @@ pip install vsa
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We support H100s (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
We support H100 (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
|
||||
First, install C++20 for ThunderKittens (if using an H100):
|
||||
First, install C++20 for ThunderKittens (if using H100):
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
|
||||
@@ -6,7 +6,7 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
|
||||
try:
|
||||
from fastvideo_kernel import video_sparse_attn
|
||||
from vsa import video_sparse_attn
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo_kernel import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -1,149 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Video DiT configuration for native FastVideo implementation.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def is_longcat_blocks(n: str, m) -> bool:
|
||||
"""FSDP shard condition for LongCat transformer blocks."""
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatVideoArchConfig(DiTArchConfig):
|
||||
"""Architecture configuration for native LongCat Video DiT."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Enable torch.compile for transformer blocks (major speedup!)
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Parameter name mapping for weight conversion
|
||||
# Maps original LongCat third_party names -> native FastVideo names
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Embedders
|
||||
r"^x_embedder\.(.*)$": r"patch_embed.\1",
|
||||
r"^t_embedder\.mlp\.0\.(.*)$": r"time_embedder.linear_1.\1",
|
||||
r"^t_embedder\.mlp\.2\.(.*)$": r"time_embedder.linear_2.\1",
|
||||
r"^y_embedder\.y_proj\.0\.(.*)$": r"caption_embedder.linear_1.\1",
|
||||
r"^y_embedder\.y_proj\.2\.(.*)$": r"caption_embedder.linear_2.\1",
|
||||
|
||||
# Transformer blocks - AdaLN modulation
|
||||
r"^blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
|
||||
r"blocks.\1.adaln_linear_1.\2",
|
||||
|
||||
# Transformer blocks - Normalization
|
||||
r"^blocks\.(\d+)\.mod_norm_attn\.(.*)$": r"blocks.\1.norm_attn.\2",
|
||||
r"^blocks\.(\d+)\.mod_norm_ffn\.(.*)$": r"blocks.\1.norm_ffn.\2",
|
||||
r"^blocks\.(\d+)\.pre_crs_attn_norm\.(.*)$":
|
||||
r"blocks.\1.norm_cross.\2",
|
||||
|
||||
# Self-attention: QKV fused -> separate (will need splitting in converter)
|
||||
# Original has attn.qkv.weight -> need to split into to_q, to_k, to_v
|
||||
r"^blocks\.(\d+)\.attn\.qkv\.(.*)$":
|
||||
r"blocks.\1.self_attn.qkv_fused.\2", # Marker for splitting
|
||||
r"^blocks\.(\d+)\.attn\.proj\.(.*)$":
|
||||
r"blocks.\1.self_attn.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn\.q_norm\.(.*)$":
|
||||
r"blocks.\1.self_attn.q_norm.\2",
|
||||
r"^blocks\.(\d+)\.attn\.k_norm\.(.*)$":
|
||||
r"blocks.\1.self_attn.k_norm.\2",
|
||||
|
||||
# Cross-attention
|
||||
r"^blocks\.(\d+)\.cross_attn\.q_linear\.(.*)$":
|
||||
r"blocks.\1.cross_attn.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.kv_linear\.(.*)$":
|
||||
r"blocks.\1.cross_attn.kv_fused.\2", # Marker for splitting
|
||||
r"^blocks\.(\d+)\.cross_attn\.proj\.(.*)$":
|
||||
r"blocks.\1.cross_attn.to_out.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q_norm\.(.*)$":
|
||||
r"blocks.\1.cross_attn.q_norm.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k_norm\.(.*)$":
|
||||
r"blocks.\1.cross_attn.k_norm.\2",
|
||||
|
||||
# FFN (SwiGLU)
|
||||
r"^blocks\.(\d+)\.ffn\.w1\.(.*)$": r"blocks.\1.ffn.w1.\2", # gate
|
||||
r"^blocks\.(\d+)\.ffn\.w2\.(.*)$": r"blocks.\1.ffn.w2.\2", # down
|
||||
r"^blocks\.(\d+)\.ffn\.w3\.(.*)$": r"blocks.\1.ffn.w3.\2", # up
|
||||
|
||||
# Final layer
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
|
||||
r"final_layer.adaln_linear.\1",
|
||||
r"^final_layer\.norm_final\.(.*)$": r"final_layer.norm.\1",
|
||||
r"^final_layer\.linear\.(.*)$": r"final_layer.proj.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: custom -> hf
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# LoRA parameter name mapping
|
||||
lora_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Model architecture parameters
|
||||
hidden_size: int = 4096
|
||||
depth: int = 48 # Number of transformer blocks
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128 # hidden_size / num_attention_heads
|
||||
|
||||
in_channels: int = 16 # Latent space channels
|
||||
out_channels: int = 16
|
||||
num_channels_latents: int = 16
|
||||
|
||||
# Patch embedding
|
||||
patch_size: tuple[int, int,
|
||||
int] = (1, 2, 2) # [T, H, W] - no temporal compression
|
||||
|
||||
# Text/caption embedding
|
||||
caption_channels: int = 4096 # UMT5 d_model
|
||||
|
||||
# Timestep embedding
|
||||
adaln_tembed_dim: int = 512
|
||||
frequency_embedding_size: int = 256
|
||||
|
||||
# FFN
|
||||
mlp_ratio: int = 4
|
||||
|
||||
# Attention backend support
|
||||
_supported_attention_backends: tuple = field(default_factory=lambda: (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
))
|
||||
|
||||
# Text padding behavior
|
||||
text_tokens_zero_pad: bool = True
|
||||
|
||||
# Block Sparse Attention (BSA)
|
||||
enable_bsa: bool = False
|
||||
bsa_params: dict | None = field(
|
||||
default_factory=lambda: {
|
||||
"sparsity": 0.9375,
|
||||
"cdf_threshold": None,
|
||||
"chunk_3d_shape_q": [4, 4, 4],
|
||||
"chunk_3d_shape_k": [4, 4, 4],
|
||||
})
|
||||
|
||||
# LoRA exclusions
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
# Ensure attention_head_dim matches
|
||||
self.attention_head_dim = self.hidden_size // self.num_attention_heads
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatVideoConfig(DiTConfig):
|
||||
"""Main configuration for LongCat Video DiT."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=LongCatVideoArchConfig)
|
||||
|
||||
prefix: str = "longcat"
|
||||
@@ -1,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}'"
|
||||
)
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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] = {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.")
|
||||
|
||||
@@ -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}."
|
||||
)
|
||||
|
||||
@@ -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
|
||||
-946
@@ -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)
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
Reference in New Issue
Block a user