Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
104a539a22 | ||
|
|
f8bfc76015 | ||
|
|
8f1e6c3336 | ||
|
|
8e7d2e7879 | ||
|
|
6ab2870942 | ||
|
|
e0ad145152 | ||
|
|
da04d08426 | ||
|
|
1f70032af5 |
@@ -0,0 +1,236 @@
|
||||
name: Publish FastVideo Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/fastvideo_kernel/pyproject.toml"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/fastvideo_kernel
|
||||
# Get current commit's version from pyproject.toml
|
||||
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
torch-cuda:
|
||||
- torch-version: '2.5.1'
|
||||
cuda-version: '12.4.1'
|
||||
torch-cuda-short: 'cu124'
|
||||
- torch-version: '2.6.0'
|
||||
cuda-version: '12.6.3'
|
||||
torch-cuda-short: 'cu126'
|
||||
- torch-version: '2.7.1'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton
|
||||
|
||||
cd csrc/fastvideo_kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/fastvideo_kernel
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}-py${{ matrix.python-version }}
|
||||
path: csrc/fastvideo_kernel/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton
|
||||
|
||||
cd csrc/fastvideo_kernel
|
||||
git submodule update --init --recursive
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/fastvideo_kernel/dist/
|
||||
@@ -30,6 +30,8 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
@@ -2,12 +2,12 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
We support H100 (via TK) and any other GPU (via triton) for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
@@ -16,6 +16,14 @@ git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you want to skip the compilation of the TK kernel and only use the Triton version, try below:
|
||||
|
||||
```bash
|
||||
SKIP_SM90_EXT=1 python setup.py install
|
||||
or
|
||||
SKIP_SM90_EXT=1 pip install --no-build-isolation .
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
@@ -30,7 +38,7 @@ sudo apt install clang-11
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
@@ -43,7 +51,7 @@ bash scripts/inference/v1_inference_wan_STA.sh
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
@@ -58,7 +66,6 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
@@ -67,7 +74,7 @@ python ../benchmarks/bench_sta.py
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
@@ -82,7 +89,7 @@ Here is a diagram of how the window is configured and passed through the FastVid
|
||||
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
@@ -51,21 +51,28 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
ext_modules = []
|
||||
|
||||
if os.environ.get("SKIP_SM90_EXT", "0") != "1":
|
||||
ext_modules.append(
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
)
|
||||
else:
|
||||
print("ENV SKIP_SM90_EXT=1, skip st_attn_cuda compile")
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
|
||||
@@ -7,12 +7,17 @@ try:
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
try:
|
||||
from st_attn.st_attn_triton import sliding_tile_attention_triton
|
||||
except ImportError:
|
||||
sliding_tile_attention_triton = None
|
||||
|
||||
def sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
seq_length = q_all.shape[2]
|
||||
dit_seq_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
@@ -46,4 +51,13 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
return hidden_states[:, :, :seq_length]
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
major, minor = torch.cuda.get_device_capability(q_all.device)
|
||||
if major == 9 and minor == 0 and sta_fwd is not None:
|
||||
return sliding_tile_attention_SM90(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
|
||||
elif sliding_tile_attention_triton is not None:
|
||||
return sliding_tile_attention_triton(q_all, k_all, v_all, window_size, text_length, has_text, dit_seq_shape)
|
||||
else:
|
||||
raise ImportError("No suitable sliding tile attention implementation found.")
|
||||
|
||||
@@ -0,0 +1,327 @@
|
||||
import math
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
def is_cuda():
|
||||
return triton.runtime.driver.active.get_current_target().backend == "cuda"
|
||||
|
||||
|
||||
def is_hip():
|
||||
target = triton.runtime.driver.active.get_current_target()
|
||||
return target.backend == 'hip'
|
||||
|
||||
def get_common_autotune_config():
|
||||
configs = [
|
||||
triton.Config({'BLOCK_Q': BLOCK_Q, 'BLOCK_KV': BLOCK_KV}, num_stages=s, num_warps=w) \
|
||||
for BLOCK_Q in [32, 64, 128]\
|
||||
for BLOCK_KV in [32, 64, 128]\
|
||||
for s in [1, 2, 3, 4]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
return configs
|
||||
|
||||
|
||||
def get_cuda_autotune_config():
|
||||
# cuda and hip can use differnt autotune configs
|
||||
return get_common_autotune_config()
|
||||
|
||||
|
||||
def get_hip_autotune_config():
|
||||
# cuda and hip can use differnt autotune configs
|
||||
return get_common_autotune_config()
|
||||
|
||||
|
||||
def get_autotune_config():
|
||||
if is_cuda():
|
||||
return get_cuda_autotune_config()
|
||||
else:
|
||||
return get_hip_autotune_config()
|
||||
|
||||
|
||||
@triton.jit
|
||||
def clamp_int(value, min_val, max_val):
|
||||
ret = tl.where(value > max_val, max_val, value)
|
||||
ret = tl.where(ret < min_val, min_val, ret)
|
||||
return ret
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_loop(
|
||||
q, k, v, kv_mask, m, l, acc, sm_scale,
|
||||
MASK_KV: tl.constexpr,
|
||||
):
|
||||
scores = tl.dot(q, k.T) #[BLOCK_Q, BLOCK_KV]
|
||||
scores = scores * sm_scale
|
||||
if MASK_KV:
|
||||
scores = tl.where(kv_mask[None, :], scores, -float('inf'))
|
||||
|
||||
current_m = tl.max(scores, axis=1)
|
||||
new_m = tl.maximum(m, current_m)
|
||||
exp_scores = tl.math.exp2(scores - new_m[:, None])
|
||||
current_l = tl.sum(exp_scores, axis=1)
|
||||
|
||||
# Update L <- L * exp(M - M') + L1, M <- M'
|
||||
alpha = tl.math.exp2(m - new_m)
|
||||
l = l * alpha + current_l
|
||||
m = new_m
|
||||
|
||||
# Update O <- O * exp(M - M') + P @ V
|
||||
acc = (acc * alpha[:, None] + tl.dot(exp_scores.to(v.type.element_ty), v))
|
||||
|
||||
return m, l, acc
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=get_autotune_config(),
|
||||
key=['head_dim'],
|
||||
)
|
||||
@triton.jit
|
||||
def triton_sta_kernel(
|
||||
Q, K, V, output,
|
||||
batch_size: int, num_heads: int, seq_len: int, head_dim: int,
|
||||
img_seq_len: int,
|
||||
text_length: int,
|
||||
canvas_t: int, canvas_h: int, canvas_w: int,
|
||||
kernel_t: int, kernel_h: int, kernel_w: int,
|
||||
tile_t: int, tile_h: int, tile_w: int,
|
||||
scale: float,
|
||||
has_text: tl.constexpr,
|
||||
text_q: tl.constexpr,
|
||||
BLOCK_Q: tl.constexpr,
|
||||
BLOCK_KV: tl.constexpr,
|
||||
BLOCK_DIM: tl.constexpr,
|
||||
):
|
||||
total_tile_size = tile_t * tile_h * tile_w
|
||||
q_block_per_tile = (total_tile_size + BLOCK_Q - 1) // BLOCK_Q
|
||||
|
||||
batch_idx = tl.program_id(0)
|
||||
head_idx = tl.program_id(1)
|
||||
if text_q:
|
||||
q_block_idx = tl.program_id(2)
|
||||
else:
|
||||
q_tile_flat = tl.program_id(2) // q_block_per_tile
|
||||
q_block_idx = tl.program_id(2) % q_block_per_tile
|
||||
|
||||
m = tl.full((BLOCK_Q,), -float('inf'), dtype=tl.float32)
|
||||
l = tl.zeros((BLOCK_Q,), dtype=tl.float32)
|
||||
acc = tl.zeros((BLOCK_Q, BLOCK_DIM), dtype=tl.float32)
|
||||
|
||||
q_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
if text_q:
|
||||
q_base_idx = img_seq_len + q_block_idx * BLOCK_Q
|
||||
else:
|
||||
q_base_idx = q_tile_flat * total_tile_size + q_block_idx * BLOCK_Q
|
||||
|
||||
q_offset_in_tile = tl.arange(0, BLOCK_Q)
|
||||
q_idx = q_base_idx + q_offset_in_tile
|
||||
q_mask = (q_block_idx * BLOCK_Q + tl.arange(0, BLOCK_Q)) < total_tile_size
|
||||
|
||||
q = tl.load(
|
||||
Q + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=q_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_Q, BLOCK_DIM]
|
||||
|
||||
# Scale sm_scale by log_2(e) and use 2^x instead of exp
|
||||
sm_scale = scale * 1.4426950408889634
|
||||
|
||||
num_tiles_t = canvas_t // tile_t
|
||||
num_tiles_h = canvas_h // tile_h
|
||||
num_tiles_w = canvas_w // tile_w
|
||||
tiles_per_hw = num_tiles_h * num_tiles_w
|
||||
|
||||
if text_q:
|
||||
kv_tile_start_t = 0
|
||||
kv_tile_end_t = num_tiles_t
|
||||
|
||||
kv_tile_start_h = 0
|
||||
kv_tile_end_h = num_tiles_h
|
||||
|
||||
kv_tile_start_w = 0
|
||||
kv_tile_end_w = num_tiles_w
|
||||
|
||||
else:
|
||||
q_tile_t = q_tile_flat // tiles_per_hw
|
||||
remaining = q_tile_flat % tiles_per_hw
|
||||
q_tile_h = remaining // num_tiles_w
|
||||
q_tile_w = remaining % num_tiles_w
|
||||
|
||||
kernel_center_t = clamp_int(q_tile_t, kernel_t // 2, (num_tiles_t - 1) - kernel_t // 2)
|
||||
kernel_center_h = clamp_int(q_tile_h, kernel_h // 2, (num_tiles_h - 1) - kernel_h // 2)
|
||||
kernel_center_w = clamp_int(q_tile_w, kernel_w // 2, (num_tiles_w - 1) - kernel_w // 2)
|
||||
|
||||
kv_tile_start_t = kernel_center_t - kernel_t // 2
|
||||
kv_tile_end_t = kernel_center_t + kernel_t // 2 + 1
|
||||
kv_tile_end_t = tl.where(kv_tile_end_t > num_tiles_t, num_tiles_t, kv_tile_end_t)
|
||||
|
||||
kv_tile_start_h = kernel_center_h - kernel_h // 2
|
||||
kv_tile_end_h = kernel_center_h + kernel_h // 2 + 1
|
||||
kv_tile_end_h = tl.where(kv_tile_end_h > num_tiles_h, num_tiles_h, kv_tile_end_h)
|
||||
|
||||
kv_tile_start_w = kernel_center_w - kernel_w // 2
|
||||
kv_tile_end_w = kernel_center_w + kernel_w // 2 + 1
|
||||
kv_tile_end_w = tl.where(kv_tile_end_w > num_tiles_w, num_tiles_w, kv_tile_end_w)
|
||||
|
||||
# for kv_img
|
||||
for kv_tile_t in tl.range(kv_tile_start_t, kv_tile_end_t):
|
||||
for kv_tile_h in tl.range(kv_tile_start_h, kv_tile_end_h):
|
||||
for kv_tile_w in tl.range(kv_tile_start_w, kv_tile_end_w):
|
||||
kv_base_idx = (kv_tile_t * num_tiles_h * num_tiles_w + kv_tile_h * num_tiles_w + kv_tile_w) * total_tile_size
|
||||
|
||||
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
|
||||
kv_offset_in_block = tl.arange(0, BLOCK_KV)
|
||||
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
|
||||
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < total_tile_size
|
||||
|
||||
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
|
||||
k = tl.load(
|
||||
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
v = tl.load(
|
||||
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
|
||||
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, False)
|
||||
|
||||
|
||||
# for kv_text
|
||||
if has_text:
|
||||
kv_base_idx = img_seq_len
|
||||
for kv_block_idx in tl.range(0, total_tile_size, BLOCK_KV):
|
||||
kv_offset_in_block = tl.arange(0, BLOCK_KV)
|
||||
kv_idx = kv_base_idx + kv_block_idx + kv_offset_in_block
|
||||
kv_mask = (kv_block_idx + tl.arange(0, BLOCK_KV)) < text_length
|
||||
|
||||
kv_offset = (batch_idx * num_heads + head_idx) * seq_len * head_dim
|
||||
|
||||
k = tl.load(
|
||||
K + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
v = tl.load(
|
||||
V + kv_offset + kv_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
mask=kv_mask[:, None],
|
||||
other=0.0
|
||||
) # [BLOCK_KV, BLOCK_DIM]
|
||||
|
||||
m, l, acc = _attn_fwd_loop(q, k, v, kv_mask, m, l, acc, sm_scale, True)
|
||||
|
||||
|
||||
output_acc = acc / l[:, None]
|
||||
tl.store(
|
||||
output + q_offset + q_idx[:, None] * head_dim + tl.arange(0, BLOCK_DIM)[None, :],
|
||||
output_acc,
|
||||
mask=q_mask[:, None]
|
||||
) # [BLOCK_Q, BLOCK_DIM]
|
||||
|
||||
|
||||
def sliding_tile_attention_triton(
|
||||
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
window_size, text_length: int,
|
||||
has_text=True, dit_seq_shape='30x48x80') -> torch.Tensor:
|
||||
seq_length = q.shape[2]
|
||||
if has_text:
|
||||
assert q.shape[2] >= 115200 and q.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
|
||||
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
|
||||
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo
|
||||
assert q.shape[2] == 82944
|
||||
elif dit_seq_shape == '18x48x80': # Wan
|
||||
assert q.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
assert q.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
if dit_seq_shape == '30x48x80': # Hunyuan
|
||||
canvas_t, canvas_h, canvas_w = 30, 48, 80
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
elif dit_seq_shape == '36x48x48': # Stepvideo
|
||||
canvas_t, canvas_h, canvas_w = 36, 48, 48
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
elif dit_seq_shape == '18x48x80': # Wan
|
||||
canvas_t, canvas_h, canvas_w = 18, 48, 80
|
||||
tile_t, tile_h, tile_w = 6, 8, 8
|
||||
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
num_tiles_t = canvas_t // tile_t
|
||||
num_tiles_h = canvas_h // tile_h
|
||||
num_tiles_w = canvas_w // tile_w
|
||||
num_tiles = num_tiles_t * num_tiles_h * num_tiles_w
|
||||
|
||||
total_tile_size = tile_t * tile_h * tile_w
|
||||
|
||||
# BLOCK_Q=128
|
||||
# BLOCK_KV=128
|
||||
BLOCK_DIM = head_dim
|
||||
|
||||
output = torch.empty_like(q)
|
||||
|
||||
# for q_img
|
||||
# kernel_size maybe different for different head
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (kernel_t, kernel_h, kernel_w) in enumerate(window_size):
|
||||
for batch in range(batch_size):
|
||||
q_head, k_head, v_head, o_head = (q[batch:batch + 1, head_index:head_index + 1],
|
||||
k[batch:batch + 1, head_index:head_index + 1],
|
||||
v[batch:batch + 1, head_index:head_index + 1],
|
||||
output[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
# triton_sta_kernel[(1, 1, num_tiles * triton.cdiv(total_tile_size, BLOCK_Q))](
|
||||
grid = lambda META: (1, 1, num_tiles * triton.cdiv(total_tile_size, META['BLOCK_Q']))
|
||||
triton_sta_kernel[grid](
|
||||
q_head, k_head, v_head, o_head,
|
||||
1, 1, seq_len, head_dim,
|
||||
img_seq_len,
|
||||
text_length,
|
||||
canvas_t, canvas_h, canvas_w,
|
||||
kernel_t, kernel_h, kernel_w,
|
||||
tile_t, tile_h, tile_w,
|
||||
scale=1.0 / (head_dim ** 0.5),
|
||||
has_text=has_text,
|
||||
text_q=False,
|
||||
# BLOCK_Q=BLOCK_Q,
|
||||
# BLOCK_KV=BLOCK_KV,
|
||||
BLOCK_DIM=BLOCK_DIM,
|
||||
)
|
||||
|
||||
# for q_text
|
||||
# kernel_t, kernel_h, kernel_w is not used, set to (3, 3, 3)
|
||||
if has_text:
|
||||
# triton_sta_kernel[(batch_size, num_heads, triton.cdiv(total_tile_size, BLOCK_Q))](
|
||||
grid = lambda META: (batch_size, num_heads, triton.cdiv(total_tile_size, META['BLOCK_Q']))
|
||||
triton_sta_kernel[grid](
|
||||
q, k, v, output,
|
||||
batch_size, num_heads, seq_len, head_dim,
|
||||
img_seq_len,
|
||||
text_length,
|
||||
canvas_t, canvas_h, canvas_w,
|
||||
3, 3, 3,
|
||||
#kernel_t, kernel_h, kernel_w,
|
||||
tile_t, tile_h, tile_w,
|
||||
scale=1.0 / (head_dim ** 0.5),
|
||||
has_text=has_text,
|
||||
text_q=True,
|
||||
# BLOCK_Q=BLOCK_Q,
|
||||
# BLOCK_KV=BLOCK_KV,
|
||||
BLOCK_DIM=BLOCK_DIM,
|
||||
)
|
||||
|
||||
if has_text:
|
||||
if pad_size > 0:
|
||||
output = output[:, :, :seq_length]
|
||||
|
||||
return output
|
||||
@@ -0,0 +1,7 @@
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
__pycache__/
|
||||
*.so
|
||||
*.pyc
|
||||
.ipynb_checkpoints/
|
||||
@@ -0,0 +1,187 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
@@ -0,0 +1,6 @@
|
||||
include LICENSE
|
||||
include README.md
|
||||
include pyproject.toml
|
||||
recursive-include src/fastvideo_kernel *.cu *.cuh *.cpp *.h
|
||||
recursive-include csrc *.cu *.cuh *.cpp *.h
|
||||
recursive-include tk *.cu *.cuh *.cpp *.h
|
||||
@@ -0,0 +1,31 @@
|
||||
# FastVideo Kernel
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
cd csrc/fastvideo_kernel
|
||||
pip install .
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from fastvideo_kernel import sliding_tile_attention, video_sparse_attn, moba_attn_varlen
|
||||
|
||||
# Example: Sliding Tile Attention
|
||||
out = sliding_tile_attention(q, k, v, window_sizes, text_len)
|
||||
|
||||
# Example: Video Sparse Attention (with Triton fallback)
|
||||
out = video_sparse_attn(q, k, v, block_sizes, topk=5)
|
||||
|
||||
# Example: VMoBA
|
||||
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
|
||||
```
|
||||
|
||||
## Requirements
|
||||
|
||||
- H100 GPU (sm_90a) for CUDA kernels
|
||||
- Triton for non-H100 fallback
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,23 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,573 @@
|
||||
// # Define TORCH_COMPILE macro
|
||||
|
||||
#include "kittens.cuh"
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <stdio.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
|
||||
return (value < min) ? min : ((value > max) ? max : value);
|
||||
}
|
||||
// #define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
__device__ __forceinline__ int abs_int(int value) {
|
||||
return (value < 0) ? -value : value;
|
||||
}
|
||||
|
||||
|
||||
constexpr int CONSUMER_WARPGROUPS = (3);
|
||||
constexpr int PRODUCER_WARPGROUPS = (1);
|
||||
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
|
||||
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
|
||||
|
||||
using namespace kittens;
|
||||
namespace cg = cooperative_groups;
|
||||
|
||||
template<int D> struct fwd_attend_ker_tile_dims {};
|
||||
template<> struct fwd_attend_ker_tile_dims<64> {
|
||||
constexpr static int tile_width = (64);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (4);
|
||||
};
|
||||
template<> struct fwd_attend_ker_tile_dims<128> {
|
||||
constexpr static int tile_width = (128);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (2);
|
||||
};
|
||||
|
||||
template<int D> struct fwd_globals {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
|
||||
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
q_gl q;
|
||||
k_gl k;
|
||||
v_gl v;
|
||||
l_gl l;
|
||||
o_gl o;
|
||||
|
||||
const int N;
|
||||
const int text_L;
|
||||
const int hr;
|
||||
};
|
||||
|
||||
|
||||
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
|
||||
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
|
||||
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
extern __shared__ int __shm[];
|
||||
tma_swizzle_allocator al((int*)&__shm[0]);
|
||||
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
|
||||
|
||||
using K = fwd_attend_ker_tile_dims<D>;
|
||||
|
||||
using q_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
using k_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using v_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
|
||||
using o_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
|
||||
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
|
||||
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
|
||||
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
|
||||
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
|
||||
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
|
||||
int img_kv_blocks;
|
||||
int kv_blocks = g.N / (K::kv_height);
|
||||
if constexpr (text_kv) {
|
||||
img_kv_blocks = kv_blocks - 3;
|
||||
} else {
|
||||
img_kv_blocks = kv_blocks;
|
||||
}
|
||||
int kv_head_idx = blockIdx.y / g.hr;
|
||||
int seq_idx;
|
||||
if constexpr (text_q) {
|
||||
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
} else {
|
||||
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
}
|
||||
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
|
||||
if (threadIdx.x == 0) {
|
||||
init_semaphore(qsmem_semaphore, 0, 1);
|
||||
for(int j = 0; j < K::stages; j++) {
|
||||
init_semaphore(k_smem_arrived[j], 0, 1);
|
||||
init_semaphore(v_smem_arrived[j], 0, 1);
|
||||
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
|
||||
}
|
||||
|
||||
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
|
||||
|
||||
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
|
||||
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
|
||||
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
|
||||
}
|
||||
|
||||
if constexpr (text_q){
|
||||
for (int j = 0; j < K::stages - 1; j++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
|
||||
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
|
||||
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
|
||||
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int count = 0;
|
||||
int j = 0;
|
||||
while (count < K::stages - 1) {
|
||||
int kt = j / 3 / (CH * CW);
|
||||
int kh = (j / 3) % (CH * CW) / CW;
|
||||
int kw = (j / 3) % CW;
|
||||
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
|
||||
if (mask){
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
|
||||
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
|
||||
count += 1;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int pipe_idx = K::stages - 1;
|
||||
|
||||
if(warpgroupid == NUM_WARPGROUPS-1) {
|
||||
warpgroup::decrease_registers<32>();
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
|
||||
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
|
||||
}
|
||||
else { kv_iters = kv_blocks-2;}
|
||||
|
||||
if(warpid == NUM_WORKERS-4) {
|
||||
if constexpr (text_q){
|
||||
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
|
||||
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int k_t_min = clamp_int(qt-DT, 0, CT-1);
|
||||
int k_t_max = clamp_int(qt+DT, 0, CT-1);
|
||||
int k_h_min = clamp_int(qh-DH, 0, CH-1);
|
||||
int k_h_max = clamp_int(qh+DH, 0, CH-1);
|
||||
int k_w_min = clamp_int(qw-DW, 0, CW-1);
|
||||
int k_w_max = clamp_int(qw+DW, 0, CW-1);
|
||||
int count = 0;
|
||||
for (int kt = k_t_min; kt <= k_t_max; kt++) {
|
||||
for (int kh = k_h_min; kh <= k_h_max; kh++) {
|
||||
for (int kw = k_w_min; kw <= k_w_max; kw++) {
|
||||
for (int j = 0; j <= 2; j++){
|
||||
if (count >= K::stages - 1) {
|
||||
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
} else {
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// for text
|
||||
for (int index = img_kv_blocks; index < kv_blocks; index++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
warpgroup::increase_registers<160>();
|
||||
|
||||
rt_fl<16, K::kv_height> att_block;
|
||||
rt_bf<16, K::kv_height> att_block_mma;
|
||||
rt_fl<16, K::tile_width> o_reg;
|
||||
|
||||
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
|
||||
|
||||
neg_infty(max_vec);
|
||||
zero(norm_vec);
|
||||
zero(o_reg);
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
|
||||
kv_iters = (kv_iters/8);
|
||||
}
|
||||
else if constexpr (text_q){
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
kv_iters = img_kv_blocks - 1;
|
||||
} else {
|
||||
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
}
|
||||
|
||||
kittens::wait(qsmem_semaphore, 0);
|
||||
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
if constexpr(text_kv) {
|
||||
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
// apply non-pad mask
|
||||
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
|
||||
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
|
||||
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
|
||||
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
}
|
||||
|
||||
div_row(o_reg, o_reg, norm_vec);
|
||||
warpgroup::store(o_smem[warpgroupid], o_reg);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
|
||||
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
|
||||
}
|
||||
|
||||
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
|
||||
log(norm_vec, norm_vec);
|
||||
add(norm_vec, norm_vec, max_vec_scaled);
|
||||
|
||||
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
|
||||
else { mul(norm_vec, norm_vec, -11.313708499f); }
|
||||
|
||||
warpgroup::store(l_smem[warpgroupid], norm_vec);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
|
||||
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
|
||||
}
|
||||
tma::store_async_wait();
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
#include "pyutils/torch_helpers.cuh"
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <iostream>
|
||||
|
||||
torch::Tensor
|
||||
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
|
||||
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
|
||||
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
|
||||
auto hr = qo_heads / kv_heads;
|
||||
|
||||
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
|
||||
|
||||
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
|
||||
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
|
||||
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
|
||||
|
||||
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
|
||||
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
|
||||
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
|
||||
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
|
||||
if (head_dim == 128) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<128>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
// Shared memory size for the kernel.
|
||||
// We use the maximum available shared memory (kittens::MAX_SHARED_MEMORY)
|
||||
// which is approximately 227KB on H100, necessary for the high-performance
|
||||
// TMA-based attention tiles with multiple stages.
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
dim3 grid_text(2, qo_heads, batch);
|
||||
if (!process_text) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, true, DT_VAL, DH_VAL, DW_VAL, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(2, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 3, 0); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 1, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(2, 2, 2); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(2, 2, 3); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(2, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(2, 0, 5); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
} else {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
}
|
||||
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
if (kernel_aspect_ratio_flag == 2){
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(3, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(1, 3, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 1, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 3, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 0, 0); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(3, 0, 3); }
|
||||
else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(3, 3, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6) { LAUNCH_IMAGE_KER(0, 0, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(0, 3, 0); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
else if (kernel_aspect_ratio_flag == 3) {
|
||||
#define LAUNCH_IMAGE_KER(DT_VAL, DH_VAL, DW_VAL) \
|
||||
cudaFuncSetAttribute( \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10>, \
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize, \
|
||||
mem_size \
|
||||
); \
|
||||
fwd_attend_ker<128, false, false, false, DT_VAL, DH_VAL, DW_VAL, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 1, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 1, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(1, 2, 2); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 3, 0); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(1, 2, 3); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(1, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 3, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3) { LAUNCH_IMAGE_KER(1, 3, 1); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1) { LAUNCH_IMAGE_KER(1, 0, 0); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 3, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 2, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 3, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7) { LAUNCH_IMAGE_KER(0, 2, 3); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9) { LAUNCH_IMAGE_KER(0, 2, 4); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 0, 5); }
|
||||
else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(1, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){ LAUNCH_IMAGE_KER(0, 1, 5); }
|
||||
else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5) { LAUNCH_IMAGE_KER(0, 3, 2); }
|
||||
else {
|
||||
TORCH_CHECK(false, "Invalid kernel size: ", kernel_t_size, "x", kernel_h_size, "x", kernel_w_size);
|
||||
}
|
||||
#undef LAUNCH_IMAGE_KER
|
||||
}
|
||||
|
||||
else {
|
||||
TORCH_CHECK(false, "Unsupported kernel_aspect_ratio_flag: ", kernel_aspect_ratio_flag);
|
||||
}
|
||||
|
||||
}
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return o;
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=61.0", "torch>=2.5.0", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.1.0"
|
||||
description = "CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
license = {text = "Apache-2.0"}
|
||||
authors = [{name = "Hao AI Lab"}]
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
]
|
||||
dependencies = [
|
||||
"torch>=2.5.0",
|
||||
"triton>=2.0.0"
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -0,0 +1,132 @@
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
ROOT = Path(__file__).parent.absolute()
|
||||
CSRC_DIR = ROOT / "csrc"
|
||||
|
||||
# Path to ThunderKittens (TK)
|
||||
def get_tk_dir():
|
||||
tk_env = os.getenv("THUNDERKITTENS_ROOT")
|
||||
if tk_env:
|
||||
return tk_env
|
||||
|
||||
# Check common locations
|
||||
possible_paths = [
|
||||
ROOT / "tk",
|
||||
ROOT / "csrc" / "tk",
|
||||
ROOT.parent / "attn" / "sliding_tile_attn" / "tk",
|
||||
ROOT.parent / "attn" / "video_sparse_attn" / "tk",
|
||||
]
|
||||
for p in possible_paths:
|
||||
if (p / "include" / "kittens.cuh").exists():
|
||||
return str(p)
|
||||
|
||||
# Default fallback
|
||||
return str(ROOT.parent / "attn" / "sliding_tile_attn" / "tk")
|
||||
|
||||
TK_DIR = get_tk_dir()
|
||||
|
||||
def get_cuda_flags(tk_root: str) -> list:
|
||||
python_include = subprocess.check_output(
|
||||
["python", "-c", "import sysconfig; print(sysconfig.get_path('include'))"]
|
||||
).decode().strip()
|
||||
|
||||
torch_includes = subprocess.check_output([
|
||||
"python", "-c",
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; "
|
||||
"print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip().split()
|
||||
|
||||
return [
|
||||
"-DNDEBUG",
|
||||
"-Xcompiler=-Wno-psabi",
|
||||
"-Xcompiler=-fno-strict-aliasing",
|
||||
"--expt-extended-lambda",
|
||||
"--expt-relaxed-constexpr",
|
||||
"-forward-unknown-to-host-compiler",
|
||||
"--use_fast_math",
|
||||
"-std=c++20",
|
||||
"-O3",
|
||||
"-Xnvlink=--verbose",
|
||||
"-Xptxas=--verbose",
|
||||
"-Xptxas=--warn-on-spills",
|
||||
f"-I{tk_root}/include",
|
||||
f"-I{tk_root}/prototype",
|
||||
f"-I{python_include}",
|
||||
"-DTORCH_COMPILE",
|
||||
"-DKITTENS_HOPPER",
|
||||
"-arch=sm_90a",
|
||||
] + torch_includes
|
||||
|
||||
def get_extensions():
|
||||
if not torch.cuda.is_available():
|
||||
return []
|
||||
|
||||
extensions = []
|
||||
cpp_flags = ["-std=c++20", "-O3"]
|
||||
|
||||
# Check if TK is available
|
||||
if not os.path.exists(os.path.join(TK_DIR, "include", "kittens.cuh")):
|
||||
print(f"Warning: ThunderKittens not found at {TK_DIR}. CUDA kernels will not be built.")
|
||||
return []
|
||||
|
||||
cuda_flags = get_cuda_flags(TK_DIR)
|
||||
|
||||
# STA Extension
|
||||
extensions.append(CUDAExtension(
|
||||
"fastvideo_kernel._C.st_attn",
|
||||
sources=[
|
||||
"csrc/st_attn.cpp",
|
||||
"csrc/st_attn_h100.cu",
|
||||
],
|
||||
extra_compile_args={
|
||||
"cxx": cpp_flags + ["-DTK_COMPILE_ST_ATTN"],
|
||||
"nvcc": cuda_flags + ["-DTK_COMPILE_ST_ATTN"]
|
||||
},
|
||||
libraries=["cuda"],
|
||||
))
|
||||
|
||||
# VSA Extension
|
||||
extensions.append(CUDAExtension(
|
||||
"fastvideo_kernel._C.vsa",
|
||||
sources=[
|
||||
"csrc/vsa.cpp",
|
||||
"csrc/block_sparse_h100.cu",
|
||||
],
|
||||
extra_compile_args={
|
||||
"cxx": cpp_flags + ["-DTK_COMPILE_BLOCK_SPARSE"],
|
||||
"nvcc": cuda_flags + ["-DTK_COMPILE_BLOCK_SPARSE"]
|
||||
},
|
||||
libraries=["cuda"],
|
||||
))
|
||||
|
||||
return extensions
|
||||
|
||||
ext_modules = []
|
||||
if not any(arg in sys.argv for arg in ["clean", "egg_info", "--version"]):
|
||||
try:
|
||||
import torch
|
||||
ext_modules = get_extensions()
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to configure CUDA extensions: {e}")
|
||||
|
||||
setup(
|
||||
name="fastvideo-kernel",
|
||||
version="0.1.0",
|
||||
description="Unified CUDA kernels for FastVideo",
|
||||
long_description=open("README.md").read(),
|
||||
long_description_content_type="text/markdown",
|
||||
license="Apache-2.0",
|
||||
author="Hao AI Lab",
|
||||
url="https://github.com/hao-ai-lab/FastVideo",
|
||||
package_dir={"": "src"},
|
||||
packages=find_packages(where="src"),
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={"build_ext": BuildExtension} if ext_modules else {},
|
||||
python_requires=">=3.10",
|
||||
install_requires=["torch>=2.5.0", "triton>=2.0.0"],
|
||||
)
|
||||
@@ -0,0 +1,21 @@
|
||||
__version__ = "0.1.0"
|
||||
|
||||
from fastvideo_kernel.ops import (
|
||||
sliding_tile_attention,
|
||||
video_sparse_attn,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.vmoba import (
|
||||
moba_attn_varlen,
|
||||
process_moba_input,
|
||||
process_moba_output,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
"__version__",
|
||||
]
|
||||
@@ -0,0 +1,103 @@
|
||||
import math
|
||||
import torch
|
||||
from .triton_kernels.block_sparse_attn_triton import triton_block_sparse_attn_forward
|
||||
from .triton_kernels.index import map_to_index
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C.st_attn import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
try:
|
||||
from fastvideo_kernel._C.vsa import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
|
||||
def sliding_tile_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
window_size: list,
|
||||
text_length: int,
|
||||
has_text: bool = True,
|
||||
seq_shape: str = "30x48x80",
|
||||
) -> torch.Tensor:
|
||||
if sta_fwd is None:
|
||||
raise RuntimeError("STA kernel not compiled. Requires H100 and ThunderKittens at build time.")
|
||||
|
||||
seq_length = q.shape[2]
|
||||
shape_map = {"30x48x80": 1, "36x48x48": 2, "18x48x80": 3}
|
||||
|
||||
if has_text:
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
q = torch.cat([q, q[:, :, -pad_size:]], dim=2)
|
||||
k = torch.cat([k, k[:, :, -pad_size:]], dim=2)
|
||||
v = torch.cat([v, v[:, :, -pad_size:]], dim=2)
|
||||
|
||||
output = torch.empty_like(q)
|
||||
flag = shape_map[seq_shape]
|
||||
|
||||
for head_idx, (t, h, w) in enumerate(window_size):
|
||||
sta_fwd(
|
||||
q[:, head_idx:head_idx+1],
|
||||
k[:, head_idx:head_idx+1],
|
||||
v[:, head_idx:head_idx+1],
|
||||
output[:, head_idx:head_idx+1],
|
||||
t, h, w, text_length, False, has_text, flag
|
||||
)
|
||||
|
||||
if has_text:
|
||||
sta_fwd(q, k, v, output, 3, 3, 3, text_length, True, True, flag)
|
||||
|
||||
return output[:, :, :seq_length]
|
||||
|
||||
|
||||
def video_sparse_attn(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
topk: int,
|
||||
block_size: int | tuple = 64,
|
||||
compress_attn_weight: torch.Tensor = None,
|
||||
) -> torch.Tensor:
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
batch, heads, seq_len, dim = q.shape
|
||||
|
||||
# Compression branch
|
||||
q_c = q.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
k_c = k.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
v_c = v.view(batch, heads, seq_len // block_elements, block_elements, dim)
|
||||
|
||||
q_c = (q_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
|
||||
k_c = (k_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
|
||||
v_c = (v_c.float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
|
||||
|
||||
scores = torch.matmul(q_c, k_c.transpose(-2, -1)) / (dim ** 0.5)
|
||||
attn = torch.softmax(scores, dim=-1)
|
||||
out_c = torch.matmul(attn, v_c)
|
||||
|
||||
out_c = out_c.view(batch, heads, seq_len // block_elements, 1, dim)
|
||||
out_c = out_c.repeat(1, 1, 1, block_elements, 1).view(batch, heads, seq_len, dim)
|
||||
|
||||
# Sparse branch
|
||||
topk_idx = torch.topk(scores, topk, dim=-1).indices
|
||||
mask = torch.zeros_like(scores, dtype=torch.bool).scatter_(-1, topk_idx, True)
|
||||
|
||||
if block_sparse_fwd is not None:
|
||||
idx, num = map_to_index(mask)
|
||||
out_s, _ = block_sparse_fwd(q, k, v, idx, num, variable_block_sizes.int())
|
||||
else:
|
||||
idx, num = map_to_index(mask)
|
||||
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
return out_c + out_s
|
||||
+449
@@ -0,0 +1,449 @@
|
||||
"""
|
||||
Fused Attention
|
||||
===============
|
||||
|
||||
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
(https://tridao.me/publications/flash2/flash2.pdf)
|
||||
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
import math # small utility needed by the sparse wrapper
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
|
||||
# the code below and commenting out the equivalent parameters is convenient for
|
||||
# re-tuning.
|
||||
configs = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64]\
|
||||
for BN in [64]\
|
||||
for s in [3, 4, 7]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(Q, K, V, sm_scale, #
|
||||
q2k_index, q2k_num, max_kv_blks, #
|
||||
variable_block_sizes,
|
||||
M, Out, #
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vk, stride_vn,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
Z, H, N_CTX, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr):
|
||||
"""
|
||||
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
|
||||
(32×64 and 64×32) – memory footprint unchanged.
|
||||
"""
|
||||
|
||||
# ----- program-id mapping -----
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(1) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
# ----- base pointers -----
|
||||
qvk_off = (b.to(tl.int64) * stride_qz +
|
||||
h.to(tl.int64) * stride_qh)
|
||||
|
||||
Q_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
K_base = tl.make_block_ptr(
|
||||
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
|
||||
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_base = tl.make_block_ptr(
|
||||
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
|
||||
|
||||
O_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
# ----- accumulators -----
|
||||
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
q = tl.load(Q_ptr)
|
||||
|
||||
# ----- sparse loop over valid K/V tiles -----
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
|
||||
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
|
||||
|
||||
k = tl.load(K_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
# mask out invalid columns
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptr)
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
# ----- epilogue -----
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(M + off_hz * N_CTX + offs_m, m_i)
|
||||
tl.store(O_ptr, acc.to(Out.type.element_ty))
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO, #
|
||||
Delta, #
|
||||
Z, H, N_CTX, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
|
||||
):
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
# Filled in by the wrapper.
|
||||
start_n, start_m, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
|
||||
step_m = BLOCK_M1
|
||||
kv_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_N1
|
||||
meta_base = ((b * H + h) * q_tiles + kv_blk)
|
||||
|
||||
q_blocks = tl.load(k2q_num + meta_base) # int32
|
||||
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + kv_blk)
|
||||
|
||||
|
||||
|
||||
for blk_idx in range(q_blocks*2):
|
||||
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
|
||||
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
|
||||
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(tl.bfloat16)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(tl.bfloat16)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
# Increment pointers.
|
||||
return dk, dv
|
||||
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr,
|
||||
# Filled in by the wrapper.
|
||||
start_m, start_n, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N2)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
|
||||
step_n = BLOCK_N2
|
||||
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M2
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + q_blk)
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
|
||||
p = tl.where(mask[None, :], p , 0.0)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd(Q, K, V, sm_scale, #
|
||||
DO, #
|
||||
DQ, DK, DV, #
|
||||
M, D,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_z, stride_h, stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr):
|
||||
LN2 = 0.6931471824645996 # = ln(2)
|
||||
|
||||
bhid = tl.program_id(2)
|
||||
off_chz = (bhid * N_CTX).to(tl.int64)
|
||||
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
|
||||
pid = tl.program_id(0)
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += adj
|
||||
K += adj
|
||||
V += adj
|
||||
DO += adj
|
||||
DQ += adj
|
||||
DK += adj
|
||||
DV += adj
|
||||
M += off_chz
|
||||
D += off_chz
|
||||
|
||||
# load scales
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
start_n = pid * BLOCK_N1
|
||||
start_m = 0
|
||||
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
|
||||
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
|
||||
num_steps = N_CTX // BLOCK_M1
|
||||
|
||||
dk, dv = _attn_bwd_dkdv( #
|
||||
dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
|
||||
start_n, start_m, num_steps #
|
||||
)
|
||||
|
||||
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dv_ptrs, dv)
|
||||
|
||||
# Write back dK.
|
||||
dk *= sm_scale
|
||||
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dk_ptrs, dk)
|
||||
|
||||
# THIS BLOCK DOES DQ:
|
||||
start_m = pid * BLOCK_M2
|
||||
end_n = 0
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
|
||||
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
num_steps = N_CTX // BLOCK_N2
|
||||
dq = _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
|
||||
start_m, end_n, num_steps #
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
|
||||
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
|
||||
|
||||
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
q, k, v, sm_scale,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
M, o,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
B, H, T,
|
||||
HEAD_DIM=D, STAGE=3
|
||||
)
|
||||
|
||||
return o, M
|
||||
|
||||
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do, #
|
||||
delta, #
|
||||
BATCH, N_HEAD, N_CTX, #
|
||||
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
|
||||
)
|
||||
|
||||
|
||||
max_q_blks = k2q_index.shape[-1]
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
|
||||
N_HEAD, N_CTX, #
|
||||
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
|
||||
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
|
||||
HEAD_DIM=D #
|
||||
)
|
||||
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
@@ -0,0 +1,868 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
|
||||
|
||||
import random
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
except ImportError:
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
_flash_attn_varlen_backward = _unsupported
|
||||
flash_attn_varlen_func = _unsupported
|
||||
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def calc_chunks(cu_seqlen, moba_chunk_size):
|
||||
"""
|
||||
Calculate chunk boundaries.
|
||||
|
||||
For vision tasks we include all chunks (even the last one which might be shorter)
|
||||
so that every chunk can be selected.
|
||||
"""
|
||||
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
|
||||
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
|
||||
cu_num_chunk = torch.ones(
|
||||
batch_num_chunk.numel() + 1,
|
||||
device=cu_seqlen.device,
|
||||
dtype=batch_num_chunk.dtype,
|
||||
)
|
||||
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
|
||||
num_chunk = cu_num_chunk[-1]
|
||||
chunk_sizes = torch.full(
|
||||
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_sizes[0] = 0
|
||||
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
|
||||
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
|
||||
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
|
||||
chunk_to_batch = torch.zeros(
|
||||
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_to_batch[cu_num_chunk[1:-1]] = 1
|
||||
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
|
||||
|
||||
# Do not filter out any chunk
|
||||
filtered_chunk_indices = torch.arange(
|
||||
num_chunk, device=cu_seqlen.device, dtype=torch.int32
|
||||
)
|
||||
num_filtered_chunk = num_chunk
|
||||
|
||||
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
|
||||
|
||||
|
||||
# --- Threshold Selection Helper Functions ---
|
||||
|
||||
def _select_threshold_query_head(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects chunks for each <query, head> pair based on threshold.
|
||||
Normalization and sorting happen along the chunk dimension (dim=0).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per <head, query> (across chunks)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
row_min = gate_min_val.amin(dim=0) # (H, S)
|
||||
row_max = gate_masked.amax(dim=0) # (H, S)
|
||||
denom = row_max - row_min
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
|
||||
|
||||
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
|
||||
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
|
||||
|
||||
# 2) compute how much more normalized weight we need beyond self
|
||||
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
|
||||
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
|
||||
|
||||
# 3) zero out the self‐chunk in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0
|
||||
|
||||
# 4) sort the other chunks by descending norm, per <head,seq>
|
||||
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per <head,seq>
|
||||
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
|
||||
|
||||
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
|
||||
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
|
||||
any_cond = cond.any(dim=0) # (H, S)
|
||||
# Find the index of the first True value along dim 0. If none, use C-1.
|
||||
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff
|
||||
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
|
||||
|
||||
# 8) scatter it back to original chunk order
|
||||
others_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
others_mask.scatter_(0, sorted_idx, sorted_mask)
|
||||
|
||||
# 9) finally, include every self‐chunk plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_block(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <query, head> pairs for each block based on threshold.
|
||||
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
HS = H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per block (across heads and queries)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_denom = block_max - block_min
|
||||
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
|
||||
|
||||
gate_norm = (gate - block_min) / block_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights *per block*
|
||||
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
|
||||
|
||||
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
|
||||
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
|
||||
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort the other <head, seq> pairs by descending norm, per block
|
||||
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
|
||||
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per block
|
||||
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
|
||||
|
||||
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
|
||||
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
|
||||
any_cond = cond_flat.any(dim=1) # (C,)
|
||||
# Find the index of the first True value along dim 1. If none, use HS-1.
|
||||
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff per block
|
||||
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
|
||||
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
|
||||
|
||||
# 8) scatter it back to original <head, seq> order per block
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
|
||||
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_overall(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query, head> triplets globally based on threshold.
|
||||
Normalization and sorting happen across all valid entries.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
CHS = C * H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization globally across all valid entries
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
overall_max = gate_masked.max() # scalar
|
||||
overall_min = gate_min_val.min() # scalar
|
||||
overall_denom = overall_max - overall_min
|
||||
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
|
||||
|
||||
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights globally
|
||||
self_norm_sum_overall = self_norm_entries.sum() # scalar
|
||||
|
||||
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
|
||||
total_norm_sum_overall = gate_norm.sum() # scalar
|
||||
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort all other entries by descending norm, globally
|
||||
others_flat = others_norm.flatten() # (C*H*S,)
|
||||
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
|
||||
|
||||
# Only sort the valid 'other' entries
|
||||
valid_others_indices = torch.where(valid_others_mask_flat)[0]
|
||||
valid_others_values = others_flat[valid_others_indices]
|
||||
|
||||
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
|
||||
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
|
||||
|
||||
# 5) cumulative‑sum the sorted valid 'other' norms globally
|
||||
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
|
||||
|
||||
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
|
||||
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
|
||||
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
|
||||
any_cond = cond_values.any() # scalar
|
||||
|
||||
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
|
||||
cutoff_idx_in_sorted = torch.where(
|
||||
any_cond,
|
||||
cond_values.float().argmax(dim=0),
|
||||
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
|
||||
)
|
||||
|
||||
# 7) build a mask selecting the top-k others based on the cutoff
|
||||
# Select the original indices corresponding to the top entries in the sorted list
|
||||
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
|
||||
|
||||
# 8) create the mask in the original flat shape
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
|
||||
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
|
||||
others_mask_flat[selected_other_indices] = True
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_head_global(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query> globally for each head based on threshold.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# 1) LSE‐style normalization per head (across chunks and sequence dims)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
|
||||
|
||||
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
denom = max_per_head - min_per_head
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
|
||||
|
||||
gate_norm = (gate - min_per_head) / denom
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 2) sum normalized self‐chunk contributions per head
|
||||
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 3) total normalized sum per head
|
||||
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 4) how much more normalized weight needed per head
|
||||
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0)
|
||||
|
||||
# 5) zero out self‐chunk entries to focus on "others"
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
|
||||
|
||||
# 6) flatten chunk and sequence dims, per head
|
||||
CS = C * S
|
||||
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
|
||||
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
|
||||
# 7) vectorized selection of “others” per head
|
||||
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
|
||||
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
|
||||
|
||||
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
|
||||
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
|
||||
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
|
||||
|
||||
has_cutoff = cond.any(dim=1) # (H,)
|
||||
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
|
||||
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
|
||||
|
||||
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
|
||||
|
||||
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
|
||||
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
|
||||
|
||||
# 8) reshape selection mask back to (C, H, S)
|
||||
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
|
||||
|
||||
# 9) include self‐chunks plus selected others, and obey valid mask
|
||||
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
class MixedAttention(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
):
|
||||
ctx.max_seqlen = max_seqlen
|
||||
ctx.moba_chunk_size = moba_chunk_size
|
||||
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
|
||||
|
||||
# Non-causal self-attention branch
|
||||
# return out, softmax_lse, S_dmask, rng_state
|
||||
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
# MOBA attention branch (non-causal)
|
||||
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
|
||||
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
|
||||
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
|
||||
|
||||
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
output_2d = output.view(-1, q.shape[2])
|
||||
|
||||
max_lse_1d = self_attn_lse_sh.view(-1)
|
||||
max_lse_1d = max_lse_1d.index_reduce(
|
||||
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
|
||||
)
|
||||
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
|
||||
moba_attn_lse = (
|
||||
moba_attn_lse.view(-1)
|
||||
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
|
||||
.reshape_as(moba_attn_lse)
|
||||
)
|
||||
|
||||
mixed_attn_se_sh = self_attn_lse_sh.exp()
|
||||
moba_attn_se = moba_attn_lse.exp()
|
||||
|
||||
mixed_attn_se_sh.view(-1).index_add_(
|
||||
0, moba_q_sh_indices, moba_attn_se.view(-1)
|
||||
)
|
||||
mixed_attn_lse_sh = mixed_attn_se_sh.log()
|
||||
|
||||
# Combine self-attention output
|
||||
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
|
||||
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
|
||||
output_2d += self_attn_out_sh.reshape_as(output_2d)
|
||||
|
||||
# Combine MOBA attention output
|
||||
mixed_attn_lse = (
|
||||
mixed_attn_lse_sh.view(-1)
|
||||
.index_select(0, moba_q_sh_indices)
|
||||
.view_as(moba_attn_lse)
|
||||
)
|
||||
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
|
||||
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
|
||||
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
|
||||
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
|
||||
output = output.to(q.dtype)
|
||||
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
|
||||
ctx.save_for_backward(
|
||||
output,
|
||||
mixed_attn_lse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, d_output):
|
||||
|
||||
max_seqlen = ctx.max_seqlen
|
||||
moba_chunk_size = ctx.moba_chunk_size
|
||||
softmax_scale = ctx.softmax_scale
|
||||
|
||||
(
|
||||
output,
|
||||
mixed_attn_vlse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
) = ctx.saved_tensors
|
||||
|
||||
d_output = d_output.contiguous()
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_output,
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
out=output,
|
||||
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
|
||||
dq=dq,
|
||||
dk=dk,
|
||||
dv=dv,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
headdim = q.shape[-1]
|
||||
d_moba_output = (
|
||||
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
moba_output = (
|
||||
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
|
||||
mixed_attn_vlse = (
|
||||
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
|
||||
)
|
||||
|
||||
dmq = torch.empty_like(moba_q)
|
||||
dmkv = torch.empty_like(moba_kv)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_moba_output,
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
out=moba_output,
|
||||
softmax_lse=mixed_attn_vlse,
|
||||
dq=dmq,
|
||||
dk=dmkv[:,0],
|
||||
dv=dmkv[:,1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
|
||||
|
||||
|
||||
def moba_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int,
|
||||
moba_chunk_size: int,
|
||||
moba_topk: int,
|
||||
select_mode: str = 'threshold', # "topk" or "threshold"
|
||||
simsum_threshold: float = 0.25,
|
||||
threshold_type: str = 'query_head',
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Accelerated MOBA attention for vision tasks with proper LSE normalization.
|
||||
|
||||
This version:
|
||||
- Splits KV into chunks.
|
||||
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
|
||||
by amplifying the diagonal (self-chunk) logits.
|
||||
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
|
||||
reduction so that attending to each query over the selected chunks is equivalent
|
||||
to the original algorithm.
|
||||
"""
|
||||
# Stack keys and values.
|
||||
kv = torch.stack((k, v), dim=1)
|
||||
seqlen, num_head, head_dim = q.shape
|
||||
|
||||
# Compute chunk boundaries.
|
||||
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
|
||||
cu_seqlens, moba_chunk_size
|
||||
)
|
||||
|
||||
self_attn_cu_seqlen = cu_chunk
|
||||
|
||||
# Update top-k selection to include the self chunk.
|
||||
moba_topk = min(moba_topk, num_filtered_chunk)
|
||||
|
||||
# --- Build filtered KV from chunks ---
|
||||
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
|
||||
max_chunk_len = int(chunk_lengths.max().item())
|
||||
|
||||
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
|
||||
indices = chunk_starts.unsqueeze(1) + range_tensor
|
||||
indices = torch.clamp(indices, max=kv.shape[0] - 1)
|
||||
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
|
||||
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
|
||||
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
|
||||
|
||||
# Compute key_gate_weight over valid tokens.
|
||||
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
|
||||
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
|
||||
key_sum = (key_values * valid_mask_exp).sum(dim=1)
|
||||
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
|
||||
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
|
||||
|
||||
# Compute gate logits between key_gate_weight and queries.
|
||||
q_float = q.float()
|
||||
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
|
||||
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
|
||||
|
||||
# Amplify the diagonal (self chunk) contributions.
|
||||
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
|
||||
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
|
||||
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
|
||||
amplification_factor = 1e9 # Example factor; adjust as needed.
|
||||
origin_gate = gate.clone()
|
||||
gate = gate.clone()
|
||||
if select_mode == "topk":
|
||||
gate[gate_self_chunk_mask] += amplification_factor
|
||||
|
||||
# Exclude positions that are outside the valid batch boundaries.
|
||||
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
|
||||
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
|
||||
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
|
||||
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
|
||||
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
|
||||
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
|
||||
|
||||
if select_mode == 'topk':
|
||||
# We amplify self‐chunk in gate already, so self entries will rank highest.
|
||||
valid_gate_mask = gate != -float("inf")
|
||||
if threshold_type == 'query_head':
|
||||
# === per‐<head,seq> top-k across chunks (original behavior) ===
|
||||
# gate: (C, H, S)
|
||||
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
|
||||
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
gate_idx_mask.scatter_(0, gate_topk_idx, True)
|
||||
gate_mask = valid_gate_mask & gate_idx_mask
|
||||
elif threshold_type == 'overall':
|
||||
# === global top-k across all (chunk, head, seq) entries ===
|
||||
C, H, S = gate.shape
|
||||
flat_gate = gate.flatten()
|
||||
flat_mask = valid_gate_mask.flatten()
|
||||
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
|
||||
# pick topk global entries
|
||||
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
|
||||
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
|
||||
others_mask_flat[idx] = True
|
||||
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
|
||||
elif threshold_type == 'head_global':
|
||||
# per-head top-k across all chunks and sequence positions
|
||||
C, H, S = gate.shape
|
||||
CS = C * S
|
||||
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
|
||||
# pick top-k indices per head
|
||||
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
|
||||
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
|
||||
gate_idx_flat.scatter_(1, topk_idx, True)
|
||||
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid threshold_type for topk: {threshold_type}. "
|
||||
"Choose 'query_head', 'block', or 'overall'."
|
||||
)
|
||||
elif select_mode == 'threshold':
|
||||
# Delegate to the specific thresholding function
|
||||
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
|
||||
if threshold_type == 'query_head':
|
||||
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'block':
|
||||
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'overall':
|
||||
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'head_global':
|
||||
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
else:
|
||||
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
|
||||
else:
|
||||
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
|
||||
|
||||
# eliminate self_chunk in MoBA branch
|
||||
gate_mask = gate_mask & ~gate_self_chunk_mask
|
||||
# if gate_mask is all false, perform flash_attn instead
|
||||
if gate_mask.sum() == 0:
|
||||
return flash_attn_varlen_func(
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
|
||||
)
|
||||
|
||||
# Determine which query positions are selected.
|
||||
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
|
||||
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
|
||||
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
|
||||
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
|
||||
|
||||
# Build cumulative sequence lengths for the selected queries.
|
||||
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
|
||||
q_zero_mask = moba_seqlen_q == 0
|
||||
valid_expert_mask = ~q_zero_mask
|
||||
if q_zero_mask.sum() > 0:
|
||||
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
|
||||
moba_cu_seqlen_q = torch.cat(
|
||||
(
|
||||
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
|
||||
moba_seqlen_q.cumsum(dim=0),
|
||||
),
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
# Rearrange gathered KV for the MOBA branch.
|
||||
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
|
||||
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
|
||||
if q_zero_mask.sum() > 0:
|
||||
experts_tensor = experts_tensor[valid_expert_mask]
|
||||
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
|
||||
|
||||
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
|
||||
mask = seq_range < valid_expert_lengths.unsqueeze(1)
|
||||
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
|
||||
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
|
||||
|
||||
moba_cu_seqlen_kv = torch.cat(
|
||||
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
|
||||
valid_expert_lengths.cumsum(dim=0)],
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
assert (
|
||||
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
|
||||
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
|
||||
|
||||
return MixedAttention.apply(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
|
||||
def process_moba_input(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
"""
|
||||
Process inputs for the attention function.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
|
||||
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
|
||||
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed input tensor.
|
||||
"""
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
|
||||
else:
|
||||
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
|
||||
if len(chunk_size) == 2:
|
||||
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
|
||||
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
|
||||
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
|
||||
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
|
||||
elif len(chunk_size) == 3:
|
||||
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
|
||||
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
|
||||
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
|
||||
else:
|
||||
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
|
||||
|
||||
return x, moba_chunk_size
|
||||
|
||||
|
||||
def process_moba_output(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
pass
|
||||
elif len(chunk_size) == 2:
|
||||
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
|
||||
elif len(chunk_size) == 3:
|
||||
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# TEST
|
||||
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
|
||||
random.seed(0)
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
|
||||
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
|
||||
max_seqlen = q.shape[1]
|
||||
q = rearrange(q, "b s ... -> (b s) ...")
|
||||
k = rearrange(k, "b s ... -> (b s) ...")
|
||||
v = rearrange(v, "b s ... -> (b s) ...")
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
|
||||
"""Speed test comparing flash_attn vs moba_attention"""
|
||||
# Get data
|
||||
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
|
||||
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
|
||||
vo_grad = torch.randn_like(q)
|
||||
|
||||
# Warmup
|
||||
warmup_iters = 3
|
||||
perf_test_iters = 10
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_flash = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_moba = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
|
||||
|
||||
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
|
||||
print(f"Speedup: {time_flash / time_moba:.2f}x")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
CUDA_VISIBLE_DEVICES=1 \
|
||||
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
"""
|
||||
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
|
||||
@@ -0,0 +1,71 @@
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import BoolTensor, IntTensor
|
||||
from torch.nn.attention.flex_attention import create_block_mask
|
||||
|
||||
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
|
||||
torch._inductor.config.realize_opcount_threshold = 100
|
||||
|
||||
|
||||
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
|
||||
"""Generates a 3D NATTEN attention mask with a given kernel size.
|
||||
|
||||
Args:
|
||||
canvas_t: The time dimension of the canvas.
|
||||
canvas_h: The height of the canvas.
|
||||
canvas_w: The width of the canvas.
|
||||
kernel_t: The time dimension of the kernel.
|
||||
kernel_h: The height of the kernel.
|
||||
kernel_w: The width of the kernel.
|
||||
"""
|
||||
canvas_t, canvas_h, canvas_w = canvas_twh
|
||||
kernel_t, kernel_h, kernel_w = kernel_twh
|
||||
tile_t_size, tile_h_size, tile_w_size = tile_twh
|
||||
total_tile_size = tile_t_size * tile_h_size * tile_w_size
|
||||
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
|
||||
tile_id = idx // total_tile_size
|
||||
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
|
||||
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
|
||||
tile_w = tile_id % canvas_tile_w
|
||||
return tile_t, tile_h, tile_w
|
||||
|
||||
def sta_mask_3d(
|
||||
b: IntTensor,
|
||||
h: IntTensor,
|
||||
q_idx: IntTensor,
|
||||
kv_idx: IntTensor,
|
||||
) -> BoolTensor:
|
||||
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
|
||||
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
|
||||
# kernel nominally attempts to center itself on the query, but kernel center
|
||||
# is clamped to a fixed distance (kernel half-length) from the canvas edge
|
||||
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
|
||||
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
|
||||
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
|
||||
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
|
||||
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
|
||||
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
|
||||
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
|
||||
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
|
||||
|
||||
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return sta_mask_3d
|
||||
|
||||
|
||||
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
|
||||
img_seq_len = img_size[0] * img_size[1] * img_size[2]
|
||||
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
|
||||
mask = create_block_mask(image_mask,
|
||||
B=None,
|
||||
H=None,
|
||||
Q_LEN=img_seq_len + text_max_len,
|
||||
KV_LEN=img_seq_len + text_max_len,
|
||||
device=device,
|
||||
_compile=True)
|
||||
return mask
|
||||
@@ -0,0 +1,63 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
|
||||
# Local support import
|
||||
from .support_flex_sta import get_sliding_tile_attention_mask
|
||||
|
||||
# USE OUR NEW PACKAGE!
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
return output
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
# Using the same parameters as the original test
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=2):
|
||||
print(f"Running correctness check: batch={b}, heads={h}, seq_len={n}, dim={d}")
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
|
||||
for kernel_size in kernel_size_ls:
|
||||
print(f"Testing kernel_size: {kernel_size}")
|
||||
for xi in tqdm(range(num_iterations)):
|
||||
torch.manual_seed(xi)
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
|
||||
pt_o = flex_test(Q, K, V, kernel_size)
|
||||
|
||||
diff = pt_o - tk_o
|
||||
abs_diff = torch.abs(diff)
|
||||
max_d = torch.max(abs_diff).item()
|
||||
avg_d = torch.sum(abs_diff).item() / (b * h * n * d)
|
||||
|
||||
if max_d > 0.1:
|
||||
print(f"Warning: Large diff detected! max={max_d}, avg={avg_d}")
|
||||
|
||||
print("\n✅ TEST COMPLETE: New package matches FlexAttention behavior.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
check_correctness(b, h, n, d, causal, mean, std, num_iterations=2)
|
||||
@@ -0,0 +1,97 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from fastvideo_kernel.vmoba import moba_attn_varlen
|
||||
|
||||
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
|
||||
"""
|
||||
Generates random data for testing the variable-length attention function.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
random.seed(42)
|
||||
torch.cuda.manual_seed_all(42)
|
||||
|
||||
# Generate sequence lengths for each item in the batch
|
||||
if batch_size > 1:
|
||||
# Ensure sequence lengths are reasonably distributed
|
||||
avg_seqlen = total_seqlen // batch_size
|
||||
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
|
||||
remaining_len = total_seqlen - sum(seqlens)
|
||||
if remaining_len > 0:
|
||||
seqlens.append(remaining_len)
|
||||
else: # Adjust if sum exceeds total_seqlen
|
||||
seqlens.append(avg_seqlen)
|
||||
current_sum = sum(seqlens)
|
||||
seqlens[-1] -= (current_sum - total_seqlen)
|
||||
# Ensure all lengths are positive
|
||||
seqlens = [max(1, s) for s in seqlens]
|
||||
# Final adjustment to match total_seqlen
|
||||
seqlens[-1] += total_seqlen - sum(seqlens)
|
||||
|
||||
else:
|
||||
seqlens = [total_seqlen]
|
||||
|
||||
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
|
||||
max_seqlen = max(seqlens) if seqlens else 0
|
||||
|
||||
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [1, 2])
|
||||
@pytest.mark.parametrize("total_seqlen", [512, 1024])
|
||||
@pytest.mark.parametrize("num_heads", [8])
|
||||
@pytest.mark.parametrize("head_dim", [64])
|
||||
@pytest.mark.parametrize("moba_chunk_size", [64])
|
||||
@pytest.mark.parametrize("moba_topk", [2, 4])
|
||||
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
|
||||
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
def test_moba_attn_varlen_forward(
|
||||
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
|
||||
):
|
||||
"""
|
||||
Tests the forward pass of moba_attn_varlen for basic correctness.
|
||||
It checks output shape, dtype, and for the presence of NaNs/Infs.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
pytest.skip("float32 is not supported in flash attention")
|
||||
|
||||
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
|
||||
batch_size, total_seqlen, num_heads, head_dim, dtype
|
||||
)
|
||||
|
||||
# Ensure chunk size is not larger than the smallest sequence length
|
||||
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
|
||||
if moba_chunk_size > min_seqlen:
|
||||
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
|
||||
|
||||
try:
|
||||
output = moba_attn_varlen(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
moba_chunk_size=moba_chunk_size,
|
||||
moba_topk=moba_topk,
|
||||
select_mode=select_mode,
|
||||
threshold_type=threshold_type,
|
||||
simsum_threshold=0.5, # A reasonable default for threshold mode
|
||||
)
|
||||
except Exception as e:
|
||||
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
|
||||
|
||||
# 1. Check output shape
|
||||
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
|
||||
|
||||
# 2. Check output dtype
|
||||
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
|
||||
|
||||
# 3. Check for NaNs or Infs in the output
|
||||
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
|
||||
@@ -55,17 +55,10 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
# Install FastVideo Kernels
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -55,17 +55,10 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
# Install FastVideo Kernels
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -55,17 +55,10 @@ RUN source $HOME/.local/bin/env && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
# Install FastVideo Kernels
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
FROM rocm/pytorch:rocm7.1_ubuntu22.04_py3.10_pytorch_release_2.9.1
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject_other.toml ./pyproject.toml
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[rocm] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install FastVideo Kernels
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/fastvideo_kernel && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r docs/requirements-mkdocs.txt
|
||||
pip install -r requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -6,7 +6,7 @@ Thank you for your interest in contributing to FastVideo. We want to make the pr
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Profiling FastVideo
|
||||
|
||||
!!! warning
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
|
||||
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down inference.
|
||||
|
||||
## Profiling with PyTorch
|
||||
|
||||
@@ -49,5 +49,5 @@ Traces can be visualized using <https://ui.perfetto.dev/>.
|
||||
### Best Practices
|
||||
|
||||
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
|
||||
- After profiling, clean up trace directories to avoid filling disks.
|
||||
- After profiling, clean up trace directories to avoid filling disk storage.
|
||||
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
|
||||
|
||||
@@ -74,9 +74,11 @@ To add a new SSIM test, follow these steps:
|
||||
generator.generate_video(prompt, ...)
|
||||
|
||||
# Compare with Reference
|
||||
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
|
||||
assert ssim_values[0] >= 0.98 # Threshold
|
||||
```
|
||||
ssim_values = compute_video_ssim_torchvision(
|
||||
reference_path, generated_path, use_ms_ssim=True
|
||||
)
|
||||
assert ssim_values[0] >= 0.98 # Threshold
|
||||
```
|
||||
|
||||
4. **Reference Videos**:
|
||||
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# 🎯 Distillation
|
||||
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computations, enabling much faster video generation.
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
|
||||
@@ -7,6 +7,11 @@ Get up and running with FastVideo in minutes!
|
||||
First, install FastVideo:
|
||||
|
||||
```bash
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
@@ -15,36 +20,53 @@ pip install fastvideo
|
||||
### Text-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import FastVideoPipeline
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-t2v-1.3B")
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
# Generate a video
|
||||
prompt = "A cat playing with a ball of yarn"
|
||||
video = pipe(prompt, num_frames=16, height=512, width=512)
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import FastVideoPipeline
|
||||
from PIL import Image
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
# Load an image
|
||||
image = Image.open("input.jpg")
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-i2v-14B-480p")
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
|
||||
# Generate a video from the image
|
||||
video = pipe(image, num_frames=16, height=480, width=480)
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
## Next Steps
|
||||
|
||||
@@ -7,7 +7,7 @@ pip install st_attn
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently, we only have an implementation for H100s.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
|
||||
@@ -30,7 +30,7 @@ path_to_your_dataset_folder/
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
To generate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
|
||||
@@ -7,9 +7,9 @@ pip install vsa
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We support H100 (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
We support H100s (via ThunderKittens) and any other GPU (via Triton) for VSA.
|
||||
|
||||
First, install C++20 for ThunderKittens (if using H100):
|
||||
First, install C++20 for ThunderKittens (if using an H100):
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,8 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
|
||||
from dataclasses import dataclass
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
@@ -46,6 +47,29 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlashAttnMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
attn_mask: torch.Tensor | None = None
|
||||
|
||||
|
||||
class FlashAttnMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> FlashAttnMetadata:
|
||||
return FlashAttnMetadata(current_timestep=current_timestep,
|
||||
attn_mask=attn_mask)
|
||||
|
||||
|
||||
class FlashAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -66,12 +90,27 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
output = flash_attn_func(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
if attn_metadata is not None and hasattr(
|
||||
attn_metadata,
|
||||
"attn_mask") and attn_metadata.attn_mask is not None:
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import flash_attn_no_pad
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
attn_mask = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0),
|
||||
value=True)
|
||||
output = flash_attn_no_pad(qkv,
|
||||
attn_mask,
|
||||
causal=False,
|
||||
dropout_p=0,
|
||||
softmax_scale=None)
|
||||
else:
|
||||
output = flash_attn_func(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
return output
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata)
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -30,6 +31,29 @@ class SDPABackend(AttentionBackend):
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
@dataclass
|
||||
class SDPAMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
attn_mask: torch.Tensor | None = None
|
||||
|
||||
|
||||
class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
return SDPAMetadata(current_timestep=current_timestep,
|
||||
attn_mask=attn_mask)
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -51,14 +75,15 @@ class SDPAImpl(AttentionImpl):
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
attn_metadata: SDPAMetadata,
|
||||
) -> torch.Tensor:
|
||||
# transpose to bs, heads, seq_len, head_dim
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
attn_mask = attn_metadata.attn_mask if attn_metadata is not None else None
|
||||
attn_kwargs = {
|
||||
"attn_mask": None,
|
||||
"attn_mask": attn_mask,
|
||||
"dropout_p": self.dropout,
|
||||
"is_causal": self.causal,
|
||||
"scale": self.softmax_scale
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from st_attn import sliding_tile_attention
|
||||
from fastvideo_kernel import sliding_tile_attention
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
|
||||
@@ -6,7 +6,7 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
|
||||
try:
|
||||
from vsa import video_sparse_attn
|
||||
from fastvideo_kernel import video_sparse_attn
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ from dataclasses import dataclass
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from fastvideo_kernel import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
|
||||
@@ -108,6 +108,7 @@ class DistributedAttention(nn.Module):
|
||||
# Since mask is [batch, full_seq_len], it's already in the correct format
|
||||
|
||||
# LOAY TODO, instead of slicing repeatedly maintain an original qkv and rewrite into that
|
||||
valid_seq_len = None
|
||||
if attention_mask is not None:
|
||||
valid_seq_len = (attention_mask[0] == 1).sum().item()
|
||||
qkv = qkv[:, :valid_seq_len, :, :]
|
||||
@@ -140,8 +141,9 @@ class DistributedAttention(nn.Module):
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
if replicated_q is not None:
|
||||
replicated_output = output[:, seq_len * world_size:]
|
||||
output = output[:, :seq_len * world_size]
|
||||
split_idx = seq_len * world_size if valid_seq_len is None else valid_seq_len
|
||||
replicated_output = output[:, split_idx:]
|
||||
output = output[:, :split_idx]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
|
||||
#
|
||||
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
|
||||
# output and results there from are provided "AS IS" without any express or implied warranties of
|
||||
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
|
||||
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
|
||||
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
|
||||
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
|
||||
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
|
||||
# of rights and permissions under this agreement.
|
||||
# See the License for the specific language governing permissions and limitations under the License.
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
def flash_attn_no_pad(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
batch_size = qkv.shape[0]
|
||||
seqlen = qkv.shape[1]
|
||||
nheads = qkv.shape[-2]
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
|
||||
x, key_padding_mask)
|
||||
|
||||
x_unpad = rearrange(x_unpad,
|
||||
"nnz (three h d) -> nnz three h d",
|
||||
three=3,
|
||||
h=nheads)
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
output = rearrange(
|
||||
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices,
|
||||
batch_size, seqlen),
|
||||
"b s (h d) -> b s h d",
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def flash_attn_no_pad_v3(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
from flash_attn_interface import flash_attn_varlen_func as flash_attn_varlen_func_v3
|
||||
|
||||
if flash_attn_varlen_func_v3 is None:
|
||||
raise ImportError("FlashAttention V3 backend not available")
|
||||
|
||||
batch_size, seqlen, _, nheads, head_dim = qkv.shape
|
||||
query, key, value = qkv.unbind(dim=2)
|
||||
|
||||
query_unpad, indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
|
||||
rearrange(query, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
key_unpad, _, cu_seqlens_k, _, _ = unpad_input(
|
||||
rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
value_unpad, _, _, _, _ = unpad_input(
|
||||
rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
|
||||
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
|
||||
|
||||
output_unpad = flash_attn_varlen_func_v3(query_unpad,
|
||||
key_unpad,
|
||||
value_unpad,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_q,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic)
|
||||
|
||||
output = rearrange(pad_input(
|
||||
rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size,
|
||||
seqlen),
|
||||
"b s (h d) -> b s h d",
|
||||
h=nheads)
|
||||
return output
|
||||
@@ -1,10 +1,13 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
|
||||
"CosmosVideoConfig", "Cosmos25VideoConfig"
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
|
||||
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
|
||||
"LongCatVideoConfig"
|
||||
]
|
||||
|
||||
@@ -23,6 +23,8 @@ class DiTArchConfig(ArchConfig):
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
in_channels: int = 0
|
||||
out_channels: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_double_block(n: str, m) -> bool:
|
||||
return "double_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_refiner_block(n: str, m) -> bool:
|
||||
return "refiner" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_txt_in(n: str, m) -> bool:
|
||||
return n.split(".")[-1] == "txt_in"
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideo15ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block])
|
||||
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^context_embedder\.proj_in\.(.*)$":
|
||||
r"txt_in.input_embedder.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
|
||||
|
||||
# 2. txt_in_2 mapping:
|
||||
r"^context_embedder_2\.(.*)$":
|
||||
r"txt_in_2.\1",
|
||||
|
||||
# 3. x_embedder mapping:
|
||||
r"^x_embedder\.proj\.(.*)$":
|
||||
r"img_in.proj.\1",
|
||||
|
||||
# 4. Top-level time_text_embed mappings:
|
||||
r"^time_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"time_in.timestep_embedder.mlp.fc_in.\1",
|
||||
r"^time_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"time_in.timestep_embedder.mlp.fc_out.\1",
|
||||
r"^time_embed\.timestep_embedder_r\.linear_1\.(.*)$":
|
||||
r"time_in.timestep_embedder_r.mlp.fc_in.\1",
|
||||
r"^time_embed\.timestep_embedder_r\.linear_2\.(.*)$":
|
||||
r"time_in.timestep_embedder_r.mlp.fc_out.\1",
|
||||
|
||||
# 5. transformer_blocks mapping:
|
||||
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
|
||||
r"double_blocks.\1.img_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
|
||||
r"double_blocks.\1.txt_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 7. Final layers mapping:
|
||||
r"^norm_out\.linear\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^proj_out\.(.*)$":
|
||||
r"final_layer.linear.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: custom -> hf
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
in_channels: int = 65
|
||||
out_channels: int = 32
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128
|
||||
num_layers: int = 54
|
||||
num_refiner_layers: int = 2
|
||||
mlp_ratio: float = 4.0
|
||||
patch_size: int = 1
|
||||
patch_size_t: int = 1
|
||||
qk_norm: str = "rms_norm"
|
||||
text_embed_dim: int = 3584
|
||||
text_embed_2_dim: int = 1472
|
||||
image_embed_dim: int = 1152
|
||||
rope_theta: float = 256.0
|
||||
rope_axes_dim: tuple[int, ...] = (16, 56, 56)
|
||||
target_size: int = 640
|
||||
task_type: str = "i2v"
|
||||
use_meanflow: bool = False
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
|
||||
self.num_channels_latents: int = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideo15Config(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=HunyuanVideo15ArchConfig)
|
||||
|
||||
prefix: str = "Hunyuan15"
|
||||
@@ -0,0 +1,149 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat Video DiT configuration for native FastVideo implementation.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def is_longcat_blocks(n: str, m) -> bool:
|
||||
"""FSDP shard condition for LongCat transformer blocks."""
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatVideoArchConfig(DiTArchConfig):
|
||||
"""Architecture configuration for native LongCat Video DiT."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Enable torch.compile for transformer blocks (major speedup!)
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_longcat_blocks])
|
||||
|
||||
# Parameter name mapping for weight conversion
|
||||
# Maps original LongCat third_party names -> native FastVideo names
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Embedders
|
||||
r"^x_embedder\.(.*)$": r"patch_embed.\1",
|
||||
r"^t_embedder\.mlp\.0\.(.*)$": r"time_embedder.linear_1.\1",
|
||||
r"^t_embedder\.mlp\.2\.(.*)$": r"time_embedder.linear_2.\1",
|
||||
r"^y_embedder\.y_proj\.0\.(.*)$": r"caption_embedder.linear_1.\1",
|
||||
r"^y_embedder\.y_proj\.2\.(.*)$": r"caption_embedder.linear_2.\1",
|
||||
|
||||
# Transformer blocks - AdaLN modulation
|
||||
r"^blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
|
||||
r"blocks.\1.adaln_linear_1.\2",
|
||||
|
||||
# Transformer blocks - Normalization
|
||||
r"^blocks\.(\d+)\.mod_norm_attn\.(.*)$": r"blocks.\1.norm_attn.\2",
|
||||
r"^blocks\.(\d+)\.mod_norm_ffn\.(.*)$": r"blocks.\1.norm_ffn.\2",
|
||||
r"^blocks\.(\d+)\.pre_crs_attn_norm\.(.*)$":
|
||||
r"blocks.\1.norm_cross.\2",
|
||||
|
||||
# Self-attention: QKV fused -> separate (will need splitting in converter)
|
||||
# Original has attn.qkv.weight -> need to split into to_q, to_k, to_v
|
||||
r"^blocks\.(\d+)\.attn\.qkv\.(.*)$":
|
||||
r"blocks.\1.self_attn.qkv_fused.\2", # Marker for splitting
|
||||
r"^blocks\.(\d+)\.attn\.proj\.(.*)$":
|
||||
r"blocks.\1.self_attn.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn\.q_norm\.(.*)$":
|
||||
r"blocks.\1.self_attn.q_norm.\2",
|
||||
r"^blocks\.(\d+)\.attn\.k_norm\.(.*)$":
|
||||
r"blocks.\1.self_attn.k_norm.\2",
|
||||
|
||||
# Cross-attention
|
||||
r"^blocks\.(\d+)\.cross_attn\.q_linear\.(.*)$":
|
||||
r"blocks.\1.cross_attn.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.kv_linear\.(.*)$":
|
||||
r"blocks.\1.cross_attn.kv_fused.\2", # Marker for splitting
|
||||
r"^blocks\.(\d+)\.cross_attn\.proj\.(.*)$":
|
||||
r"blocks.\1.cross_attn.to_out.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q_norm\.(.*)$":
|
||||
r"blocks.\1.cross_attn.q_norm.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k_norm\.(.*)$":
|
||||
r"blocks.\1.cross_attn.k_norm.\2",
|
||||
|
||||
# FFN (SwiGLU)
|
||||
r"^blocks\.(\d+)\.ffn\.w1\.(.*)$": r"blocks.\1.ffn.w1.\2", # gate
|
||||
r"^blocks\.(\d+)\.ffn\.w2\.(.*)$": r"blocks.\1.ffn.w2.\2", # down
|
||||
r"^blocks\.(\d+)\.ffn\.w3\.(.*)$": r"blocks.\1.ffn.w3.\2", # up
|
||||
|
||||
# Final layer
|
||||
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
|
||||
r"final_layer.adaln_linear.\1",
|
||||
r"^final_layer\.norm_final\.(.*)$": r"final_layer.norm.\1",
|
||||
r"^final_layer\.linear\.(.*)$": r"final_layer.proj.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: custom -> hf
|
||||
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# LoRA parameter name mapping
|
||||
lora_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Model architecture parameters
|
||||
hidden_size: int = 4096
|
||||
depth: int = 48 # Number of transformer blocks
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128 # hidden_size / num_attention_heads
|
||||
|
||||
in_channels: int = 16 # Latent space channels
|
||||
out_channels: int = 16
|
||||
num_channels_latents: int = 16
|
||||
|
||||
# Patch embedding
|
||||
patch_size: tuple[int, int,
|
||||
int] = (1, 2, 2) # [T, H, W] - no temporal compression
|
||||
|
||||
# Text/caption embedding
|
||||
caption_channels: int = 4096 # UMT5 d_model
|
||||
|
||||
# Timestep embedding
|
||||
adaln_tembed_dim: int = 512
|
||||
frequency_embedding_size: int = 256
|
||||
|
||||
# FFN
|
||||
mlp_ratio: int = 4
|
||||
|
||||
# Attention backend support
|
||||
_supported_attention_backends: tuple = field(default_factory=lambda: (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
))
|
||||
|
||||
# Text padding behavior
|
||||
text_tokens_zero_pad: bool = True
|
||||
|
||||
# Block Sparse Attention (BSA)
|
||||
enable_bsa: bool = False
|
||||
bsa_params: dict | None = field(
|
||||
default_factory=lambda: {
|
||||
"sparsity": 0.9375,
|
||||
"cdf_threshold": None,
|
||||
"chunk_3d_shape_q": [4, 4, 4],
|
||||
"chunk_3d_shape_k": [4, 4, 4],
|
||||
})
|
||||
|
||||
# LoRA exclusions
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
# Ensure attention_head_dim matches
|
||||
self.attention_head_dim = self.hidden_size // self.num_attention_heads
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatVideoConfig(DiTConfig):
|
||||
"""Main configuration for LongCat Video DiT."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=LongCatVideoArchConfig)
|
||||
|
||||
prefix: str = "longcat"
|
||||
@@ -6,9 +6,11 @@ from fastvideo.configs.models.encoders.clip import (
|
||||
CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig"
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
|
||||
"Qwen2_5_VLConfig"
|
||||
]
|
||||
|
||||
@@ -72,6 +72,7 @@ class EncoderConfig(ModelConfig):
|
||||
@dataclass
|
||||
class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||
is_chat_model: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("embed_tokens")
|
||||
|
||||
|
||||
def _is_final_norm(n: str, m) -> bool:
|
||||
return n.endswith("norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Qwen2_5_VLArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 152064
|
||||
hidden_size: int = 8192
|
||||
intermediate_size: int = 29568
|
||||
num_hidden_layers: int = 80
|
||||
num_attention_heads: int = 64
|
||||
num_key_value_heads: int = 8
|
||||
hidden_act: str = "silu"
|
||||
max_position_embeddings: int = 32768
|
||||
initializer_range: float = 0.02
|
||||
rms_norm_eps: float = 1e-05
|
||||
use_cache: bool = True
|
||||
tie_word_embeddings: bool = False
|
||||
rope_theta: float = 1000000.0
|
||||
use_sliding_window: bool = False
|
||||
sliding_window: int | None = 4096
|
||||
max_window_layers: int = 80
|
||||
layer_types: list = field(default_factory=list)
|
||||
attention_dropout: float = 0.0
|
||||
rope_scaling: dict | None = None
|
||||
bos_token_id: int | None = None
|
||||
eos_token_id: int | None = None
|
||||
pad_token_id: int | None = None
|
||||
vision_token_id: int = 151654
|
||||
model_type: str = "qwen2_5_vl_text"
|
||||
dtype: str = "bfloat16"
|
||||
|
||||
stacked_params_mapping: list[tuple[str, str, str
|
||||
| int]] = field(default_factory=lambda: [
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0),
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_norm])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.sliding_window = self.sliding_window if self.use_sliding_window else None
|
||||
# for backward compatibility
|
||||
if self.num_key_value_heads is None:
|
||||
self.num_key_value_heads = self.num_attention_heads
|
||||
if self.layer_types is None:
|
||||
self.layer_types = [
|
||||
"sliding_attention" if self.sliding_window is not None
|
||||
and i >= self.max_window_layers else "full_attention"
|
||||
for i in range(self.num_hidden_layers)
|
||||
]
|
||||
if self.rope_scaling is not None and "type" in self.rope_scaling:
|
||||
if self.rope_scaling["type"] == "mrope":
|
||||
self.rope_scaling["type"] = "default"
|
||||
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
|
||||
|
||||
self.tokenizer_kwargs = {
|
||||
"add_generation_prompt": True,
|
||||
"tokenize": True,
|
||||
"return_dict": True,
|
||||
"padding": "max_length",
|
||||
"max_length": 1000 + 108,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Qwen2_5_VLConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=Qwen2_5_VLArchConfig)
|
||||
prefix: str = "qwen2_5_vl"
|
||||
is_chat_model: bool = True
|
||||
@@ -40,6 +40,8 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
eos_token_id: int = 1
|
||||
classifier_dropout: float = 0.0
|
||||
text_len: int = 512
|
||||
dtype: str | None = None
|
||||
gradient_checkpointing: bool = False
|
||||
stacked_params_mapping: list[tuple[str, str,
|
||||
str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
@@ -68,6 +70,7 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
self.hidden_size = self.d_model
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
@@ -8,4 +9,5 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"StepVideoVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15VAEArchConfig(VAEArchConfig):
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 32
|
||||
block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024)
|
||||
layers_per_block: int = 2
|
||||
spatial_compression_ratio: int = 16
|
||||
temporal_compression_ratio: int = 4
|
||||
downsample_match_channel: bool = True
|
||||
upsample_match_channel: bool = True
|
||||
scaling_factor: float = 1.03682
|
||||
|
||||
def __post_init__(self):
|
||||
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
|
||||
1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15VAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=Hunyuan15VAEArchConfig)
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.pipelines.base import (PipelineConfig,
|
||||
SlidingTileAttnConfig)
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
@@ -11,8 +12,8 @@ from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"SelfForcingWanT2V480PConfig", "CosmosConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
import re
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
|
||||
Qwen2_5_VLConfig, T5Config)
|
||||
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
PROMPT_TEMPLATE_TOKEN_LENGTH = 108
|
||||
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = "You are a helpful assistant. Describe the video by detailing the following aspects: \
|
||||
1. The main content and theme of the video. \
|
||||
2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects. \
|
||||
3. Actions, events, behaviors temporal relationships, physical movement changes of the objects. \
|
||||
4. background environment, light, style and atmosphere. \
|
||||
5. camera angles, movements, and transitions used in the video."
|
||||
|
||||
|
||||
def extract_glyph_texts(prompt: str) -> str | None:
|
||||
"""
|
||||
Extract glyph texts from prompt using regex pattern.
|
||||
|
||||
Args:
|
||||
prompt: Input prompt string
|
||||
|
||||
Returns:
|
||||
List of extracted glyph texts
|
||||
"""
|
||||
pattern = r"\"(.*?)\"|“(.*?)”"
|
||||
matches = re.findall(pattern, prompt)
|
||||
result = [match[0] or match[1] for match in matches]
|
||||
result = list(dict.fromkeys(result)) if len(result) > 1 else result
|
||||
|
||||
if result:
|
||||
formatted_result = ". ".join([f'Text "{text}"'
|
||||
for text in result]) + ". "
|
||||
else:
|
||||
formatted_result = None
|
||||
|
||||
return formatted_result
|
||||
|
||||
|
||||
def format_text_input(prompt: str, system_message: str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Apply text to template.
|
||||
|
||||
Args:
|
||||
prompt (List[str]): Input text.
|
||||
system_message (str): System message.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: List of chat conversation.
|
||||
"""
|
||||
|
||||
template = [{
|
||||
"role": "system",
|
||||
"content": system_message
|
||||
}, {
|
||||
"role": "user",
|
||||
"content": prompt if prompt else " "
|
||||
}]
|
||||
|
||||
return template
|
||||
|
||||
|
||||
def qwen_preprocess_text(prompt: str) -> list[dict[str, Any]]:
|
||||
output = format_text_input(prompt, PROMPT_TEMPLATE_ENCODE_VIDEO)
|
||||
return output
|
||||
|
||||
|
||||
def qwen_postprocess_text(
|
||||
outputs: BaseEncoderOutput,
|
||||
mask: torch.tensor) -> tuple[torch.tensor, torch.tensor]:
|
||||
assert outputs.hidden_states is not None
|
||||
output = outputs.hidden_states[-3]
|
||||
output = output[:, PROMPT_TEMPLATE_TOKEN_LENGTH:]
|
||||
mask = mask[:, PROMPT_TEMPLATE_TOKEN_LENGTH:]
|
||||
return output, mask
|
||||
|
||||
|
||||
def byt5_preprocess_text(prompt: str) -> str | None:
|
||||
prompts = [prompt] if isinstance(prompt, str) else prompt
|
||||
glyph_texts = [extract_glyph_texts(p) for p in prompts]
|
||||
return glyph_texts[0]
|
||||
|
||||
|
||||
def byt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15T2V480PConfig(PipelineConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=HunyuanVideo15Config)
|
||||
# VAE
|
||||
vae_config: VAEConfig = field(default_factory=Hunyuan15VAEConfig)
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (Qwen2_5_VLConfig(), T5Config()))
|
||||
preprocess_text_funcs: tuple[Callable[[Any], Any], ...] = field(
|
||||
default_factory=lambda: (qwen_preprocess_text, byt5_preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
|
||||
default_factory=lambda: (qwen_postprocess_text, byt5_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", "fp32"))
|
||||
text_encoder_crop_start: int = PROMPT_TEMPLATE_TOKEN_LENGTH
|
||||
text_encoder_max_lengths: tuple[int, ...] = field(
|
||||
default_factory=lambda: (1000 + PROMPT_TEMPLATE_TOKEN_LENGTH, 256))
|
||||
|
||||
vae_tiling: bool = True
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
flow_shift: int = 9
|
||||
@@ -0,0 +1,355 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
import html
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatDiTArchConfig(DiTArchConfig):
|
||||
"""Extended DiTArchConfig with LongCat-specific fields.
|
||||
|
||||
NOTE: This is for Phase 1 wrapper compatibility. For native model (Phase 2),
|
||||
use LongCatVideoConfig from fastvideo.configs.models.dits.longcat instead.
|
||||
"""
|
||||
# LongCat-specific architecture parameters
|
||||
adaln_tembed_dim: int = 512
|
||||
caption_channels: int = 4096
|
||||
depth: int = 48
|
||||
enable_bsa: bool = False
|
||||
enable_flashattn3: bool = False
|
||||
enable_flashattn2: bool = True
|
||||
enable_xformers: bool = False
|
||||
frequency_embedding_size: int = 256
|
||||
in_channels: int = 16
|
||||
mlp_ratio: int = 4
|
||||
num_heads: int = 32
|
||||
out_channels: int = 16
|
||||
text_tokens_zero_pad: bool = True
|
||||
patch_size: list[int] = field(default_factory=lambda: [1, 2, 2])
|
||||
cp_split_hw: list[int] | None = None
|
||||
bsa_params: dict | None = None
|
||||
|
||||
|
||||
def longcat_preprocess_text(prompt: str) -> str:
|
||||
"""Clean and preprocess text like original LongCat implementation.
|
||||
|
||||
This function applies the same text cleaning pipeline as the original
|
||||
LongCat-Video implementation to ensure identical tokenization results.
|
||||
|
||||
Steps:
|
||||
1. basic_clean: Fix unicode issues and unescape HTML entities
|
||||
2. whitespace_clean: Normalize whitespace to single spaces
|
||||
|
||||
Args:
|
||||
prompt: Raw input text prompt
|
||||
|
||||
Returns:
|
||||
Cleaned and normalized text prompt
|
||||
"""
|
||||
# basic_clean: fix unicode and HTML entities
|
||||
text = ftfy.fix_text(prompt)
|
||||
text = html.unescape(html.unescape(text))
|
||||
text = text.strip()
|
||||
|
||||
# whitespace_clean: normalize whitespace
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
text = text.strip()
|
||||
|
||||
return text
|
||||
|
||||
|
||||
def umt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""
|
||||
Postprocess UMT5/T5 encoder outputs to fixed length 512 embeddings.
|
||||
"""
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
assert torch.isnan(hidden_state).sum() == 0
|
||||
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT2V480PConfig(PipelineConfig):
|
||||
"""Configuration for LongCat pipeline (480p) aligned to LongCat-Video modules.
|
||||
|
||||
Components expected by loaders:
|
||||
- tokenizer: AutoTokenizer
|
||||
- text_encoder: UMT5EncoderModel
|
||||
- transformer: LongCatVideoTransformer3DModel (Phase 1 wrapper)
|
||||
OR LongCatTransformer3DModel (Phase 2 native)
|
||||
- vae: AutoencoderKLWan (Wan VAE, 4x8 compression)
|
||||
- scheduler: FlowMatchEulerDiscreteScheduler
|
||||
"""
|
||||
|
||||
# DiT config with LongCat-specific arch_config
|
||||
# NOTE: For Phase 1 wrapper, uses LongCatDiTArchConfig
|
||||
# For Phase 2 native model, can use LongCatVideoConfig directly
|
||||
dit_config: DiTConfig = field(
|
||||
default_factory=lambda: DiTConfig(arch_config=LongCatDiTArchConfig()))
|
||||
|
||||
# VAE config: Wan VAE with encoder+decoder enabled
|
||||
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Precision defaults
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", ))
|
||||
|
||||
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
|
||||
text_encoder_configs: tuple[T5Config, ...] = field(
|
||||
default_factory=lambda: (T5Config(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (longcat_preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(umt5_postprocess_text, ))
|
||||
|
||||
# LongCat-specific runtime toggles (consumed by pipeline/stages)
|
||||
enable_kv_cache: bool = True
|
||||
offload_kv_cache: bool = False
|
||||
enable_bsa: bool = False
|
||||
use_distill: bool = False
|
||||
enhance_hf: bool = False
|
||||
# Optional BSA parameter dict (kept for backward/phase-1 compatibility).
|
||||
# `LongCatPipeline.initialize_pipeline()` uses this as a base and then applies
|
||||
# CLI overrides (bsa_sparsity / bsa_chunk_{q,k} / bsa_cdf_threshold).
|
||||
bsa_params: dict | None = None
|
||||
# BSA runtime overrides (preferred over bsa_params if provided via CLI)
|
||||
bsa_sparsity: float | None = None
|
||||
bsa_cdf_threshold: float | None = None
|
||||
bsa_chunk_q: list[int] | None = None
|
||||
bsa_chunk_k: list[int] | None = None
|
||||
t_thresh: float | None = None # refine stage default controlled by sampling args
|
||||
|
||||
# LongCat does not need flow_shift
|
||||
flow_shift: float | None = None
|
||||
dmd_denoising_steps: list[int] | None = None
|
||||
|
||||
def __post_init__(self):
|
||||
# LongCat inference requires vae encoder and decoder
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatT2V704PConfig(LongCatT2V480PConfig):
|
||||
"""Configuration for LongCat pipeline (704p) with BSA enabled by default.
|
||||
|
||||
Uses the same resolution and BSA parameters as original LongCat refinement stage.
|
||||
BSA parameters configured in transformer config.json with chunk_3d_shape=[4,4,4]:
|
||||
- Input: 704×1280×96
|
||||
- VAE (8x): 88×160×96
|
||||
- Patch [1,2,2]: 44×80×96
|
||||
- chunk [4,4,4]: 96%4=0, 44%4=0, 80%4=0 ✅
|
||||
|
||||
This configuration matches the original LongCat refinement stage parameters.
|
||||
"""
|
||||
|
||||
# Enable BSA by default for 704p
|
||||
enable_bsa: bool = True
|
||||
|
||||
|
||||
ASPECT_RATIO_627 = {
|
||||
'0.26': ([320, 1216], 1),
|
||||
'0.31': ([352, 1120], 1),
|
||||
'0.38': ([384, 1024], 1),
|
||||
'0.43': ([416, 960], 1),
|
||||
'0.52': ([448, 864], 1),
|
||||
'0.58': ([480, 832], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'0.74': ([544, 736], 1),
|
||||
'0.86': ([576, 672], 1),
|
||||
'0.95': ([608, 640], 1),
|
||||
'1.05': ([640, 608], 1),
|
||||
'1.17': ([672, 576], 1),
|
||||
'1.29': ([704, 544], 1),
|
||||
'1.35': ([736, 544], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'1.67': ([800, 480], 1),
|
||||
'1.73': ([832, 480], 1),
|
||||
'2.00': ([896, 448], 1),
|
||||
'2.31': ([960, 416], 1),
|
||||
'2.58': ([992, 384], 1),
|
||||
'2.75': ([1056, 384], 1),
|
||||
'3.09': ([1088, 352], 1),
|
||||
'3.70': ([1184, 320], 1),
|
||||
'3.80': ([1216, 320], 1),
|
||||
'3.90': ([1248, 320], 1),
|
||||
'4.00': ([1280, 320], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_627_F64 = {
|
||||
'0.26': ([320, 1216], 1),
|
||||
'0.38': ([384, 1024], 1),
|
||||
'0.50': ([448, 896], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'0.82': ([576, 704], 1),
|
||||
'1.00': ([640, 640], 1),
|
||||
'1.22': ([704, 576], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'1.86': ([832, 448], 1),
|
||||
'2.00': ([896, 448], 1),
|
||||
'2.50': ([960, 384], 1),
|
||||
'2.83': ([1088, 384], 1),
|
||||
'3.60': ([1152, 320], 1),
|
||||
'3.80': ([1216, 320], 1),
|
||||
'4.00': ([1280, 320], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_627_F128 = {
|
||||
'0.25': ([256, 1024], 1),
|
||||
'0.38': ([384, 1024], 1),
|
||||
'0.43': ([384, 896], 1),
|
||||
'0.57': ([512, 896], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'1.00': ([640, 640], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'1.75': ([896, 512], 1),
|
||||
'2.33': ([896, 384], 1),
|
||||
'2.67': ([1024, 384], 1),
|
||||
'4.00': ([1024, 256], 1),
|
||||
}
|
||||
|
||||
ASPECT_RATIO_627_F256 = {
|
||||
'0.25': ([256, 1024], 1),
|
||||
'0.33': ([256, 768], 1),
|
||||
'0.50': ([256, 512], 1),
|
||||
'0.67': ([512, 768], 1),
|
||||
'1.00': ([512, 512], 1),
|
||||
'1.50': ([768, 512], 1),
|
||||
'2.00': ([512, 256], 1),
|
||||
'3.00': ([768, 256], 1),
|
||||
'4.00': ([1024, 256], 1),
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960 = {
|
||||
'0.25': ([480, 1920], 1),
|
||||
'0.29': ([512, 1792], 1),
|
||||
'0.32': ([544, 1696], 1),
|
||||
'0.36': ([576, 1600], 1),
|
||||
'0.40': ([608, 1504], 1),
|
||||
'0.49': ([672, 1376], 1),
|
||||
'0.54': ([704, 1312], 1),
|
||||
'0.59': ([736, 1248], 1),
|
||||
'0.69': ([800, 1152], 1),
|
||||
'0.74': ([832, 1120], 1),
|
||||
'0.82': ([864, 1056], 1),
|
||||
'0.88': ([896, 1024], 1),
|
||||
'0.94': ([928, 992], 1),
|
||||
'1.00': ([960, 960], 1),
|
||||
'1.07': ([992, 928], 1),
|
||||
'1.14': ([1024, 896], 1),
|
||||
'1.22': ([1056, 864], 1),
|
||||
'1.31': ([1088, 832], 1),
|
||||
'1.35': ([1120, 832], 1),
|
||||
'1.44': ([1152, 800], 1),
|
||||
'1.70': ([1248, 736], 1),
|
||||
'2.00': ([1344, 672], 1),
|
||||
'2.05': ([1376, 672], 1),
|
||||
'2.47': ([1504, 608], 1),
|
||||
'2.53': ([1536, 608], 1),
|
||||
'2.83': ([1632, 576], 1),
|
||||
'3.06': ([1664, 544], 1),
|
||||
'3.12': ([1696, 544], 1),
|
||||
'3.62': ([1856, 512], 1),
|
||||
'3.93': ([1888, 480], 1),
|
||||
'4.00': ([1920, 480], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960_F64 = {
|
||||
'0.22': ([448, 2048], 1),
|
||||
'0.29': ([512, 1792], 1),
|
||||
'0.36': ([576, 1600], 1),
|
||||
'0.45': ([640, 1408], 1),
|
||||
'0.55': ([704, 1280], 1),
|
||||
'0.63': ([768, 1216], 1),
|
||||
'0.76': ([832, 1088], 1),
|
||||
'0.88': ([896, 1024], 1),
|
||||
'1.00': ([960, 960], 1),
|
||||
'1.14': ([1024, 896], 1),
|
||||
'1.31': ([1088, 832], 1),
|
||||
'1.50': ([1152, 768], 1),
|
||||
'1.58': ([1216, 768], 1),
|
||||
'1.82': ([1280, 704], 1),
|
||||
'1.91': ([1344, 704], 1),
|
||||
'2.20': ([1408, 640], 1),
|
||||
'2.30': ([1472, 640], 1),
|
||||
'2.67': ([1536, 576], 1),
|
||||
'2.89': ([1664, 576], 1),
|
||||
'3.62': ([1856, 512], 1),
|
||||
'3.75': ([1920, 512], 1)
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960_F128 = {
|
||||
'0.20': ([384, 1920], 1),
|
||||
'0.27': ([512, 1920], 1),
|
||||
'0.33': ([512, 1536], 1),
|
||||
'0.42': ([640, 1536], 1),
|
||||
'0.50': ([640, 1280], 1),
|
||||
'0.60': ([768, 1280], 1),
|
||||
'0.67': ([768, 1152], 1),
|
||||
'0.78': ([896, 1152], 1),
|
||||
'1.00': ([1024, 1024], 1),
|
||||
'1.29': ([1152, 896], 1),
|
||||
'1.50': ([1152, 768], 1),
|
||||
'1.67': ([1280, 768], 1),
|
||||
'2.00': ([1280, 640], 1),
|
||||
'2.40': ([1536, 640], 1),
|
||||
'3.00': ([1536, 512], 1),
|
||||
'3.75': ([1920, 512], 1),
|
||||
'5.00': ([1920, 384], 1),
|
||||
}
|
||||
|
||||
ASPECT_RATIO_960_F256 = {
|
||||
'0.33': ([512, 1536], 1),
|
||||
'0.60': ([768, 1280], 1),
|
||||
'1.00': ([1024, 1024], 1),
|
||||
'1.67': ([1280, 768], 1),
|
||||
'3.00': ([1536, 512], 1),
|
||||
}
|
||||
|
||||
|
||||
def get_bucket_config(resolution, scale_factor_spatial):
|
||||
if resolution == '480p':
|
||||
if scale_factor_spatial == 16 or scale_factor_spatial == 32:
|
||||
return ASPECT_RATIO_627
|
||||
elif scale_factor_spatial == 64:
|
||||
return ASPECT_RATIO_627_F64
|
||||
elif scale_factor_spatial == 128:
|
||||
return ASPECT_RATIO_627_F128
|
||||
elif scale_factor_spatial == 256:
|
||||
return ASPECT_RATIO_627_F256
|
||||
elif resolution == '720p':
|
||||
if scale_factor_spatial == 16 or scale_factor_spatial == 32:
|
||||
return ASPECT_RATIO_960
|
||||
elif scale_factor_spatial == 64:
|
||||
return ASPECT_RATIO_960_F64
|
||||
elif scale_factor_spatial == 128:
|
||||
return ASPECT_RATIO_960_F128
|
||||
elif scale_factor_spatial == 256:
|
||||
return ASPECT_RATIO_960_F256
|
||||
|
||||
raise ValueError(
|
||||
f"Unsupported resolution '{resolution}' or scale_factor_spatial '{scale_factor_spatial}'"
|
||||
)
|
||||
@@ -7,7 +7,9 @@ from collections.abc import Callable
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
@@ -27,6 +29,10 @@ logger = init_logger(__name__)
|
||||
PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
|
||||
Hunyuan15T2V480PConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
|
||||
Hunyuan15T2V720PConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
|
||||
@@ -56,6 +62,8 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan":
|
||||
lambda id: "hunyuan" in id.lower(),
|
||||
"hunyuan15":
|
||||
lambda id: "hunyuan15" in id.lower(),
|
||||
"matrixgame":
|
||||
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
|
||||
"wanpipeline":
|
||||
@@ -70,14 +78,19 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
lambda id: "stepvideo" in id.lower(),
|
||||
"cosmos":
|
||||
lambda id: "cosmos" in id.lower(),
|
||||
"longcat":
|
||||
lambda id: "longcat" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"longcat": LongCatT2V480PConfig,
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"matrixgame": MatrixGameI2V480PConfig,
|
||||
"hunyuan15":
|
||||
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
|
||||
"wanpipeline":
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
@@ -123,6 +136,7 @@ def get_pipeline_config_cls_from_name(
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
|
||||
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
|
||||
return pipeline_config_cls
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
|
||||
|
||||
@@ -3,6 +3,7 @@ from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -27,6 +28,17 @@ class SamplingParam:
|
||||
keyboard_cond: Any | None = None # Shape: (B, T, K)
|
||||
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
|
||||
|
||||
# Refine inputs (LongCat 480p->720p upscaling)
|
||||
# Path-based refine (load stage1 video from disk, e.g. MP4)
|
||||
refine_from: str | None = None # Path to stage1 video (480p output from distill)
|
||||
t_thresh: float = 0.5 # Threshold for timestep scheduling in refinement
|
||||
spatial_refine_only: bool = False # If True, only spatial (no temporal doubling)
|
||||
num_cond_frames: int = 0 # Number of conditioning frames
|
||||
# In-memory refine input (for two-stage pipeline where stage1 frames are already in memory)
|
||||
# This mirrors LongCat's demo where a list of frames (e.g. np.ndarray or PIL.Image)
|
||||
# is passed directly to the refinement pipeline instead of reloading from disk.
|
||||
stage1_video: Any | None = None
|
||||
|
||||
# Text inputs
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
@@ -50,6 +62,7 @@ class SamplingParam:
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -215,6 +228,31 @@ class SamplingParam:
|
||||
default=SamplingParam.video_path,
|
||||
help="Path to input video for video-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--refine-from",
|
||||
type=str,
|
||||
default=SamplingParam.refine_from,
|
||||
help="Path to stage1 video for refinement (LongCat 480p->720p)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t-thresh",
|
||||
type=float,
|
||||
default=SamplingParam.t_thresh,
|
||||
help=
|
||||
"Threshold for timestep scheduling in refinement (default: 0.5)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial-refine-only",
|
||||
action=StoreBoolean,
|
||||
default=SamplingParam.spatial_refine_only,
|
||||
help="Only perform spatial super-resolution (no temporal doubling)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-cond-frames",
|
||||
type=int,
|
||||
default=SamplingParam.num_cond_frames,
|
||||
help="Number of conditioning frames for refinement",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--moba-config-path",
|
||||
type=str,
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import numpy as np
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_480P_SamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 121
|
||||
height: int = 480
|
||||
width: int = 848
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 6.0
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
sigmas: list[float] | None = field(
|
||||
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
|
||||
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
@@ -5,6 +5,7 @@ from typing import Any
|
||||
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
|
||||
@@ -35,6 +36,10 @@ logger = init_logger(__name__)
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
|
||||
Hunyuan15_480P_SamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
|
||||
Hunyuan15_720P_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
|
||||
# Wan2.1
|
||||
@@ -86,6 +91,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan":
|
||||
lambda id: "hunyuan" in id.lower(),
|
||||
"hunyuan15":
|
||||
lambda id: "hunyuan15" in id.lower(),
|
||||
"wanpipeline":
|
||||
lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo":
|
||||
@@ -105,6 +112,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
"hunyuan":
|
||||
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"hunyuan15":
|
||||
Hunyuan15_480P_SamplingParam, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
|
||||
"wanpipeline":
|
||||
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
|
||||
|
||||
@@ -319,6 +319,62 @@ class FastVideoArgs:
|
||||
"Path to a text file containing prompts (one per line) for batch processing",
|
||||
)
|
||||
|
||||
# LoRA parameters (inference-time adapter loading)
|
||||
parser.add_argument(
|
||||
"--lora-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.lora_path,
|
||||
help=
|
||||
"Path to a LoRA adapter (directory or HF repo id). If set, LoRA will be applied at inference.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora-nickname",
|
||||
type=str,
|
||||
default=FastVideoArgs.lora_nickname,
|
||||
help=
|
||||
"Nickname to refer to the loaded LoRA adapter (useful for swapping).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora-target-modules",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=FastVideoArgs.lora_target_modules,
|
||||
help=
|
||||
"Optional list of module name substrings to restrict LoRA injection (e.g. q_proj k_proj v_proj).",
|
||||
)
|
||||
|
||||
# BSA runtime control (LongCat)
|
||||
parser.add_argument(
|
||||
"--enable-bsa",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Enable Block Sparse Attention (BSA) at runtime (overrides config).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bsa-sparsity",
|
||||
type=float,
|
||||
help="BSA sparsity (e.g., 0.9375).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bsa-cdf-threshold",
|
||||
type=float,
|
||||
help="BSA CDF threshold (optional).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bsa-chunk-q",
|
||||
nargs=3,
|
||||
type=int,
|
||||
metavar=("T", "H", "W"),
|
||||
help="BSA chunk_3d_shape_q as three ints, e.g., 4 4 4.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bsa-chunk-k",
|
||||
nargs=3,
|
||||
type=int,
|
||||
metavar=("T", "H", "W"),
|
||||
help="BSA chunk_3d_shape_k as three ints, e.g., 4 4 4.",
|
||||
)
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
"--STA-mode",
|
||||
|
||||
@@ -87,6 +87,7 @@ _ACTIVATION_REGISTRY = {
|
||||
"gelu_pytorch_tanh": lambda: nn.GELU(approximate="tanh"),
|
||||
"relu": nn.ReLU,
|
||||
"silu": nn.SiLU,
|
||||
"swish": nn.SiLU,
|
||||
"quick_gelu": QuickGELU,
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
3D Rotary Position Embedding (RoPE) for video transformers.
|
||||
|
||||
Reference: https://arxiv.org/pdf/2104.09864.pdf
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange, repeat
|
||||
|
||||
|
||||
def broadcast(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor:
|
||||
"""
|
||||
Broadcast and concatenate tensors along a dimension.
|
||||
"""
|
||||
num_tensors = len(tensors)
|
||||
shape_lens = set(len(t.shape) for t in tensors)
|
||||
assert len(
|
||||
shape_lens) == 1, "tensors must all have the same number of dimensions"
|
||||
|
||||
shape_len = list(shape_lens)[0]
|
||||
dim = (dim + shape_len) if dim < 0 else dim
|
||||
|
||||
dims = list(zip(*[list(t.shape) for t in tensors], strict=False))
|
||||
expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim]
|
||||
|
||||
assert all(
|
||||
len(set(t[1])) <= 2 for t in
|
||||
expandable_dims), "invalid dimensions for broadcastable concatenation"
|
||||
|
||||
max_dims = [(t[0], max(t[1])) for t in expandable_dims]
|
||||
expanded_dims = [(t[0], (t[1], ) * num_tensors) for t in max_dims]
|
||||
expanded_dims.insert(dim, (dim, dims[dim]))
|
||||
expandable_shapes = list(zip(*[t[1] for t in expanded_dims], strict=False))
|
||||
tensors = [
|
||||
t[0].expand(*t[1])
|
||||
for t in zip(tensors, expandable_shapes, strict=False)
|
||||
]
|
||||
|
||||
return torch.cat(tensors, dim=dim)
|
||||
|
||||
|
||||
def rotate_half(x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Rotate half the hidden dims of the input.
|
||||
"""
|
||||
x = rearrange(x, "... (d r) -> ... d r", r=2)
|
||||
x1, x2 = x.unbind(dim=-1)
|
||||
x = torch.stack((-x2, x1), dim=-1)
|
||||
return rearrange(x, "... d r -> ... (d r)")
|
||||
|
||||
|
||||
class RotaryPositionalEmbedding3D(nn.Module):
|
||||
"""
|
||||
3D Rotary Positional Embedding for video transformers.
|
||||
|
||||
Splits the head dimension across temporal, height, and width dimensions,
|
||||
computing separate rotary embeddings for each and concatenating them.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
head_dim: int,
|
||||
base: float = 10000.0,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
head_dim: Dimension of each attention head
|
||||
base: Base value for exponential frequency
|
||||
"""
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
assert self.head_dim % 8 == 0, "head_dim must be a multiple of 8 for 3D RoPE"
|
||||
self.base = base
|
||||
|
||||
# Cache for precomputed frequencies
|
||||
self.freqs_dict: dict[tuple, torch.Tensor] = {}
|
||||
|
||||
def register_grid_size(self, grid_size: tuple[int, int, int]) -> None:
|
||||
"""
|
||||
Precompute and register frequencies for a given grid size.
|
||||
|
||||
Args:
|
||||
grid_size: (T, H, W) tuple of grid dimensions
|
||||
"""
|
||||
if grid_size not in self.freqs_dict:
|
||||
self.freqs_dict[grid_size] = self.precompute_freqs_3d(grid_size)
|
||||
|
||||
def precompute_freqs_3d(self, grid_size: tuple[int, int,
|
||||
int]) -> torch.Tensor:
|
||||
"""
|
||||
Precompute 3D rotary frequencies.
|
||||
|
||||
Args:
|
||||
grid_size: (num_frames, height, width)
|
||||
|
||||
Returns:
|
||||
freqs: [T*H*W, head_dim] tensor of frequencies
|
||||
"""
|
||||
num_frames, height, width = grid_size
|
||||
|
||||
# Split head_dim across 3 dimensions
|
||||
# Temporal gets the remainder to ensure exact division
|
||||
dim_t = self.head_dim - 4 * (self.head_dim // 6)
|
||||
dim_h = 2 * (self.head_dim // 6)
|
||||
dim_w = 2 * (self.head_dim // 6)
|
||||
|
||||
# Compute frequency bands for each dimension
|
||||
freqs_t = 1.0 / (self.base**(
|
||||
torch.arange(0, dim_t, 2)[:(dim_t // 2)].float() / dim_t))
|
||||
freqs_h = 1.0 / (self.base**(
|
||||
torch.arange(0, dim_h, 2)[:(dim_h // 2)].float() / dim_h))
|
||||
freqs_w = 1.0 / (self.base**(
|
||||
torch.arange(0, dim_w, 2)[:(dim_w // 2)].float() / dim_w))
|
||||
|
||||
# Create position grids
|
||||
grid_t = torch.arange(num_frames, dtype=torch.float32)
|
||||
grid_h = torch.arange(height, dtype=torch.float32)
|
||||
grid_w = torch.arange(width, dtype=torch.float32)
|
||||
|
||||
# Compute frequencies for each position
|
||||
freqs_t = torch.einsum("..., f -> ... f", grid_t, freqs_t)
|
||||
freqs_h = torch.einsum("..., f -> ... f", grid_h, freqs_h)
|
||||
freqs_w = torch.einsum("..., f -> ... f", grid_w, freqs_w)
|
||||
|
||||
# Duplicate for complex pair representation
|
||||
freqs_t = repeat(freqs_t, "... n -> ... (n r)", r=2)
|
||||
freqs_h = repeat(freqs_h, "... n -> ... (n r)", r=2)
|
||||
freqs_w = repeat(freqs_w, "... n -> ... (n r)", r=2)
|
||||
|
||||
# Broadcast and concatenate across all 3 dimensions
|
||||
freqs = broadcast(
|
||||
[
|
||||
freqs_t[:, None, None, :], # [T, 1, 1, dim_t]
|
||||
freqs_h[None, :, None, :], # [1, H, 1, dim_h]
|
||||
freqs_w[None, None, :, :], # [1, 1, W, dim_w]
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
# Flatten spatial dimensions: [T, H, W, head_dim] -> [T*H*W, head_dim]
|
||||
freqs = rearrange(freqs, "T H W D -> (T H W) D")
|
||||
|
||||
return freqs
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
grid_size: tuple[int, int, int],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply 3D rotary positional embedding to queries and keys.
|
||||
|
||||
Args:
|
||||
q: Query tensor [B, num_heads, seq_len, head_dim]
|
||||
k: Key tensor [B, num_heads, seq_len, head_dim]
|
||||
grid_size: (T, H, W) tuple of grid dimensions
|
||||
|
||||
Returns:
|
||||
(q_rotated, k_rotated): Rotated query and key tensors
|
||||
"""
|
||||
# Register grid size if not cached
|
||||
if grid_size not in self.freqs_dict:
|
||||
self.register_grid_size(grid_size)
|
||||
|
||||
# Get cached frequencies
|
||||
freqs_cis = self.freqs_dict[grid_size].to(q.device)
|
||||
|
||||
# Cast to float32 for precision
|
||||
q_, k_ = q.float(), k.float()
|
||||
freqs_cis = freqs_cis.float()
|
||||
|
||||
# Compute cos and sin
|
||||
cos = freqs_cis.cos()
|
||||
sin = freqs_cis.sin()
|
||||
|
||||
# Reshape for broadcasting: [1, 1, seq_len, head_dim]
|
||||
cos = rearrange(cos, "n d -> 1 1 n d")
|
||||
sin = rearrange(sin, "n d -> 1 1 n d")
|
||||
|
||||
# Apply rotation
|
||||
q_ = (q_ * cos) + (rotate_half(q_) * sin)
|
||||
k_ = (k_ * cos) + (rotate_half(k_) * sin)
|
||||
|
||||
# Cast back to original dtype
|
||||
return q_.type_as(q), k_.type_as(k)
|
||||
|
||||
|
||||
def apply_rotary_emb_3d(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
rope_module: RotaryPositionalEmbedding3D,
|
||||
grid_size: tuple[int, int, int],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Convenience function to apply 3D RoPE.
|
||||
|
||||
Args:
|
||||
q: Query tensor [B, num_heads, seq_len, head_dim]
|
||||
k: Key tensor [B, num_heads, seq_len, head_dim]
|
||||
rope_module: RotaryPositionalEmbedding3D module
|
||||
grid_size: (T, H, W) grid dimensions
|
||||
|
||||
Returns:
|
||||
(q_rotated, k_rotated): Rotated tensors
|
||||
"""
|
||||
return rope_module(q, k, grid_size)
|
||||
@@ -0,0 +1,853 @@
|
||||
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Any, Dict, Optional, List
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather_with_unpad,
|
||||
sequence_model_parallel_shard)
|
||||
from fastvideo.configs.models.dits import HunyuanVideo15Config
|
||||
from fastvideo.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
# TODO(will-PY-refactor): RMSNorm ....
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
|
||||
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
|
||||
TimestepEmbedder, unpatchify)
|
||||
from fastvideo.models.dits.base import CachableDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.attention.backends.abstract import AttentionMetadata
|
||||
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.distributed.utils import create_attention_mask_for_padding
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class HunyuanRMSNorm(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
elementwise_affine=True,
|
||||
eps: float = 1e-6,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
"""
|
||||
Initialize the RMSNorm normalization layer.
|
||||
|
||||
Args:
|
||||
dim (int): The dimension of the input tensor.
|
||||
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
|
||||
|
||||
Attributes:
|
||||
eps (float): A small value added to the denominator for numerical stability.
|
||||
weight (nn.Parameter): Learnable scaling parameter.
|
||||
|
||||
"""
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
if elementwise_affine:
|
||||
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
|
||||
|
||||
def _norm(self, x) -> torch.Tensor:
|
||||
"""
|
||||
Apply the RMSNorm normalization to the input tensor.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The input tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The normalized tensor.
|
||||
|
||||
"""
|
||||
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
Forward pass through the RMSNorm layer.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The input tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The output tensor after applying RMSNorm.
|
||||
|
||||
"""
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
if hasattr(self, "weight"):
|
||||
output = output * self.weight
|
||||
return output
|
||||
|
||||
|
||||
class HunyuanVideo15TimeEmbedding(nn.Module):
|
||||
r"""
|
||||
Time embedding for HunyuanVideo 1.5.
|
||||
|
||||
Supports standard timestep embedding and optional reference timestep embedding for MeanFlow-based super-resolution
|
||||
models.
|
||||
|
||||
Args:
|
||||
embedding_dim (`int`):
|
||||
The dimension of the output embedding.
|
||||
"""
|
||||
|
||||
def __init__(self, embedding_dim: int, use_meanflow: bool = False):
|
||||
super().__init__()
|
||||
|
||||
self.timestep_embedder = TimestepEmbedder(hidden_size=embedding_dim)
|
||||
|
||||
self.use_meanflow = use_meanflow
|
||||
self.time_proj_r = None
|
||||
self.timestep_embedder_r = None
|
||||
if use_meanflow:
|
||||
self.timestep_embedder_r = TimestepEmbedder(hidden_size=embedding_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
timestep_r: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
timesteps_emb = self.timestep_embedder(timestep)
|
||||
|
||||
if timestep_r is not None:
|
||||
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
|
||||
timesteps_emb = timesteps_emb + timesteps_emb_r
|
||||
|
||||
return timesteps_emb
|
||||
|
||||
|
||||
class HunyuanVideo15ByT5TextProjection(nn.Module):
|
||||
def __init__(self, in_features: int, hidden_size: int, out_features: int):
|
||||
super().__init__()
|
||||
self.norm = nn.LayerNorm(in_features)
|
||||
self.linear_1 = nn.Linear(in_features, hidden_size)
|
||||
self.linear_2 = nn.Linear(hidden_size, hidden_size)
|
||||
self.linear_3 = nn.Linear(hidden_size, out_features)
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.norm(encoder_hidden_states)
|
||||
hidden_states = self.linear_1(hidden_states)
|
||||
hidden_states = self.act_fn(hidden_states)
|
||||
hidden_states = self.linear_2(hidden_states)
|
||||
hidden_states = self.act_fn(hidden_states)
|
||||
hidden_states = self.linear_3(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideo15ImageProjection(nn.Module):
|
||||
def __init__(self, in_channels: int, hidden_size: int):
|
||||
super().__init__()
|
||||
self.norm_in = nn.LayerNorm(in_channels)
|
||||
self.linear_1 = nn.Linear(in_channels, in_channels)
|
||||
self.act_fn = nn.GELU()
|
||||
self.linear_2 = nn.Linear(in_channels, hidden_size)
|
||||
self.norm_out = nn.LayerNorm(hidden_size)
|
||||
|
||||
def forward(self, image_embeds: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.norm_in(image_embeds)
|
||||
hidden_states = self.linear_1(hidden_states)
|
||||
hidden_states = self.act_fn(hidden_states)
|
||||
hidden_states = self.linear_2(hidden_states)
|
||||
hidden_states = self.norm_out(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
A multimodal DiT block with separate modulation for text and image/video,
|
||||
using distributed attention and linear layers.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: torch.dtype | None = None,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.num_attention_heads = num_attention_heads
|
||||
head_dim = hidden_size // num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
# Image modulation components
|
||||
self.img_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mod",
|
||||
)
|
||||
|
||||
# Fused operations for image stream
|
||||
self.img_attn_norm = LayerNormScaleShift(hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.img_mlp_residual = ScaleResidual()
|
||||
|
||||
# Image attention components
|
||||
self.img_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_qkv")
|
||||
|
||||
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
self.img_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_proj")
|
||||
|
||||
self.img_mlp = MLP(hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mlp")
|
||||
|
||||
# Text modulation components
|
||||
self.txt_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.txt_mod",
|
||||
)
|
||||
|
||||
# Fused operations for text stream
|
||||
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.txt_mlp_residual = ScaleResidual()
|
||||
|
||||
# Text attention components
|
||||
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
# QK norm layers for text
|
||||
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
self.txt_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
txt: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple,
|
||||
seq_attention_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# Process modulation vectors
|
||||
img_mod_outputs = self.img_mod(vec)
|
||||
(
|
||||
img_attn_shift,
|
||||
img_attn_scale,
|
||||
img_attn_gate,
|
||||
img_mlp_shift,
|
||||
img_mlp_scale,
|
||||
img_mlp_gate,
|
||||
) = torch.chunk(img_mod_outputs, 6, dim=-1)
|
||||
|
||||
txt_mod_outputs = self.txt_mod(vec)
|
||||
(
|
||||
txt_attn_shift,
|
||||
txt_attn_scale,
|
||||
txt_attn_gate,
|
||||
txt_mlp_shift,
|
||||
txt_mlp_scale,
|
||||
txt_mlp_gate,
|
||||
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
|
||||
|
||||
# Prepare image for attention using fused operation
|
||||
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
|
||||
# Get QKV for image
|
||||
img_qkv, _ = self.img_attn_qkv(img_attn_input)
|
||||
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
|
||||
self.num_attention_heads, -1)
|
||||
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
|
||||
2]
|
||||
|
||||
# Apply QK-Norm if needed
|
||||
img_q = self.img_attn_q_norm(img_q).to(img_v)
|
||||
img_k = self.img_attn_k_norm(img_k).to(img_v)
|
||||
|
||||
# Prepare text for attention using fused operation
|
||||
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
|
||||
|
||||
# Get QKV for text
|
||||
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
|
||||
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
|
||||
self.num_attention_heads, -1)
|
||||
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
|
||||
2]
|
||||
# Apply QK-Norm if needed
|
||||
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
|
||||
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
|
||||
|
||||
# seq_len = txt_q.shape[1] + img_q.shape[1]
|
||||
# attention_mask = F.pad(encoder_attention_mask, (seq_len - encoder_attention_mask.shape[1], 0), value=True)
|
||||
# attention_mask = attention_mask.bool()
|
||||
# self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1)
|
||||
# self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)
|
||||
# attention_mask = (self_attn_mask_1 & self_attn_mask_2).bool()
|
||||
|
||||
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
|
||||
attn_metadata = FlashAttnMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
attn_mask=encoder_attention_mask,
|
||||
)
|
||||
# Run distributed attention
|
||||
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
|
||||
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis, attention_mask=seq_attention_mask)
|
||||
|
||||
img_attn_out, _ = self.img_attn_proj(
|
||||
img_attn.view(batch_size, image_seq_len, -1))
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
|
||||
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
|
||||
|
||||
# Process image MLP
|
||||
img_mlp_out = self.img_mlp(img_mlp_input)
|
||||
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
|
||||
|
||||
# Process text attention output
|
||||
txt_attn_out, _ = self.txt_attn_proj(
|
||||
txt_attn.reshape(batch_size, text_seq_len, -1))
|
||||
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
|
||||
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
|
||||
|
||||
# Process text MLP
|
||||
txt_mlp_out = self.txt_mlp(txt_mlp_input)
|
||||
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class HunyuanVideo15Transformer3DModel(CachableDiT):
|
||||
r"""
|
||||
A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5).
|
||||
"""
|
||||
|
||||
# shard single stream, double stream blocks, and refiner_blocks
|
||||
_fsdp_shard_conditions = HunyuanVideo15Config()._fsdp_shard_conditions
|
||||
_compile_conditions = HunyuanVideo15Config()._compile_conditions
|
||||
_supported_attention_backends = HunyuanVideo15Config(
|
||||
)._supported_attention_backends
|
||||
param_names_mapping = HunyuanVideo15Config().param_names_mapping
|
||||
reverse_param_names_mapping = HunyuanVideo15Config(
|
||||
).reverse_param_names_mapping
|
||||
lora_param_names_mapping = HunyuanVideo15Config().lora_param_names_mapping
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: HunyuanVideo15Config,
|
||||
hf_config: dict[str, Any],
|
||||
) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.out_channels = config.out_channels or config.in_channels
|
||||
self.patch_size = (config.patch_size_t, config.patch_size, config.patch_size)
|
||||
|
||||
# 1. Latent and condition embedders
|
||||
self.img_in = PatchEmbed(self.patch_size,
|
||||
config.in_channels,
|
||||
self.hidden_size,
|
||||
prefix=f"{config.prefix}.img_in")
|
||||
self.image_embedder = HunyuanVideo15ImageProjection(config.image_embed_dim, self.hidden_size)
|
||||
|
||||
self.txt_in = SingleTokenRefiner(config.text_embed_dim,
|
||||
self.hidden_size,
|
||||
config.num_attention_heads,
|
||||
depth=config.num_refiner_layers,
|
||||
dtype=None,
|
||||
prefix=f"{config.prefix}.txt_in")
|
||||
|
||||
self.txt_in_2 = HunyuanVideo15ByT5TextProjection(config.text_embed_2_dim, 2048, self.hidden_size)
|
||||
|
||||
self.time_in = HunyuanVideo15TimeEmbedding(self.hidden_size, use_meanflow=config.use_meanflow)
|
||||
|
||||
self.cond_type_embed = nn.Embedding(3, self.hidden_size)
|
||||
|
||||
# 3. Dual stream transformer blocks
|
||||
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
MMDoubleStreamBlock(
|
||||
hidden_size=self.hidden_size,
|
||||
num_attention_heads=config.num_attention_heads,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
dtype=None,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{config.prefix}.double_blocks.{i}"
|
||||
)
|
||||
for i in range(config.num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 5. Output projection
|
||||
self.final_layer = FinalLayer(self.hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
prefix=f"{config.prefix}.final_layer")
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: List[torch.Tensor],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: List[torch.Tensor],
|
||||
encoder_attention_mask: List[torch.Tensor],
|
||||
guidance: Optional[torch.Tensor] = None,
|
||||
timestep_r: Optional[torch.LongTensor] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
|
||||
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# 1. RoPE
|
||||
# Get rotary embeddings
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||
(post_patch_num_frames, post_patch_height, post_patch_width), self.hidden_size,
|
||||
self.num_attention_heads, self.config.rope_axes_dim, self.config.rope_theta)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# 2. Conditional embeddings
|
||||
temb = self.time_in(timestep, timestep_r=timestep_r)
|
||||
|
||||
hidden_states = self.img_in(hidden_states)
|
||||
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
|
||||
|
||||
current_seq_len = hidden_states.shape[1]
|
||||
sp_world_size = get_sp_world_size()
|
||||
padded_seq_len = current_seq_len * sp_world_size
|
||||
|
||||
if padded_seq_len > original_seq_len:
|
||||
seq_attention_mask = create_attention_mask_for_padding(
|
||||
seq_len=original_seq_len,
|
||||
padded_seq_len=padded_seq_len,
|
||||
batch_size=batch_size,
|
||||
device=hidden_states.device,
|
||||
)
|
||||
else:
|
||||
seq_attention_mask = None
|
||||
|
||||
# qwen text embedding
|
||||
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep, encoder_attention_mask)
|
||||
|
||||
encoder_hidden_states_cond_emb = self.cond_type_embed(
|
||||
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
|
||||
)
|
||||
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
|
||||
|
||||
# byt5 text embedding
|
||||
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
|
||||
|
||||
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
|
||||
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
|
||||
)
|
||||
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
|
||||
|
||||
# image embed
|
||||
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
|
||||
is_t2v = torch.all(encoder_hidden_states_image == 0)
|
||||
if is_t2v:
|
||||
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
|
||||
encoder_attention_mask_3 = torch.zeros(
|
||||
(batch_size, encoder_hidden_states_3.shape[1]),
|
||||
dtype=encoder_attention_mask.dtype,
|
||||
device=encoder_attention_mask.device,
|
||||
)
|
||||
else:
|
||||
encoder_attention_mask_3 = torch.ones(
|
||||
(batch_size, encoder_hidden_states_3.shape[1]),
|
||||
dtype=encoder_attention_mask.dtype,
|
||||
device=encoder_attention_mask.device,
|
||||
)
|
||||
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
|
||||
2
|
||||
* torch.ones_like(
|
||||
encoder_hidden_states_3[:, :, 0],
|
||||
dtype=torch.long,
|
||||
)
|
||||
)
|
||||
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
|
||||
|
||||
# reorder and combine text tokens: combine valid tokens first, then padding
|
||||
encoder_attention_mask = encoder_attention_mask.bool()
|
||||
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
|
||||
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
|
||||
new_encoder_hidden_states = []
|
||||
new_encoder_attention_mask = []
|
||||
|
||||
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
encoder_hidden_states_2,
|
||||
encoder_attention_mask_2,
|
||||
encoder_hidden_states_3,
|
||||
encoder_attention_mask_3,
|
||||
):
|
||||
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
|
||||
new_encoder_hidden_states.append(
|
||||
torch.cat(
|
||||
[
|
||||
image[image_mask], # valid image
|
||||
text_2[text_mask_2], # valid byt5
|
||||
text[text_mask], # valid mllm
|
||||
image[~image_mask], # invalid image (zeroed)
|
||||
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
|
||||
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
# Apply same reordering to attention masks
|
||||
new_encoder_attention_mask.append(
|
||||
torch.cat(
|
||||
[
|
||||
image_mask[image_mask],
|
||||
text_mask_2[text_mask_2],
|
||||
text_mask[text_mask],
|
||||
image_mask[~image_mask],
|
||||
text_mask_2[~text_mask_2],
|
||||
text_mask[~text_mask],
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
|
||||
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
|
||||
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.double_blocks:
|
||||
hidden_states, encoder_hidden_states = self._gradient_checkpointing_func(
|
||||
block,
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
temb,
|
||||
freqs_cis,
|
||||
seq_attention_mask
|
||||
)
|
||||
|
||||
else:
|
||||
for block in self.double_blocks:
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
temb,
|
||||
freqs_cis,
|
||||
seq_attention_mask
|
||||
)
|
||||
|
||||
# Final layer processing
|
||||
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
|
||||
hidden_states = self.final_layer(hidden_states, temb)
|
||||
# Unpatchify to get original shape
|
||||
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
|
||||
|
||||
return hidden_states
|
||||
|
||||
class SingleTokenRefiner(nn.Module):
|
||||
"""
|
||||
A token refiner that processes text embeddings with attention to improve
|
||||
their representation for cross-attention with image features.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
depth=2,
|
||||
qkv_bias=True,
|
||||
dtype=None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
# Input projection
|
||||
# self.input_embedder = ReplicatedLinear(
|
||||
# in_channels,
|
||||
# hidden_size,
|
||||
# bias=True,
|
||||
# params_dtype=dtype,
|
||||
# prefix=f"{prefix}.input_embedder")
|
||||
self.input_embedder = nn.Linear(in_channels, hidden_size, bias=True)
|
||||
|
||||
# Timestep embedding
|
||||
self.t_embedder = TimestepEmbedder(hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.t_embedder")
|
||||
|
||||
# Context embedding
|
||||
self.c_embedder = MLP(in_channels,
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
act_type="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.c_embedder")
|
||||
|
||||
# Refiner blocks
|
||||
self.refiner_blocks = nn.ModuleList([
|
||||
IndividualTokenRefinerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.refiner_blocks.{i}",
|
||||
) for i in range(depth)
|
||||
])
|
||||
|
||||
def forward(self, x, t, mask=None):
|
||||
# Get timestep embeddings
|
||||
timestep_aware_representations = self.t_embedder(t)
|
||||
|
||||
# Get context-aware representations
|
||||
original_dtype = x.dtype
|
||||
if mask is None:
|
||||
context_aware_representations = x.mean(dim=1)
|
||||
else:
|
||||
mask_float = mask.float().unsqueeze(-1) # [B, L, 1]
|
||||
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
|
||||
|
||||
context_aware_representations = self.c_embedder(
|
||||
context_aware_representations)
|
||||
c = timestep_aware_representations + context_aware_representations
|
||||
# Project input
|
||||
x = self.input_embedder(x)
|
||||
# Process through refiner blocks
|
||||
for block in self.refiner_blocks:
|
||||
x = block(x, c, mask)
|
||||
return x
|
||||
|
||||
class IndividualTokenRefinerBlock(nn.Module):
|
||||
"""
|
||||
A transformer block for refining individual tokens with self-attention.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=True,
|
||||
dtype=None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
# Normalization and attention
|
||||
self.norm1 = nn.LayerNorm(hidden_size,
|
||||
eps=1e-6,
|
||||
elementwise_affine=True,
|
||||
dtype=dtype)
|
||||
|
||||
self.self_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=qkv_bias,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.self_attn_qkv")
|
||||
|
||||
self.self_attn_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
bias=qkv_bias,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.self_attn_proj")
|
||||
|
||||
# MLP
|
||||
self.norm2 = nn.LayerNorm(hidden_size,
|
||||
eps=1e-6,
|
||||
elementwise_affine=True,
|
||||
dtype=dtype)
|
||||
self.mlp = MLP(hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
act_type="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.mlp")
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.adaLN_modulation")
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=hidden_size // num_attention_heads,
|
||||
# TODO: remove hardcode; remove STA
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
|
||||
def forward(self, x, c, mask=None):
|
||||
if mask is not None:
|
||||
mask = mask.clone().bool()
|
||||
mask[:, 0] = True # Prevent attention weights from becoming NaN
|
||||
|
||||
# Get modulation parameters
|
||||
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=-1)
|
||||
# Self-attention
|
||||
norm_x = self.norm1(x)
|
||||
qkv, _ = self.self_attn_qkv(norm_x)
|
||||
|
||||
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
|
||||
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
|
||||
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
|
||||
|
||||
# Run scaled dot product attention
|
||||
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
|
||||
attn_metadata = FlashAttnMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
attn_mask=mask,
|
||||
)
|
||||
# Run distributed attention
|
||||
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
|
||||
attn_output = self.attn(q, k, v) # [B, L, H, D]
|
||||
attn_output = attn_output.reshape(batch_size, seq_len,
|
||||
-1) # [B, L, H*D]
|
||||
|
||||
# Project and apply residual connection with gating
|
||||
attn_out, _ = self.self_attn_proj(attn_output)
|
||||
x = x + attn_out * gate_msa.unsqueeze(1)
|
||||
|
||||
# MLP
|
||||
mlp_out = self.mlp(self.norm2(x))
|
||||
x = x + mlp_out * gate_mlp.unsqueeze(1)
|
||||
|
||||
return x
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of DiT that projects features to pixel space.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
hidden_size,
|
||||
patch_size,
|
||||
out_channels,
|
||||
dtype=None,
|
||||
prefix: str = "") -> None:
|
||||
super().__init__()
|
||||
|
||||
# Normalization
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
eps=1e-6,
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
|
||||
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
|
||||
|
||||
self.linear = ReplicatedLinear(hidden_size,
|
||||
output_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear")
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.adaLN_modulation")
|
||||
|
||||
def forward(self, x, c):
|
||||
# What the heck HF? Why you change the scale and shift order here???
|
||||
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
|
||||
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
x, _ = self.linear(x)
|
||||
return x
|
||||
@@ -0,0 +1,866 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Native LongCat Video DiT implementation using FastVideo conventions.
|
||||
|
||||
This is a Phase 2 reimplementation that replaces the third_party wrapper
|
||||
with native FastVideo layers for better performance and integration.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.configs.models.dits import LongCatVideoConfig
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.layernorm import RMSNorm, FP32LayerNorm
|
||||
from fastvideo.layers.activation import get_act_fn
|
||||
from fastvideo.layers.rotary_embedding_3d import RotaryPositionalEmbedding3D
|
||||
from fastvideo.attention.layer import DistributedAttention, LocalAttention
|
||||
from fastvideo.models.dits.base import CachableDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.third_party.longcat_video.block_sparse_attention.bsa_interface import flash_attn_bsa_3d
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Embeddings
|
||||
# ============================================================================
|
||||
|
||||
class PatchEmbed3D(nn.Module):
|
||||
"""
|
||||
3D patch embedding using Conv3d.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
in_channels: int = 16,
|
||||
embed_dim: int = 4096,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.in_channels = in_channels
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
self.proj = nn.Conv3d(
|
||||
in_channels,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: [B, C, T, H, W]
|
||||
Returns:
|
||||
[B, N, C] where N = (T/pt) * (H/ph) * (W/pw)
|
||||
"""
|
||||
# Padding if needed
|
||||
_, _, T, H, W = x.shape
|
||||
if W % self.patch_size[2] != 0:
|
||||
x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
|
||||
if H % self.patch_size[1] != 0:
|
||||
x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
|
||||
if T % self.patch_size[0] != 0:
|
||||
x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - T % self.patch_size[0]))
|
||||
|
||||
x = self.proj(x) # [B, C, T', H', W']
|
||||
x = x.flatten(2).transpose(1, 2) # [B, N, C]
|
||||
return x
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Sinusoidal timestep embedding + MLP projection.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
frequency_embedding_size: int = 256,
|
||||
adaln_tembed_dim: int = 512,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
# Use FastVideo's ReplicatedLinear
|
||||
self.linear_1 = ReplicatedLinear(
|
||||
frequency_embedding_size,
|
||||
adaln_tembed_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
self.act = nn.SiLU()
|
||||
self.linear_2 = ReplicatedLinear(
|
||||
adaln_tembed_dim,
|
||||
adaln_tembed_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000) -> torch.Tensor:
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
"""
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half
|
||||
)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
def forward(self, t: torch.Tensor, latent_shape: tuple | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
t: [B] or [B, T] timesteps
|
||||
latent_shape: (T, H, W) for temporal expansion
|
||||
Returns:
|
||||
[B, T, C]
|
||||
"""
|
||||
# Sinusoidal embedding in FP32
|
||||
t_freq = self.timestep_embedding(t.flatten(), self.frequency_embedding_size)
|
||||
|
||||
# Cast to model dtype before MLP
|
||||
# Handle LoRA wrapper if present
|
||||
linear_layer = self.linear_1.base_layer if hasattr(self.linear_1, 'base_layer') else self.linear_1
|
||||
target_dtype = linear_layer.weight.dtype
|
||||
if t_freq.dtype != target_dtype:
|
||||
t_freq = t_freq.to(target_dtype)
|
||||
|
||||
# MLP projection
|
||||
t_emb, _ = self.linear_1(t_freq)
|
||||
t_emb = self.act(t_emb)
|
||||
t_emb, _ = self.linear_2(t_emb)
|
||||
|
||||
# Reshape if needed
|
||||
if latent_shape is not None and len(t.shape) > 1:
|
||||
B = t.shape[0]
|
||||
T = latent_shape[0]
|
||||
t_emb = t_emb.reshape(B, T, -1)
|
||||
|
||||
return t_emb
|
||||
|
||||
|
||||
class CaptionEmbedder(nn.Module):
|
||||
"""
|
||||
Caption embedding with MLP projection and optional text compaction.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
caption_channels: int = 4096,
|
||||
hidden_size: int = 4096,
|
||||
text_tokens_zero_pad: bool = True,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.text_tokens_zero_pad = text_tokens_zero_pad
|
||||
|
||||
# Two-layer MLP using ReplicatedLinear
|
||||
self.linear_1 = ReplicatedLinear(
|
||||
caption_channels,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
self.act = nn.SiLU()
|
||||
self.linear_2 = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
encoder_hidden_states: [B, N_text, C_text] or [B, 1, N_text, C_text]
|
||||
encoder_attention_mask: [B, N_text] or [B, 1, 1, N_text]
|
||||
Returns:
|
||||
y: [B, N_text, C] - standard padded representation (like other models)
|
||||
"""
|
||||
# Handle extra dimension from wrapper
|
||||
if len(encoder_hidden_states.shape) == 4:
|
||||
encoder_hidden_states = encoder_hidden_states.squeeze(1)
|
||||
|
||||
# Project
|
||||
y, _ = self.linear_1(encoder_hidden_states)
|
||||
y = self.act(y)
|
||||
y, _ = self.linear_2(y) # [B, N_text, C]
|
||||
|
||||
# Handle attention masking - just zero out padded tokens if requested
|
||||
if encoder_attention_mask is not None:
|
||||
# Remove extra dimensions
|
||||
if len(encoder_attention_mask.shape) == 4:
|
||||
encoder_attention_mask = encoder_attention_mask.squeeze(1).squeeze(1)
|
||||
elif len(encoder_attention_mask.shape) == 3:
|
||||
encoder_attention_mask = encoder_attention_mask.squeeze(1)
|
||||
|
||||
# Zero out padded tokens if requested
|
||||
if self.text_tokens_zero_pad:
|
||||
y = y * encoder_attention_mask.unsqueeze(-1)
|
||||
|
||||
# Return standard format [B, N_text, C] - no compaction!
|
||||
return y
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Attention Modules (Placeholders for now)
|
||||
# ============================================================================
|
||||
|
||||
class LongCatSelfAttention(nn.Module):
|
||||
"""
|
||||
Self-attention with 3D RoPE support and optional BSA.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
config: LongCatVideoConfig,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
# Separate Q/K/V projections (not fused like original)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
|
||||
# Per-head RMS normalization
|
||||
self.q_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
|
||||
self.k_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
|
||||
|
||||
# Output projection
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
|
||||
# 3D RoPE
|
||||
self.rope_3d = RotaryPositionalEmbedding3D(head_dim=self.head_dim)
|
||||
|
||||
# BSA configuration
|
||||
self.enable_bsa = getattr(config, 'enable_bsa', False)
|
||||
self.bsa_params = getattr(config, 'bsa_params', None)
|
||||
|
||||
# FastVideo attention backend (used when BSA is disabled)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
supported_attention_backends=config._supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # [B, N, C]
|
||||
latent_shape: tuple, # (T, H, W)
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass with 3D RoPE and optional BSA.
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T, H, W = latent_shape
|
||||
|
||||
# Project to Q/K/V
|
||||
q, _ = self.to_q(x)
|
||||
k, _ = self.to_k(x)
|
||||
v, _ = self.to_v(x)
|
||||
|
||||
# Reshape to heads: [B, N, num_heads, head_dim]
|
||||
q = q.view(B, N, self.num_heads, self.head_dim)
|
||||
k = k.view(B, N, self.num_heads, self.head_dim)
|
||||
v = v.view(B, N, self.num_heads, self.head_dim)
|
||||
|
||||
# Per-head RMS normalization
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
# For RoPE: need [B, num_heads, N, head_dim]
|
||||
q_rope = q.transpose(1, 2)
|
||||
k_rope = k.transpose(1, 2)
|
||||
|
||||
# Apply 3D RoPE
|
||||
q_rope, k_rope = self.rope_3d(q_rope, k_rope, grid_size=latent_shape)
|
||||
|
||||
# Transpose back: [B, N, num_heads, head_dim] or [B, H, N, D] for BSA
|
||||
q = q_rope.transpose(1, 2)
|
||||
k = k_rope.transpose(1, 2)
|
||||
|
||||
# === Attention: BSA or standard ===
|
||||
if self.enable_bsa and T > 1: # Only use BSA for multi-frame videos
|
||||
# BSA expects [B, H, S, D] format
|
||||
q_bsa = q.transpose(1, 2).contiguous() # [B, num_heads, N, head_dim]
|
||||
k_bsa = k.transpose(1, 2).contiguous()
|
||||
v_bsa = v.transpose(1, 2).contiguous()
|
||||
|
||||
# Handle SP split: BSA operates on per-rank spatial dimensions
|
||||
# Replicate LongCat's cp_split_hw logic exactly
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
sp_size = get_sp_world_size()
|
||||
if sp_size > 1:
|
||||
# Calculate optimal 2D split (same as LongCat's get_optimal_split)
|
||||
factors = []
|
||||
for i in range(1, int(sp_size**0.5) + 1):
|
||||
if sp_size % i == 0:
|
||||
factors.append([i, sp_size // i])
|
||||
cp_split_hw = min(factors, key=lambda x: abs(x[0] - x[1]))
|
||||
|
||||
# Split H and W dimensions by their respective factors
|
||||
T_bsa, H_bsa, W_bsa = latent_shape
|
||||
assert H_bsa % cp_split_hw[0] == 0 and W_bsa % cp_split_hw[1] == 0, \
|
||||
f"H {H_bsa} must be divisible by {cp_split_hw[0]}, W {W_bsa} must be divisible by {cp_split_hw[1]}"
|
||||
H_bsa = H_bsa // cp_split_hw[0]
|
||||
W_bsa = W_bsa // cp_split_hw[1]
|
||||
latent_shape_bsa = (T_bsa, H_bsa, W_bsa)
|
||||
else:
|
||||
latent_shape_bsa = latent_shape
|
||||
|
||||
# Call BSA with per-rank latent shape
|
||||
out = flash_attn_bsa_3d(
|
||||
q_bsa, k_bsa, v_bsa,
|
||||
latent_shape_q=latent_shape_bsa,
|
||||
latent_shape_k=latent_shape_bsa,
|
||||
**self.bsa_params
|
||||
) # [B, num_heads, N, head_dim]
|
||||
|
||||
# Transpose back: [B, N, num_heads, head_dim]
|
||||
out = out.transpose(1, 2)
|
||||
else:
|
||||
# Standard attention: [B, N, num_heads, head_dim]
|
||||
out, _ = self.attn(q, k, v)
|
||||
|
||||
# Reshape and project out
|
||||
out = out.reshape(B, N, C)
|
||||
out, _ = self.to_out(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class LongCatCrossAttention(nn.Module):
|
||||
"""
|
||||
Cross-attention for text conditioning (standard implementation like other models).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
config: LongCatVideoConfig,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
# Separate Q/K/V projections
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
|
||||
# Per-head RMS normalization
|
||||
self.q_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
|
||||
self.k_norm = RMSNorm(self.head_dim, eps=1e-6, dtype=dtype or torch.float32)
|
||||
|
||||
# Output projection
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True, params_dtype=dtype)
|
||||
|
||||
# Cross-attention uses LocalAttention (FastVideo standard)
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=config.arch_config._supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # [B, N_img, C]
|
||||
context: torch.Tensor, # [B, N_text, C]
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass for cross-attention (standard implementation).
|
||||
|
||||
Args:
|
||||
x: Image tokens [B, N_img, C]
|
||||
context: Text tokens [B, N_text, C] (standard padded format)
|
||||
"""
|
||||
B, N_img, C = x.shape
|
||||
|
||||
# Project Q, K, V (standard cross-attention like WanVideo/StepVideo/Cosmos)
|
||||
q, _ = self.to_q(x)
|
||||
k, _ = self.to_k(context)
|
||||
v, _ = self.to_v(context)
|
||||
|
||||
N_text = context.shape[1]
|
||||
|
||||
# Reshape to heads
|
||||
q = q.view(B, N_img, self.num_heads, self.head_dim)
|
||||
k = k.view(B, N_text, self.num_heads, self.head_dim)
|
||||
v = v.view(B, N_text, self.num_heads, self.head_dim)
|
||||
|
||||
# Per-head RMS normalization
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
# Run cross-attention using FastVideo's LocalAttention
|
||||
# LocalAttention handles different q and k/v sequence lengths automatically
|
||||
out = self.attn(q, k, v) # [B, N_img, num_heads, head_dim]
|
||||
|
||||
# Reshape and project out
|
||||
out = out.reshape(B, N_img, C)
|
||||
out, _ = self.to_out(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Feed-Forward Network
|
||||
# ============================================================================
|
||||
|
||||
class LongCatSwiGLUFFN(nn.Module):
|
||||
"""
|
||||
SwiGLU feed-forward network using FastVideo's ReplicatedLinear.
|
||||
|
||||
FFN(x) = down(gate(x) * SiLU(up(x)))
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# Three projections for SwiGLU (no bias as per original)
|
||||
self.w1 = ReplicatedLinear(dim, hidden_dim, bias=False, params_dtype=dtype) # gate
|
||||
self.w3 = ReplicatedLinear(dim, hidden_dim, bias=False, params_dtype=dtype) # up
|
||||
self.w2 = ReplicatedLinear(hidden_dim, dim, bias=False, params_dtype=dtype) # down
|
||||
|
||||
self.act = nn.SiLU()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass: SiLU(w1(x)) * w3(x) -> w2 (matching original LongCat)
|
||||
"""
|
||||
w1_out, _ = self.w1(x)
|
||||
w3_out, _ = self.w3(x)
|
||||
combined = self.act(w1_out) * w3_out
|
||||
out, _ = self.w2(combined)
|
||||
return out
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Modulation Utilities
|
||||
# ============================================================================
|
||||
|
||||
def modulate_fp32(norm: nn.Module, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply modulation in FP32 for numerical stability (matching original LongCat).
|
||||
|
||||
shift and scale should already be FP32 from torch.amp.autocast context.
|
||||
"""
|
||||
# Ensure modulation params are FP32 (should be from autocast)
|
||||
assert shift.dtype == torch.float32 and scale.dtype == torch.float32, \
|
||||
f"shift and scale must be FP32, got {shift.dtype} and {scale.dtype}"
|
||||
|
||||
orig_dtype = x.dtype
|
||||
|
||||
# Normalize and modulate in FP32
|
||||
x_norm = norm(x.to(torch.float32))
|
||||
x_mod = x_norm * (scale + 1) + shift
|
||||
|
||||
return x_mod.to(orig_dtype)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Transformer Block
|
||||
# ============================================================================
|
||||
|
||||
class LongCatTransformerBlock(nn.Module):
|
||||
"""
|
||||
Single-stream transformer block with:
|
||||
- AdaLN modulation (FP32)
|
||||
- Self-attention
|
||||
- Cross-attention
|
||||
- SwiGLU FFN
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: int,
|
||||
adaln_tembed_dim: int,
|
||||
config: LongCatVideoConfig,
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.num_heads = num_heads
|
||||
|
||||
# AdaLN modulation (6 parameters: scale/shift for attn & ffn, gate for residual)
|
||||
self.adaln_linear_1 = ReplicatedLinear(
|
||||
adaln_tembed_dim,
|
||||
6 * hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
self.adaln_act = nn.SiLU()
|
||||
|
||||
# Normalization layers (CRITICAL: Use LayerNorm not RMSNorm like original!)
|
||||
# Original LongCat uses LayerNorm_FP32 with elementwise_affine=False
|
||||
self.norm_attn = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
|
||||
self.norm_ffn = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
|
||||
# Cross-attention norm has elementwise_affine=True (has weight and bias)
|
||||
self.norm_cross = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=True)
|
||||
|
||||
# Self-attention
|
||||
self.self_attn = LongCatSelfAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_heads,
|
||||
config=config,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
# Cross-attention
|
||||
self.cross_attn = LongCatCrossAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_heads,
|
||||
config=config,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
# SwiGLU FFN
|
||||
ffn_hidden_dim = int(hidden_size * mlp_ratio * 2 / 3)
|
||||
# Round up to nearest multiple of 256
|
||||
ffn_hidden_dim = 256 * ((ffn_hidden_dim + 255) // 256)
|
||||
|
||||
self.ffn = LongCatSwiGLUFFN(
|
||||
dim=hidden_size,
|
||||
hidden_dim=ffn_hidden_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # [B, N, C]
|
||||
context: torch.Tensor, # [B, N_text, C]
|
||||
t: torch.Tensor, # [B, T, C_t]
|
||||
latent_shape: tuple, # (T, H, W)
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass with AdaLN modulation.
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T, H, W = latent_shape
|
||||
x_orig_dtype = x.dtype # Save for later casting
|
||||
|
||||
# === AdaLN Modulation (CRITICAL: FP32 for stability like original) ===
|
||||
# Use autocast to compute modulation params in FP32
|
||||
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
|
||||
t_mod = self.adaln_act(t)
|
||||
mod_params, _ = self.adaln_linear_1(t_mod)
|
||||
# Ensure FP32 output (needed when LoRA is applied)
|
||||
if mod_params.dtype != torch.float32:
|
||||
mod_params = mod_params.float()
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = \
|
||||
mod_params.unsqueeze(2).chunk(6, dim=-1) # [B, T, 1, C]
|
||||
|
||||
# === Self-Attention ===
|
||||
x_norm = modulate_fp32(self.norm_attn, x.view(B, T, -1, C), shift_msa, scale_msa)
|
||||
x_norm = x_norm.view(B, N, C)
|
||||
|
||||
attn_out = self.self_attn(x_norm, latent_shape=latent_shape)
|
||||
|
||||
# Residual with gating (CRITICAL: FP32 like original, then cast back)
|
||||
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
|
||||
x = x + (gate_msa * attn_out.view(B, T, -1, C)).view(B, N, C)
|
||||
x = x.to(x_orig_dtype)
|
||||
|
||||
# === Cross-Attention ===
|
||||
x_norm_cross = self.norm_cross(x)
|
||||
cross_out = self.cross_attn(x_norm_cross, context)
|
||||
x = x + cross_out
|
||||
|
||||
# === FFN ===
|
||||
x_norm_ffn = modulate_fp32(self.norm_ffn, x.view(B, T, -1, C), shift_mlp, scale_mlp)
|
||||
x_norm_ffn = x_norm_ffn.view(B, N, C)
|
||||
|
||||
ffn_out = self.ffn(x_norm_ffn)
|
||||
|
||||
# Residual with gating (CRITICAL: FP32 like original, then cast back)
|
||||
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
|
||||
x = x + (gate_mlp * ffn_out.view(B, T, -1, C)).view(B, N, C)
|
||||
x = x.to(x_orig_dtype)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Final Layer
|
||||
# ============================================================================
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
Final output projection with AdaLN modulation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
out_channels: int,
|
||||
adaln_tembed_dim: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
dtype: torch.dtype | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# AdaLN for final layer (2 parameters: scale and shift)
|
||||
self.adaln_linear = ReplicatedLinear(
|
||||
adaln_tembed_dim,
|
||||
2 * hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
self.adaln_act = nn.SiLU()
|
||||
|
||||
# CRITICAL: Use LayerNorm not RMSNorm! (matches original)
|
||||
self.norm = FP32LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
|
||||
|
||||
# Output projection
|
||||
num_patch = patch_size[0] * patch_size[1] * patch_size[2]
|
||||
self.proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
num_patch * out_channels,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # [B, N, C]
|
||||
t: torch.Tensor, # [B, T, C_t]
|
||||
latent_shape: tuple,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Returns: [B, N, out_channels * patch_size^3]
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T, _, _ = latent_shape
|
||||
|
||||
# AdaLN modulation (FP32 for stability like original)
|
||||
with torch.amp.autocast(device_type='cuda', dtype=torch.float32):
|
||||
t_mod = self.adaln_act(t)
|
||||
mod_params, _ = self.adaln_linear(t_mod)
|
||||
# Ensure FP32 output (needed when LoRA is applied)
|
||||
if mod_params.dtype != torch.float32:
|
||||
mod_params = mod_params.float()
|
||||
shift, scale = mod_params.unsqueeze(2).chunk(2, dim=-1)
|
||||
|
||||
# Modulate
|
||||
x = modulate_fp32(self.norm, x.view(B, T, -1, C), shift, scale)
|
||||
x = x.reshape(B, N, C)
|
||||
|
||||
# Project
|
||||
x, _ = self.proj(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Main Model
|
||||
# ============================================================================
|
||||
|
||||
class LongCatTransformer3DModel(CachableDiT):
|
||||
"""
|
||||
Native LongCat Video Transformer using FastVideo layers.
|
||||
|
||||
This is a Phase 2 implementation that replaces third_party dependencies.
|
||||
"""
|
||||
|
||||
# FSDP sharding: shard at each transformer block
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "blocks" in n and n.split(".")[-1].isdigit(),
|
||||
]
|
||||
# torch.compile optimization: compile each transformer block for speedup
|
||||
_compile_conditions = [
|
||||
lambda n, m: "blocks" in n and n.split(".")[-1].isdigit(),
|
||||
]
|
||||
|
||||
# Parameter name mapping (for weight conversion)
|
||||
param_names_mapping = {} # Will be defined in config
|
||||
reverse_param_names_mapping = {}
|
||||
lora_param_names_mapping = {}
|
||||
|
||||
# Supported attention backends
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
def __init__(self, config: LongCatVideoConfig, hf_config: dict[str, Any]):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
# Extract architecture parameters
|
||||
self.hidden_size = config.hidden_size # 4096
|
||||
self.num_attention_heads = config.num_attention_heads # 32
|
||||
self.depth = config.depth # 48
|
||||
self.mlp_ratio = config.mlp_ratio # 4
|
||||
self.in_channels = config.in_channels # 16
|
||||
self.out_channels = config.out_channels # 16
|
||||
self.num_channels_latents = config.in_channels
|
||||
self.patch_size = config.patch_size # [1, 2, 2]
|
||||
|
||||
# Embeddings
|
||||
self.patch_embed = PatchEmbed3D(
|
||||
patch_size=self.patch_size,
|
||||
in_channels=self.in_channels,
|
||||
embed_dim=self.hidden_size,
|
||||
)
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
frequency_embedding_size=config.frequency_embedding_size,
|
||||
adaln_tembed_dim=config.adaln_tembed_dim,
|
||||
)
|
||||
|
||||
self.caption_embedder = CaptionEmbedder(
|
||||
caption_channels=config.caption_channels,
|
||||
hidden_size=self.hidden_size,
|
||||
text_tokens_zero_pad=getattr(config, 'text_tokens_zero_pad', True),
|
||||
)
|
||||
|
||||
# Transformer blocks (48 blocks)
|
||||
self.blocks = nn.ModuleList([
|
||||
LongCatTransformerBlock(
|
||||
hidden_size=self.hidden_size,
|
||||
num_heads=self.num_attention_heads,
|
||||
mlp_ratio=self.mlp_ratio,
|
||||
adaln_tembed_dim=config.adaln_tembed_dim,
|
||||
config=config,
|
||||
)
|
||||
for _ in range(self.depth)
|
||||
])
|
||||
|
||||
# Output projection
|
||||
self.final_layer = FinalLayer(
|
||||
hidden_size=self.hidden_size,
|
||||
out_channels=self.out_channels,
|
||||
adaln_tembed_dim=config.adaln_tembed_dim,
|
||||
patch_size=self.patch_size,
|
||||
)
|
||||
|
||||
def enable_bsa(self):
|
||||
"""Enable BSA for all self-attention layers."""
|
||||
for block in self.blocks:
|
||||
block.self_attn.enable_bsa = True
|
||||
|
||||
def disable_bsa(self):
|
||||
"""Disable BSA for all self-attention layers."""
|
||||
for block in self.blocks:
|
||||
block.self_attn.enable_bsa = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor, # [B, C, T, H, W]
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor], # [B, N_text, C_text]
|
||||
timestep: torch.LongTensor, # [B] or [B, T]
|
||||
encoder_attention_mask: torch.Tensor | None = None, # [B, N_text]
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
|
||||
guidance: float | None = None, # Unused, for API compatibility
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass with FastVideo parameter ordering.
|
||||
|
||||
NOTE: This follows FastVideo convention:
|
||||
(hidden_states, encoder_hidden_states, timestep)
|
||||
"""
|
||||
B, _, T, H, W = hidden_states.shape
|
||||
|
||||
N_t = T // self.patch_size[0]
|
||||
N_h = H // self.patch_size[1]
|
||||
N_w = W // self.patch_size[2]
|
||||
|
||||
# Handle list of encoder outputs (take first one)
|
||||
if isinstance(encoder_hidden_states, list):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
# 1. Patch embedding
|
||||
x = self.patch_embed(hidden_states) # [B, N, C]
|
||||
|
||||
# 2. Timestep embedding
|
||||
# Expand timestep from [B] to [B, T] if needed
|
||||
if timestep.ndim == 1:
|
||||
timestep = timestep.unsqueeze(1).expand(-1, N_t) # [B, T]
|
||||
|
||||
t = self.time_embedder(timestep.flatten(), latent_shape=(N_t, N_h, N_w))
|
||||
if t.ndim == 2:
|
||||
t = t.reshape(B, N_t, -1) # [B, T, C_t]
|
||||
|
||||
# 3. Caption embedding (standard format, no compaction)
|
||||
context = self.caption_embedder(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask
|
||||
) # [B, N_text, C]
|
||||
|
||||
# 4. Transformer blocks
|
||||
for i, block in enumerate(self.blocks):
|
||||
x = block(
|
||||
x, context, t,
|
||||
latent_shape=(N_t, N_h, N_w)
|
||||
)
|
||||
|
||||
# 5. Output projection
|
||||
output = self.final_layer(x, t, latent_shape=(N_t, N_h, N_w))
|
||||
|
||||
# Reshape to [B, C_out, T, H, W]
|
||||
output = self.unpatchify(output, N_t, N_h, N_w)
|
||||
|
||||
# Cast to float32 for better accuracy (as per original)
|
||||
output = output.to(torch.float32)
|
||||
|
||||
return output
|
||||
|
||||
def unpatchify(self, x: torch.Tensor, N_t: int, N_h: int, N_w: int) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: [B, N, C] where C = T_p * H_p * W_p * C_out
|
||||
Returns:
|
||||
[B, C_out, T, H, W]
|
||||
"""
|
||||
T_p, H_p, W_p = self.patch_size
|
||||
x = rearrange(
|
||||
x,
|
||||
"B (N_t N_h N_w) (T_p H_p W_p C_out) -> B C_out (N_t T_p) (N_h H_p) (N_w W_p)",
|
||||
N_t=N_t,
|
||||
N_h=N_h,
|
||||
N_w=N_w,
|
||||
T_p=T_p,
|
||||
H_p=H_p,
|
||||
W_p=W_p,
|
||||
C_out=self.out_channels,
|
||||
)
|
||||
return x
|
||||
|
||||
@@ -0,0 +1,387 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from transformers: https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py
|
||||
|
||||
import math
|
||||
from typing import Any, Optional, Tuple, Union, List, Callable, Iterable
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, Qwen2_5_VLConfig
|
||||
from fastvideo.distributed import get_tp_rank, get_tp_world_size
|
||||
from fastvideo.layers.activation import get_act_fn, SiluAndMul
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.linear import MergedColumnParallelLinear, QKVParallelLinear, RowParallelLinear
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.loader.weight_utils import default_weight_loader
|
||||
from fastvideo.models.mask_utils import sdpa_mask
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
|
||||
"""
|
||||
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
|
||||
num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
|
||||
"""
|
||||
batch, num_key_value_heads, slen, head_dim = hidden_states.shape
|
||||
if n_rep == 1:
|
||||
return hidden_states
|
||||
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
|
||||
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
|
||||
|
||||
def sdpa_attention_forward(
|
||||
module: torch.nn.Module,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor],
|
||||
dropout: float = 0.0,
|
||||
scaling: Optional[float] = None,
|
||||
is_causal: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> tuple[torch.Tensor, None]:
|
||||
if kwargs.get("output_attentions", False) or kwargs.get("head_mask") is not None:
|
||||
logger.warning_once(
|
||||
"`sdpa` attention does not support `output_attentions=True` or `head_mask`."
|
||||
" Please set your attention to `eager` if you want any of these features."
|
||||
)
|
||||
|
||||
if hasattr(module, "num_key_value_groups"):
|
||||
key = repeat_kv(key, module.num_key_value_groups)
|
||||
value = repeat_kv(value, module.num_key_value_groups)
|
||||
|
||||
if attention_mask is not None and attention_mask.ndim == 4:
|
||||
attention_mask = attention_mask[:, :, :, : key.shape[-2]]
|
||||
# If attention_mask is not None, convert it to boolean type
|
||||
if attention_mask is not None and attention_mask.dtype != torch.bool:
|
||||
attention_mask = attention_mask.bool()
|
||||
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=attention_mask,
|
||||
dropout_p=dropout,
|
||||
scale=scaling,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
attn_output = attn_output.transpose(1, 2).contiguous()
|
||||
|
||||
return attn_output, None
|
||||
|
||||
def rotate_half(x):
|
||||
"""Rotates half the hidden dims of the input."""
|
||||
x1 = x[..., : x.shape[-1] // 2]
|
||||
x2 = x[..., x.shape[-1] // 2 :]
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
|
||||
def apply_multimodal_rotary_pos_emb(q, k, cos, sin, mrope_section, unsqueeze_dim=1):
|
||||
mrope_section = [s * 2 for s in mrope_section]
|
||||
cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
|
||||
unsqueeze_dim
|
||||
)
|
||||
sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
|
||||
unsqueeze_dim
|
||||
)
|
||||
|
||||
q_embed = (q * cos) + (rotate_half(q) * sin)
|
||||
k_embed = (k * cos) + (rotate_half(k) * sin)
|
||||
return q_embed, k_embed
|
||||
|
||||
|
||||
class Qwen2_5_VLRotaryEmbedding(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, device=None):
|
||||
super().__init__()
|
||||
self.max_seq_len_cached = config.max_position_embeddings
|
||||
self.original_max_seq_len = config.max_position_embeddings
|
||||
self.config = config
|
||||
|
||||
self.rope_type = config.rope_scaling.get("rope_type", "default")
|
||||
self.base = config.rope_theta
|
||||
|
||||
# Simplified initialization
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
dim = head_dim
|
||||
self.attention_scaling = 1.0
|
||||
|
||||
inv_freq = 1.0 / (
|
||||
self.base ** (torch.arange(0, dim, 2, dtype=torch.int64).float().to(device) / dim)
|
||||
)
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
self.original_inv_freq = inv_freq
|
||||
|
||||
def forward(self, x, position_ids):
|
||||
# In contrast to other models, Qwen2_5_VL has different position ids for the grids
|
||||
# So we expand the inv_freq to shape (3, ...)
|
||||
inv_freq_expanded = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1)
|
||||
position_ids_expanded = position_ids[:, :, None, :].float() # shape (3, bs, 1, positions)
|
||||
|
||||
device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
|
||||
with torch.autocast(device_type=device_type, enabled=False): # Force float32
|
||||
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3)
|
||||
emb = torch.cat((freqs, freqs), dim=-1)
|
||||
cos = emb.cos() * self.attention_scaling
|
||||
sin = emb.sin() * self.attention_scaling
|
||||
|
||||
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
|
||||
|
||||
|
||||
class Qwen2_5_VLMLP(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, quant_config: QuantizationConfig | None = None, prefix: str = ""):
|
||||
super().__init__()
|
||||
self.gate_up_proj = MergedColumnParallelLinear(
|
||||
input_size=config.hidden_size,
|
||||
output_sizes=[config.intermediate_size] * 2,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.gate_up_proj",
|
||||
)
|
||||
self.down_proj = RowParallelLinear(
|
||||
input_size=config.intermediate_size,
|
||||
output_size=config.hidden_size,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.down_proj",
|
||||
)
|
||||
self.act_fn = SiluAndMul()
|
||||
|
||||
def forward(self, x):
|
||||
x, _ = self.gate_up_proj(x)
|
||||
x = self.act_fn(x)
|
||||
x, _ = self.down_proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class Qwen2_5_VLAttention(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int, quant_config: QuantizationConfig | None = None, prefix: str = ""):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.layer_idx = layer_idx
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_heads = config.num_attention_heads
|
||||
self.head_dim = self.hidden_size // self.num_heads
|
||||
self.num_key_value_heads = config.num_key_value_heads
|
||||
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
|
||||
|
||||
tp_size = get_tp_world_size()
|
||||
self.total_num_heads = self.num_heads
|
||||
assert self.total_num_heads % tp_size == 0
|
||||
self.num_heads = self.total_num_heads // tp_size
|
||||
|
||||
self.total_num_kv_heads = self.num_key_value_heads
|
||||
if self.total_num_kv_heads >= tp_size:
|
||||
assert self.total_num_kv_heads % tp_size == 0
|
||||
self.num_kv_heads = self.total_num_kv_heads // tp_size
|
||||
else:
|
||||
assert tp_size % self.total_num_kv_heads == 0
|
||||
self.num_kv_heads = 1
|
||||
|
||||
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
|
||||
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.scaling = self.head_dim**-0.5
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=self.hidden_size,
|
||||
head_size=self.head_dim,
|
||||
total_num_heads=self.total_num_heads,
|
||||
total_num_kv_heads=self.total_num_kv_heads,
|
||||
bias=True,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
)
|
||||
|
||||
self.o_proj = RowParallelLinear(
|
||||
input_size=self.total_num_heads * self.head_dim,
|
||||
output_size=self.hidden_size,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
self.layer_type = config.layer_types[layer_idx] if config.layer_types else "full_attention"
|
||||
self.sliding_window = config.sliding_window if self.layer_type == "sliding_attention" else None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
output_attentions: bool = False,
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
bsz, q_len, _ = hidden_states.size()
|
||||
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
query_states, key_states, value_states = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
|
||||
query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
|
||||
key_states = key_states.view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
|
||||
value_states = value_states.view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
|
||||
|
||||
cos, sin = position_embeddings
|
||||
query_states, key_states = apply_multimodal_rotary_pos_emb(
|
||||
query_states, key_states, cos, sin, self.config.rope_scaling["mrope_section"]
|
||||
)
|
||||
|
||||
attn_output = sdpa_attention_forward(self, query_states, key_states, value_states, attention_mask, dropout=self.config.attention_dropout, scaling=self.scaling, is_causal=False)[0].reshape(bsz, q_len, -1)
|
||||
attn_output, _ = self.o_proj(attn_output)
|
||||
|
||||
return attn_output
|
||||
|
||||
class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int, quant_config: QuantizationConfig | None = None, prefix: str = ""):
|
||||
super().__init__()
|
||||
self.self_attn = Qwen2_5_VLAttention(config, layer_idx, quant_config=quant_config, prefix=f"{prefix}.self_attn")
|
||||
self.mlp = Qwen2_5_VLMLP(config, quant_config=quant_config, prefix=f"{prefix}.mlp")
|
||||
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
output_attentions: Optional[bool] = False,
|
||||
) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
|
||||
hidden_states = self.self_attn(
|
||||
hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
position_embeddings=position_embeddings,
|
||||
output_attentions=output_attentions,
|
||||
)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
outputs = (hidden_states,)
|
||||
|
||||
return outputs
|
||||
|
||||
class Qwen2_5_VLTextModel(TextEncoder):
|
||||
def __init__(self, config: Qwen2_5_VLConfig):
|
||||
super().__init__(config)
|
||||
|
||||
quant_config = None
|
||||
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([
|
||||
Qwen2_5_VLDecoderLayer(config, layer_idx, quant_config=quant_config, prefix=f"{config.prefix}.layers.{layer_idx}")
|
||||
for layer_idx in range(config.num_hidden_layers)
|
||||
])
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config)
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.embed_tokens(input_ids)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.get_input_embeddings(input_ids)
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
if position_ids is None:
|
||||
seq_length = hidden_states.shape[1]
|
||||
cache_position = torch.arange(seq_length, device=hidden_states.device)
|
||||
position_ids = cache_position.view(1, 1, -1).expand(3, hidden_states.shape[0], -1)
|
||||
|
||||
mask_kwargs = {
|
||||
"batch_size": hidden_states.shape[0],
|
||||
"cache_position": cache_position,
|
||||
"kv_length": attention_mask.shape[-1],
|
||||
"kv_offset": 0,
|
||||
"attention_mask": attention_mask
|
||||
}
|
||||
|
||||
position_embeddings = self.rotary_emb(hidden_states, position_ids)
|
||||
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
|
||||
for decoder_layer in self.layers:
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
layer_outputs = decoder_layer(
|
||||
hidden_states,
|
||||
attention_mask=sdpa_mask(**mask_kwargs),
|
||||
position_ids=position_ids,
|
||||
position_embeddings=position_embeddings,
|
||||
output_attentions=False,
|
||||
)
|
||||
hidden_states = layer_outputs[0]
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=hidden_states,
|
||||
hidden_states=all_hidden_states,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
for param_name, weight_name, shard_id in self.config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
# if name.endswith(".bias") and name not in params_dict:
|
||||
# continue
|
||||
if name not in params_dict:
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
# if name.endswith(".bias") and name not in params_dict:
|
||||
# continue
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
@@ -16,6 +16,7 @@ import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from torch.distributed import init_device_mesh
|
||||
from transformers import AutoImageProcessor, AutoTokenizer
|
||||
from transformers import UMT5EncoderModel
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.configs.models import EncoderConfig
|
||||
@@ -405,18 +406,16 @@ class VAELoader(ComponentLoader):
|
||||
target_device = get_local_torch_device()
|
||||
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]):
|
||||
fastvideo_args.pipeline_config.vae_precision] if fastvideo_args.pipeline_config.vae_precision else torch.bfloat16):
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(target_device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(model_path), "*.safetensors"))
|
||||
# TODO(PY)
|
||||
assert len(
|
||||
safetensors_list
|
||||
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
loaded = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
vae.load_state_dict(
|
||||
loaded, strict=False) # We might only load encoder or decoder
|
||||
|
||||
@@ -478,8 +477,6 @@ class TransformerLoader(ComponentLoader):
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s, default_dtype: %s", cls_name,
|
||||
default_dtype)
|
||||
assert fastvideo_args.hsdp_shard_dim is not None
|
||||
model = maybe_load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
import torch
|
||||
from typing import Callable, Optional
|
||||
|
||||
def and_masks(*mask_functions: Callable) -> Callable:
|
||||
"""Returns a mask function that is the intersection of provided mask functions"""
|
||||
if not all(callable(arg) for arg in mask_functions):
|
||||
raise RuntimeError(f"All inputs should be callable mask_functions: {mask_functions}")
|
||||
|
||||
def and_mask(batch_idx, head_idx, q_idx, kv_idx):
|
||||
result = q_idx.new_ones((), dtype=torch.bool)
|
||||
for mask in mask_functions:
|
||||
result = result & mask(batch_idx, head_idx, q_idx, kv_idx).to(result.device)
|
||||
return result
|
||||
|
||||
return and_mask
|
||||
|
||||
def causal_mask_function(batch_idx: int, head_idx: int, q_idx: int, kv_idx: int) -> bool:
|
||||
"""
|
||||
This creates a basic lower-diagonal causal mask.
|
||||
"""
|
||||
return kv_idx <= q_idx
|
||||
|
||||
def padding_mask_function(padding_mask: torch.Tensor) -> Callable:
|
||||
"""
|
||||
This return the mask_function function corresponding to a 2D padding mask.
|
||||
"""
|
||||
|
||||
def inner_mask(batch_idx: int, head_idx: int, q_idx: int, kv_idx: int) -> bool:
|
||||
# Note that here the mask should ALWAYS be at least of the max `kv_index` size in the dimension 1. This is because
|
||||
# we cannot pad it here in the mask_function as we don't know the final size, and we cannot try/except, as it is not
|
||||
# vectorizable on accelerator devices
|
||||
return padding_mask[batch_idx, kv_idx]
|
||||
|
||||
return inner_mask
|
||||
|
||||
def prepare_padding_mask(
|
||||
attention_mask: Optional[torch.Tensor], kv_length: int, kv_offset: int
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
From the 2D attention mask, prepare the correct padding mask to use by potentially padding it.
|
||||
"""
|
||||
local_padding_mask = attention_mask
|
||||
if attention_mask is not None:
|
||||
# Pad it if necessary
|
||||
if (padding_length := kv_length + kv_offset - attention_mask.shape[-1]) > 0:
|
||||
local_padding_mask = torch.nn.functional.pad(attention_mask, (0, padding_length))
|
||||
return local_padding_mask
|
||||
|
||||
def _non_vmap_expansion_sdpa(
|
||||
batch_indices: torch.Tensor, head_indices: torch.Tensor, q_indices: torch.Tensor, kv_indices: torch.Tensor
|
||||
):
|
||||
"""
|
||||
Used to broadcast our mask_functions over the all 4 dimensions (b_idx, h_idx, q_idx, kv_idx) of the inputs.
|
||||
Allows the usage of any index-based mask function without relying on vmap.
|
||||
|
||||
NOTE: This is limited to index based functions only and is not guaranteed to work otherwise.
|
||||
|
||||
Reference:
|
||||
- https://github.com/huggingface/optimum-onnx/blob/c123e8f4fab61b54a8e0e31ce74462bcacca576e/optimum/exporters/onnx/model_patcher.py#L362-L365
|
||||
"""
|
||||
batch_indices = batch_indices[:, None, None, None]
|
||||
head_indices = head_indices[None, :, None, None]
|
||||
q_indices = q_indices[None, None, :, None]
|
||||
kv_indices = kv_indices[None, None, None, :]
|
||||
return batch_indices, head_indices, q_indices, kv_indices
|
||||
|
||||
def sdpa_mask(
|
||||
batch_size: int,
|
||||
cache_position: torch.Tensor,
|
||||
kv_length: int,
|
||||
kv_offset: int = 0,
|
||||
mask_function: Callable = causal_mask_function,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
local_size: Optional[int] = None,
|
||||
allow_is_causal_skip: bool = True,
|
||||
allow_is_bidirectional_skip: bool = False,
|
||||
allow_torch_fix: bool = True,
|
||||
use_vmap: bool = False,
|
||||
**kwargs,
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
Create a 4D boolean mask of shape `(batch_size, 1, query_length, kv_length)` where a value of True indicates that
|
||||
the element should take part in the attention computation, and False that it should not.
|
||||
This function can only be used with torch>=2.5, as the context manager is otherwise not available.
|
||||
|
||||
Args:
|
||||
batch_size (`int`):
|
||||
The batch size of the input sequence.
|
||||
cache_position (`torch.Tensor`):
|
||||
A tensor of shape (query_length,) indicating the current indices of the input sequence elements.
|
||||
kv_length (`int`):
|
||||
The size that the key and value states will have during the attention computation.
|
||||
kv_offset (`int`, optional):
|
||||
An optional offset to indicate at which first position the key and values states will refer to.
|
||||
mask_function (`Callable`):
|
||||
The mask factory function describing the mask pattern.
|
||||
attention_mask (`torch.Tensor`, optional):
|
||||
The 2D attention mask corresponding to padded tokens of shape (batch_size, number_of_seen_tokens+q_length)
|
||||
local_size (`int`, optional):
|
||||
The size of the local attention, if we do not use full attention. This is used only if `allow_is_causal_skip=True`
|
||||
to try to skip mask creation if possible.
|
||||
allow_is_causal_skip (`bool`, optional):
|
||||
Whether to allow to return `None` for the mask under conditions where we can use the `is_causal` argument in
|
||||
`torch.sdpa` instead. Default to `True`.
|
||||
allow_is_bidirectional_skip (`bool`, optional):
|
||||
Whether to allow to return `None` for the mask under conditions where we do not have to add any bias,
|
||||
i.e. full attention without any padding. Default to `False`.
|
||||
allow_torch_fix (`bool`, optional):
|
||||
Whether to update the mask in case a query is not attending to any tokens, to solve a bug in torch's older
|
||||
versions. We need an arg to skip it when using eager. By default `True`.
|
||||
use_vmap (`bool`, optional):
|
||||
Whether to use `vmap` during the mask construction or not. Allows powerful custom patterns that may not be
|
||||
index-based (for the cost of speed performance). By default `False`.
|
||||
|
||||
|
||||
## Creating a simple causal mask:
|
||||
|
||||
To create the following causal mask:
|
||||
|
||||
0 ■ ⬚ ⬚ ⬚ ⬚
|
||||
1 ■ ■ ⬚ ⬚ ⬚
|
||||
2 ■ ■ ■ ⬚ ⬚
|
||||
3 ■ ■ ■ ■ ⬚
|
||||
4 ■ ■ ■ ■ ■
|
||||
|
||||
You can do
|
||||
|
||||
```python
|
||||
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5)
|
||||
>>> tensor([[[[ True, False, False, False, False],
|
||||
[ True, True, False, False, False],
|
||||
[ True, True, True, False, False],
|
||||
[ True, True, True, True, False],
|
||||
[ True, True, True, True, True]]]])
|
||||
```
|
||||
|
||||
## Creating a sliding window mask:
|
||||
|
||||
To create the following sliding window mask (`sliding_window=3`):
|
||||
|
||||
0 ■ ⬚ ⬚ ⬚ ⬚
|
||||
1 ■ ■ ⬚ ⬚ ⬚
|
||||
2 ■ ■ ■ ⬚ ⬚
|
||||
3 ⬚ ■ ■ ■ ⬚
|
||||
4 ⬚ ⬚ ■ ■ ■
|
||||
|
||||
You can do
|
||||
|
||||
```python
|
||||
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5, mask_function=sliding_window_causal_mask_function(3))
|
||||
>>> tensor([[[[ True, False, False, False, False],
|
||||
[ True, True, False, False, False],
|
||||
[ True, True, True, False, False],
|
||||
[False, True, True, True, False],
|
||||
[False, False, True, True, True]]]])
|
||||
```
|
||||
|
||||
## Creating a chunked attention mask
|
||||
|
||||
To create the following chunked attention mask (`chunk_size=3`):
|
||||
|
||||
0 ■ ⬚ ⬚ ⬚ ⬚
|
||||
1 ■ ■ ⬚ ⬚ ⬚
|
||||
2 ■ ■ ■ ⬚ ⬚
|
||||
3 ⬚ ⬚ ⬚ ■ ⬚
|
||||
4 ⬚ ⬚ ⬚ ■ ■
|
||||
|
||||
You can do
|
||||
|
||||
```python
|
||||
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5, mask_function=chunked_causal_mask_function(3, torch.zeros(1, dtype=int)))
|
||||
>>> tensor([[[[ True, False, False, False, False],
|
||||
[ True, True, False, False, False],
|
||||
[ True, True, True, False, False],
|
||||
[False, False, False, True, False],
|
||||
[False, False, False, True, True]]]])
|
||||
```
|
||||
|
||||
"""
|
||||
q_length = cache_position.shape[0]
|
||||
|
||||
# Potentially pad the 2D mask
|
||||
padding_mask = prepare_padding_mask(attention_mask, kv_length, kv_offset)
|
||||
|
||||
# Potentially add the padding 2D mask
|
||||
if padding_mask is not None:
|
||||
mask_function = and_masks(mask_function, padding_mask_function(padding_mask))
|
||||
|
||||
batch_arange = torch.arange(batch_size, device=cache_position.device)
|
||||
head_arange = torch.arange(1, device=cache_position.device)
|
||||
# Similar to `kv_arange = torch.arange(start=kv_offset, end=kv_offset + kv_length, device=cache_position.device)`
|
||||
# but without data-dependent slicing (i.e. torch.compile friendly)
|
||||
kv_arange = torch.arange(kv_length, device=cache_position.device) + kv_offset
|
||||
|
||||
# Actual mask creation
|
||||
# Apply mask function element-wise through broadcasting
|
||||
attention_mask = mask_function(*_non_vmap_expansion_sdpa(batch_arange, head_arange, cache_position, kv_arange))
|
||||
# Expand the mask to match batch size and query length if they weren't used in the mask function
|
||||
attention_mask = attention_mask.expand(batch_size, -1, q_length, kv_length)
|
||||
|
||||
return attention_mask
|
||||
@@ -24,10 +24,14 @@ logger = init_logger(__name__)
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"HunyuanVideoTransformer3DModel":
|
||||
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
|
||||
"HunyuanVideo15Transformer3DModel":
|
||||
("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel")
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -45,6 +49,7 @@ _TEXT_ENCODER_MODELS = {
|
||||
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
|
||||
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
|
||||
"BertModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
@@ -56,6 +61,7 @@ _IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
_VAE_MODELS = {
|
||||
"AutoencoderKLHunyuanVideo":
|
||||
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,703 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from diffusers
|
||||
|
||||
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
|
||||
from fastvideo.layers.activation import get_act_fn
|
||||
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
|
||||
from fastvideo.models.vaes.common import ParallelTiledVAE
|
||||
|
||||
class HunyuanVideo15CausalConv3d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size: Union[int, Tuple[int, int, int]] = 3,
|
||||
stride: Union[int, Tuple[int, int, int]] = 1,
|
||||
padding: Union[int, Tuple[int, int, int]] = 0,
|
||||
dilation: Union[int, Tuple[int, int, int]] = 1,
|
||||
bias: bool = True,
|
||||
pad_mode: str = "replicate",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
|
||||
|
||||
self.pad_mode = pad_mode
|
||||
self.time_causal_padding = (
|
||||
kernel_size[0] // 2,
|
||||
kernel_size[0] // 2,
|
||||
kernel_size[1] // 2,
|
||||
kernel_size[1] // 2,
|
||||
kernel_size[2] - 1,
|
||||
0,
|
||||
)
|
||||
|
||||
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode)
|
||||
return self.conv(hidden_states)
|
||||
|
||||
|
||||
class HunyuanVideo15RMS_norm(nn.Module):
|
||||
r"""
|
||||
A custom RMS normalization layer.
|
||||
|
||||
Args:
|
||||
dim (int): The number of dimensions to normalize over.
|
||||
channel_first (bool, optional): Whether the input tensor has channels as the first dimension.
|
||||
Default is True.
|
||||
images (bool, optional): Whether the input represents image data. Default is True.
|
||||
bias (bool, optional): Whether to include a learnable bias term. Default is False.
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None:
|
||||
super().__init__()
|
||||
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
|
||||
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
|
||||
|
||||
self.channel_first = channel_first
|
||||
self.scale = dim**0.5
|
||||
self.gamma = nn.Parameter(torch.ones(shape))
|
||||
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
|
||||
|
||||
def forward(self, x):
|
||||
return F.normalize(x, dim=(1 if self.channel_first else -1)) * self.scale * self.gamma + self.bias
|
||||
|
||||
|
||||
class HunyuanVideo15AttnBlock(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = HunyuanVideo15RMS_norm(in_channels, images=False)
|
||||
|
||||
self.to_q = nn.Conv3d(in_channels, in_channels, kernel_size=1)
|
||||
self.to_k = nn.Conv3d(in_channels, in_channels, kernel_size=1)
|
||||
self.to_v = nn.Conv3d(in_channels, in_channels, kernel_size=1)
|
||||
self.proj_out = nn.Conv3d(in_channels, in_channels, kernel_size=1)
|
||||
|
||||
@staticmethod
|
||||
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
|
||||
"""Prepare a causal attention mask for 3D videos.
|
||||
|
||||
Args:
|
||||
n_frame (int): Number of frames (temporal length).
|
||||
n_hw (int): Product of height and width.
|
||||
dtype: Desired mask dtype.
|
||||
device: Device for the mask.
|
||||
batch_size (int, optional): If set, expands for batch.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Causal attention mask.
|
||||
"""
|
||||
seq_len = n_frame * n_hw
|
||||
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
|
||||
for i in range(seq_len):
|
||||
i_frame = i // n_hw
|
||||
mask[i, : (i_frame + 1) * n_hw] = 0
|
||||
if batch_size is not None:
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
return mask
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
identity = x
|
||||
|
||||
x = self.norm(x)
|
||||
|
||||
query = self.to_q(x)
|
||||
key = self.to_k(x)
|
||||
value = self.to_v(x)
|
||||
|
||||
batch_size, channels, frames, height, width = query.shape
|
||||
|
||||
query = query.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
|
||||
key = key.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
|
||||
value = value.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
|
||||
|
||||
attention_mask = self.prepare_causal_attention_mask(
|
||||
frames, height * width, query.dtype, query.device, batch_size=batch_size
|
||||
)
|
||||
|
||||
x = nn.functional.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask)
|
||||
|
||||
# batch_size, 1, frames * height * width, channels
|
||||
|
||||
x = x.squeeze(1).reshape(batch_size, frames, height, width, channels).permute(0, 4, 1, 2, 3)
|
||||
x = self.proj_out(x)
|
||||
|
||||
return x + identity
|
||||
|
||||
|
||||
class HunyuanVideo15Upsample(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, add_temporal_upsample: bool = True):
|
||||
super().__init__()
|
||||
factor = 2 * 2 * 2 if add_temporal_upsample else 1 * 2 * 2
|
||||
self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels * factor, kernel_size=3)
|
||||
|
||||
self.add_temporal_upsample = add_temporal_upsample
|
||||
self.repeats = factor * out_channels // in_channels
|
||||
|
||||
@staticmethod
|
||||
def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2):
|
||||
"""
|
||||
Convert (b, r1*r2*r3*c, f, h, w) -> (b, c, r1*f, r2*h, r3*w)
|
||||
|
||||
Args:
|
||||
tensor: Input tensor of shape (b, r1*r2*r3*c, f, h, w)
|
||||
r1: temporal upsampling factor
|
||||
r2: height upsampling factor
|
||||
r3: width upsampling factor
|
||||
"""
|
||||
b, packed_c, f, h, w = tensor.shape
|
||||
factor = r1 * r2 * r3
|
||||
c = packed_c // factor
|
||||
|
||||
tensor = tensor.view(b, r1, r2, r3, c, f, h, w)
|
||||
tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3)
|
||||
return tensor.reshape(b, c, f * r1, h * r2, w * r3)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
r1 = 2 if self.add_temporal_upsample else 1
|
||||
h = self.conv(x)
|
||||
if self.add_temporal_upsample:
|
||||
h_first = h[:, :, :1, :, :]
|
||||
h_first = self._dcae_upsample_rearrange(h_first, r1=1, r2=2, r3=2)
|
||||
h_first = h_first[:, : h_first.shape[1] // 2]
|
||||
h_next = h[:, :, 1:, :, :]
|
||||
h_next = self._dcae_upsample_rearrange(h_next, r1=r1, r2=2, r3=2)
|
||||
h = torch.cat([h_first, h_next], dim=2)
|
||||
|
||||
# shortcut computation
|
||||
x_first = x[:, :, :1, :, :]
|
||||
x_first = self._dcae_upsample_rearrange(x_first, r1=1, r2=2, r3=2)
|
||||
x_first = x_first.repeat_interleave(repeats=self.repeats // 2, dim=1)
|
||||
|
||||
x_next = x[:, :, 1:, :, :]
|
||||
x_next = self._dcae_upsample_rearrange(x_next, r1=r1, r2=2, r3=2)
|
||||
x_next = x_next.repeat_interleave(repeats=self.repeats, dim=1)
|
||||
shortcut = torch.cat([x_first, x_next], dim=2)
|
||||
|
||||
else:
|
||||
h = self._dcae_upsample_rearrange(h, r1=r1, r2=2, r3=2)
|
||||
shortcut = x.repeat_interleave(repeats=self.repeats, dim=1)
|
||||
shortcut = self._dcae_upsample_rearrange(shortcut, r1=r1, r2=2, r3=2)
|
||||
return h + shortcut
|
||||
|
||||
|
||||
class HunyuanVideo15Downsample(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, add_temporal_downsample: bool = True):
|
||||
super().__init__()
|
||||
factor = 2 * 2 * 2 if add_temporal_downsample else 1 * 2 * 2
|
||||
self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels // factor, kernel_size=3)
|
||||
|
||||
self.add_temporal_downsample = add_temporal_downsample
|
||||
self.group_size = factor * in_channels // out_channels
|
||||
|
||||
@staticmethod
|
||||
def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2):
|
||||
"""
|
||||
Convert (b, c, r1*f, r2*h, r3*w) -> (b, r1*r2*r3*c, f, h, w)
|
||||
|
||||
This packs spatial/temporal dimensions into channels (opposite of upsample)
|
||||
"""
|
||||
b, c, packed_f, packed_h, packed_w = tensor.shape
|
||||
f, h, w = packed_f // r1, packed_h // r2, packed_w // r3
|
||||
|
||||
tensor = tensor.view(b, c, f, r1, h, r2, w, r3)
|
||||
tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6)
|
||||
return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
r1 = 2 if self.add_temporal_downsample else 1
|
||||
h = self.conv(x)
|
||||
if self.add_temporal_downsample:
|
||||
h_first = h[:, :, :1, :, :]
|
||||
h_first = self._dcae_downsample_rearrange(h_first, r1=1, r2=2, r3=2)
|
||||
h_first = torch.cat([h_first, h_first], dim=1)
|
||||
h_next = h[:, :, 1:, :, :]
|
||||
h_next = self._dcae_downsample_rearrange(h_next, r1=r1, r2=2, r3=2)
|
||||
h = torch.cat([h_first, h_next], dim=2)
|
||||
|
||||
# shortcut computation
|
||||
x_first = x[:, :, :1, :, :]
|
||||
x_first = self._dcae_downsample_rearrange(x_first, r1=1, r2=2, r3=2)
|
||||
B, C, T, H, W = x_first.shape
|
||||
x_first = x_first.view(B, h.shape[1], self.group_size // 2, T, H, W).mean(dim=2)
|
||||
x_next = x[:, :, 1:, :, :]
|
||||
x_next = self._dcae_downsample_rearrange(x_next, r1=r1, r2=2, r3=2)
|
||||
B, C, T, H, W = x_next.shape
|
||||
x_next = x_next.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
|
||||
shortcut = torch.cat([x_first, x_next], dim=2)
|
||||
else:
|
||||
h = self._dcae_downsample_rearrange(h, r1=r1, r2=2, r3=2)
|
||||
shortcut = self._dcae_downsample_rearrange(x, r1=r1, r2=2, r3=2)
|
||||
B, C, T, H, W = shortcut.shape
|
||||
shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
|
||||
|
||||
return h + shortcut
|
||||
|
||||
|
||||
class HunyuanVideo15ResnetBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
non_linearity: str = "swish",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
self.nonlinearity = get_act_fn(non_linearity)
|
||||
|
||||
self.norm1 = HunyuanVideo15RMS_norm(in_channels, images=False)
|
||||
self.conv1 = HunyuanVideo15CausalConv3d(in_channels, out_channels, kernel_size=3)
|
||||
|
||||
self.norm2 = HunyuanVideo15RMS_norm(out_channels, images=False)
|
||||
self.conv2 = HunyuanVideo15CausalConv3d(out_channels, out_channels, kernel_size=3)
|
||||
|
||||
self.conv_shortcut = None
|
||||
if in_channels != out_channels:
|
||||
self.conv_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
residual = self.conv_shortcut(residual)
|
||||
|
||||
return hidden_states + residual
|
||||
|
||||
|
||||
class HunyuanVideo15MidBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
num_layers: int = 1,
|
||||
add_attention: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.add_attention = add_attention
|
||||
|
||||
# There is always at least one resnet
|
||||
resnets = [
|
||||
HunyuanVideo15ResnetBlock(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
)
|
||||
]
|
||||
attentions = []
|
||||
|
||||
for _ in range(num_layers):
|
||||
if self.add_attention:
|
||||
attentions.append(HunyuanVideo15AttnBlock(in_channels))
|
||||
else:
|
||||
attentions.append(None)
|
||||
|
||||
resnets.append(
|
||||
HunyuanVideo15ResnetBlock(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
)
|
||||
)
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.resnets[0](hidden_states)
|
||||
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
hidden_states = attn(hidden_states)
|
||||
hidden_states = resnet(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideo15DownBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
num_layers: int = 1,
|
||||
downsample_out_channels: Optional[int] = None,
|
||||
add_temporal_downsample: int = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
resnets = []
|
||||
|
||||
for i in range(num_layers):
|
||||
in_channels = in_channels if i == 0 else out_channels
|
||||
resnets.append(
|
||||
HunyuanVideo15ResnetBlock(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if downsample_out_channels is not None:
|
||||
self.downsamplers = nn.ModuleList(
|
||||
[
|
||||
HunyuanVideo15Downsample(
|
||||
out_channels,
|
||||
out_channels=downsample_out_channels,
|
||||
add_temporal_downsample=add_temporal_downsample,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states)
|
||||
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = downsampler(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideo15UpBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
num_layers: int = 1,
|
||||
upsample_out_channels: Optional[int] = None,
|
||||
add_temporal_upsample: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
resnets = []
|
||||
|
||||
for i in range(num_layers):
|
||||
input_channels = in_channels if i == 0 else out_channels
|
||||
|
||||
resnets.append(
|
||||
HunyuanVideo15ResnetBlock(
|
||||
in_channels=input_channels,
|
||||
out_channels=out_channels,
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if upsample_out_channels is not None:
|
||||
self.upsamplers = nn.ModuleList(
|
||||
[
|
||||
HunyuanVideo15Upsample(
|
||||
out_channels,
|
||||
out_channels=upsample_out_channels,
|
||||
add_temporal_upsample=add_temporal_upsample,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.upsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
|
||||
|
||||
else:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = upsampler(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideo15Encoder3D(nn.Module):
|
||||
r"""
|
||||
3D vae encoder for HunyuanImageRefiner.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 64,
|
||||
block_out_channels: Tuple[int, ...] = (128, 256, 512, 1024, 1024),
|
||||
layers_per_block: int = 2,
|
||||
temporal_compression_ratio: int = 4,
|
||||
spatial_compression_ratio: int = 16,
|
||||
downsample_match_channel: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.group_size = block_out_channels[-1] // self.out_channels
|
||||
|
||||
self.conv_in = HunyuanVideo15CausalConv3d(in_channels, block_out_channels[0], kernel_size=3)
|
||||
self.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
input_channel = block_out_channels[0]
|
||||
for i in range(len(block_out_channels)):
|
||||
add_spatial_downsample = i < np.log2(spatial_compression_ratio)
|
||||
output_channel = block_out_channels[i]
|
||||
if not add_spatial_downsample:
|
||||
down_block = HunyuanVideo15DownBlock3D(
|
||||
num_layers=layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
downsample_out_channels=None,
|
||||
add_temporal_downsample=False,
|
||||
)
|
||||
input_channel = output_channel
|
||||
else:
|
||||
add_temporal_downsample = i >= np.log2(spatial_compression_ratio // temporal_compression_ratio)
|
||||
downsample_out_channels = block_out_channels[i + 1] if downsample_match_channel else output_channel
|
||||
down_block = HunyuanVideo15DownBlock3D(
|
||||
num_layers=layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
downsample_out_channels=downsample_out_channels,
|
||||
add_temporal_downsample=add_temporal_downsample,
|
||||
)
|
||||
input_channel = downsample_out_channels
|
||||
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[-1])
|
||||
|
||||
self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False)
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.conv_in(hidden_states)
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for down_block in self.down_blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(down_block, hidden_states)
|
||||
|
||||
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
|
||||
else:
|
||||
for down_block in self.down_blocks:
|
||||
hidden_states = down_block(hidden_states)
|
||||
|
||||
hidden_states = self.mid_block(hidden_states)
|
||||
|
||||
batch_size, _, frame, height, width = hidden_states.shape
|
||||
short_cut = hidden_states.view(batch_size, -1, self.group_size, frame, height, width).mean(dim=2)
|
||||
|
||||
hidden_states = self.norm_out(hidden_states)
|
||||
hidden_states = self.conv_act(hidden_states)
|
||||
hidden_states = self.conv_out(hidden_states)
|
||||
|
||||
hidden_states += short_cut
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideo15Decoder3D(nn.Module):
|
||||
r"""
|
||||
Causal decoder for 3D video-like data used for HunyuanImage-1.5 Refiner.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 32,
|
||||
out_channels: int = 3,
|
||||
block_out_channels: Tuple[int, ...] = (1024, 1024, 512, 256, 128),
|
||||
layers_per_block: int = 2,
|
||||
spatial_compression_ratio: int = 16,
|
||||
temporal_compression_ratio: int = 4,
|
||||
upsample_match_channel: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.repeat = block_out_channels[0] // self.in_channels
|
||||
|
||||
self.conv_in = HunyuanVideo15CausalConv3d(self.in_channels, block_out_channels[0], kernel_size=3)
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
# mid
|
||||
self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[0])
|
||||
|
||||
# up
|
||||
input_channel = block_out_channels[0]
|
||||
for i in range(len(block_out_channels)):
|
||||
output_channel = block_out_channels[i]
|
||||
|
||||
add_spatial_upsample = i < np.log2(spatial_compression_ratio)
|
||||
add_temporal_upsample = i < np.log2(temporal_compression_ratio)
|
||||
if add_spatial_upsample or add_temporal_upsample:
|
||||
upsample_out_channels = block_out_channels[i + 1] if upsample_match_channel else output_channel
|
||||
up_block = HunyuanVideo15UpBlock3D(
|
||||
num_layers=self.layers_per_block + 1,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
upsample_out_channels=upsample_out_channels,
|
||||
add_temporal_upsample=add_temporal_upsample,
|
||||
)
|
||||
input_channel = upsample_out_channels
|
||||
else:
|
||||
up_block = HunyuanVideo15UpBlock3D(
|
||||
num_layers=self.layers_per_block + 1,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
upsample_out_channels=None,
|
||||
add_temporal_upsample=False,
|
||||
)
|
||||
input_channel = output_channel
|
||||
|
||||
self.up_blocks.append(up_block)
|
||||
|
||||
# out
|
||||
self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False)
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.conv_in(hidden_states) + hidden_states.repeat_interleave(repeats=self.repeat, dim=1)
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
|
||||
|
||||
for up_block in self.up_blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(up_block, hidden_states)
|
||||
else:
|
||||
hidden_states = self.mid_block(hidden_states)
|
||||
|
||||
for up_block in self.up_blocks:
|
||||
hidden_states = up_block(hidden_states)
|
||||
|
||||
# post-process
|
||||
hidden_states = self.norm_out(hidden_states)
|
||||
hidden_states = self.conv_act(hidden_states)
|
||||
hidden_states = self.conv_out(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
|
||||
r"""
|
||||
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used for
|
||||
HunyuanVideo-1.5.
|
||||
|
||||
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
|
||||
for all models (such as downloading or saving).
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Hunyuan15VAEConfig,
|
||||
) -> None:
|
||||
nn.Module.__init__(self)
|
||||
ParallelTiledVAE.__init__(self, config)
|
||||
|
||||
if config.load_encoder:
|
||||
self.encoder = HunyuanVideo15Encoder3D(
|
||||
in_channels=config.in_channels,
|
||||
out_channels=config.latent_channels * 2,
|
||||
block_out_channels=config.block_out_channels,
|
||||
layers_per_block=config.layers_per_block,
|
||||
temporal_compression_ratio=config.temporal_compression_ratio,
|
||||
spatial_compression_ratio=config.spatial_compression_ratio,
|
||||
downsample_match_channel=config.downsample_match_channel,
|
||||
)
|
||||
|
||||
if config.load_decoder:
|
||||
self.decoder = HunyuanVideo15Decoder3D(
|
||||
in_channels=config.latent_channels,
|
||||
out_channels=config.out_channels,
|
||||
block_out_channels=list(reversed(config.block_out_channels)),
|
||||
layers_per_block=config.layers_per_block,
|
||||
temporal_compression_ratio=config.temporal_compression_ratio,
|
||||
spatial_compression_ratio=config.spatial_compression_ratio,
|
||||
upsample_match_channel=config.upsample_match_channel,
|
||||
)
|
||||
|
||||
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
|
||||
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
|
||||
# intermediate tiles together, the memory requirement can be lowered.
|
||||
self.use_tiling = False
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
self.tile_sample_min_width = 256
|
||||
self.tile_sample_min_num_frames = 2000 # Fill in a random large number, as hy1.5 vae does not use temporal tiling
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.encoder(x)
|
||||
return x
|
||||
|
||||
def _decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
sample_posterior: bool = False,
|
||||
return_dict: bool = True,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Args:
|
||||
sample (`torch.Tensor`): Input sample.
|
||||
sample_posterior (`bool`, *optional*, defaults to `False`):
|
||||
Whether to sample from the posterior.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
|
||||
"""
|
||||
x = sample
|
||||
posterior = self.encode(x).latent_dist
|
||||
if sample_posterior:
|
||||
z = posterior.sample(generator=generator)
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z)
|
||||
return dec
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Hunyuan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Hunyuan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
Hy15ImageEncodingStage)
|
||||
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class HunyuanVideo15Pipeline(ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
|
||||
"transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage_primary",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2")
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2")
|
||||
],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=Hy15ImageEncodingStage(image_encoder=None,
|
||||
image_processor=None))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = HunyuanVideo15Pipeline
|
||||
@@ -0,0 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LongCat pipeline module."""
|
||||
|
||||
from fastvideo.pipelines.basic.longcat.longcat_pipeline import LongCatPipeline
|
||||
|
||||
__all__ = ["LongCatPipeline"]
|
||||
@@ -0,0 +1,145 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat video diffusion pipeline implementation (Phase 1: Wrapper).
|
||||
|
||||
This module contains a wrapper implementation of the LongCat video diffusion pipeline
|
||||
using FastVideo's modular pipeline architecture with the original LongCat modules.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
|
||||
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
LongCat video diffusion pipeline with LoRA support.
|
||||
|
||||
Phase 1 implementation using wrapper modules from third_party/longcat_video.
|
||||
This validates the pipeline infrastructure before full FastVideo integration.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize LongCat-specific components."""
|
||||
|
||||
# Enable BSA (Block Sparse Attention) if configured
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
transformer = self.get_module("transformer", None)
|
||||
if transformer is None:
|
||||
raise RuntimeError(
|
||||
"Transformer module not found during initializing LongCat pipeline."
|
||||
)
|
||||
# If user toggles BSA via CLI/config
|
||||
if pipeline_config.enable_bsa:
|
||||
# Build effective BSA params:
|
||||
# 1) from explicit CLI overrides if provided
|
||||
# 2) else from pipeline_config.bsa_params
|
||||
# 3) else fall back to reasonable defaults
|
||||
bsa_params_cfg = pipeline_config.bsa_params
|
||||
sparsity = pipeline_config.bsa_sparsity
|
||||
cdf_threshold = pipeline_config.bsa_cdf_threshold
|
||||
chunk_q = pipeline_config.bsa_chunk_q
|
||||
chunk_k = pipeline_config.bsa_chunk_k
|
||||
|
||||
effective_bsa_params = dict(bsa_params_cfg) if isinstance(
|
||||
bsa_params_cfg, dict) else {}
|
||||
if sparsity is not None:
|
||||
effective_bsa_params['sparsity'] = sparsity
|
||||
if cdf_threshold is not None:
|
||||
effective_bsa_params['cdf_threshold'] = cdf_threshold
|
||||
if chunk_q is not None:
|
||||
effective_bsa_params['chunk_3d_shape_q'] = chunk_q
|
||||
if chunk_k is not None:
|
||||
effective_bsa_params['chunk_3d_shape_k'] = chunk_k
|
||||
# Provide defaults if still missing
|
||||
effective_bsa_params.setdefault('sparsity', 0.9375)
|
||||
effective_bsa_params.setdefault('chunk_3d_shape_q', [4, 4, 4])
|
||||
effective_bsa_params.setdefault('chunk_3d_shape_k', [4, 4, 4])
|
||||
|
||||
if hasattr(transformer, 'enable_bsa'):
|
||||
logger.info(
|
||||
"Enabling Block Sparse Attention (BSA) for LongCat transformer"
|
||||
)
|
||||
transformer.enable_bsa()
|
||||
# Propagate params to all attention modules
|
||||
if hasattr(transformer, 'blocks'):
|
||||
try:
|
||||
for blk in transformer.blocks:
|
||||
if hasattr(blk, 'self_attn'):
|
||||
blk.self_attn.bsa_params = effective_bsa_params
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to set BSA params on all blocks: %s", e)
|
||||
logger.info("BSA parameters in effect: %s",
|
||||
effective_bsa_params)
|
||||
else:
|
||||
logger.warning(
|
||||
"BSA is enabled in config but transformer does not support it"
|
||||
)
|
||||
else:
|
||||
# Explicitly disable if present
|
||||
if hasattr(transformer, 'disable_bsa'):
|
||||
transformer.disable_bsa()
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
# Add refine initialization stage (will be skipped if not refining)
|
||||
self.add_stage(stage_name="longcat_refine_init_stage",
|
||||
stage=LongCatRefineInitStage(vae=self.get_module("vae")))
|
||||
|
||||
# First prepare generic timesteps (for non-refine paths)
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
# Then override timesteps for refinement (will be a no-op if not refining),
|
||||
# matching LongCat's generate_refine schedule.
|
||||
self.add_stage(stage_name="longcat_refine_timestep_stage",
|
||||
stage=LongCatRefineTimestepStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=LongCatDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"),
|
||||
pipeline=self))
|
||||
|
||||
|
||||
EntryClass = LongCatPipeline
|
||||
@@ -287,6 +287,7 @@ class ComposedPipelineBase(ABC):
|
||||
# remove keys that are not pipeline modules
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
model_index.pop("workload_type", None)
|
||||
if "boundary_ratio" in model_index and model_index[
|
||||
"boundary_ratio"] is not None:
|
||||
logger.info(
|
||||
|
||||
@@ -91,6 +91,14 @@ class ForwardBatch:
|
||||
video_path: str | None = None
|
||||
video_latent: torch.Tensor | None = None
|
||||
|
||||
# Refine inputs (LongCat)
|
||||
refine_from: str | None = None
|
||||
t_thresh: float = 0.5
|
||||
spatial_refine_only: bool = False
|
||||
num_cond_frames: int = 0
|
||||
stage1_video: list[
|
||||
PIL.Image.Image] | None = None # Loaded frames from refine_from
|
||||
|
||||
# Primary encoder embeddings
|
||||
prompt_embeds: list[torch.Tensor] = field(default_factory=list)
|
||||
negative_prompt_embeds: list[torch.Tensor] | None = None
|
||||
|
||||
@@ -25,9 +25,11 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanCausalDMDPipeline": "wan",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
"HunyuanVideo15Pipeline": "hunyuan15",
|
||||
"Cosmos2VideoToWorldPipeline": "cosmos",
|
||||
"MatrixGamePipeline": "matrixgame",
|
||||
"MatrixGameCausalDMDPipeline": "matrixgame",
|
||||
"LongCatPipeline": "longcat",
|
||||
}
|
||||
|
||||
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
|
||||
|
||||
@@ -16,7 +16,7 @@ from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
|
||||
from fastvideo.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.pipelines.stages.image_encoding import (
|
||||
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
|
||||
ImageVAEEncodingStage, VideoVAEEncodingStage)
|
||||
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage)
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.pipelines.stages.latent_preparation import (
|
||||
CosmosLatentPreparationStage, LatentPreparationStage)
|
||||
@@ -44,6 +44,7 @@ __all__ = [
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
"MatrixGameImageEncodingStage",
|
||||
"Hy15ImageEncodingStage",
|
||||
"RefImageEncodingStage",
|
||||
"ImageVAEEncodingStage",
|
||||
"VideoVAEEncodingStage",
|
||||
|
||||
@@ -187,6 +187,21 @@ class DecodingStage(PipelineStage):
|
||||
# Convert to CPU float32 for compatibility
|
||||
frames = frames.cpu().float()
|
||||
|
||||
# Crop padding if this is a LongCat refinement
|
||||
if hasattr(batch, 'num_cond_frames_added') and hasattr(
|
||||
batch, 'new_frame_size_before_padding'):
|
||||
num_cond_frames_added = batch.num_cond_frames_added
|
||||
new_frame_size = batch.new_frame_size_before_padding
|
||||
if num_cond_frames_added > 0 or frames.shape[2] != new_frame_size:
|
||||
# frames is [B, C, T, H, W], crop temporal dimension
|
||||
frames = frames[:, :,
|
||||
num_cond_frames_added:num_cond_frames_added +
|
||||
new_frame_size, :, :]
|
||||
logger.info(
|
||||
"Cropped LongCat refinement padding: %s:%s, final shape: %s",
|
||||
num_cond_frames_added,
|
||||
num_cond_frames_added + new_frame_size, frames.shape)
|
||||
|
||||
# Update batch with decoded image
|
||||
batch.output = frames
|
||||
|
||||
|
||||
@@ -100,6 +100,33 @@ class ImageEncodingStage(PipelineStage):
|
||||
return result
|
||||
|
||||
|
||||
class Hy15ImageEncodingStage(ImageEncodingStage):
|
||||
"""
|
||||
Stage for encoding image prompts into embeddings for HunyuanVideo1.5 models.
|
||||
"""
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify image encoding stage inputs."""
|
||||
return VerificationResult()
|
||||
|
||||
def forward(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into image encoder hidden states.
|
||||
"""
|
||||
if batch.pil_image is None:
|
||||
batch.image_embeds = [
|
||||
torch.zeros(1, 729, 1152, device=get_local_torch_device())
|
||||
]
|
||||
|
||||
raw_latent_shape = list(batch.raw_latent_shape)
|
||||
raw_latent_shape[1] = 1
|
||||
batch.video_latent = torch.zeros(tuple(raw_latent_shape),
|
||||
device=get_local_torch_device())
|
||||
return batch
|
||||
|
||||
|
||||
class MatrixGameImageEncodingStage(ImageEncodingStage):
|
||||
CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073]
|
||||
CLIP_STD = [0.26862954, 0.26130258, 0.27577711]
|
||||
|
||||
@@ -136,16 +136,26 @@ class LatentPreparationStage(PipelineStage):
|
||||
)
|
||||
# Generate or use provided latents
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
else:
|
||||
# Pre-initialized latents:
|
||||
# - For LongCat refine (refine_from or stage1_video present), we should not re-scale by init_noise_sigma.
|
||||
# - For other models, keep the original behavior.
|
||||
latents = latents.to(device)
|
||||
is_longcat_refine = (batch.refine_from
|
||||
is not None) or (batch.stage1_video
|
||||
is not None)
|
||||
if (not is_longcat_refine) and hasattr(self.scheduler,
|
||||
"init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
|
||||
# Scale the initial noise if needed
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
# Update batch with prepared latents
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = bcthw_shape
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat-specific denoising stage implementing CFG-zero optimized guidance.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatDenoisingStage(DenoisingStage):
|
||||
"""
|
||||
LongCat denoising stage with CFG-zero optimized guidance scale.
|
||||
|
||||
Implements:
|
||||
1. Optimized CFG scale from CFG-zero paper
|
||||
2. Negation of noise prediction before scheduler step (flow matching convention)
|
||||
3. Batched CFG computation (unlike standard FastVideo separate passes)
|
||||
"""
|
||||
|
||||
def optimized_scale(self, positive_flat, negative_flat) -> torch.Tensor:
|
||||
"""
|
||||
Calculate optimized scale from CFG-zero paper.
|
||||
|
||||
st_star = (v_cond^T * v_uncond) / ||v_uncond||^2
|
||||
|
||||
Args:
|
||||
positive_flat: Conditional prediction, flattened [B, -1]
|
||||
negative_flat: Unconditional prediction, flattened [B, -1]
|
||||
|
||||
Returns:
|
||||
st_star: Optimized scale [B, 1]
|
||||
"""
|
||||
# Calculate dot product
|
||||
dot_product = torch.sum(positive_flat * negative_flat,
|
||||
dim=1,
|
||||
keepdim=True)
|
||||
# Squared norm of uncondition
|
||||
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
|
||||
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
|
||||
st_star = dot_product / squared_norm
|
||||
return st_star
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run LongCat denoising loop with optimized CFG.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
from fastvideo.models.model_loader import TransformerLoader
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
# Get transformer dtype
|
||||
if hasattr(self.transformer, 'module'):
|
||||
transformer_dtype = next(self.transformer.module.parameters()).dtype
|
||||
else:
|
||||
transformer_dtype = next(self.transformer.parameters()).dtype
|
||||
|
||||
target_dtype = transformer_dtype
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Extract batch parameters
|
||||
latents = batch.latents
|
||||
timesteps = batch.timesteps
|
||||
prompt_embeds = batch.prompt_embeds[0] # LongCat uses single encoder
|
||||
prompt_attention_mask = batch.prompt_attention_mask[
|
||||
0] if batch.prompt_attention_mask else None
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_classifier_free_guidance = batch.do_classifier_free_guidance
|
||||
|
||||
# Get negative prompts if doing CFG
|
||||
if do_classifier_free_guidance:
|
||||
negative_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
|
||||
if batch.negative_attention_mask
|
||||
else None)
|
||||
# Concatenate for batched processing
|
||||
prompt_embeds_combined = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
if prompt_attention_mask is not None:
|
||||
prompt_attention_mask_combined = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask],
|
||||
dim=0)
|
||||
else:
|
||||
prompt_attention_mask_combined = None
|
||||
else:
|
||||
prompt_embeds_combined = prompt_embeds
|
||||
prompt_attention_mask_combined = prompt_attention_mask
|
||||
|
||||
# Denoising loop
|
||||
num_inference_steps = len(timesteps)
|
||||
with tqdm(total=num_inference_steps,
|
||||
desc="LongCat Denoising") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# Expand latents for CFG
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = latent_model_input.to(target_dtype)
|
||||
|
||||
# Expand timestep to match batch size
|
||||
timestep = t.expand(
|
||||
latent_model_input.shape[0]).to(target_dtype)
|
||||
|
||||
# Run transformer with context
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=batch,
|
||||
), torch.autocast(device_type='cuda',
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds_combined,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask_combined,
|
||||
)
|
||||
|
||||
# Apply CFG with optimized scale
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
|
||||
B = noise_pred_cond.shape[0]
|
||||
positive = noise_pred_cond.reshape(B, -1)
|
||||
negative = noise_pred_uncond.reshape(B, -1)
|
||||
|
||||
# Calculate optimized scale (CFG-zero)
|
||||
st_star = self.optimized_scale(positive, negative)
|
||||
|
||||
# Reshape for broadcasting
|
||||
st_star = st_star.view(B, 1, 1, 1, 1)
|
||||
|
||||
# Apply optimized CFG formula
|
||||
noise_pred = (
|
||||
noise_pred_uncond * st_star + guidance_scale *
|
||||
(noise_pred_cond - noise_pred_uncond * st_star))
|
||||
|
||||
# CRITICAL: Negate noise prediction for flow matching scheduler
|
||||
noise_pred = -noise_pred
|
||||
|
||||
# Compute previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
# Update batch with denoised latents
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -0,0 +1,310 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat refinement initialization stage.
|
||||
|
||||
This stage prepares the latent variables for LongCat's 480p->720p refinement by:
|
||||
1. Loading the stage1 (480p) video
|
||||
2. Upsampling it to 720p resolution
|
||||
3. Encoding it with VAE
|
||||
4. Mixing with noise according to t_thresh
|
||||
"""
|
||||
|
||||
import math
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vision_utils import load_video
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.configs.pipelines.longcat import get_bucket_config
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatRefineInitStage(PipelineStage):
|
||||
"""
|
||||
Stage for initializing LongCat refinement from a stage1 (480p) video.
|
||||
|
||||
This replicates the logic from LongCatVideoPipeline.generate_refine():
|
||||
- Load stage1_video frames
|
||||
- Upsample spatially and temporally
|
||||
- VAE encode and normalize
|
||||
- Mix with noise according to t_thresh
|
||||
"""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Initialize latents for refinement.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with initialized latents for refinement.
|
||||
"""
|
||||
refine_from = batch.refine_from
|
||||
in_memory_stage1 = batch.stage1_video
|
||||
|
||||
# Only run for refinement tasks: either a path (refine_from) or in-memory video is provided
|
||||
if refine_from is None and in_memory_stage1 is None:
|
||||
# Not a refinement task, skip
|
||||
return batch
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 1. Obtain stage1 frames (either from disk or from in-memory input)
|
||||
# ------------------------------------------------------------------
|
||||
if in_memory_stage1 is not None:
|
||||
# User provided stage1 frames directly (e.g., from distilled stage output)
|
||||
if len(in_memory_stage1) == 0:
|
||||
raise ValueError(
|
||||
"stage1_video is empty; expected a non-empty list of frames"
|
||||
)
|
||||
|
||||
if isinstance(in_memory_stage1[0], Image.Image):
|
||||
pil_images = in_memory_stage1
|
||||
else:
|
||||
# Assume numpy arrays or torch tensors with shape [H, W, C]
|
||||
pil_images = [
|
||||
Image.fromarray(np.array(frame))
|
||||
for frame in in_memory_stage1
|
||||
]
|
||||
|
||||
logger.info(
|
||||
"Initializing LongCat refinement from in-memory stage1_video (%s frames)",
|
||||
len(pil_images))
|
||||
else:
|
||||
# Path-based refine: load video from disk (original design)
|
||||
logger.info("Initializing LongCat refinement from file: %s",
|
||||
refine_from)
|
||||
stage1_video_path = Path(refine_from)
|
||||
if not stage1_video_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Stage1 video not found: {refine_from}")
|
||||
|
||||
# Load video frames as PIL Images
|
||||
pil_images, original_fps = load_video(str(stage1_video_path),
|
||||
return_fps=True)
|
||||
logger.info("Loaded stage1 video: %s frames @ %s fps",
|
||||
len(pil_images), original_fps)
|
||||
|
||||
# Store in batch for reference (use PIL images, same as official demo)
|
||||
batch.stage1_video = pil_images
|
||||
|
||||
# Get parameters from batch
|
||||
num_frames = len(pil_images)
|
||||
spatial_refine_only = batch.spatial_refine_only
|
||||
t_thresh = batch.t_thresh
|
||||
num_cond_frames = batch.num_cond_frames if hasattr(
|
||||
batch, 'num_cond_frames') else 0
|
||||
|
||||
# Calculate new frame count (temporal upsampling if not spatial_refine_only)
|
||||
new_num_frames = num_frames if spatial_refine_only else 2 * num_frames
|
||||
logger.info(
|
||||
"Refine mode: %s",
|
||||
'spatial only' if spatial_refine_only else 'spatial + temporal')
|
||||
|
||||
# Update batch.num_frames to reflect the upsampled count
|
||||
batch.num_frames = new_num_frames
|
||||
|
||||
# Use bucket system to select resolution (exactly like LongCat)
|
||||
# Calculate scale_factor_spatial considering SP split
|
||||
sp_size = fastvideo_args.sp_size if fastvideo_args.sp_size > 0 else 1
|
||||
vae_scale_factor_spatial = 8 # VAE spatial downsampling
|
||||
patch_size_spatial = 2 # LongCat patch size
|
||||
bsa_latent_granularity = 4
|
||||
scale_factor_spatial = vae_scale_factor_spatial * patch_size_spatial * bsa_latent_granularity # 64
|
||||
|
||||
# Calculate optimal split like LongCat (cp_split_hw logic)
|
||||
# For sp_size=1: [1,1], max=1
|
||||
# For sp_size=2: [1,2], max=2
|
||||
# For sp_size=4: [2,2], max=2
|
||||
# For sp_size=8: [2,4], max=4
|
||||
if sp_size > 1:
|
||||
# Get optimal 2D split factors (mimic context_parallel_util.get_optimal_split)
|
||||
factors = []
|
||||
for i in range(1, int(sp_size**0.5) + 1):
|
||||
if sp_size % i == 0:
|
||||
factors.append([i, sp_size // i])
|
||||
cp_split_hw = min(factors, key=lambda x: abs(x[0] - x[1]))
|
||||
scale_factor_spatial *= max(cp_split_hw)
|
||||
logger.info("SP split: sp_size=%s, cp_split_hw=%s, max_split=%s",
|
||||
sp_size, cp_split_hw, max(cp_split_hw))
|
||||
else:
|
||||
cp_split_hw = [1, 1]
|
||||
|
||||
# Get bucket config and find closest bucket for the input aspect ratio
|
||||
bucket_config = get_bucket_config('720p', scale_factor_spatial)
|
||||
|
||||
# Get input aspect ratio from stage1 video
|
||||
input_height, input_width = pil_images[0].height, pil_images[0].width
|
||||
input_ratio = input_height / input_width
|
||||
|
||||
# Find closest bucket
|
||||
closest_ratio = min(bucket_config.keys(),
|
||||
key=lambda x: abs(float(x) - input_ratio))
|
||||
height, width = bucket_config[closest_ratio][0]
|
||||
|
||||
logger.info("Input aspect ratio: %.2f (%sx%s)", input_ratio,
|
||||
input_width, input_height)
|
||||
logger.info("Matched bucket ratio: %s -> resolution: %sx%s",
|
||||
closest_ratio, width, height)
|
||||
logger.info("Target: %sx%s @ %s frames (sp_size=%s, scale_factor=%s)",
|
||||
width, height, new_num_frames, sp_size,
|
||||
scale_factor_spatial)
|
||||
|
||||
# Override batch height/width with bucket-selected resolution
|
||||
batch.height = height
|
||||
batch.width = width
|
||||
|
||||
# Convert PIL images to tensor [T, C, H, W]
|
||||
stage1_video_tensor = torch.stack([
|
||||
torch.from_numpy(np.array(img)).permute(2, 0, 1) # HWC -> CHW
|
||||
for img in pil_images
|
||||
]).float() # [T, C, H, W]
|
||||
|
||||
device = batch.prompt_embeds[0].device
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
stage1_video_tensor = stage1_video_tensor.to(device=device, dtype=dtype)
|
||||
|
||||
# Replicate LongCat's exact preprocessing (lines 1227-1235 in pipeline_longcat_video.py)
|
||||
# First: spatial interpolation to target (height, width) on [T, C, H, W]
|
||||
video_down = F.interpolate(stage1_video_tensor,
|
||||
size=(height, width),
|
||||
mode='bilinear',
|
||||
align_corners=True)
|
||||
|
||||
# Rearrange to [C, T, H, W] and add batch dimension -> [1, C, T, H, W]
|
||||
video_down = video_down.permute(1, 0, 2,
|
||||
3).unsqueeze(0) # [1, C, T, H, W]
|
||||
video_down = video_down / 255.0 # Normalize to [0, 1]
|
||||
|
||||
# Then: temporal+spatial interpolation to (new_num_frames, height, width)
|
||||
video_up = F.interpolate(video_down,
|
||||
size=(new_num_frames, height, width),
|
||||
mode='trilinear',
|
||||
align_corners=True)
|
||||
|
||||
# Rescale to [-1, 1] for VAE
|
||||
video_up = video_up * 2.0 - 1.0
|
||||
|
||||
logger.info("Upsampled video shape: %s", video_up.shape)
|
||||
|
||||
# Padding logic (exactly like LongCat lines 1237-1255)
|
||||
# Only pad temporal dimension to ensure BSA compatibility
|
||||
vae_scale_factor_temporal = 4
|
||||
num_noise_frames = video_up.shape[2] - num_cond_frames
|
||||
|
||||
num_cond_latents = 0
|
||||
num_cond_frames_added = 0
|
||||
if num_cond_frames > 0:
|
||||
num_cond_latents = 1 + math.ceil(
|
||||
(num_cond_frames - 1) / vae_scale_factor_temporal)
|
||||
num_cond_latents = math.ceil(
|
||||
num_cond_latents /
|
||||
bsa_latent_granularity) * bsa_latent_granularity
|
||||
num_cond_frames_added = 1 + (
|
||||
num_cond_latents -
|
||||
1) * vae_scale_factor_temporal - num_cond_frames
|
||||
num_cond_frames = num_cond_frames + num_cond_frames_added
|
||||
|
||||
num_noise_latents = math.ceil(num_noise_frames /
|
||||
vae_scale_factor_temporal)
|
||||
num_noise_latents = math.ceil(
|
||||
num_noise_latents / bsa_latent_granularity) * bsa_latent_granularity
|
||||
num_noise_frames_added = num_noise_latents * vae_scale_factor_temporal - num_noise_frames
|
||||
|
||||
if num_cond_frames_added > 0 or num_noise_frames_added > 0:
|
||||
logger.info(
|
||||
"Padding temporal dimension for BSA: cond_frames+=%s, noise_frames+=%s",
|
||||
num_cond_frames_added, num_noise_frames_added)
|
||||
pad_front = video_up[:, :, 0:1].repeat(1, 1, num_cond_frames_added,
|
||||
1, 1)
|
||||
pad_back = video_up[:, :, -1:].repeat(1, 1, num_noise_frames_added,
|
||||
1, 1)
|
||||
video_up = torch.cat([pad_front, video_up, pad_back], dim=2)
|
||||
logger.info("Padded video shape: %s", video_up.shape)
|
||||
|
||||
# Update batch with actual frame count after padding
|
||||
batch.num_frames = video_up.shape[2]
|
||||
|
||||
# Store padding info for later cropping (CRITICAL for correct output!)
|
||||
batch.num_cond_frames_added = num_cond_frames_added
|
||||
batch.num_noise_frames_added = num_noise_frames_added
|
||||
batch.new_frame_size_before_padding = new_num_frames
|
||||
|
||||
# Store num_cond_latents for denoising stage
|
||||
if num_cond_latents > 0:
|
||||
batch.num_cond_latents = num_cond_latents
|
||||
logger.info("Will use num_cond_latents=%s during denoising",
|
||||
num_cond_latents)
|
||||
|
||||
logger.info("Padding info: cond+=%s, noise+=%s, original=%s",
|
||||
num_cond_frames_added, num_noise_frames_added,
|
||||
new_num_frames)
|
||||
|
||||
# VAE encode
|
||||
logger.info("Encoding stage1 video with VAE...")
|
||||
vae_dtype = next(self.vae.parameters()).dtype
|
||||
vae_device = next(self.vae.parameters()).device
|
||||
video_up = video_up.to(dtype=vae_dtype, device=vae_device)
|
||||
|
||||
with torch.no_grad():
|
||||
latent_dist = self.vae.encode(video_up)
|
||||
# Extract tensor from latent distribution
|
||||
if hasattr(latent_dist, 'latent_dist'):
|
||||
# Nested distribution wrapper
|
||||
latent_up = latent_dist.latent_dist.sample()
|
||||
elif hasattr(latent_dist, 'sample'):
|
||||
# DiagonalGaussianDistribution or similar
|
||||
latent_up = latent_dist.sample()
|
||||
elif hasattr(latent_dist, 'latents'):
|
||||
# Direct latents tensor
|
||||
latent_up = latent_dist.latents
|
||||
else:
|
||||
# Assume it's already a tensor
|
||||
latent_up = latent_dist
|
||||
|
||||
# Normalize latents using VAE config (exactly like LongCat)
|
||||
if hasattr(self.vae.config, 'latents_mean') and hasattr(
|
||||
self.vae.config, 'latents_std'):
|
||||
latents_mean = torch.tensor(self.vae.config.latents_mean).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latent_up.device,
|
||||
latent_up.dtype)
|
||||
# LongCat uses: 1.0 / latents_std (equivalent to dividing by latents_std)
|
||||
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(
|
||||
1, self.vae.config.z_dim, 1, 1, 1).to(latent_up.device,
|
||||
latent_up.dtype)
|
||||
# LongCat: (latents - mean) * (1/std)
|
||||
latent_up = (latent_up - latents_mean) * latents_std
|
||||
|
||||
logger.info("Encoded latent shape: %s", latent_up.shape)
|
||||
|
||||
# Mix with noise according to t_thresh
|
||||
# latent_up = (1 - t_thresh) * latent_up + t_thresh * noise
|
||||
noise = torch.randn_like(latent_up).contiguous()
|
||||
latent_up = (1 - t_thresh) * latent_up + t_thresh * noise
|
||||
|
||||
logger.info("Applied t_thresh=%s noise mixing", t_thresh)
|
||||
|
||||
# Store in batch
|
||||
batch.latents = latent_up.to(dtype)
|
||||
batch.raw_latent_shape = latent_up.shape
|
||||
|
||||
logger.info("LongCat refinement initialization complete")
|
||||
|
||||
return batch
|
||||
@@ -0,0 +1,104 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LongCat refinement timestep preparation stage.
|
||||
|
||||
This stage prepares special timesteps for LongCat refinement that start from t_thresh.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LongCatRefineTimestepStage(PipelineStage):
|
||||
"""
|
||||
Stage for preparing timesteps specific to LongCat refinement.
|
||||
|
||||
For refinement, we need to start from t_thresh instead of t=1.0, so we:
|
||||
1. Generate normal timesteps for num_inference_steps
|
||||
2. Filter to only keep timesteps < t_thresh * 1000
|
||||
3. Prepend t_thresh * 1000 as the first timestep
|
||||
"""
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare refinement-specific timesteps.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with refinement timesteps.
|
||||
"""
|
||||
# Only apply if this is a refinement task
|
||||
# Trigger when either a refine_from path or in-memory stage1_video is provided
|
||||
if batch.refine_from is None and batch.stage1_video is None:
|
||||
return batch
|
||||
|
||||
device = get_local_torch_device()
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
t_thresh = batch.t_thresh
|
||||
|
||||
logger.info("Preparing LongCat refinement timesteps (t_thresh=%s)",
|
||||
t_thresh)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 1) Match LongCatVideoPipeline.get_timesteps_sigmas (non-distill):
|
||||
# sigmas = linspace(1, 0.001, num_inference_steps) on CPU
|
||||
# ------------------------------------------------------------------
|
||||
base_sigmas = torch.linspace(
|
||||
1.0,
|
||||
0.001,
|
||||
num_inference_steps,
|
||||
dtype=torch.float32,
|
||||
device=
|
||||
"cpu", # scheduler.set_timesteps expects CPU-convertible sigmas
|
||||
)
|
||||
# Let the scheduler build its internal timestep schedule from sigmas
|
||||
self.scheduler.set_timesteps(num_inference_steps,
|
||||
sigmas=base_sigmas,
|
||||
device=device)
|
||||
base_timesteps = self.scheduler.timesteps
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 2) Apply t_thresh cropping exactly like generate_refine:
|
||||
# timesteps = [t_thresh*1000] + [t for t in base_timesteps if t < t_thresh*1000]
|
||||
# sigmas = timesteps / 1000 (with trailing zero)
|
||||
# ------------------------------------------------------------------
|
||||
t_thresh_value = t_thresh * 1000.0
|
||||
t_thresh_tensor = torch.tensor(t_thresh_value,
|
||||
dtype=base_timesteps.dtype,
|
||||
device=device)
|
||||
filtered_timesteps = base_timesteps[base_timesteps < t_thresh_tensor]
|
||||
|
||||
timesteps = torch.cat(
|
||||
[t_thresh_tensor.unsqueeze(0), filtered_timesteps])
|
||||
|
||||
# Update scheduler with these custom timesteps and corresponding sigmas
|
||||
self.scheduler.timesteps = timesteps
|
||||
sigmas = torch.cat([timesteps / 1000.0, torch.zeros(1, device=device)])
|
||||
self.scheduler.sigmas = sigmas
|
||||
|
||||
logger.info("Refinement timesteps: %s steps starting from t=%s",
|
||||
len(timesteps), t_thresh)
|
||||
logger.info("First few timesteps: %s", timesteps[:5].tolist())
|
||||
|
||||
# Store in batch so downstream stages (denoising) use the same schedule
|
||||
batch.timesteps = timesteps
|
||||
|
||||
return batch
|
||||
@@ -6,6 +6,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
"""
|
||||
|
||||
import torch
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -100,9 +101,9 @@ class TextEncodingStage(PipelineStage):
|
||||
"""Verify text encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
|
||||
result.add_check(
|
||||
"negative_prompt", batch.negative_prompt, lambda x: not batch.
|
||||
do_classifier_free_guidance or V.string_not_empty(x))
|
||||
# result.add_check(
|
||||
# "negative_prompt", batch.negative_prompt, lambda x: not batch.
|
||||
# do_classifier_free_guidance or V.string_not_empty(x))
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.is_list)
|
||||
@@ -203,20 +204,45 @@ class TextEncodingStage(PipelineStage):
|
||||
preprocess_func = preprocess_funcs[i]
|
||||
postprocess_func = postprocess_funcs[i]
|
||||
|
||||
processed_texts: list[str] = []
|
||||
for prompt_str in texts:
|
||||
processed_texts.append(preprocess_func(prompt_str))
|
||||
|
||||
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
|
||||
if max_length is not None:
|
||||
tok_kwargs["max_length"] = max_length
|
||||
elif hasattr(fastvideo_args.pipeline_config,
|
||||
"text_encoder_max_lengths"):
|
||||
tok_kwargs[
|
||||
"max_length"] = fastvideo_args.pipeline_config.text_encoder_max_lengths[
|
||||
i]
|
||||
|
||||
if truncation is not None:
|
||||
tok_kwargs["truncation"] = truncation
|
||||
if padding is not None:
|
||||
tok_kwargs["padding"] = padding
|
||||
|
||||
text_inputs = tokenizer(processed_texts,
|
||||
**tok_kwargs).to(target_device)
|
||||
processed_texts: list[str] = []
|
||||
for prompt_str in texts:
|
||||
processed_text = preprocess_func(prompt_str)
|
||||
if processed_text is not None:
|
||||
processed_texts.append(processed_text)
|
||||
else:
|
||||
# Assuming batch_size = 1
|
||||
prompt_embeds = torch.zeros((1, tok_kwargs["max_length"],
|
||||
encoder_config.hidden_size),
|
||||
device=target_device)
|
||||
attention_mask = torch.zeros((1, tok_kwargs["max_length"]),
|
||||
device=target_device,
|
||||
dtype=torch.int64)
|
||||
embeds_list.append(prompt_embeds)
|
||||
attn_masks_list.append(attention_mask)
|
||||
return self.return_embeds(embeds_list, attn_masks_list,
|
||||
return_type,
|
||||
return_attention_mask, indices)
|
||||
|
||||
if encoder_config.is_chat_model:
|
||||
text_inputs = tokenizer.apply_chat_template(
|
||||
processed_texts, **tok_kwargs).to(target_device)
|
||||
else:
|
||||
text_inputs = tokenizer(processed_texts,
|
||||
**tok_kwargs).to(target_device)
|
||||
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
@@ -228,13 +254,29 @@ class TextEncodingStage(PipelineStage):
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
prompt_embeds = postprocess_func(outputs)
|
||||
try:
|
||||
prompt_embeds = postprocess_func(outputs)
|
||||
except Exception:
|
||||
prompt_embeds, attention_mask = postprocess_func(
|
||||
outputs, attention_mask)
|
||||
|
||||
if dtype is not None:
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype)
|
||||
embeds_list.append(prompt_embeds)
|
||||
if return_attention_mask:
|
||||
attn_masks_list.append(attention_mask)
|
||||
|
||||
return self.return_embeds(embeds_list, attn_masks_list, return_type,
|
||||
return_attention_mask, indices)
|
||||
|
||||
def return_embeds(
|
||||
self,
|
||||
embeds_list: list[torch.Tensor],
|
||||
attn_masks_list: list[torch.Tensor],
|
||||
return_type: str = "list",
|
||||
return_attention_mask: bool = False,
|
||||
indices: list[int] | None = None,
|
||||
) -> Any:
|
||||
# Shape results according to return_type
|
||||
if return_type == "list":
|
||||
if return_attention_mask:
|
||||
|
||||
@@ -7,6 +7,8 @@ This module contains implementations of timestep preparation stages for diffusio
|
||||
|
||||
import inspect
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -71,7 +73,12 @@ class TimestepPreparationStage(PipelineStage):
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(timesteps=timesteps,
|
||||
# Convert timesteps to CPU if it's a tensor (for numpy conversion in scheduler)
|
||||
if isinstance(timesteps, torch.Tensor):
|
||||
timesteps_for_scheduler = timesteps.cpu()
|
||||
else:
|
||||
timesteps_for_scheduler = timesteps
|
||||
scheduler.set_timesteps(timesteps=timesteps_for_scheduler,
|
||||
device=device,
|
||||
**extra_set_timesteps_kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
|
||||
@@ -126,7 +126,7 @@ class CudaPlatformBase(Platform):
|
||||
logger.info("Selected backend: %s", selected_backend)
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
from fastvideo_kernel import sliding_tile_attention # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sliding_tile_attn import ( # noqa: F401
|
||||
SlidingTileAttentionBackend)
|
||||
@@ -169,7 +169,7 @@ class CudaPlatformBase(Platform):
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
from fastvideo_kernel import video_sparse_attn # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.video_sparse_attn import ( # noqa: F401
|
||||
VideoSparseAttentionBackend)
|
||||
@@ -188,8 +188,7 @@ class CudaPlatformBase(Platform):
|
||||
|
||||
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
||||
try:
|
||||
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
|
||||
moba_attn_varlen)
|
||||
from fastvideo_kernel import moba_attn_varlen # noqa: F401
|
||||
from fastvideo.attention.backends.vmoba import ( # noqa: F401
|
||||
VMOBAAttentionBackend)
|
||||
logger.info("Using Video MOBA Attention backend.")
|
||||
|
||||
@@ -58,6 +58,13 @@ class RocmPlatform(Platform):
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_torch_device(cls):
|
||||
"""
|
||||
Return torch.cuda
|
||||
"""
|
||||
return torch.cuda
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
@@ -71,8 +78,23 @@ class RocmPlatform(Platform):
|
||||
elif selected_backend in (AttentionBackendEnum.FLASH_ATTN, None):
|
||||
pass
|
||||
|
||||
elif selected_backend in (AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN):
|
||||
elif selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sliding_tile_attn import ( # noqa: F401
|
||||
SlidingTileAttentionBackend)
|
||||
logger.info("Using Sliding Tile Attention backend.")
|
||||
|
||||
return "fastvideo.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.error(
|
||||
"Failed to import Sliding Tile Attention backend: %s",
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Sliding Tile Attention backend is not installed. ") from e
|
||||
|
||||
elif selected_backend in (AttentionBackendEnum.SAGE_ATTN):
|
||||
raise ValueError(
|
||||
f"{selected_backend.name} is not supported on {cls.device_name}."
|
||||
)
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoConfig, AutoTokenizer, T5EncoderModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.configs.models.encoders import T5Config
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def t5_model_paths():
|
||||
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
model_path = maybe_download_model(base_model_path)
|
||||
text_encoder_path = os.path.join(model_path, "text_encoder_2")
|
||||
tokenizer_path = os.path.join(model_path, "tokenizer_2")
|
||||
return text_encoder_path, tokenizer_path
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_t5_encoder(t5_model_paths):
|
||||
# Initialize the two model implementations
|
||||
text_encoder_path, tokenizer_path = t5_model_paths
|
||||
hf_config = AutoConfig.from_pretrained(text_encoder_path)
|
||||
print(hf_config)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision_str = "fp32"
|
||||
precision = PRECISION_TO_TYPE[precision_str]
|
||||
model1 = T5EncoderModel.from_pretrained(text_encoder_path).to(
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
|
||||
|
||||
args = FastVideoArgs(model_path=text_encoder_path,
|
||||
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
|
||||
text_encoder_precisions=(precision_str,)),
|
||||
pin_cpu_memory=False)
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load(text_encoder_path, args)
|
||||
model2 = model2.to(precision)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Check number of parameters
|
||||
logger.info("Model1 has %s parameters", len(params1))
|
||||
logger.info("Model2 has %s parameters", len(params2))
|
||||
|
||||
# check if embed_tokens are the same
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight"]
|
||||
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
|
||||
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
"Once upon a time", "The quick brown fox jumps over",
|
||||
"In a galaxy far, far away"
|
||||
]
|
||||
|
||||
logger.info("Testing T5 encoder with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info("Testing prompt: %s", prompt)
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens = tokenizer(prompt,
|
||||
padding="max_length",
|
||||
max_length=512,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt").to(device)
|
||||
|
||||
# Get outputs from HuggingFace implementation
|
||||
# filter out padding input_ids
|
||||
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
|
||||
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
|
||||
outputs1 = model1(input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask.float())[0]
|
||||
print("--------------------------------")
|
||||
logger.info("Testing model2")
|
||||
|
||||
# Get outputs from our implementation
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
).last_hidden_state
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
|
||||
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info("Maximum difference in last hidden states: %s",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %s",
|
||||
mean_diff_hidden.item())
|
||||
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert max_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
@@ -0,0 +1,150 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import pytest
|
||||
import torch
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoConfig, AutoTokenizer, Qwen2_5_VLTextModel
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.configs.models.encoders import Qwen2_5_VLConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29505"
|
||||
|
||||
@pytest.fixture
|
||||
def qwen_model_path():
|
||||
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
model_path = maybe_download_model(base_model_path)
|
||||
text_encoder_path = os.path.join(model_path, "text_encoder")
|
||||
tokenizer_path = os.path.join(model_path, "tokenizer")
|
||||
return text_encoder_path, tokenizer_path
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_qwen2_5_encoder(qwen_model_path):
|
||||
text_encoder_path, tokenizer_path = qwen_model_path
|
||||
hf_config = AutoConfig.from_pretrained(text_encoder_path)
|
||||
print(hf_config)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
# Qwen2.5-VL default dtype is usually bf16
|
||||
precision_str = "fp32"
|
||||
precision = PRECISION_TO_TYPE[precision_str]
|
||||
|
||||
logger.info(f"Using precision: {precision_str}")
|
||||
|
||||
# Load HF model (Base model)
|
||||
model1 = Qwen2_5_VLTextModel.from_pretrained(text_encoder_path).to(
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
|
||||
# Load FastVideo model
|
||||
args = FastVideoArgs(model_path=text_encoder_path,
|
||||
pipeline_config=PipelineConfig(text_encoder_configs=(Qwen2_5_VLConfig(),),
|
||||
text_encoder_precisions=(precision_str,)),
|
||||
pin_cpu_memory=False)
|
||||
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load(text_encoder_path, args)
|
||||
model2 = model2.to(precision)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
logger.info("Model1 has %s parameters", len(params1))
|
||||
logger.info("Model2 has %s parameters", len(params2))
|
||||
|
||||
# Check common layers like Norms which are likely not merged/sharded in a way that changes name significantly
|
||||
# or simple linear layers if names match.
|
||||
# Note: FastVideo uses QKVParallelLinear, so q_proj, k_proj, v_proj are merged.
|
||||
# HF Qwen2_5_VL uses separate projections? No, usually they are separate nn.Linear in HF.
|
||||
|
||||
weights_to_check = [
|
||||
"norm.weight",
|
||||
"layers.{}.self_attn.o_proj.weight",
|
||||
"layers.{}.input_layernorm.weight",
|
||||
"layers.{}.post_attention_layernorm.weight",
|
||||
"layers.{}.mlp.down_proj.weight"
|
||||
]
|
||||
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights_to_check:
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
|
||||
|
||||
# Check shape
|
||||
assert p1.shape == p2.shape, f"Shape mismatch for {w}: {p1.shape} vs {p2.shape}"
|
||||
|
||||
# Check values
|
||||
assert_close(p1, p2, atol=1e-7, rtol=1e-7, msg=f"Weight mismatch for {w}")
|
||||
|
||||
# Test with sample prompts
|
||||
prompts = [
|
||||
"Hello world",
|
||||
"The quick brown fox jumps over the lazy dog."
|
||||
]
|
||||
|
||||
logger.info("Testing with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info(f"Prompt: {prompt}")
|
||||
tokens = tokenizer(prompt, return_tensors="pt", padding="max_length", max_length=1000, truncation=True).to(device)
|
||||
|
||||
# HF Forward
|
||||
# AutoModel for Qwen2.5-VL usually returns BaseModelOutputWithPast
|
||||
outputs1 = model1(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True
|
||||
).hidden_states[-3]
|
||||
|
||||
# FastVideo Forward
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True
|
||||
).hidden_states[-3]
|
||||
|
||||
# Compare
|
||||
# Filter padding for comparison if needed, but here we just check raw output matching
|
||||
|
||||
# Check shapes
|
||||
assert outputs1.shape == outputs2.shape, f"Output shape mismatch: {outputs1.shape} vs {outputs2.shape}"
|
||||
|
||||
diff = torch.abs(outputs1 - outputs2)
|
||||
max_diff = diff.max().item()
|
||||
mean_diff = diff.mean().item()
|
||||
|
||||
logger.info(f"Max diff: {max_diff}")
|
||||
logger.info(f"Mean diff: {mean_diff}")
|
||||
|
||||
# Thresholds
|
||||
# Qwen2.5-VL RoPE is complex, if our implementation is slightly off (e.g. float32 conversion logic in RoPE),
|
||||
# differences might appear. But should be small.
|
||||
if precision_str == "bf16":
|
||||
atol = 5e-2 # relaxed for bf16
|
||||
else:
|
||||
atol = 1e-3
|
||||
|
||||
if max_diff > atol:
|
||||
logger.warning(f"Max diff {max_diff} > {atol}. Checking if it's acceptable...")
|
||||
# If mean diff is small, maybe just outliers
|
||||
assert mean_diff < atol, f"Mean diff {mean_diff} too high"
|
||||
else:
|
||||
logger.info("Outputs match within tolerance.")
|
||||
|
||||
@@ -4,6 +4,7 @@ app = modal.App()
|
||||
|
||||
import os
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
image_version = os.getenv("IMAGE_VERSION")
|
||||
image_tag = f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{image_version}"
|
||||
print(f"Using image: {image_tag}")
|
||||
@@ -74,9 +75,15 @@ def run_vae_tests():
|
||||
def run_transformer_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
@app.function(
|
||||
gpu="L40S:4",
|
||||
image=image,
|
||||
timeout=2700,
|
||||
secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})],
|
||||
volumes={"/root/data": model_vol}
|
||||
)
|
||||
def run_ssim_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
run_test("export MODEL_PATH='/root/data/weights' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests():
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers import AutoencoderKLHunyuanVideo15
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.logger import init_logger
|
||||
# from fastvideo.models.vaes.hunyuanvae import (
|
||||
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.loader.component_loader import VAELoader
|
||||
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
"data", BASE_MODEL_PATH))
|
||||
VAE_PATH = os.path.join(MODEL_PATH, "vae")
|
||||
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_hunyuan_vae():
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.float32
|
||||
precision_str = "fp32"
|
||||
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=Hunyuan15VAEConfig(), vae_precision=precision_str))
|
||||
args.device = device
|
||||
args.vae_cpu_offload = False
|
||||
|
||||
model1 = AutoencoderKLHunyuanVideo15.from_pretrained(
|
||||
VAE_PATH, torch_dtype=precision).to(device).eval()
|
||||
model1.enable_tiling()
|
||||
|
||||
loader = VAELoader()
|
||||
model2 = loader.load(VAE_PATH, args)
|
||||
|
||||
model2.enable_tiling()
|
||||
|
||||
batch_size = 1
|
||||
|
||||
# Video input [B, C, T, H, W]
|
||||
input_tensor = torch.randn(batch_size,
|
||||
3,
|
||||
81,
|
||||
512,
|
||||
512,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
latent1 = model1.encode(input_tensor, return_dict=False)[0].mode()
|
||||
latent2 = model2.encode(input_tensor).mode()
|
||||
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
max_diff_encode = torch.max(torch.abs(latent1.float() - latent2.float()))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1.float() - latent2.float()))
|
||||
logger.info("Maximum difference between encoded latents: %s",
|
||||
max_diff_encode.item())
|
||||
logger.info("Mean difference between encoded latents: %s",
|
||||
mean_diff_encode.item())
|
||||
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {max_diff_encode.item()}, mean diff = {mean_diff_encode.item()}"
|
||||
|
||||
# Test decoding
|
||||
latent1 = latent1 / model1.config.scaling_factor
|
||||
latent2 = latent2 / model2.config.scaling_factor
|
||||
|
||||
with torch.no_grad():
|
||||
video1 = model1.decode(latent1, return_dict=False)[0]
|
||||
video2 = model2.decode(latent2)
|
||||
|
||||
assert video1.shape == video2.shape, f"Video shapes don't match: {video1.shape} vs {video2.shape}"
|
||||
max_diff_decode = torch.max(torch.abs(video1.float() - video2.float()))
|
||||
mean_diff_decode = torch.mean(torch.abs(video1.float() - video2.float()))
|
||||
logger.info("Maximum difference between decoded videos: %s",
|
||||
max_diff_decode.item())
|
||||
logger.info("Mean difference between decoded videos: %s",
|
||||
mean_diff_decode.item())
|
||||
assert max_diff_decode < 1e-5, f"Decoded videos differ significantly: max diff = {max_diff_decode.item()}, mean diff = {mean_diff_decode.item()}"
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
"""Block-sparse attention kernels for LongCat."""
|
||||
|
||||
|
||||
@@ -0,0 +1,656 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import math
|
||||
|
||||
from .common import _attn_fwd_gating, _attn_bwd_preprocess, configs_gating_preset
|
||||
from .flash_attn_bsa_varlen_mask import (
|
||||
_attn_fwd_bsa_varlen, _attn_fwd_bsa_varlen_align, _attn_bwd_dkdv_bsa_varlen_wrapper, _attn_bwd_dq_bsa_varlen_wrapper, _attn_bwd_dq_bsa_varlen_align_wrapper,
|
||||
configs_fwd_bsa_varlen_preset, configs_fwd_bsa_varlen_align_preset, configs_bwd_dkdv_bsa_varlen_preset, configs_bwd_dq_bsa_varlen_preset, configs_bwd_dq_bsa_varlen_align_preset
|
||||
)
|
||||
|
||||
torch._dynamo.config.cache_size_limit = 32
|
||||
|
||||
def is_cuda():
|
||||
return triton.runtime.driver.active.get_current_target().backend == "cuda"
|
||||
|
||||
def supports_tma():
|
||||
return is_cuda() and torch.cuda.get_device_capability()[0] >= 9
|
||||
|
||||
HAS_TMA_DESC = "nv_tma_desc_type" in dir(tl)
|
||||
|
||||
if HAS_TMA_DESC:
|
||||
print("TMA benchmarks will be running with experimental grid constant TMA descriptor.", )
|
||||
else:
|
||||
print("TMA benchmarks will be running without grid constant TMA descriptor.", )
|
||||
|
||||
|
||||
# TmaAutoTuneHelper used in htyu's PR #5622
|
||||
class TmaAutoTuneHelper:
|
||||
|
||||
# duck typing wrapper to implement the same interface as TmaDescKernelParam in Triton PR #4498
|
||||
class KernelParamWrapper:
|
||||
|
||||
def __init__(self, desc):
|
||||
self.desc = desc
|
||||
|
||||
def tma_desc_cpu_ptr(self):
|
||||
return self.desc.data_ptr()
|
||||
|
||||
TMA_SIZE = 128
|
||||
|
||||
def __init__(self):
|
||||
self.fill_1d_tma_descriptor_inner = (triton.runtime.driver.active.utils.fill_1d_tma_descriptor)
|
||||
self.fill_2d_tma_descriptor_inner = (triton.runtime.driver.active.utils.fill_2d_tma_descriptor)
|
||||
if HAS_TMA_DESC:
|
||||
self.descriptors = {}
|
||||
else:
|
||||
self.cuda_descriptors = {}
|
||||
|
||||
# Call this method outside of the lambda function for grid size
|
||||
def init_tma_descriptor(self, name):
|
||||
if HAS_TMA_DESC:
|
||||
self.descriptors[name] = torch.empty(TmaAutoTuneHelper.TMA_SIZE, device="cpu", dtype=torch.int8)
|
||||
else:
|
||||
self.cuda_descriptors[name] = torch.empty(TmaAutoTuneHelper.TMA_SIZE, device="cuda", dtype=torch.int8)
|
||||
|
||||
# Call this method inside the lambda function for grid size
|
||||
def fill_1d_tma_descriptor(self, name, ptr, dim, block_dim, element_size):
|
||||
if HAS_TMA_DESC:
|
||||
desc_x = self.descriptors[name]
|
||||
assert desc_x.data_ptr() % 64 == 0
|
||||
self.fill_1d_tma_descriptor_inner(ptr, dim, block_dim, element_size, desc_x.data_ptr())
|
||||
else:
|
||||
desc_x = self.cuda_descriptors[name]
|
||||
buf_x = torch.empty_like(desc_x, device="cpu", pin_memory=True)
|
||||
self.fill_1d_tma_descriptor_inner(ptr, dim, block_dim, element_size, buf_x.data_ptr())
|
||||
desc_x.copy_(buf_x, non_blocking=True)
|
||||
|
||||
# Call this method inside the lambda function for grid size
|
||||
def fill_2d_tma_descriptor(self, name, ptr, dim1, dim0, block_dim1, block_dim0, element_size):
|
||||
if HAS_TMA_DESC:
|
||||
desc_x = self.descriptors[name]
|
||||
assert desc_x.data_ptr() % 64 == 0
|
||||
self.fill_2d_tma_descriptor_inner(ptr, dim1, dim0, block_dim1, block_dim0, element_size, desc_x.data_ptr())
|
||||
else:
|
||||
desc_x = self.cuda_descriptors[name]
|
||||
buf_x = torch.empty_like(desc_x, device="cpu", pin_memory=True)
|
||||
self.fill_2d_tma_descriptor_inner(ptr, dim1, dim0, block_dim1, block_dim0, element_size, buf_x.data_ptr())
|
||||
desc_x.copy_(buf_x, non_blocking=True)
|
||||
|
||||
def get_tma_descriptor_kernel_param(self, name):
|
||||
if HAS_TMA_DESC:
|
||||
assert self.descriptors[name] is not None
|
||||
return self.KernelParamWrapper(self.descriptors[name])
|
||||
else:
|
||||
assert self.cuda_descriptors[name] is not None
|
||||
return self.cuda_descriptors[name]
|
||||
|
||||
|
||||
@triton.jit
|
||||
def create_mask_from_indices_kernel(
|
||||
block_indices,
|
||||
block_mask,
|
||||
stride_bz, stride_bh, stride_bm, stride_bs,
|
||||
stride_mz, stride_mh, stride_mm, stride_mn,
|
||||
H,
|
||||
):
|
||||
i_zh, i_m, i_s = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
i_z, i_h = i_zh // H, i_zh % H
|
||||
|
||||
off_b = i_z.to(tl.int64) * stride_bz + i_h.to(tl.int64) * stride_bh + i_m.to(tl.int64) * stride_bm + i_s.to(tl.int64) * stride_bs
|
||||
|
||||
b_i = tl.load(block_indices + off_b)
|
||||
|
||||
off_m = i_z.to(tl.int64) * stride_mz + i_h.to(tl.int64) * stride_mh + i_m.to(tl.int64) * stride_mm + b_i.to(tl.int64) * stride_mn
|
||||
|
||||
b_m = 1
|
||||
tl.store(block_mask + off_m, b_m.to(block_mask.dtype.element_ty))
|
||||
|
||||
def create_mask_from_indices_triton(
|
||||
block_indices,
|
||||
N_cols
|
||||
):
|
||||
B, H, N_rows, S = block_indices.shape
|
||||
block_mask = torch.zeros((B, H, N_rows, N_cols), dtype=torch.bool, device=block_indices.device)
|
||||
create_mask_from_indices_kernel[(B * H, N_rows, S)](
|
||||
block_indices,
|
||||
block_mask,
|
||||
block_indices.stride(0), block_indices.stride(1), block_indices.stride(2), block_indices.stride(3),
|
||||
block_mask.stride(0), block_mask.stride(1), block_mask.stride(2), block_mask.stride(3),
|
||||
H,
|
||||
)
|
||||
return block_mask
|
||||
|
||||
@torch.compile
|
||||
def create_mask_from_indices_varlen(block_indices, N_cols_mask):
|
||||
|
||||
B, H, M, _ = block_indices.shape
|
||||
device = block_indices.device
|
||||
|
||||
mask = torch.zeros((B, H, M, N_cols_mask), dtype=torch.bool, device=device)
|
||||
|
||||
valid = block_indices < N_cols_mask
|
||||
|
||||
b_idx = torch.arange(B, device=device)[:, None, None, None].expand_as(block_indices)
|
||||
h_idx = torch.arange(H, device=device)[None, :, None, None].expand_as(block_indices)
|
||||
m_idx = torch.arange(M, device=device)[None, None, :, None].expand_as(block_indices)
|
||||
|
||||
valid_coords = (b_idx[valid], h_idx[valid], m_idx[valid], block_indices[valid])
|
||||
|
||||
mask[valid_coords] = True
|
||||
|
||||
return mask
|
||||
|
||||
@torch.compile
|
||||
def create_indices_k_from_indices_q_varlen(
|
||||
block_indices,
|
||||
N_cols_mask # indicate the number of the last dimension of the bool mask, since this information cannot be determined by block_indices, which may contain invalid elements
|
||||
):
|
||||
block_mask_qk = create_mask_from_indices_varlen(block_indices, N_cols_mask)
|
||||
B, H, M, N = block_mask_qk.shape
|
||||
block_mask_kq = block_mask_qk.permute(0, 1, 3, 2)
|
||||
indices = torch.arange(M, device=block_indices.device).view(1, 1, 1, -1).expand_as(block_mask_kq)
|
||||
block_indices_k = torch.where(block_mask_kq, indices, M)
|
||||
block_indices_k, _ = torch.sort(block_indices_k, dim=-1)
|
||||
|
||||
block_indices_k_lens = (block_indices_k < M).sum(dim=-1)
|
||||
|
||||
return block_indices_k, block_indices_k_lens
|
||||
|
||||
|
||||
@torch.compile
|
||||
def mean_pooling_compression(
|
||||
x: torch.Tensor,
|
||||
block_size: int
|
||||
) -> torch.Tensor:
|
||||
B, H, S = x.shape[:3]
|
||||
num_block = math.ceil(S / block_size)
|
||||
if S % block_size != 0:
|
||||
x = F.pad(x, (0, 0, 0, num_block * block_size - S))
|
||||
x_cmp = x.view(B, H, num_block, block_size, -1).mean(dim=3)
|
||||
return x_cmp
|
||||
|
||||
@torch.compile
|
||||
def cal_score(q, k):
|
||||
k_transposed = k.transpose(-1, -2) # [b, h, d, s_k]
|
||||
score = torch.matmul(q, k_transposed) # [b, h, s_q, s_k]
|
||||
return score
|
||||
|
||||
def cal_score_triton(q, k):
|
||||
B, H, s_q, D = q.shape
|
||||
s_k = k.shape[2]
|
||||
|
||||
score = torch.empty(B, H, s_q, s_k, device=q.device, dtype=q.dtype)
|
||||
|
||||
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_gating_preset['default']
|
||||
|
||||
grid = lambda args: (triton.cdiv(s_q, args["BLOCK_M"]), B * H, 1)
|
||||
_attn_fwd_gating[grid](
|
||||
q, k, score,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
score.stride(0), score.stride(1), score.stride(2), score.stride(3),
|
||||
H, s_q, s_k,
|
||||
HEAD_DIM=D,
|
||||
**kernel_config
|
||||
)
|
||||
return score
|
||||
|
||||
@torch.compile
|
||||
def get_select_indices_topk(q, k, sparsity):
|
||||
score = cal_score(q, k)
|
||||
block_indices, block_indices_lens = get_select_indices_topk_from_score(score, sparsity)
|
||||
return block_indices, block_indices_lens
|
||||
|
||||
@torch.compile
|
||||
def get_select_indices_topk_from_score(score, sparsity):
|
||||
num_selected = int((1 - sparsity) * score.shape[-1])
|
||||
block_indices = torch.topk(score, num_selected)[1]
|
||||
|
||||
block_indices_lens = torch.full(
|
||||
(block_indices.shape[0], block_indices.shape[1], block_indices.shape[2]),
|
||||
num_selected,
|
||||
dtype=torch.int32,
|
||||
device=block_indices.device
|
||||
)
|
||||
|
||||
return block_indices, block_indices_lens
|
||||
|
||||
@torch.compile
|
||||
def get_select_indices_cdf(q, k, cdf_threshold):
|
||||
score = cal_score(q, k)
|
||||
head_dim = q.shape[-1]
|
||||
block_indices, block_indices_lens = get_select_indices_cdf_from_score(score, cdf_threshold, 1 / head_dim**0.5)
|
||||
return block_indices, block_indices_lens
|
||||
|
||||
@torch.compile
|
||||
def get_select_indices_cdf_from_score(score, cdf_threshold, sm_scale):
|
||||
weights = torch.softmax(score * sm_scale, dim=-1)
|
||||
|
||||
B, H, Sq, Sk = weights.shape
|
||||
cdf_threshold = torch.full((H,), cdf_threshold, device=weights.device).view(1, H, 1, 1).expand(B, -1, Sq, -1)
|
||||
weights_sorted = torch.sort(weights, dim=-1, descending=True)
|
||||
cdf = torch.cumsum(weights_sorted.values, dim=-1)
|
||||
num_selected = torch.searchsorted(cdf, cdf_threshold, right=True)
|
||||
|
||||
return weights_sorted.indices, num_selected.squeeze(-1)
|
||||
|
||||
@torch.compile
|
||||
def get_select_indices_cdf_topk(q, k, sparsity, cdf_threshold):
|
||||
score = cal_score(q, k)
|
||||
head_dim = q.shape[-1]
|
||||
block_indices, block_indices_lens = get_select_indices_cdf_topk_from_score(score, sparsity, cdf_threshold, 1 / head_dim**0.5)
|
||||
return block_indices, block_indices_lens
|
||||
|
||||
@torch.compile
|
||||
def get_select_indices_cdf_topk_from_score(score, sparsity, cdf_threshold, sm_scale):
|
||||
weights = torch.softmax(score * sm_scale, dim=-1)
|
||||
|
||||
B, H, Sq, Sk = weights.shape
|
||||
cdf_threshold = torch.full((H,), cdf_threshold, device=weights.device).view(1, H, 1, 1).expand(B, -1, Sq, -1)
|
||||
weights_sorted = torch.sort(weights, dim=-1, descending=True)
|
||||
cdf = torch.cumsum(weights_sorted.values, dim=-1)
|
||||
num_selected = torch.searchsorted(cdf, cdf_threshold, right=True)
|
||||
|
||||
# max(cdf, topk)
|
||||
num_selected_topk = int((1 - sparsity) * score.shape[-1])
|
||||
num_selected[num_selected < num_selected_topk] = num_selected_topk
|
||||
|
||||
return weights_sorted.indices, num_selected.squeeze(-1)
|
||||
|
||||
def get_select_indices(q, k, sparsity, cdf_threshold):
|
||||
if sparsity is not None and cdf_threshold is None:
|
||||
block_indices, block_indices_lens = get_select_indices_topk(q, k, sparsity)
|
||||
elif sparsity is None and cdf_threshold is not None:
|
||||
block_indices, block_indices_lens = get_select_indices_cdf(q, k, cdf_threshold)
|
||||
elif sparsity is not None and cdf_threshold is not None:
|
||||
block_indices, block_indices_lens = get_select_indices_cdf_topk(q, k, sparsity, cdf_threshold)
|
||||
else:
|
||||
raise ValueError
|
||||
return block_indices, block_indices_lens
|
||||
|
||||
def get_select_indices_from_score(score, sparsity, cdf_threshold):
|
||||
if sparsity is not None and cdf_threshold is None:
|
||||
block_indices, block_indices_lens = get_select_indices_topk_from_score(score, sparsity)
|
||||
elif sparsity is None and cdf_threshold is not None:
|
||||
block_indices, block_indices_lens = get_select_indices_cdf_from_score(score, cdf_threshold)
|
||||
elif sparsity is not None and cdf_threshold is not None:
|
||||
block_indices, block_indices_lens = get_select_indices_cdf_topk_from_score(score, sparsity, cdf_threshold)
|
||||
else:
|
||||
raise ValueError
|
||||
return block_indices, block_indices_lens
|
||||
|
||||
def attn_fwd_bsa_varlen_triton(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
sm_scale,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
chunk_size_q,
|
||||
chunk_size_k,
|
||||
sparsity
|
||||
):
|
||||
|
||||
B, H, Seq, D = q.shape
|
||||
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
|
||||
grid = lambda args: (triton.cdiv(q.shape[2], args["BLOCK_M"]), q.shape[0] * q.shape[1], 1)
|
||||
|
||||
config_key = 'BLOCK_N_LG=64' if chunk_size_k == 64 else 'default'
|
||||
if chunk_size_k > 128:
|
||||
fwd_func = _attn_fwd_bsa_varlen
|
||||
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_fwd_bsa_varlen_preset[config_key]
|
||||
else:
|
||||
fwd_func = _attn_fwd_bsa_varlen_align
|
||||
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_fwd_bsa_varlen_align_preset[config_key]
|
||||
|
||||
block_indices = block_indices.contiguous()
|
||||
block_indices_lens = block_indices_lens.contiguous()
|
||||
|
||||
fwd_func[grid](
|
||||
q, k, v, sm_scale, M, o,
|
||||
block_indices, # [B, H, M_COMPRESS, S]
|
||||
block_indices_lens, # [B, H, M_COMPRESS, S_MAX]
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
block_indices.stride(0), block_indices.stride(1), block_indices.stride(2), block_indices.stride(3),
|
||||
block_indices_lens.stride(0), block_indices_lens.stride(1), block_indices_lens.stride(2),
|
||||
H, Seq,
|
||||
D,
|
||||
BLOCK_M=chunk_size_q,
|
||||
BLOCK_N_LG=chunk_size_k,
|
||||
SPARSITY=sparsity,
|
||||
**kernel_config
|
||||
)
|
||||
|
||||
LN2 = 0.6931471824645996
|
||||
lse = M * LN2 # convert back to natural units (M is of base 2)
|
||||
|
||||
return o, lse
|
||||
|
||||
def attn_bwd_bsa_varlen_triton(
|
||||
do,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
o,
|
||||
dq,
|
||||
dk,
|
||||
dv,
|
||||
sm_scale,
|
||||
M,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
chunk_size_q,
|
||||
chunk_size_k,
|
||||
sparsity
|
||||
):
|
||||
RCP_LN2 = 1.4426950408889634
|
||||
M = M * RCP_LN2 # ln -> log2
|
||||
|
||||
do = do.contiguous()
|
||||
# assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
BATCH, N_HEAD, N_CTX, HEAD_DIM = q.shape
|
||||
N_CTX_KV = k.shape[-2]
|
||||
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2) # reciprocal
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
|
||||
if min(chunk_size_q, chunk_size_k) >= 128:
|
||||
PRE_BLOCK = 128
|
||||
else:
|
||||
PRE_BLOCK = min(chunk_size_q, chunk_size_k)
|
||||
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do,
|
||||
delta,
|
||||
N_CTX,
|
||||
BLOCK_M=PRE_BLOCK,
|
||||
HEAD_DIM=HEAD_DIM
|
||||
)
|
||||
|
||||
block_indices_k, block_indices_k_lens = create_indices_k_from_indices_q_varlen(
|
||||
block_indices=block_indices,
|
||||
N_cols_mask=N_CTX_KV // chunk_size_k
|
||||
)
|
||||
|
||||
block_indices = block_indices.contiguous()
|
||||
block_indices_lens = block_indices_lens.contiguous()
|
||||
block_indices_k = block_indices_k.contiguous()
|
||||
block_indices_k_lens = block_indices_k_lens.contiguous()
|
||||
|
||||
config_key = 'BLOCK_N_DQ_LG=64' if chunk_size_k == 64 else 'default'
|
||||
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_bwd_dkdv_bsa_varlen_preset[config_key]
|
||||
|
||||
grid_dkdv = lambda args: (triton.cdiv(arg_k.shape[2], args["BLOCK_N"]), 1, arg_k.shape[0] * arg_k.shape[1])
|
||||
_attn_bwd_dkdv_bsa_varlen_wrapper[grid_dkdv](
|
||||
q, arg_k, v, sm_scale, # softmax scale
|
||||
do,
|
||||
dk, dv,
|
||||
M, # lse (log2)
|
||||
delta,
|
||||
block_indices_k,
|
||||
block_indices_k_lens,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
dk.stride(0), dk.stride(1), dk.stride(2), dk.stride(3),
|
||||
dv.stride(0), dv.stride(1), dv.stride(2), dv.stride(3),
|
||||
do.stride(0), do.stride(1), do.stride(2), do.stride(3),
|
||||
M.stride(0), M.stride(1), M.stride(2),
|
||||
delta.stride(0), delta.stride(1), delta.stride(2),
|
||||
block_indices_k.stride(0), block_indices_k.stride(1), block_indices_k.stride(2), block_indices_k.stride(3),
|
||||
block_indices_k_lens.stride(0), block_indices_k_lens.stride(1), block_indices_k_lens.stride(2),
|
||||
N_HEAD, N_CTX,
|
||||
BLOCK_M=chunk_size_q,
|
||||
BLOCK_N_DQ_LG=chunk_size_k,
|
||||
HEAD_DIM=HEAD_DIM,
|
||||
SPARSITY=sparsity,
|
||||
**kernel_config
|
||||
)
|
||||
|
||||
config_key = 'BLOCK_N_DQ_LG=64' if chunk_size_k == 64 else 'default'
|
||||
if chunk_size_k > 128:
|
||||
bwd_dq_func = _attn_bwd_dq_bsa_varlen_wrapper
|
||||
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_bwd_dq_bsa_varlen_preset[config_key]
|
||||
else:
|
||||
bwd_dq_func = _attn_bwd_dq_bsa_varlen_align_wrapper
|
||||
kernel_config = {} if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1' else configs_bwd_dq_bsa_varlen_align_preset[config_key]
|
||||
|
||||
grid_dq = lambda args: (triton.cdiv(q.shape[2], args["BLOCK_M"]), 1, q.shape[0] * q.shape[1])
|
||||
bwd_dq_func[grid_dq](
|
||||
q, arg_k, v,
|
||||
do,
|
||||
dq,
|
||||
M, # lse (log2)
|
||||
delta,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
dq.stride(0), dq.stride(1), dq.stride(2), dq.stride(3),
|
||||
do.stride(0), do.stride(1), do.stride(2), do.stride(3),
|
||||
M.stride(0), M.stride(1), M.stride(2),
|
||||
delta.stride(0), delta.stride(1), delta.stride(2),
|
||||
block_indices.stride(0), block_indices.stride(1), block_indices.stride(2), block_indices.stride(3),
|
||||
block_indices_lens.stride(0), block_indices_lens.stride(1), block_indices_lens.stride(2),
|
||||
N_HEAD, N_CTX,
|
||||
BLOCK_M=chunk_size_q,
|
||||
BLOCK_N_DQ_LG=chunk_size_k,
|
||||
HEAD_DIM=HEAD_DIM,
|
||||
SPARSITY=sparsity,
|
||||
**kernel_config
|
||||
)
|
||||
|
||||
@torch.compile
|
||||
def make_block_indices_varlen_cp_list(block_indices, cp_size, num_blocks_k_full):
|
||||
"""
|
||||
Args:
|
||||
block_indices: [B, H, num_blocks_q_per_cp_rank, num_blocks_k_full]
|
||||
|
||||
Return:
|
||||
a list of [block_indices, block_indices_lens] for k from each cp_rank
|
||||
- each block_indices starts from zero
|
||||
- block_indices_lens indicates the valid number of elements in the last dimension of block_indices
|
||||
"""
|
||||
res = []
|
||||
num_blocks_per_rank = num_blocks_k_full // cp_size
|
||||
for i in range(cp_size):
|
||||
block_indices_tmp = block_indices.clone()
|
||||
min_block_idx = i * num_blocks_per_rank
|
||||
block_indices_tmp -= min_block_idx
|
||||
block_indices_tmp[block_indices_tmp < 0] = num_blocks_per_rank # block_indices_tmp < 0 indicate invalid indices, set them to num_blocks_per_rank in order to sort them to the tail, so that the first N elements of the block_indices indicated by block_indices_lens are valid
|
||||
|
||||
block_indices_tmp, _ = torch.sort(block_indices_tmp, dim=-1)
|
||||
|
||||
block_indices_tmp_lens = (block_indices_tmp < num_blocks_per_rank).sum(dim=-1)
|
||||
|
||||
res.append([block_indices_tmp, block_indices_tmp_lens])
|
||||
|
||||
return res
|
||||
|
||||
@torch.compile
|
||||
def flash_attn_fwd_softmax_lse_correction(
|
||||
softmax_lse: torch.Tensor,
|
||||
softmax_lse_per_step: torch.Tensor,
|
||||
):
|
||||
"""Merge softmax stats of each step in Attention with context parallelism"""
|
||||
max_scale = torch.max(softmax_lse, softmax_lse_per_step)
|
||||
min_scale = torch.min(softmax_lse, softmax_lse_per_step)
|
||||
lse_diff = min_scale - max_scale
|
||||
lse_diff = lse_diff.nan_to_num(nan=0.) # handle cases: tensor(-inf) - tensor(-inf) = tensor(nan); In the current cp implementation, it is possible that lses of 2 cp ranks are both -inf, if no block is selected from both cp ranks. In such cases, the finally corrected lse should remain -inf.
|
||||
new_scale = max_scale + torch.log1p(torch.exp(lse_diff)) # a + ln(1 + e^(b - a)) = ln(e^a) + ln(1 + e^(b - a)) = ln(e^a + e^b)
|
||||
softmax_lse.copy_(new_scale)
|
||||
|
||||
@torch.compile
|
||||
def flash_attn_fwd_out_correction_init(
|
||||
out_init_step: torch.Tensor, # b h s d
|
||||
softmax_lse: torch.Tensor, # b h s
|
||||
softmax_lse_init_step: torch.Tensor,
|
||||
):
|
||||
"""Merge partial outputs of the first step in Attention with context parallelism"""
|
||||
softmax_lse_corrected_exp = torch.exp(softmax_lse_init_step - softmax_lse)
|
||||
softmax_lse_corrected_exp = softmax_lse_corrected_exp.unsqueeze(-1)
|
||||
out_corrected = out_init_step * softmax_lse_corrected_exp
|
||||
return out_corrected.to(out_init_step.dtype)
|
||||
|
||||
|
||||
@torch.compile
|
||||
def flash_attn_fwd_out_correction(
|
||||
out: torch.Tensor,
|
||||
out_per_step: torch.Tensor,
|
||||
softmax_lse: torch.Tensor,
|
||||
softmax_lse_per_step: torch.Tensor,
|
||||
):
|
||||
"""Merge partial outputs of each step in Attention with context parallelism"""
|
||||
softmax_lse_corrected_exp = torch.exp(softmax_lse_per_step - softmax_lse)
|
||||
softmax_lse_corrected_exp = softmax_lse_corrected_exp.unsqueeze(-1)
|
||||
out_corrected = out_per_step * softmax_lse_corrected_exp
|
||||
out.add_(out_corrected)
|
||||
|
||||
@torch.compile
|
||||
def topk_sort(score, num_chunks_selected):
|
||||
block_indices = torch.topk(score, num_chunks_selected)[1]
|
||||
block_indices, _ = torch.sort(block_indices, dim=-1)
|
||||
return block_indices
|
||||
|
||||
class _attention_bsa(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, chunk_size_q, chunk_size_k, sparsity, cdf_threshold, sm_scale, use_tma=False):
|
||||
# shape constraints
|
||||
HEAD_DIM_Q, HEAD_DIM_K = q.shape[-1], k.shape[-1]
|
||||
# when v is in float8_e5m2 it is transposed.
|
||||
HEAD_DIM_V = v.shape[-1]
|
||||
assert HEAD_DIM_Q == HEAD_DIM_K and HEAD_DIM_K == HEAD_DIM_V
|
||||
assert HEAD_DIM_K in {16, 32, 64, 128, 256}
|
||||
|
||||
# ---------------------- gating ----------------------
|
||||
q_cmp = mean_pooling_compression(q, chunk_size_q)
|
||||
k_cmp = mean_pooling_compression(k, chunk_size_k)
|
||||
block_indices, block_indices_lens = get_select_indices(q_cmp, k_cmp, sparsity, cdf_threshold)
|
||||
|
||||
# ---------------------- bsa ----------------------
|
||||
|
||||
o, lse = attn_fwd_bsa_varlen_triton(
|
||||
q, k, v,
|
||||
sm_scale, block_indices, block_indices_lens,
|
||||
chunk_size_q, chunk_size_k,
|
||||
sparsity
|
||||
)
|
||||
|
||||
ctx.save_for_backward(q, k, v, o, lse, block_indices, block_indices_lens)
|
||||
ctx.sm_scale = sm_scale
|
||||
ctx.HEAD_DIM = HEAD_DIM_K
|
||||
ctx.chunk_size_q = chunk_size_q
|
||||
ctx.chunk_size_k = chunk_size_k
|
||||
ctx.use_tma = use_tma
|
||||
ctx.sparsity = sparsity
|
||||
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, do):
|
||||
q, k, v, o, lse, block_indices, block_indices_lens = ctx.saved_tensors
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
|
||||
attn_bwd_bsa_varlen_triton(
|
||||
do,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
o,
|
||||
dq,
|
||||
dk,
|
||||
dv,
|
||||
ctx.sm_scale,
|
||||
lse,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
ctx.chunk_size_q,
|
||||
ctx.chunk_size_k,
|
||||
ctx.sparsity
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, None, None, None, None, None
|
||||
|
||||
flash_attn_bsa = _attention_bsa.apply
|
||||
|
||||
def rearrange_THW_to_3d_block(x, Nt, Nh, Nw, t, h, w, D):
|
||||
B, H, _, D = x.shape
|
||||
x = x.view(B, H, Nt, t, Nh, h, Nw, w, D)
|
||||
x = x.permute(0, 1, 2, 4, 6, 3, 5, 7, 8) # B H Nt Nh Nw t h w D
|
||||
return x.contiguous().view(B, H, Nt * Nh * Nw * t * h * w, D)
|
||||
|
||||
def rearrange_3d_block_to_THW(x, Nt, Nh, Nw, t, h, w, D):
|
||||
B, H, _, D = x.shape
|
||||
x = x.view(B, H, Nt, Nh, Nw, t, h, w, D)
|
||||
x = x.permute(0, 1, 2, 5, 3, 6, 4, 7, 8) # B H Nt t Nh h Nw w D
|
||||
return x.contiguous().view(B, H, Nt * t * Nh * h * Nw * w, D)
|
||||
|
||||
def flash_attn_bsa_3d(
|
||||
q: torch.Tensor, # [B, H, Sq, D]
|
||||
k: torch.Tensor, # [B, H, Skv, D]
|
||||
v: torch.Tensor, # [B, H, Skv, D]
|
||||
latent_shape_q,
|
||||
latent_shape_k,
|
||||
# bsa_params
|
||||
sparsity=0.875,
|
||||
cdf_threshold=None,
|
||||
chunk_3d_shape_q=[4, 4, 8],
|
||||
chunk_3d_shape_k=[4, 4, 8],
|
||||
) -> torch.Tensor:
|
||||
_, _, Sq, head_dim_q = q.shape
|
||||
_, _, Sk, head_dim_k = k.shape
|
||||
|
||||
assert head_dim_q == head_dim_k
|
||||
head_dim = head_dim_q
|
||||
|
||||
Tq, Hq, Wq = latent_shape_q
|
||||
Tk, Hk, Wk = latent_shape_k
|
||||
|
||||
assert Tq * Hq * Wq == Sq
|
||||
assert Tk * Hk * Wk == Sk
|
||||
|
||||
tq, hq, wq = chunk_3d_shape_q
|
||||
tk, hk, wk = chunk_3d_shape_k
|
||||
|
||||
assert Tq % tq == 0 and Hq % hq == 0 and Wq % wq == 0
|
||||
assert Tk % tk == 0 and Hk % hk == 0 and Wk % wk == 0
|
||||
|
||||
Ntq = Tq // tq
|
||||
Nhq = Hq // hq
|
||||
Nwq = Wq // wq
|
||||
|
||||
Ntk = Tk // tk
|
||||
Nhk = Hk // hk
|
||||
Nwk = Wk // wk
|
||||
|
||||
q = rearrange_THW_to_3d_block(q, Ntq, Nhq, Nwq, tq, hq, wq, q.shape[-1])
|
||||
k = rearrange_THW_to_3d_block(k, Ntk, Nhk, Nwk, tk, hk, wk, k.shape[-1])
|
||||
v = rearrange_THW_to_3d_block(v, Ntk, Nhk, Nwk, tk, hk, wk, v.shape[-1])
|
||||
|
||||
chunk_size_q = tq * hq * wq
|
||||
chunk_size_k = tk * hk * wk
|
||||
|
||||
output = flash_attn_bsa(q, k, v, chunk_size_q, chunk_size_k, sparsity, cdf_threshold, 1 / head_dim**0.5)
|
||||
|
||||
output = rearrange_3d_block_to_THW(output, Ntq, Nhq, Nwq, tq, hq, wq, output.shape[-1])
|
||||
return output
|
||||
@@ -0,0 +1,111 @@
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import os
|
||||
|
||||
if os.environ.get('TRITON_AUTOTUNE_ENBALE', '0') == '1':
|
||||
autotune = triton.autotune
|
||||
else:
|
||||
def autotune(*args, **kwargs):
|
||||
def decorator(func):
|
||||
return func
|
||||
return decorator
|
||||
|
||||
configs_gating_preset = {
|
||||
'default': {
|
||||
'BLOCK_M': 64,
|
||||
'BLOCK_N': 64,
|
||||
'num_stages': 3,
|
||||
'num_warps': 8,
|
||||
}
|
||||
}
|
||||
|
||||
configs_gating = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64, 128] \
|
||||
for BN in [32, 64] \
|
||||
for s in [2, 3, 4, 5] \
|
||||
for w in [4, 8] \
|
||||
]
|
||||
|
||||
gating_reevaluate_keys = ["M", "N"] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
|
||||
@autotune(configs_gating, key=gating_reevaluate_keys)
|
||||
@triton.jit
|
||||
def _attn_fwd_gating(
|
||||
Q, K, Out,
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
H, M, N,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
):
|
||||
|
||||
tl.static_assert(BLOCK_N <= HEAD_DIM)
|
||||
start_m = tl.program_id(0)
|
||||
off_hz = tl.program_id(1)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
q_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
k_offset = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
|
||||
o_offset = off_z.to(tl.int64) * stride_oz + off_h.to(tl.int64) * stride_oh
|
||||
|
||||
# block pointers
|
||||
Q_block_ptr = tl.make_block_ptr(
|
||||
base=Q + q_offset,
|
||||
shape=(M, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
|
||||
K_block_ptr = tl.make_block_ptr(
|
||||
base=K + k_offset,
|
||||
shape=(HEAD_DIM, N),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N),
|
||||
order=(0, 1),
|
||||
)
|
||||
O_block_ptr = tl.make_block_ptr(
|
||||
base=Out + o_offset,
|
||||
shape=(M, N),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, BLOCK_N),
|
||||
order=(1, 0),
|
||||
)
|
||||
|
||||
# load q: it will stay in SRAM throughout
|
||||
q = tl.load(Q_block_ptr, boundary_check=(0,))
|
||||
for start_n in range(0, N, BLOCK_N):
|
||||
start_n = tl.multiple_of(start_n, BLOCK_N)
|
||||
# -- compute qk ----
|
||||
k = tl.load(K_block_ptr, boundary_check=(1,))
|
||||
qk = tl.dot(q, k)
|
||||
|
||||
tl.store(O_block_ptr, qk.to(Out.type.element_ty), boundary_check=(0, 1))
|
||||
|
||||
K_block_ptr = tl.advance(K_block_ptr, (0, BLOCK_N))
|
||||
O_block_ptr = tl.advance(O_block_ptr, (0, BLOCK_N))
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(
|
||||
O, DO,
|
||||
Delta, # output
|
||||
N_CTX,
|
||||
BLOCK_M: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr
|
||||
):
|
||||
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
@@ -0,0 +1,43 @@
|
||||
import torch
|
||||
|
||||
|
||||
def p2p_communicate(
|
||||
rank, send_tensor, send_dst, recv_tensor, recv_src, cp_group, batch_p2p_comm
|
||||
):
|
||||
"""Point-to-point communications of KV and dKV in Attention with context parallelism"""
|
||||
send_recv_ops = []
|
||||
|
||||
if batch_p2p_comm: # int(os.getenv("NVTE_BATCH_MHA_P2P_COMM", "0")) or (cp_size == 2) 为啥呢
|
||||
if rank % 2 == 0:
|
||||
send_op = torch.distributed.P2POp(
|
||||
torch.distributed.isend, send_tensor, send_dst, cp_group
|
||||
)
|
||||
recv_op = torch.distributed.P2POp(
|
||||
torch.distributed.irecv, recv_tensor, recv_src, cp_group
|
||||
)
|
||||
send_recv_ops.append(send_op)
|
||||
send_recv_ops.append(recv_op)
|
||||
else:
|
||||
recv_op = torch.distributed.P2POp(
|
||||
torch.distributed.irecv, recv_tensor, recv_src, cp_group
|
||||
)
|
||||
send_op = torch.distributed.P2POp(
|
||||
torch.distributed.isend, send_tensor, send_dst, cp_group
|
||||
)
|
||||
send_recv_ops.append(recv_op)
|
||||
send_recv_ops.append(send_op)
|
||||
send_recv_reqs = torch.distributed.batch_isend_irecv(send_recv_ops)
|
||||
else:
|
||||
if rank % 2 == 0:
|
||||
send_op = torch.distributed.isend(send_tensor, send_dst, cp_group)
|
||||
recv_op = torch.distributed.irecv(recv_tensor, recv_src, cp_group)
|
||||
send_recv_ops.append(send_op)
|
||||
send_recv_ops.append(recv_op)
|
||||
else:
|
||||
recv_op = torch.distributed.irecv(recv_tensor, recv_src, cp_group)
|
||||
send_op = torch.distributed.isend(send_tensor, send_dst, cp_group)
|
||||
send_recv_ops.append(recv_op)
|
||||
send_recv_ops.append(send_op)
|
||||
send_recv_reqs = send_recv_ops
|
||||
|
||||
return send_recv_reqs
|
||||
+946
@@ -0,0 +1,946 @@
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import os
|
||||
|
||||
from .common import autotune
|
||||
"""
|
||||
TRITON_REEVALUATE_KEY=1
|
||||
- autotune whenever params in reevaluate keys change
|
||||
- use in benchmark script to fine the best config
|
||||
|
||||
TRITON_AUTOTUNE_ENBALE=1
|
||||
- if set to 0, autotune will not work, and the related params must be passed to the function call.
|
||||
"""
|
||||
|
||||
configs_fwd_bsa_varlen_preset = {
|
||||
'default': {
|
||||
'BLOCK_N': 64,
|
||||
'num_stages': 3,
|
||||
'num_warps': 8,
|
||||
},
|
||||
'BLOCK_N_LG=64': {
|
||||
'BLOCK_N': 64,
|
||||
'num_stages': 3,
|
||||
'num_warps': 4,
|
||||
},
|
||||
}
|
||||
configs_fwd_bsa_varlen = [
|
||||
triton.Config({'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BN in [32, 64, 128] \
|
||||
for s in [2, 3, 4, 5] \
|
||||
for w in [4, 8] \
|
||||
]
|
||||
|
||||
fwd_bsa_reevaluate_varlen_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
|
||||
@autotune(list(configs_fwd_bsa_varlen), key=fwd_bsa_reevaluate_varlen_keys)
|
||||
@triton.jit
|
||||
def _attn_fwd_bsa_varlen(
|
||||
Q, K, V, sm_scale, M, Out,
|
||||
block_indices, # [B, H, M_COMPRESS, S_MAX]
|
||||
block_indices_lens, # [B, H, M_COMPRESS]
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vn, stride_vk,
|
||||
stride_oz, stride_oh, stride_om, stride_ok,
|
||||
stride_bz, stride_bh, stride_bm, stride_bs,
|
||||
stride_lz, stride_lh, stride_lm,
|
||||
H, N_CTX,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N_LG: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
|
||||
):
|
||||
|
||||
start_m = tl.program_id(0)
|
||||
off_hz = tl.program_id(1)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
|
||||
q_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
k_offset = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
|
||||
v_offset = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
|
||||
o_offset = off_z.to(tl.int64) * stride_oz + off_h.to(tl.int64) * stride_oh
|
||||
b_offset = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
|
||||
l_offset = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
|
||||
|
||||
# block pointers
|
||||
Q_block_ptr = tl.make_block_ptr(
|
||||
base=Q + q_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
V_block_ptr = tl.make_block_ptr(
|
||||
base=V + v_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vn, stride_vk),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
KT_block_ptr = tl.make_block_ptr(
|
||||
base=K + k_offset,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N),
|
||||
order=(0, 1),
|
||||
)
|
||||
O_block_ptr = tl.make_block_ptr(
|
||||
base=Out + o_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_ok),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
block_indices += b_offset + start_m * stride_bm
|
||||
block_indices_lens += l_offset + start_m * stride_lm
|
||||
# initialize offsets
|
||||
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
|
||||
# initialize pointer to m and l
|
||||
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
# load scales
|
||||
qk_scale = sm_scale
|
||||
qk_scale *= 1.44269504 # 1/ln2; exp2(x/ln2) == exp2(ln(e^x) / ln2) == exp2(log2(e^x)) == exp(x)
|
||||
# load q: it will stay in SRAM throughout
|
||||
q = tl.load(Q_block_ptr)
|
||||
S = tl.load(block_indices_lens)
|
||||
for i in range(S):
|
||||
block_id = tl.load(block_indices + i * stride_bs).to(tl.int32)
|
||||
lo, hi = block_id * BLOCK_N_LG, (block_id + 1) * BLOCK_N_LG
|
||||
lo = tl.multiple_of(lo, BLOCK_N)
|
||||
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
|
||||
V_block_ptr_i = tl.advance(V_block_ptr, (lo, 0))
|
||||
|
||||
# loop over k, v and update accumulator
|
||||
for start_n in range(lo, hi, BLOCK_N):
|
||||
start_n = tl.multiple_of(start_n, BLOCK_N)
|
||||
# -- compute qk ----
|
||||
kT = tl.load(KT_block_ptr_i)
|
||||
qkT = tl.dot(q, kT)
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qkT, 1) * qk_scale)
|
||||
qkT = qkT * qk_scale - m_ij[:, None]
|
||||
p = tl.math.exp2(qkT)
|
||||
|
||||
# -- update m_i and l_i
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_ij = tl.sum(p, 1)
|
||||
# -- update output accumulator --
|
||||
acc = acc * alpha[:, None]
|
||||
# update acc
|
||||
v = tl.load(V_block_ptr_i)
|
||||
acc = tl.dot(p.to(v.dtype), v, acc)
|
||||
# update m_i and l_i
|
||||
# place this at the end of the loop to reduce register pressure: https://github.com/triton-lang/triton/commit/ee6abd9
|
||||
l_i = l_i * alpha + l_ij
|
||||
m_i = m_ij
|
||||
V_block_ptr_i = tl.advance(V_block_ptr_i, (BLOCK_N, 0))
|
||||
KT_block_ptr_i = tl.advance(KT_block_ptr_i, (0, BLOCK_N))
|
||||
|
||||
# epilogue
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
m_ptrs = M + off_hz * N_CTX + offs_m
|
||||
tl.store(m_ptrs, m_i)
|
||||
tl.store(O_block_ptr, acc.to(Out.type.element_ty))
|
||||
|
||||
configs_fwd_bsa_varlen_align_preset = {
|
||||
'default': {
|
||||
'num_stages': 3,
|
||||
'num_warps': 8,
|
||||
},
|
||||
'BLOCK_N_LG=64': {
|
||||
'num_stages': 3,
|
||||
'num_warps': 4,
|
||||
},
|
||||
}
|
||||
configs_fwd_bsa_varlen_align = [
|
||||
triton.Config({}, num_stages=s, num_warps=w) \
|
||||
for s in [2, 3, 4, 5] \
|
||||
for w in [4, 8] \
|
||||
]
|
||||
|
||||
fwd_bsa_reevaluate_varlen_align_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
|
||||
@autotune(list(configs_fwd_bsa_varlen_align), key=fwd_bsa_reevaluate_varlen_align_keys)
|
||||
@triton.jit
|
||||
def _attn_fwd_bsa_varlen_align(
|
||||
Q, K, V, sm_scale, M, Out,
|
||||
block_indices, # [B, H, M_COMPRESS, S_MAX]
|
||||
block_indices_lens, # [B, H, M_COMPRESS]
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vn, stride_vk,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
stride_bz, stride_bh, stride_bm, stride_bs,
|
||||
stride_lz, stride_lh, stride_lm,
|
||||
H, N_CTX,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N_LG: tl.constexpr,
|
||||
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
|
||||
):
|
||||
|
||||
start_m = tl.program_id(0)
|
||||
off_hz = tl.program_id(1)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
|
||||
q_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
k_offset = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
|
||||
v_offset = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
|
||||
o_offset = off_z.to(tl.int64) * stride_oz + off_h.to(tl.int64) * stride_oh
|
||||
b_offset = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
|
||||
l_offset = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
|
||||
|
||||
# block pointers
|
||||
Q_block_ptr = tl.make_block_ptr(
|
||||
base=Q + q_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
V_block_ptr = tl.make_block_ptr(
|
||||
base=V + v_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vn, stride_vk),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N_LG, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
KT_block_ptr = tl.make_block_ptr(
|
||||
base=K + k_offset,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N_LG),
|
||||
order=(0, 1),
|
||||
)
|
||||
O_block_ptr = tl.make_block_ptr(
|
||||
base=Out + o_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
block_indices += b_offset + start_m * stride_bm
|
||||
block_indices_lens += l_offset + start_m * stride_lm
|
||||
# initialize offsets
|
||||
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
|
||||
# initialize pointer to m and l
|
||||
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
# load scales
|
||||
qk_scale = sm_scale
|
||||
qk_scale *= 1.44269504 # 1/ln2; exp2(x/ln2) == exp2(ln(e^x) / ln2) == exp2(log2(e^x)) == exp(x);乘1/ln2后,exp2(x/ln2) == exp(x),exp2速度更快
|
||||
# load q: it will stay in SRAM throughout
|
||||
q = tl.load(Q_block_ptr)
|
||||
S = tl.load(block_indices_lens)
|
||||
for i in range(S):
|
||||
block_id = tl.load(block_indices + i * stride_bs).to(tl.int32)
|
||||
lo = block_id * BLOCK_N_LG
|
||||
lo = tl.multiple_of(lo, BLOCK_N_LG)
|
||||
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
|
||||
V_block_ptr_i = tl.advance(V_block_ptr, (lo, 0))
|
||||
|
||||
# -- compute qk ----
|
||||
kT = tl.load(KT_block_ptr_i)
|
||||
qkT = tl.dot(q, kT)
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qkT, 1) * qk_scale)
|
||||
qkT = qkT * qk_scale - m_ij[:, None]
|
||||
p = tl.math.exp2(qkT)
|
||||
|
||||
# -- update m_i and l_i
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_ij = tl.sum(p, 1)
|
||||
# -- update output accumulator --
|
||||
acc = acc * alpha[:, None]
|
||||
# update acc
|
||||
v = tl.load(V_block_ptr_i)
|
||||
acc = tl.dot(p.to(v.dtype), v, acc) # 没除se,fa2引入的优化
|
||||
# update m_i and l_i
|
||||
# place this at the end of the loop to reduce register pressure: https://github.com/triton-lang/triton/commit/ee6abd9
|
||||
l_i = l_i * alpha + l_ij # 当前总se
|
||||
m_i = m_ij
|
||||
|
||||
|
||||
# epilogue
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
m_ptrs = M + off_hz * N_CTX + offs_m
|
||||
tl.store(m_ptrs, m_i)
|
||||
tl.store(O_block_ptr, acc.to(Out.type.element_ty))
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv_bsa_varlen(
|
||||
dk, dv,
|
||||
k, v,
|
||||
Q, DO,
|
||||
M, D,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
# shared by Q/K/V/DO.
|
||||
# stride_tok, stride_d,
|
||||
stride_qm, stride_qk,
|
||||
stride_dom, stride_dok,
|
||||
stride_mm,
|
||||
stride_dm,
|
||||
stride_bm,
|
||||
N_CTX,
|
||||
BLOCK_M1: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
):
|
||||
|
||||
QT_block_ptr = tl.make_block_ptr(
|
||||
base=Q,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_qk, stride_qm),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_M1),
|
||||
order=(0, 1),
|
||||
)
|
||||
|
||||
DO_block_ptr = tl.make_block_ptr(
|
||||
base=DO,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dom, stride_dok),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_M1, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
|
||||
S = tl.load(block_indices_lens)
|
||||
for i in range(S):
|
||||
block_id = tl.load(block_indices + i * stride_bm).to(tl.int32)
|
||||
start_m = block_id * BLOCK_M1
|
||||
start_m = tl.multiple_of(start_m, BLOCK_M1)
|
||||
|
||||
QT_block_ptr_i = tl.advance(QT_block_ptr, (0, start_m))
|
||||
DO_block_ptr_i = tl.advance(DO_block_ptr, (start_m, 0))
|
||||
|
||||
qT = tl.load(QT_block_ptr_i)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1) * stride_mm
|
||||
m = tl.load(M + offs_m)
|
||||
kqT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(kqT - m[None, :])
|
||||
|
||||
do = tl.load(DO_block_ptr_i)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(v.dtype)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
offs_d = start_m + tl.arange(0, BLOCK_M1) * stride_dm
|
||||
Di = tl.load(D + offs_d)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(v.dtype)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
|
||||
|
||||
return dk, dv
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq_bsa_varlen(
|
||||
dq,
|
||||
q, do,
|
||||
m, d,
|
||||
K, V,
|
||||
N_CTX,
|
||||
BLOCK_N2: tl.constexpr,
|
||||
BLOCK_N_LG: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
stride_bn,
|
||||
# stride_tok, stride_d,
|
||||
stride_kn, stride_kk,
|
||||
stride_vn, stride_vk,
|
||||
):
|
||||
|
||||
VT_block_ptr = tl.make_block_ptr(
|
||||
base=V,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N2),
|
||||
order=(0, 1),
|
||||
)
|
||||
KT_block_ptr = tl.make_block_ptr(
|
||||
base=K,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N2),
|
||||
order=(0, 1),
|
||||
)
|
||||
|
||||
S = tl.load(block_indices_lens)
|
||||
for i in range(S):
|
||||
block_id = tl.load(block_indices + i * stride_bn).to(tl.int32)
|
||||
lo, hi = block_id * BLOCK_N_LG, (block_id + 1) * BLOCK_N_LG
|
||||
lo = tl.multiple_of(lo, BLOCK_N2)
|
||||
|
||||
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
|
||||
VT_block_ptr_i = tl.advance(VT_block_ptr, (0, lo))
|
||||
|
||||
for start_n in range(lo, hi, BLOCK_N2):
|
||||
start_n = tl.multiple_of(start_n, BLOCK_N2)
|
||||
|
||||
kT = tl.load(KT_block_ptr_i)
|
||||
vT = tl.load(VT_block_ptr_i)
|
||||
qkT = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qkT - m)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - d)
|
||||
ds = ds.to(kT.dtype) # https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/flash_attn_triton.py: Converting ds to q.dtype here reduces register pressure and makes it much faster for BLOCK_HEADDIM=128
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
KT_block_ptr_i = tl.advance(KT_block_ptr_i, (0, BLOCK_N2))
|
||||
VT_block_ptr_i = tl.advance(VT_block_ptr_i, (0, BLOCK_N2))
|
||||
|
||||
return dq
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq_bsa_varlen_align(
|
||||
dq,
|
||||
q, do,
|
||||
m, d,
|
||||
K, V,
|
||||
N_CTX,
|
||||
BLOCK_N_LG: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
stride_bn,
|
||||
stride_kn, stride_kk,
|
||||
stride_vn, stride_vk,
|
||||
):
|
||||
|
||||
VT_block_ptr = tl.make_block_ptr(
|
||||
base=V,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N_LG),
|
||||
order=(0, 1),
|
||||
)
|
||||
KT_block_ptr = tl.make_block_ptr(
|
||||
base=K,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N_LG),
|
||||
order=(0, 1),
|
||||
)
|
||||
|
||||
S = tl.load(block_indices_lens)
|
||||
for i in range(S):
|
||||
block_id = tl.load(block_indices + i * stride_bn).to(tl.int32)
|
||||
lo = block_id * BLOCK_N_LG
|
||||
lo = tl.multiple_of(lo, BLOCK_N_LG)
|
||||
|
||||
KT_block_ptr_i = tl.advance(KT_block_ptr, (0, lo))
|
||||
VT_block_ptr_i = tl.advance(VT_block_ptr, (0, lo))
|
||||
|
||||
kT = tl.load(KT_block_ptr_i)
|
||||
vT = tl.load(VT_block_ptr_i)
|
||||
qkT = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qkT - m)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - d)
|
||||
ds = ds.to(kT.dtype) # https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/flash_attn_triton.py: Converting ds to q.dtype here reduces register pressure and makes it much faster for BLOCK_HEADDIM=128
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
|
||||
|
||||
return dq
|
||||
|
||||
configs_bwd_dkdv_bsa_varlen_preset = {
|
||||
'default': {
|
||||
'BLOCK_N': 128,
|
||||
'num_stages': 2,
|
||||
'num_warps': 8,
|
||||
},
|
||||
'BLOCK_N_DQ_LG=64': {
|
||||
'BLOCK_N': 64,
|
||||
'num_stages': 2,
|
||||
'num_warps': 4,
|
||||
}
|
||||
}
|
||||
configs_bwd_dkdv_bsa_varlen = [
|
||||
triton.Config({'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BN in [32, 64, 128] \
|
||||
for s in [2, 3, 4, 5] \
|
||||
for w in [4, 8] \
|
||||
]
|
||||
bwd_dkdv_bsa_varlen_reevaluate_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_DQ_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
|
||||
@autotune(list(configs_bwd_dkdv_bsa_varlen), key=bwd_dkdv_bsa_varlen_reevaluate_keys)
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv_bsa_varlen_wrapper(
|
||||
Q, K, V, sm_scale, # softmax scale
|
||||
DO,
|
||||
DK, DV,
|
||||
M, # lse (log2)
|
||||
D,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
# stride_z, stride_h, stride_tok, stride_d, # shared by Q/K/V/DO.
|
||||
# qkv
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vn, stride_vk,
|
||||
# dk dv do
|
||||
stride_dkz, stride_dkh, stride_dkn, stride_dkk,
|
||||
stride_dvz, stride_dvh, stride_dvn, stride_dvk,
|
||||
stride_doz, stride_doh, stride_dom, stride_dok,
|
||||
# m, d
|
||||
stride_mz, stride_mh, stride_mm,
|
||||
stride_dz, stride_dh, stride_dm,
|
||||
#
|
||||
stride_bz, stride_bh, stride_bn, stride_bm, # block_indices
|
||||
stride_lz, stride_lh, stride_ln, # block_indices_lens
|
||||
#
|
||||
H, N_CTX,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
BLOCK_N_DQ_LG: tl.constexpr, # logical block size
|
||||
HEAD_DIM: tl.constexpr,
|
||||
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
|
||||
):
|
||||
|
||||
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_N == 0)
|
||||
start_n = tl.program_id(0)
|
||||
|
||||
off_hz = tl.program_id(2)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
|
||||
off_q = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
off_k = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
|
||||
off_v = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
|
||||
|
||||
off_dk = off_z.to(tl.int64) * stride_dkz + off_h.to(tl.int64) * stride_dkh
|
||||
off_dv = off_z.to(tl.int64) * stride_dvz + off_h.to(tl.int64) * stride_dvh
|
||||
off_do = off_z.to(tl.int64) * stride_doz + off_h.to(tl.int64) * stride_doh
|
||||
|
||||
off_m = off_z.to(tl.int64) * stride_mz + off_h.to(tl.int64) * stride_mh
|
||||
off_d = off_z.to(tl.int64) * stride_dz + off_h.to(tl.int64) * stride_dh
|
||||
|
||||
off_block_incides = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
|
||||
off_block_incides_lens = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += off_q
|
||||
K += off_k
|
||||
V += off_v
|
||||
DO += off_do
|
||||
DK += off_dk
|
||||
DV += off_dv
|
||||
|
||||
M += off_m
|
||||
D += off_d
|
||||
block_indices += off_block_incides
|
||||
block_indices_lens += off_block_incides_lens
|
||||
|
||||
# ---------------------------------------- [DKDV] ----------------------------------------
|
||||
|
||||
dv = tl.zeros([BLOCK_N, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
|
||||
K_block_ptr = tl.make_block_ptr(
|
||||
base=K,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_kn, stride_kk),
|
||||
offsets=(start_n * BLOCK_N, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
V_block_ptr = tl.make_block_ptr(
|
||||
base=V,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vn, stride_vk),
|
||||
offsets=(start_n * BLOCK_N, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
DK_block_ptr = tl.make_block_ptr(
|
||||
base=DK,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dkn, stride_dkk),
|
||||
offsets=(start_n * BLOCK_N, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
DV_block_ptr = tl.make_block_ptr(
|
||||
base=DV,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dvn, stride_dvk),
|
||||
offsets=(start_n * BLOCK_N, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
|
||||
k = tl.load(K_block_ptr)
|
||||
v = tl.load(V_block_ptr)
|
||||
|
||||
k_compress_idx = start_n * BLOCK_N // BLOCK_N_DQ_LG
|
||||
block_indices_i = block_indices + k_compress_idx * stride_bn
|
||||
block_indices_lens_i = block_indices_lens + k_compress_idx * stride_ln
|
||||
|
||||
dk, dv = _attn_bwd_dkdv_bsa_varlen(
|
||||
dk, dv,
|
||||
k, v,
|
||||
Q, DO,
|
||||
M, D,
|
||||
block_indices_i,
|
||||
block_indices_lens_i,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_qm, stride_qk,
|
||||
stride_dom, stride_dok,
|
||||
stride_mm,
|
||||
stride_dm,
|
||||
#
|
||||
stride_bm,
|
||||
N_CTX,
|
||||
BLOCK_M,
|
||||
HEAD_DIM,
|
||||
)
|
||||
|
||||
# Write back dk
|
||||
dk *= sm_scale # S = scale * QKT; dK = scale * QdST
|
||||
tl.store(DK_block_ptr, dk.to(k.dtype))
|
||||
|
||||
# Write back dv
|
||||
tl.store(DV_block_ptr, dv.to(v.dtype))
|
||||
|
||||
configs_bwd_dq_bsa_varlen_preset = {
|
||||
'default': {
|
||||
'BLOCK_N_DQ': 64,
|
||||
'num_stages': 2,
|
||||
'num_warps': 8,
|
||||
},
|
||||
'BLOCK_N_DQ_LG=64': {
|
||||
'BLOCK_N_DQ': 64,
|
||||
'num_stages': 2,
|
||||
'num_warps': 4,
|
||||
},
|
||||
}
|
||||
configs_bwd_dq_bsa_varlen = [
|
||||
triton.Config({'BLOCK_N_DQ': BN}, num_stages=s, num_warps=w) \
|
||||
for BN in [32, 64, 128] \
|
||||
for s in [2, 3, 4, 5] \
|
||||
for w in [4, 8] \
|
||||
]
|
||||
bwd_dq_bsa_varlen_reevaluate_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_DQ_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
|
||||
@autotune(list(configs_bwd_dq_bsa_varlen), key=bwd_dq_bsa_varlen_reevaluate_keys)
|
||||
@triton.jit
|
||||
def _attn_bwd_dq_bsa_varlen_wrapper(
|
||||
Q, K, V, # softmax scale
|
||||
DO,
|
||||
DQ,
|
||||
M, # lse (log2)
|
||||
D,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
# stride_z, stride_h, stride_tok, stride_d, # shared by Q/K/V/DO.
|
||||
# qkv
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vn, stride_vk,
|
||||
# dq do
|
||||
stride_dqz, stride_dqh, stride_dqm, stride_dqk,
|
||||
stride_doz, stride_doh, stride_dom, stride_dok,
|
||||
# m, d
|
||||
stride_mz, stride_mh, stride_mm,
|
||||
stride_dz, stride_dh, stride_dm,
|
||||
#
|
||||
stride_bz, stride_bh, stride_bm, stride_bn, # block_indices
|
||||
stride_lz, stride_lh, stride_lm, # block_indices_lens
|
||||
#
|
||||
H, N_CTX,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N_DQ_LG: tl.constexpr, # logical block size
|
||||
BLOCK_N_DQ: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
|
||||
):
|
||||
|
||||
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_N_DQ == 0)
|
||||
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_M == 0)
|
||||
|
||||
LN2: tl.constexpr = 0.6931471824645996 # = ln(2)
|
||||
|
||||
start_m = tl.program_id(0)
|
||||
|
||||
off_hz = tl.program_id(2)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
|
||||
off_q = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
off_k = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
|
||||
off_v = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
|
||||
|
||||
off_dq = off_z.to(tl.int64) * stride_dqz + off_h.to(tl.int64) * stride_dqh
|
||||
off_do = off_z.to(tl.int64) * stride_doz + off_h.to(tl.int64) * stride_doh
|
||||
|
||||
off_m = off_z.to(tl.int64) * stride_mz + off_h.to(tl.int64) * stride_mh
|
||||
off_d = off_z.to(tl.int64) * stride_dz + off_h.to(tl.int64) * stride_dh
|
||||
|
||||
off_block_incides = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
|
||||
off_block_incides_lens = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += off_q
|
||||
K += off_k
|
||||
V += off_v
|
||||
DO += off_do
|
||||
DQ += off_dq
|
||||
|
||||
M += off_m
|
||||
D += off_d
|
||||
block_indices += off_block_incides
|
||||
block_indices_lens += off_block_incides_lens
|
||||
|
||||
# ---------------------------------------- [DQ] ----------------------------------------
|
||||
|
||||
Q_block_ptr = tl.make_block_ptr(
|
||||
base=Q,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
DO_block_ptr = tl.make_block_ptr(
|
||||
base=DO,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dom, stride_dok),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
DQ_block_ptr = tl.make_block_ptr(
|
||||
base=DQ,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dqm, stride_dqk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
|
||||
q = tl.load(Q_block_ptr)
|
||||
do = tl.load(DO_block_ptr)
|
||||
|
||||
dq = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
start_m = start_m * BLOCK_M
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M) * stride_mm
|
||||
offs_d = start_m + tl.arange(0, BLOCK_M) * stride_dm
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
d = tl.load(D + offs_d)
|
||||
d = d[:, None]
|
||||
|
||||
block_indices_m = block_indices + (start_m // BLOCK_M) * stride_bm
|
||||
block_indices_lens_m = block_indices_lens + (start_m // BLOCK_M) * stride_lm
|
||||
|
||||
dq = _attn_bwd_dq_bsa_varlen(
|
||||
dq,
|
||||
q, do,
|
||||
m, d,
|
||||
K, V,
|
||||
N_CTX,
|
||||
BLOCK_N_DQ,
|
||||
BLOCK_N_DQ_LG,
|
||||
HEAD_DIM,
|
||||
block_indices_m,
|
||||
block_indices_lens_m,
|
||||
stride_bn,
|
||||
stride_kn, stride_kk,
|
||||
stride_vn, stride_vk,
|
||||
)
|
||||
|
||||
# Write back dQ.
|
||||
dq *= LN2
|
||||
tl.store(DQ_block_ptr, dq.to(q.dtype))
|
||||
|
||||
configs_bwd_dq_bsa_varlen_align_preset = {
|
||||
'default': {
|
||||
'num_stages': 2,
|
||||
'num_warps': 8,
|
||||
},
|
||||
'BLOCK_N_DQ_LG=64': {
|
||||
'num_stages': 2,
|
||||
'num_warps': 4,
|
||||
},
|
||||
}
|
||||
configs_bwd_dq_bsa_varlen_align = [
|
||||
triton.Config({}, num_stages=s, num_warps=w) \
|
||||
for s in [2, 3, 4, 5] \
|
||||
for w in [4, 8] \
|
||||
]
|
||||
bwd_dq_bsa_varlen_align_reevaluate_keys = ['N_CTX', 'BLOCK_M', 'BLOCK_N_DQ_LG', 'SPARSITY'] if os.environ.get('TRITON_REEVALUATE_KEY', '0') == '1' else []
|
||||
@autotune(list(configs_bwd_dq_bsa_varlen_align), key=bwd_dq_bsa_varlen_align_reevaluate_keys)
|
||||
@triton.jit
|
||||
def _attn_bwd_dq_bsa_varlen_align_wrapper(
|
||||
Q, K, V, # softmax scale
|
||||
DO,
|
||||
DQ,
|
||||
M, # lse (log2)
|
||||
D,
|
||||
block_indices,
|
||||
block_indices_lens,
|
||||
# qkv
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vn, stride_vk,
|
||||
# dq do
|
||||
stride_dqz, stride_dqh, stride_dqm, stride_dqk,
|
||||
stride_doz, stride_doh, stride_dom, stride_dok,
|
||||
# m, d
|
||||
stride_mz, stride_mh, stride_mm,
|
||||
stride_dz, stride_dh, stride_dm,
|
||||
#
|
||||
stride_bz, stride_bh, stride_bm, stride_bn, # block_indices
|
||||
stride_lz, stride_lh, stride_lm, # block_indices_lens
|
||||
#
|
||||
H, N_CTX,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N_DQ_LG: tl.constexpr, # logical block size
|
||||
HEAD_DIM: tl.constexpr,
|
||||
SPARSITY: tl.constexpr, # not used; just for trigger reevaluate for benchmarking
|
||||
):
|
||||
|
||||
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_N_DQ_LG == 0)
|
||||
tl.static_assert(BLOCK_N_DQ_LG % BLOCK_M == 0)
|
||||
|
||||
LN2: tl.constexpr = 0.6931471824645996 # = ln(2)
|
||||
|
||||
start_m = tl.program_id(0)
|
||||
|
||||
off_hz = tl.program_id(2)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
|
||||
off_q = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
off_k = off_z.to(tl.int64) * stride_kz + off_h.to(tl.int64) * stride_kh
|
||||
off_v = off_z.to(tl.int64) * stride_vz + off_h.to(tl.int64) * stride_vh
|
||||
|
||||
off_dq = off_z.to(tl.int64) * stride_dqz + off_h.to(tl.int64) * stride_dqh
|
||||
off_do = off_z.to(tl.int64) * stride_doz + off_h.to(tl.int64) * stride_doh
|
||||
|
||||
off_m = off_z.to(tl.int64) * stride_mz + off_h.to(tl.int64) * stride_mh
|
||||
off_d = off_z.to(tl.int64) * stride_dz + off_h.to(tl.int64) * stride_dh
|
||||
|
||||
off_block_incides = off_z.to(tl.int64) * stride_bz + off_h.to(tl.int64) * stride_bh
|
||||
off_block_incides_lens = off_z.to(tl.int64) * stride_lz + off_h.to(tl.int64) * stride_lh
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += off_q
|
||||
K += off_k
|
||||
V += off_v
|
||||
DO += off_do
|
||||
DQ += off_dq
|
||||
|
||||
M += off_m
|
||||
D += off_d
|
||||
block_indices += off_block_incides
|
||||
block_indices_lens += off_block_incides_lens
|
||||
|
||||
# ---------------------------------------- [DQ] ----------------------------------------
|
||||
|
||||
Q_block_ptr = tl.make_block_ptr(
|
||||
base=Q,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
DO_block_ptr = tl.make_block_ptr(
|
||||
base=DO,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dom, stride_dok),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
DQ_block_ptr = tl.make_block_ptr(
|
||||
base=DQ,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_dqm, stride_dqk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
|
||||
q = tl.load(Q_block_ptr)
|
||||
do = tl.load(DO_block_ptr)
|
||||
|
||||
dq = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
start_m = start_m * BLOCK_M
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M) * stride_mm
|
||||
offs_d = start_m + tl.arange(0, BLOCK_M) * stride_dm
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
d = tl.load(D + offs_d)
|
||||
d = d[:, None]
|
||||
|
||||
block_indices_m = block_indices + (start_m // BLOCK_M) * stride_bm
|
||||
block_indices_lens_m = block_indices_lens + (start_m // BLOCK_M) * stride_lm
|
||||
|
||||
dq = _attn_bwd_dq_bsa_varlen_align(
|
||||
dq,
|
||||
q, do,
|
||||
m, d,
|
||||
K, V,
|
||||
N_CTX,
|
||||
BLOCK_N_DQ_LG,
|
||||
HEAD_DIM,
|
||||
block_indices_m,
|
||||
block_indices_lens_m,
|
||||
stride_bn,
|
||||
stride_kn, stride_kk,
|
||||
stride_vn, stride_vk,
|
||||
)
|
||||
|
||||
# Write back dQ.
|
||||
dq *= LN2
|
||||
tl.store(DQ_block_ptr, dq.to(q.dtype))
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user