Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9f0f714ed | ||
|
|
554a8f045e | ||
|
|
6d0eb1b956 | ||
|
|
20f5d1a779 | ||
|
|
5da1c9aadf | ||
|
|
8a77cf22c9 | ||
|
|
d869d90d12 | ||
|
|
554ee17de5 | ||
|
|
0be4fc62c9 |
@@ -0,0 +1,70 @@
|
||||
name: Publish FastVideo to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- 'pyproject.toml' # Trigger when pyproject.toml changes
|
||||
|
||||
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@v3
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
# Get current commit's version
|
||||
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-publish-main:
|
||||
needs: check-version-change
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install build dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install build twine wheel
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
python -m build
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: dist/
|
||||
@@ -0,0 +1,221 @@
|
||||
name: Publish Sliding Tile Attention Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
|
||||
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@v3
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.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")
|
||||
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'
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
|
||||
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
torch-version: ['2.5.1', '2.6.0']
|
||||
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
|
||||
|
||||
steps:
|
||||
- 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.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.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.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-version }}+cu${{ matrix.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-version }} --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 wheel
|
||||
run: |
|
||||
# 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
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
|
||||
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)
|
||||
# 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 }}
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels]
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
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
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
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: |
|
||||
# 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
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
@@ -23,8 +23,11 @@ jobs:
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install torch
|
||||
pip install packaging ninja
|
||||
# remove st-attn dependency because no cuda environment
|
||||
sed -i '/st_attn/d' pyproject.toml
|
||||
pip install -e .
|
||||
pip install pytest
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
+1
-1
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
path = sta_kernel/thunderkitten/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -160,7 +160,7 @@ def distill_one_step(
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
teacher_output = cond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
|
||||
@@ -813,7 +813,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
t, l, h = map(int, key.split('_'))
|
||||
result[t][l][h] = value
|
||||
return result
|
||||
torch._dynamo.config.cache_size_limit = 128
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
# if is_progress_bar:
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
|
||||
@@ -6,68 +6,12 @@ try:
|
||||
from st_attn import sliding_tile_attention
|
||||
except ImportError:
|
||||
print("Could not load Sliding Tile Attention.")
|
||||
sliding_tile_attention = None
|
||||
from functools import lru_cache
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
sliding_tile_attention = None
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from csrc.sliding_tile_attention.test.flex_sta_ref import get_sliding_tile_attention_mask
|
||||
@lru_cache(maxsize=32)
|
||||
def get_compiled_flex_attention(strategy, tile_size, image_size, text_length, device):
|
||||
"""
|
||||
Create and compile flex attention with a specific sliding block mask.
|
||||
This function is cached to avoid recompiling for the same parameters.
|
||||
|
||||
Args:
|
||||
strategy (tuple): A tuple (t, h, w) defining the strategy
|
||||
tile_size (tuple): A tuple (ts_t, ts_h, ts_w) defining the tile size
|
||||
image_size (tuple): A tuple (n_t, n_h, n_w) defining the image size
|
||||
text_length (int): The text length
|
||||
device (str): The device to use
|
||||
|
||||
Returns:
|
||||
function: A compiled flex attention function with the specified mask
|
||||
"""
|
||||
# Convert strategy to the required format (ceil(t*3/2), h*2, w)
|
||||
adjusted_strategy = strategy
|
||||
|
||||
# Get the sliding block attention mask
|
||||
mask = get_sliding_tile_attention_mask(
|
||||
adjusted_strategy,
|
||||
tile_size,
|
||||
image_size,
|
||||
text_length,
|
||||
device
|
||||
)
|
||||
|
||||
def flex_attn_with_mask(q, k, v, scale=None):
|
||||
return flex_attention(q, k, v, block_mask=mask, scale=scale)
|
||||
|
||||
# Compile the wrapper function
|
||||
compiled_flex_attn = torch.compile(flex_attn_with_mask)
|
||||
|
||||
return compiled_flex_attn
|
||||
|
||||
def flex_sliding_tile_attention(q_all, k_all, v_all, strategy, tile_size,
|
||||
image_size, text_length, scale=None):
|
||||
device = q_all.device
|
||||
|
||||
# Get the compiled flex attention function (cached if called with same parameters)
|
||||
compiled_flex_attn = get_compiled_flex_attention(
|
||||
strategy,
|
||||
tile_size,
|
||||
image_size,
|
||||
text_length,
|
||||
device
|
||||
)
|
||||
|
||||
|
||||
# Apply the compiled flex attention
|
||||
output = compiled_flex_attn(q_all, k_all, v_all, scale=scale)
|
||||
|
||||
|
||||
return output
|
||||
|
||||
def attention(
|
||||
q,
|
||||
@@ -148,32 +92,7 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
start_head = current_rank * head_num
|
||||
windows = [mask_strategy[head_idx + start_head] for head_idx in range(head_num)]
|
||||
|
||||
if sliding_tile_attention is not None:
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
|
||||
else:
|
||||
print("Sliding Tile Attention not available. Using Flex Sliding Tile Attention.")
|
||||
hidden_states = torch.empty_like(query)
|
||||
strategy_to_heads = {}
|
||||
for head_index in range(head_num):
|
||||
strategy = tuple(windows[head_index]) # Convert list to tuple for dict key
|
||||
if strategy not in strategy_to_heads:
|
||||
strategy_to_heads[strategy] = []
|
||||
strategy_to_heads[strategy].append(head_index)
|
||||
for strategy, heads in strategy_to_heads.items():
|
||||
# Gather all heads with this strategy
|
||||
query_heads = torch.cat([query[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
|
||||
key_heads = torch.cat([key[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
|
||||
value_heads = torch.cat([value[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
|
||||
|
||||
# Process all heads with this strategy at once
|
||||
# processed_heads = selected_attn_processor[processor_idx](query_heads, key_heads, value_heads)
|
||||
processed_heads = flex_sliding_tile_attention(query_heads, key_heads, value_heads, strategy, (6, 8, 8), (30, 48, 80), text_length)
|
||||
|
||||
# Distribute results back to the correct positions
|
||||
for i, head_idx in enumerate(heads):
|
||||
hidden_states[:, head_idx:head_idx + 1, :, :] = processed_heads[:, i:i + 1, :, :]
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2)
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
|
||||
else:
|
||||
query = torch.cat([query, encoder_query], dim=1)
|
||||
key = torch.cat([key, encoder_key], dim=1)
|
||||
@@ -202,4 +121,4 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
b, s, a, d = attn.shape
|
||||
attn = attn.reshape(b, s, -1)
|
||||
|
||||
return attn
|
||||
return attn
|
||||
|
||||
@@ -523,6 +523,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
if mask_strategy is None:
|
||||
mask_strategy = [[None] * self.heads_num for _ in range(len(self.double_blocks) + len(self.single_blocks))]
|
||||
img = x = hidden_states
|
||||
text_mask = encoder_attention_mask
|
||||
t = timestep
|
||||
|
||||
@@ -207,4 +207,4 @@ if __name__ == "__main__":
|
||||
# process for vae sequence parallel
|
||||
if args.vae_sp and not args.vae_tiling:
|
||||
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
|
||||
main(args)
|
||||
main(args)
|
||||
|
||||
+5
-2
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.0.1"
|
||||
version = "0.0.1.dev1"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.8"
|
||||
@@ -39,7 +39,10 @@ dependencies = [
|
||||
"gpustat", "watch",
|
||||
|
||||
# Kernel & Packaging
|
||||
"wheel"
|
||||
"wheel",
|
||||
|
||||
# Sliding Tile Atteniton Kernel
|
||||
"st_attn>=0.0.1"
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -37,4 +37,4 @@ torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
--output_path outputs_video/hunyuan/vae_sp/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
|
||||
--vae-sp
|
||||
--vae-sp
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config.py
|
||||
@@ -7,6 +7,13 @@ from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.2"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
@@ -44,8 +51,11 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
setup(name='st_attn',
|
||||
version="0.0.0",
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
@@ -56,4 +66,11 @@ setup(name='st_attn',
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
cmdclass={'build_ext': BuildExtension})
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
@@ -0,0 +1,252 @@
|
||||
# Copyright (c) Tile-AI Corporation.
|
||||
# Licensed under the MIT License.
|
||||
import math
|
||||
import torch
|
||||
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
import torch.nn.functional as F
|
||||
|
||||
def get_sta_mask(x, canvas_size=(32, 48, 80), tile_size=(4, 8, 8), kernel_size=(1, 1, 1), has_text=False):
|
||||
bsz, num_head, downsample_len, _ = x.shape
|
||||
device = x.device
|
||||
CT, CH, CW = canvas_size
|
||||
TT, TH, TW = tile_size
|
||||
NT, NH, NW = CT // TT, CH // TH, CW // TW
|
||||
KT, KH, KW = kernel_size
|
||||
DT, DH, DW = KT // 2, KH // 2, KW // 2
|
||||
|
||||
dense_mask = torch.full([bsz, num_head, downsample_len, downsample_len],
|
||||
False,
|
||||
dtype=torch.bool,
|
||||
device=device)
|
||||
indices = torch.arange(downsample_len, device=device)
|
||||
q_t = (indices // (NT * NH * NW)).unsqueeze(-1)
|
||||
q_h = ((indices // (NW)) % NH).unsqueeze(-1)
|
||||
q_w = ((indices // 1) % NW).unsqueeze(-1)
|
||||
|
||||
q_t = torch.clamp(q_t, DT, NT-DT-1)
|
||||
q_h = torch.clamp(q_h, DH, NH-DH-1)
|
||||
q_w = torch.clamp(q_w, DW, NW-DW-1)
|
||||
|
||||
k_indices = torch.arange(downsample_len, device=device)
|
||||
k_t = (k_indices // (NT * NH * NW)).unsqueeze(0)
|
||||
k_h = ((k_indices // (NW)) % NH).unsqueeze(0)
|
||||
k_w = ((k_indices // 1) % NW).unsqueeze(0)
|
||||
|
||||
t_dist = torch.abs(q_t - k_t)
|
||||
h_dist = torch.abs(q_h - k_h)
|
||||
w_dist = torch.abs(q_w - k_w)
|
||||
mask = (t_dist <= DT) & (h_dist <= DH) & (w_dist <= DW)
|
||||
|
||||
for b in range(bsz):
|
||||
for h in range(num_head):
|
||||
dense_mask[b, h] = mask
|
||||
|
||||
# for text mask
|
||||
if has_text:
|
||||
text_start = downsample_len - 3
|
||||
dense_mask[:, :, text_start:, :text_start] = True
|
||||
dense_mask[:, :, text_start:, text_start:] = torch.tril(
|
||||
torch.ones(3, 3, device=device, dtype=torch.bool)
|
||||
)
|
||||
return dense_mask
|
||||
|
||||
def blocksparse_flashattn(batch, heads, seq_len, dim, downsample_len, is_causal):
|
||||
block_M = 64
|
||||
block_N = 64
|
||||
num_stages = 1
|
||||
threads = 128
|
||||
scale = (1.0 / dim)**0.5 * 1.44269504 # log2(e)
|
||||
shape = [batch, heads, seq_len, dim]
|
||||
block_mask_shape = [batch, heads, downsample_len, downsample_len]
|
||||
|
||||
dtype = "float16"
|
||||
accum_dtype = "float"
|
||||
block_mask_dtype = "bool"
|
||||
|
||||
def kernel_func(block_M, block_N, num_stages, threads):
|
||||
|
||||
@T.macro
|
||||
def MMA0(
|
||||
K: T.Buffer(shape, dtype),
|
||||
Q_shared: T.Buffer([block_M, dim], dtype),
|
||||
K_shared: T.Buffer([block_N, dim], dtype),
|
||||
acc_s: T.Buffer([block_M, block_N], accum_dtype),
|
||||
k: T.int32,
|
||||
bx: T.int32,
|
||||
by: T.int32,
|
||||
bz: T.int32,
|
||||
):
|
||||
T.copy(K[bz, by, k * block_N:(k + 1) * block_N, :], K_shared)
|
||||
if is_causal:
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
acc_s[i, j] = T.if_then_else(bx * block_M + i >= k * block_N + j, 0,
|
||||
-T.infinity(acc_s.dtype))
|
||||
else:
|
||||
T.clear(acc_s)
|
||||
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow)
|
||||
|
||||
@T.macro
|
||||
def MMA1(
|
||||
V: T.Buffer(shape, dtype),
|
||||
V_shared: T.Buffer([block_M, dim], dtype),
|
||||
acc_s_cast: T.Buffer([block_M, block_N], dtype),
|
||||
acc_o: T.Buffer([block_M, dim], accum_dtype),
|
||||
k: T.int32,
|
||||
by: T.int32,
|
||||
bz: T.int32,
|
||||
):
|
||||
T.copy(V[bz, by, k * block_N:(k + 1) * block_N, :], V_shared)
|
||||
T.gemm(acc_s_cast, V_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)
|
||||
|
||||
@T.macro
|
||||
def Softmax(
|
||||
acc_s: T.Buffer([block_M, block_N], accum_dtype),
|
||||
acc_s_cast: T.Buffer([block_M, block_N], dtype),
|
||||
scores_max: T.Buffer([block_M], accum_dtype),
|
||||
scores_max_prev: T.Buffer([block_M], accum_dtype),
|
||||
scores_scale: T.Buffer([block_M], accum_dtype),
|
||||
scores_sum: T.Buffer([block_M], accum_dtype),
|
||||
logsum: T.Buffer([block_M], accum_dtype),
|
||||
):
|
||||
T.copy(scores_max, scores_max_prev)
|
||||
T.fill(scores_max, -T.infinity(accum_dtype))
|
||||
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
|
||||
# To do causal softmax, we need to set the scores_max to 0 if it is -inf
|
||||
# This process is called Check_inf in FlashAttention3 code, and it only need to be done
|
||||
# in the first ceil_div(kBlockM, kBlockN) steps.
|
||||
# for i in T.Parallel(block_M):
|
||||
# scores_max[i] = T.if_then_else(scores_max[i] == -T.infinity(accum_dtype), 0, scores_max[i])
|
||||
for i in T.Parallel(block_M):
|
||||
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
# Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
# max * log_2(e)) This allows the compiler to use the ffma
|
||||
# instruction instead of fadd and fmul separately.
|
||||
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
|
||||
T.reduce_sum(acc_s, scores_sum, dim=1)
|
||||
for i in T.Parallel(block_M):
|
||||
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
|
||||
T.copy(acc_s, acc_s_cast)
|
||||
|
||||
@T.macro
|
||||
def Rescale(
|
||||
acc_o: T.Buffer([block_M, dim], accum_dtype),
|
||||
scores_scale: T.Buffer([block_M], accum_dtype),
|
||||
):
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] *= scores_scale[i]
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
Q: T.Buffer(shape, dtype),
|
||||
K: T.Buffer(shape, dtype),
|
||||
V: T.Buffer(shape, dtype),
|
||||
BlockSparseMask: T.Buffer(block_mask_shape, block_mask_dtype),
|
||||
Output: T.Buffer(shape, dtype),
|
||||
):
|
||||
with T.Kernel(
|
||||
T.ceildiv(seq_len, block_M), heads, batch, threads=threads) as (bx, by, bz):
|
||||
Q_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
K_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
V_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
O_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
acc_s = T.alloc_fragment([block_M, block_N], accum_dtype)
|
||||
acc_s_cast = T.alloc_fragment([block_M, block_N], dtype)
|
||||
acc_o = T.alloc_fragment([block_M, dim], accum_dtype)
|
||||
scores_max = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_max_prev = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_scale = T.alloc_fragment([block_M], accum_dtype)
|
||||
scores_sum = T.alloc_fragment([block_M], accum_dtype)
|
||||
logsum = T.alloc_fragment([block_M], accum_dtype)
|
||||
block_mask = T.alloc_local([downsample_len], block_mask_dtype)
|
||||
|
||||
T.copy(Q[bz, by, bx * block_M:(bx + 1) * block_M, :], Q_shared)
|
||||
T.fill(acc_o, 0)
|
||||
T.fill(logsum, 0)
|
||||
T.fill(scores_max, -T.infinity(accum_dtype))
|
||||
|
||||
for vj in T.serial(downsample_len):
|
||||
block_mask[vj] = BlockSparseMask[bz, by, bx, vj]
|
||||
|
||||
loop_range = (
|
||||
T.min(T.ceildiv(seq_len, block_N), T.ceildiv(
|
||||
(bx + 1) * block_M, block_N)) if is_causal else T.ceildiv(seq_len, block_N))
|
||||
|
||||
for k in T.Pipelined(loop_range, num_stages=num_stages):
|
||||
if block_mask[k] != 0:
|
||||
MMA0(K, Q_shared, K_shared, acc_s, k, bx, by, bz)
|
||||
Softmax(acc_s, acc_s_cast, scores_max, scores_max_prev, scores_scale,
|
||||
scores_sum, logsum)
|
||||
Rescale(acc_o, scores_scale)
|
||||
MMA1(V, V_shared, acc_s_cast, acc_o, k, by, bz)
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] /= logsum[i]
|
||||
T.copy(acc_o, O_shared)
|
||||
T.copy(O_shared, Output[bz, by, bx * block_M:(bx + 1) * block_M, :])
|
||||
|
||||
return main
|
||||
|
||||
return kernel_func(block_M, block_N, num_stages, threads)
|
||||
|
||||
|
||||
def test_sta_attention():
|
||||
# Config
|
||||
BATCH, N_HEADS, SEQ_LEN, D_HEAD = 1, 24, 2048, 128
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Create inputs
|
||||
q = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
k = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
v = torch.randn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, device='cuda', dtype=torch.float16)
|
||||
|
||||
sm_scale = 1.0 / (D_HEAD**0.5)
|
||||
|
||||
# Create sparse mask (downsampled to block level)
|
||||
canvas_size = (32, 48, 80)
|
||||
tile_size = (4, 8, 8)
|
||||
delta_size = (2, 2, 2)
|
||||
BLOCK = tile_size[0] * tile_size[1] * tile_size[2]
|
||||
downsample_factor = BLOCK
|
||||
downsample_len = math.ceil(SEQ_LEN / downsample_factor)
|
||||
x_ds = torch.randn([BATCH, N_HEADS, downsample_len, downsample_len],
|
||||
device='cuda',
|
||||
dtype=torch.bfloat16)
|
||||
block_mask = get_sta_mask(x_ds, canvas_size, tile_size, delta_size, has_text=True)
|
||||
# print mask density
|
||||
print("mask density", block_mask.sum() / block_mask.numel())
|
||||
print("block_mask", block_mask)
|
||||
|
||||
# Run Triton kernel
|
||||
program = blocksparse_flashattn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, downsample_len, is_causal=True)
|
||||
kernel = tilelang.compile(program, out_idx=[4])
|
||||
|
||||
cuda_source = kernel.get_kernel_source()
|
||||
print("Generated CUDA kernel:\n", cuda_source)
|
||||
|
||||
tilelang_output = kernel(q, k, v, block_mask)
|
||||
|
||||
if True:
|
||||
# Compute reference
|
||||
# Expand block mask to full attention matrix
|
||||
full_mask = torch.kron(block_mask.float(), torch.ones(BLOCK, BLOCK, device='cuda'))
|
||||
full_mask = full_mask[..., :SEQ_LEN, :SEQ_LEN].bool()
|
||||
full_mask = full_mask & torch.tril(torch.ones_like(full_mask)) # Apply causal
|
||||
|
||||
# PyTorch reference implementation
|
||||
attn = torch.einsum('bhsd,bhtd->bhst', q, k) * sm_scale
|
||||
attn = attn.masked_fill(~full_mask, float('-inf'))
|
||||
attn = F.softmax(attn, dim=-1)
|
||||
ref_output = torch.einsum('bhst,bhtd->bhsd', attn, v)
|
||||
|
||||
print("ref_output", ref_output)
|
||||
print("tilelang_output", tilelang_output)
|
||||
|
||||
# Verify accuracy
|
||||
torch.testing.assert_close(tilelang_output, ref_output, atol=1e-2, rtol=1e-2)
|
||||
print("Pass topk sparse attention test with qlen == klen")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_sta_attention()
|
||||
Reference in New Issue
Block a user