[Feature][VSA]Update STA publish workflow (#498)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
This commit is contained in:
co-authored by
JerryZhou54
parent
66012d3a4c
commit
2a46902ecb
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -136,19 +136,21 @@ jobs:
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -163,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
path: csrc/attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -229,17 +231,19 @@ jobs:
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
packages-dir: csrc/attn/dist/
|
||||
|
||||
@@ -28,4 +28,4 @@ jobs:
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
pytest --ignore csrc/attn/test
|
||||
|
||||
@@ -27,7 +27,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
csrc/attn/tk/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
+12
-12
@@ -7,26 +7,26 @@
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
Install STA:
|
||||
## Environment Setup
|
||||
First, set up your CUDA environment:
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Install Sliding Tile Attention (STA)
|
||||
```bash
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
## Install Video Sparse Attention (VSA)
|
||||
```bash
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
|
||||
@@ -2,8 +2,8 @@ import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from st_attn import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
|
||||
from st_attn import BLOCK_M, BLOCK_N
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
@@ -9,14 +9,14 @@ 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), (36, 48, 48), 39, 'cuda', 0)
|
||||
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):
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
|
||||
@@ -37,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
'max_diff': 0
|
||||
},
|
||||
}
|
||||
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
from tqdm import tqdm
|
||||
for kernel_size in tqdm(kernel_size_ls):
|
||||
for _ in range(num_iterations):
|
||||
@@ -74,7 +74,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 24, 128
|
||||
n = 82944 # Sequence length
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
Submodule
+1
Submodule csrc/attn/tk added at 1719fb7264
+283
-24
@@ -2,7 +2,7 @@ import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
|
||||
from typing import Tuple
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
@@ -12,33 +12,116 @@ except ImportError:
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
|
||||
|
||||
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
|
||||
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
# assert all elements in q2k_block_sparse_num can be devisible by 2
|
||||
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
return o, lse
|
||||
|
||||
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return grad_q, grad_k, grad_v
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
class BlockSparseAttentionFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return o
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements % 64 == 0 and block_elements >= 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
|
||||
grad_q, grad_k, grad_v = block_sparse_attention_backward(
|
||||
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
return grad_q, grad_k, grad_v, None, None, None, None
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
|
||||
block_attn_score, topk)
|
||||
|
||||
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
k2q_block_sparse_index,
|
||||
k2q_block_sparse_num)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
@torch._dynamo.disable
|
||||
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
@@ -61,6 +144,19 @@ def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
|
||||
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
|
||||
"""
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
|
||||
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
|
||||
"""
|
||||
# assert all elements in q2k_block_sparse_num can be devisible by 2
|
||||
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
return o, lse
|
||||
|
||||
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return grad_q, grad_k, grad_v
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
@@ -128,6 +224,169 @@ def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K
|
||||
|
||||
return mask
|
||||
|
||||
@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: tl.constexpr,
|
||||
):
|
||||
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: tl.constexpr,
|
||||
):
|
||||
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.static_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
|
||||
|
||||
class BlockSparseAttentionFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
|
||||
grad_q, grad_k, grad_v = block_sparse_attention_backward(
|
||||
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
return grad_q, grad_k, grad_v, None, None, None, None
|
||||
|
||||
|
||||
class DummyOperator(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x):
|
||||
@@ -207,4 +466,4 @@ class BlockSparseAttnTorch:
|
||||
|
||||
o = DummyOperator.apply(output)
|
||||
o.register_hook(self.recompute)
|
||||
return o
|
||||
return o
|
||||
|
||||
@@ -287,16 +287,10 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
forward_batch.mask_search_final_result_pos[timestep].append(
|
||||
layer_loss_save)
|
||||
else:
|
||||
# windows = [
|
||||
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
|
||||
# for head_idx in range(head_num)
|
||||
# ]
|
||||
windows = [
|
||||
STA_param[head_idx + start_head] for head_idx in range(head_num)
|
||||
]
|
||||
# if has_text is False:
|
||||
# from IPython import embed
|
||||
# embed()
|
||||
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, windows, text_length, has_text,
|
||||
self.dit_seq_shape_str).transpose(1, 2)
|
||||
|
||||
@@ -1,18 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, List, Optional, Type, cast
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from einops import rearrange
|
||||
|
||||
try:
|
||||
from vsa import block_sparse_attn
|
||||
except ImportError: # noqa: E722
|
||||
block_sparse_attn = None
|
||||
from typing import Tuple
|
||||
from vsa import video_sparse_attn
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
@@ -181,8 +174,7 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
(1 - attn_metadata.VSA_sparsity) *
|
||||
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
|
||||
|
||||
# Cast to Any to bypass type checking for untyped function
|
||||
hidden_states = cast(Any, sparse_attn_c_s_p)(
|
||||
hidden_states = video_sparse_attn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
@@ -191,267 +183,3 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
compress_attn_weight=gate_compress).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
|
||||
def sparse_attn_c_s_p(q, k, v, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements % 64 == 0 and block_elements >= 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
|
||||
block_attn_score, topk)
|
||||
|
||||
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
k2q_block_sparse_index,
|
||||
k2q_block_sparse_num)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
|
||||
@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: tl.constexpr,
|
||||
):
|
||||
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: tl.constexpr,
|
||||
):
|
||||
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.static_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
|
||||
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
@@ -607,9 +607,8 @@ class TrainingArgs(FastVideoArgs):
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
value = getattr(args, attr, default_value)
|
||||
if value is not None:
|
||||
kwargs[attr] = value
|
||||
if getattr(args, attr, default_value) is not None:
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user