Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6b119a9f76 | ||
|
|
15fb1d9a57 | ||
|
|
66060dd22f | ||
|
|
9fe8704537 | ||
|
|
31bbdcd2a4 |
@@ -1,70 +0,0 @@
|
||||
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/
|
||||
@@ -1,221 +0,0 @@
|
||||
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,11 +23,8 @@ 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 = sta_kernel/thunderkitten/tk
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
import json
|
||||
|
||||
# [1, 24, 46336, 128]
|
||||
|
||||
# q = torch.randn(1, 24, 46336, 128)
|
||||
|
||||
# for startegy (2, 6, 1)
|
||||
# theortical speed up is 10.0
|
||||
# actual speed up is 7.6770869514877855
|
||||
# for startegy (1, 6, 10)
|
||||
# theortical speed up is 2.0
|
||||
# actual speed up is 1.81064261469606
|
||||
# for startegy (2, 3, 3)
|
||||
# theortical speed up is 6.666666666666667
|
||||
# actual speed up is 5.410431320071153
|
||||
# for startegy (2, 6, 10)
|
||||
# theortical speed up is 1.0
|
||||
# actual speed up is 0.9610114048271594
|
||||
# for startegy (2, 1, 10)
|
||||
# theortical speed up is 6.0
|
||||
# actual speed up is 5.1464607264824584
|
||||
# for startegy (2, 3, 5)
|
||||
# theortical speed up is 4.0
|
||||
# actual speed up is 3.5396014498556014
|
||||
|
||||
|
||||
initial_path = "/workspace/codefolder/FastVideo/assets/mask_strategy_hunyuan.json"
|
||||
|
||||
actuall_speed_up = 0
|
||||
|
||||
with open(initial_path, "r") as f:
|
||||
data = json.load(f)
|
||||
for key in data.keys():
|
||||
t, h, w = data[key]
|
||||
data[key] = [min(2,t), h, w]
|
||||
|
||||
t, h, w = data[key]
|
||||
|
||||
if t == 2 and h == 6 and w == 1:
|
||||
actuall_speed_up += 7.5
|
||||
elif t == 1 and h == 6 and w == 10:
|
||||
actuall_speed_up += 1.77
|
||||
elif t == 2 and h == 3 and w == 3:
|
||||
actuall_speed_up += 5.32
|
||||
elif t == 2 and h == 6 and w == 10:
|
||||
actuall_speed_up += 0.94
|
||||
elif t == 2 and h == 1 and w == 10:
|
||||
actuall_speed_up += 5.07
|
||||
elif t == 2 and h == 3 and w == 5:
|
||||
actuall_speed_up += 3.47
|
||||
|
||||
print(actuall_speed_up/len(data.keys()))
|
||||
|
||||
save_path = "/workspace/codefolder/FastVideo/assets/test_mask_strategy_hunyuan.json"
|
||||
|
||||
|
||||
# print unique values for new data
|
||||
unique_values = set()
|
||||
for key in data.keys():
|
||||
unique_values.add(tuple(data[key]))
|
||||
print(unique_values)
|
||||
|
||||
with open(save_path, "w") as f:
|
||||
json.dump(data, f, indent=4)
|
||||
@@ -7,13 +7,6 @@ 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',
|
||||
@@ -51,11 +44,8 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
setup(name='st_attn',
|
||||
version="0.0.0",
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
@@ -66,11 +56,4 @@ setup(name=PACKAGE_NAME,
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
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"])
|
||||
cmdclass={'build_ext': BuildExtension})
|
||||
@@ -0,0 +1,238 @@
|
||||
# def mask(b, h, q_idx, kv_idx):
|
||||
# return kv_idx < text_length + img_seq_len
|
||||
from torch.nn.attention.flex_attention import create_block_mask, or_masks
|
||||
from torch import IntTensor, BoolTensor
|
||||
import torch
|
||||
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
|
||||
import math
|
||||
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
|
||||
torch._inductor.config.realize_opcount_threshold = 100
|
||||
def generate_sba_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 natten_mask_mod_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
|
||||
|
||||
natten_mask_mod_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return natten_mask_mod_3d
|
||||
|
||||
|
||||
|
||||
def generate_baseline_3d_window_mask(
|
||||
canvas_twh,
|
||||
kernel_twh,
|
||||
img_seq_len,
|
||||
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
|
||||
def get_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
|
||||
t = idx // (canvas_h * canvas_w)
|
||||
x = (idx % (canvas_h * canvas_w)) // canvas_w
|
||||
y = idx % canvas_w
|
||||
return t, x, y
|
||||
|
||||
def natten_mask_mod_3d(
|
||||
b: IntTensor,
|
||||
h: IntTensor,
|
||||
q_idx: IntTensor,
|
||||
kv_idx: IntTensor,
|
||||
) -> BoolTensor:
|
||||
q_t, q_x, q_y = get_t_x_y(q_idx)
|
||||
kv_t, kv_x, kv_y = get_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.clamp(kernel_t // 2, (canvas_t - 1) - kernel_t // 2)
|
||||
kernel_center_x = q_x.clamp(kernel_h // 2, (canvas_h - 1) - kernel_h // 2)
|
||||
kernel_center_y = q_y.clamp(kernel_w // 2, (canvas_w - 1) - kernel_w // 2)
|
||||
time_mask = (kernel_center_t - kv_t).abs() <= kernel_t // 2
|
||||
hori_mask = (kernel_center_x - kv_x).abs() <= kernel_h // 2
|
||||
vert_mask = (kernel_center_y - kv_y).abs() <= kernel_w // 2
|
||||
no_pad_mask = kv_idx < text_length + img_seq_len
|
||||
return time_mask & hori_mask & vert_mask & no_pad_mask
|
||||
|
||||
natten_mask_mod_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return natten_mask_mod_3d
|
||||
|
||||
|
||||
def generate_tiled_3d_window_mask(
|
||||
canvas_twh,
|
||||
kernel_twh,
|
||||
img_seq_len,
|
||||
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, tile_h, tile_w = 4, 8, 8
|
||||
n_tile_t, n_tile_h, n_tile_w = canvas_t // tile_t, canvas_h // tile_h, canvas_w // tile_w
|
||||
|
||||
def get_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
|
||||
tile_id = idx // (tile_t * tile_h * tile_w)
|
||||
t_t, t_x, t_y = tile_id // (n_tile_h * n_tile_w), (tile_id % (n_tile_h * n_tile_w)) // n_tile_w, tile_id % n_tile_w
|
||||
t_offset = idx % (tile_t * tile_h * tile_w)
|
||||
i_t, i_x, i_y = t_offset // (tile_h * tile_w), (t_offset % (tile_h * tile_w)) // tile_w, t_offset % tile_w
|
||||
return t_t * tile_t + i_t, t_x * tile_h + i_x, t_y * tile_w + i_y
|
||||
|
||||
def natten_mask_mod_3d(
|
||||
b: IntTensor,
|
||||
h: IntTensor,
|
||||
q_idx: IntTensor,
|
||||
kv_idx: IntTensor,
|
||||
) -> BoolTensor:
|
||||
q_t, q_x, q_y = get_t_x_y(q_idx)
|
||||
kv_t, kv_x, kv_y = get_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.clamp(kernel_t // 2, (canvas_t - 1) - kernel_t // 2)
|
||||
kernel_center_x = q_x.clamp(kernel_h // 2, (canvas_h - 1) - kernel_h // 2)
|
||||
kernel_center_y = q_y.clamp(kernel_w // 2, (canvas_w - 1) - kernel_w // 2)
|
||||
time_mask = (kernel_center_t - kv_t).abs() <= kernel_t // 2
|
||||
hori_mask = (kernel_center_x - kv_x).abs() <= kernel_h // 2
|
||||
vert_mask = (kernel_center_y - kv_y).abs() <= kernel_w // 2
|
||||
no_pad_mask = kv_idx < text_length + img_seq_len
|
||||
return time_mask & hori_mask & vert_mask & no_pad_mask
|
||||
|
||||
natten_mask_mod_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return natten_mask_mod_3d
|
||||
|
||||
|
||||
def generate_text_mask(img_seq_len, text_length):
|
||||
def text_mask(b, h, q_idx, kv_idx):
|
||||
mask1 = kv_idx < text_length + img_seq_len
|
||||
mask2 = kv_idx >= img_seq_len
|
||||
return mask1 & mask2
|
||||
return text_mask
|
||||
|
||||
|
||||
def get_sliding_block_attention_mask(kernel_size, tile_size, img_size, text_length, device):
|
||||
img_seq_len = img_size[0] * img_size[1] * img_size[2]
|
||||
image_mask = generate_sba_mask(img_size, kernel_size, tile_size, text_length)
|
||||
mask = create_block_mask(image_mask, B=None, H=None, Q_LEN=img_seq_len+256 , KV_LEN=img_seq_len+256, device=device, _compile=True, BLOCK_SIZE=128)
|
||||
return mask
|
||||
|
||||
def get_baseline_sliding_window_mask(kernel_size, img_seq_len, text_length, device):
|
||||
image_mask = generate_baseline_3d_window_mask( (32, 48, 80), kernel_size, img_seq_len, text_length)
|
||||
text_mask = generate_text_mask(img_seq_len=img_seq_len, text_length=text_length)
|
||||
mask = or_masks(image_mask, text_mask)
|
||||
mask = create_block_mask(mask, B=None, H=None, Q_LEN=img_seq_len + 256 , KV_LEN=img_seq_len + 256, device=device, _compile=True)
|
||||
return mask
|
||||
|
||||
def get_tiled_sliding_window_mask(kernel_size, img_seq_len, text_length, device):
|
||||
image_mask = generate_tiled_3d_window_mask((32, 48, 80), kernel_size, img_seq_len, text_length)
|
||||
text_mask = generate_text_mask(img_seq_len=img_seq_len, text_length=text_length)
|
||||
mask = or_masks(image_mask, text_mask)
|
||||
mask = create_block_mask(mask, B=None, H=None, Q_LEN=img_seq_len + 256 , KV_LEN=img_seq_len + 256, device=device, _compile=True)
|
||||
return mask
|
||||
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True):
|
||||
seq_length = q_all.shape[2]
|
||||
# if has_text:
|
||||
# assert q_all.shape[
|
||||
# 2] == 115456, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
# assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
# target_size = math.ceil(seq_length / 384) * 384
|
||||
# pad_size = target_size - seq_length
|
||||
# if pad_size > 0:
|
||||
# q_all = torch.cat([q_all, q_all[:, :, -pad_size:]], dim=2)
|
||||
# k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
# v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
# else:
|
||||
# assert q_all.shape[2] == 82944
|
||||
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
for batch in range(q_all.shape[0]):
|
||||
q_head, k_head, v_head, o_head = (q_all[batch:batch + 1, head_index:head_index + 1],
|
||||
k_all[batch:batch + 1,
|
||||
head_index:head_index + 1], v_all[batch:batch + 1,
|
||||
head_index:head_index + 1],
|
||||
hidden_states[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
|
||||
if __name__ == "__main__":
|
||||
# benchmark speed
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
flex_attention = torch.compile(flex_attention)
|
||||
import time
|
||||
device = torch.device("cuda")
|
||||
kernel_size_ls = [(8, 6, 10), (4, 6, 10), (4, 6, 5), (4, 3, 5), (3, 3, 3)]
|
||||
tile_size = (4, 8, 8)
|
||||
|
||||
# random input
|
||||
q = torch.randn(1, 24, 123136, 128, device=device, dtype=torch.bfloat16)
|
||||
k = torch.randn(1, 24, 123136, 128, device=device, dtype=torch.bfloat16)
|
||||
v = torch.randn(1, 24, 123136, 128, device=device, dtype=torch.bfloat16)
|
||||
|
||||
for kernel_size in kernel_size_ls:
|
||||
mask = get_sliding_block_attention_mask(kernel_size, tile_size, (32, 48, 80), 39, device)
|
||||
|
||||
@@ -160,7 +160,7 @@ def distill_one_step(
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
teacher_output = cond_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,29 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
t, l, h = map(int, key.split('_'))
|
||||
result[t][l][h] = value
|
||||
return result
|
||||
|
||||
|
||||
|
||||
selected_strategies = [(2, 6, 1), (1, 6, 10), (2, 3, 3), (2, 6, 10), (2, 1, 10), (2, 3, 5)]
|
||||
text_length = prompt_mask.sum()
|
||||
# selected_attn_processor = []
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
from functools import partial
|
||||
from csrc.sliding_tile_attention.test.sba import get_sliding_block_attention_mask
|
||||
|
||||
torch._dynamo.config.cache_size_limit = 128
|
||||
# for ms in selected_strategies:
|
||||
# mask = get_sliding_block_attention_mask(ms, (6, 8, 8), (12, 48, 80), text_length, self.transformer.device)
|
||||
# attn_processor = torch.compile(partial(flex_attention, block_mask=mask))
|
||||
# selected_attn_processor.append(attn_processor)
|
||||
|
||||
# warmup for processors
|
||||
# warmup_q = torch.randn(1, 24, 46336, 128)
|
||||
# warmup_k = torch.randn(1, 24, 46336, 128)
|
||||
# warmup_v = torch.randn(1, 24, 46336, 128)
|
||||
|
||||
# for processor in selected_attn_processor:
|
||||
# processor(warmup_q, warmup_k, warmup_v)
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
# if is_progress_bar:
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -849,6 +871,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
mask_strategy=mask_strategy[i],
|
||||
guidance=guidance_expand,
|
||||
return_dict=False,
|
||||
selected_attn_processor=[],
|
||||
)[0]
|
||||
|
||||
# perform guidance
|
||||
@@ -904,14 +927,20 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
latents = (latents / self.vae.config.scaling_factor + self.vae.config.shift_factor)
|
||||
else:
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
|
||||
self.transformer = self.transformer.to('cpu')
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
if enable_tiling:
|
||||
print("tiling is enabled")
|
||||
self.vae.enable_tiling()
|
||||
if enable_vae_sp:
|
||||
self.vae.enable_parallel()
|
||||
|
||||
image = self.vae.decode(latents, return_dict=False, generator=generator)[0]
|
||||
|
||||
self.transformer = self.transformer.to(device)
|
||||
|
||||
if expand_temporal_dim or image.shape[2] == 1:
|
||||
image = image.squeeze(2)
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ from fastvideo.models.hunyuan.text_encoder import TextEncoder
|
||||
from fastvideo.models.hunyuan.utils.data_utils import align_to
|
||||
from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
from fastvideo.models.hunyuan.modules.fp8 import convert_fp8_linear
|
||||
|
||||
class Inference(object):
|
||||
|
||||
@@ -76,7 +76,7 @@ class Inference(object):
|
||||
|
||||
# =========================== Build main model ===========================
|
||||
logger.info("Building model...")
|
||||
factor_kwargs = {"device": device, "dtype": PRECISION_TO_TYPE[args.precision]}
|
||||
factor_kwargs = {"device": 'cpu', "dtype": PRECISION_TO_TYPE[args.precision]}
|
||||
in_channels = args.latent_channels
|
||||
out_channels = args.latent_channels
|
||||
|
||||
@@ -86,6 +86,11 @@ class Inference(object):
|
||||
out_channels=out_channels,
|
||||
factor_kwargs=factor_kwargs,
|
||||
)
|
||||
|
||||
if args.use_fp8:
|
||||
print("loading fp8 model")
|
||||
convert_fp8_linear(model, args.dit_weight, original_dtype=PRECISION_TO_TYPE[args.precision])
|
||||
|
||||
model = model.to(device)
|
||||
model = Inference.load_state_dict(args, model, pretrained_model_path)
|
||||
if args.enable_torch_compile:
|
||||
|
||||
@@ -1,18 +1,78 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
import time
|
||||
# try:
|
||||
# from st_attn import sliding_tile_attention
|
||||
# except ImportError:
|
||||
# print("Could not load Sliding Tile Attention.")
|
||||
# sliding_tile_attention = None
|
||||
|
||||
try:
|
||||
from st_attn import sliding_tile_attention
|
||||
except ImportError:
|
||||
print("Could not load Sliding Tile Attention.")
|
||||
sliding_tile_attention = None
|
||||
|
||||
import tk_4090_cuda as tk
|
||||
from functools import lru_cache, partial
|
||||
from csrc.sliding_tile_attention.test.sba import get_sliding_block_attention_mask
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
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
|
||||
|
||||
|
||||
@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_block_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 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,
|
||||
k,
|
||||
@@ -35,10 +95,10 @@ def attention(
|
||||
|
||||
|
||||
def tile(x, sp_size):
|
||||
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
|
||||
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=12 // sp_size, h=48, w=80)
|
||||
return rearrange(x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
n_t=5,
|
||||
n_t=2,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
@@ -49,16 +109,16 @@ def tile(x, sp_size):
|
||||
def untile(x, sp_size):
|
||||
x = rearrange(x,
|
||||
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
|
||||
n_t=5,
|
||||
n_t=2,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
return rearrange(x, "b (t sp h w) head d -> b (sp t h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
|
||||
return rearrange(x, "b (t sp h w) head d -> b (sp t h w) head d", sp=sp_size, t=12 // sp_size, h=48, w=80)
|
||||
|
||||
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=None):
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=None, selected_attn_processor=None):
|
||||
query, encoder_query = q
|
||||
key, encoder_key = k
|
||||
value, encoder_value = v
|
||||
@@ -91,21 +151,110 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
current_rank = nccl_info.rank_within_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [mask_strategy[head_idx + start_head] for head_idx in range(head_num)]
|
||||
|
||||
# Initialize the output tensor
|
||||
hidden_states = torch.empty_like(query)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_start = time.time()
|
||||
|
||||
# Group heads by their mask strategy
|
||||
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)
|
||||
|
||||
# Create a mapping from strategy to processor index
|
||||
strategy_to_processor = {
|
||||
(2, 6, 1): 0,
|
||||
(1, 6, 10): 1,
|
||||
(2, 3, 3): 2,
|
||||
(2, 6, 10): 3,
|
||||
(2, 1, 10): 4,
|
||||
(2, 3, 5): 5
|
||||
}
|
||||
|
||||
# Process heads with the same strategy together
|
||||
for strategy, heads in strategy_to_heads.items():
|
||||
torch.cuda.synchronize()
|
||||
strategy_start = time.time()
|
||||
|
||||
# Get processor index for this strategy
|
||||
processor_idx = strategy_to_processor.get(strategy, 3) # Default to processor 3
|
||||
|
||||
# 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 = sliding_tile_attention(query_heads, key_heads, value_heads, strategy, (6, 8, 8), (12, 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, :, :]
|
||||
|
||||
torch.cuda.synchronize()
|
||||
strategy_end = time.time()
|
||||
print(f"Time taken for strategy {strategy} with {len(heads)} heads: {strategy_end - strategy_start}")
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_end = time.time()
|
||||
print(f"Time taken for optimized sliding tile attention: {time_end - time_start}")
|
||||
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
|
||||
hidden_states = hidden_states.transpose(1, 2)
|
||||
else:
|
||||
import copy
|
||||
init_query = copy.deepcopy(query)
|
||||
init_key = copy.deepcopy(key)
|
||||
init_value = copy.deepcopy(value)
|
||||
|
||||
query = torch.cat([query, encoder_query], dim=1)
|
||||
key = torch.cat([key, encoder_key], dim=1)
|
||||
value = torch.cat([value, encoder_value], dim=1)
|
||||
# B, S, 3, H, D
|
||||
result = torch.empty_like(query)
|
||||
torch.cuda.synchronize()
|
||||
start_time = time.time()
|
||||
output = tk.attention_fwd_4090(query, key, value, result, text_length)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
end_time = time.time()
|
||||
print(f"Time taken for tk attention: {end_time - start_time}")
|
||||
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
|
||||
|
||||
|
||||
|
||||
torch.cuda.synchronize()
|
||||
flash_start_time = time.time()
|
||||
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
torch.cuda.synchronize()
|
||||
flash_end_time = time.time()
|
||||
print(f"Time taken for flash attention: {flash_end_time - flash_start_time}")
|
||||
|
||||
|
||||
query = torch.cat([tile(init_query, nccl_info.sp_size), encoder_query], dim=1).transpose(1, 2)
|
||||
key = torch.cat([tile(init_key, nccl_info.sp_size), encoder_key], dim=1).transpose(1, 2)
|
||||
value = torch.cat([tile(init_value, nccl_info.sp_size), encoder_value], dim=1).transpose(1, 2)
|
||||
torch.cuda.synchronize()
|
||||
flex_start_time = time.time()
|
||||
mask = get_sliding_block_attention_mask((2,6,10), (6, 8, 8), (12, 48, 80), text_length, 'cuda')
|
||||
flex_hidden_states = flex_attention(query, key, value, block_mask=mask).transpose(1, 2)
|
||||
torch.cuda.synchronize()
|
||||
flex_end_time = time.time()
|
||||
print(f"Time taken for flex attention: {flex_end_time - flex_start_time}")
|
||||
|
||||
|
||||
# hidden_states = flex_hidden_states
|
||||
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes((sequence_length, encoder_sequence_length),
|
||||
dim=1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
hidden_states = untile(hidden_states, nccl_info.sp_size)
|
||||
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
|
||||
_bits = torch.tensor(bits)
|
||||
_mantissa_bit = torch.tensor(mantissa_bit)
|
||||
_sign_bits = torch.tensor(sign_bits)
|
||||
M = torch.clamp(torch.round(_mantissa_bit), 1, _bits - _sign_bits)
|
||||
E = _bits - _sign_bits - M
|
||||
bias = 2 ** (E - 1) - 1
|
||||
mantissa = 1
|
||||
for i in range(mantissa_bit - 1):
|
||||
mantissa += 1 / (2 ** (i+1))
|
||||
maxval = mantissa * 2 ** (2**E - 1 - bias)
|
||||
return maxval
|
||||
|
||||
def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
|
||||
"""
|
||||
Default is E4M3.
|
||||
"""
|
||||
bits = torch.tensor(bits)
|
||||
mantissa_bit = torch.tensor(mantissa_bit)
|
||||
sign_bits = torch.tensor(sign_bits)
|
||||
M = torch.clamp(torch.round(mantissa_bit), 1, bits - sign_bits)
|
||||
E = bits - sign_bits - M
|
||||
bias = 2 ** (E - 1) - 1
|
||||
mantissa = 1
|
||||
for i in range(mantissa_bit - 1):
|
||||
mantissa += 1 / (2 ** (i+1))
|
||||
maxval = mantissa * 2 ** (2**E - 1 - bias)
|
||||
minval = - maxval
|
||||
minval = - maxval if sign_bits == 1 else torch.zeros_like(maxval)
|
||||
input_clamp = torch.min(torch.max(x, minval), maxval)
|
||||
log_scales = torch.clamp((torch.floor(torch.log2(torch.abs(input_clamp)) + bias)).detach(), 1.0)
|
||||
log_scales = 2.0 ** (log_scales - M - bias.type(x.dtype))
|
||||
# dequant
|
||||
qdq_out = torch.round(input_clamp / log_scales) * log_scales
|
||||
return qdq_out, log_scales
|
||||
|
||||
def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
|
||||
for i in range(len(x.shape) - 1):
|
||||
scale = scale.unsqueeze(-1)
|
||||
new_x = x / scale
|
||||
quant_dequant_x, log_scales = quantize_to_fp8(new_x, bits=bits, mantissa_bit=mantissa_bit, sign_bits=sign_bits)
|
||||
return quant_dequant_x, scale, log_scales
|
||||
|
||||
def fp8_activation_dequant(qdq_out, scale, dtype):
|
||||
qdq_out = qdq_out.type(dtype)
|
||||
quant_dequant_x = qdq_out * scale.to(dtype)
|
||||
return quant_dequant_x
|
||||
|
||||
def fp8_linear_forward(cls, original_dtype, input):
|
||||
weight_dtype = cls.weight.dtype
|
||||
#####
|
||||
if cls.weight.dtype != torch.float8_e4m3fn:
|
||||
maxval = get_fp_maxval()
|
||||
scale = torch.max(torch.abs(cls.weight.flatten())) / maxval
|
||||
linear_weight, scale, log_scales = fp8_tensor_quant(cls.weight, scale)
|
||||
linear_weight = linear_weight.to(torch.float8_e4m3fn)
|
||||
weight_dtype = linear_weight.dtype
|
||||
else:
|
||||
scale = cls.fp8_scale.to(cls.weight.device)
|
||||
linear_weight = cls.weight
|
||||
#####
|
||||
|
||||
if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0:
|
||||
if True or len(input.shape) == 3:
|
||||
cls_dequant = fp8_activation_dequant(linear_weight, scale, original_dtype)
|
||||
if cls.bias != None:
|
||||
output = F.linear(input, cls_dequant, cls.bias)
|
||||
else:
|
||||
output = F.linear(input, cls_dequant)
|
||||
return output
|
||||
else:
|
||||
return cls.original_forward(input.to(original_dtype))
|
||||
else:
|
||||
return cls.original_forward(input)
|
||||
|
||||
def convert_fp8_linear(module, dit_weight_path, original_dtype, params_to_keep={}):
|
||||
setattr(module, "fp8_matmul_enabled", True)
|
||||
|
||||
# loading fp8 mapping file
|
||||
fp8_map_path = dit_weight_path.replace('.pt', '_map.pt')
|
||||
if os.path.exists(fp8_map_path):
|
||||
fp8_map = torch.load(fp8_map_path, map_location=lambda storage, loc: storage)
|
||||
else:
|
||||
raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.")
|
||||
|
||||
fp8_layers = []
|
||||
for key, layer in module.named_modules():
|
||||
if isinstance(layer, nn.Linear) and ('double_blocks' in key or 'single_blocks' in key):
|
||||
fp8_layers.append(key)
|
||||
original_forward = layer.forward
|
||||
layer.weight = torch.nn.Parameter(layer.weight.to(torch.float8_e4m3fn))
|
||||
setattr(layer, "fp8_scale", fp8_map[key].to(dtype=original_dtype))
|
||||
setattr(layer, "original_forward", original_forward)
|
||||
setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input))
|
||||
|
||||
@@ -110,6 +110,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
freqs_cis: tuple = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
mask_strategy=None,
|
||||
selected_attn_processor=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
(
|
||||
img_mod1_shift,
|
||||
@@ -171,6 +172,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
mask_strategy=mask_strategy,
|
||||
selected_attn_processor=selected_attn_processor,
|
||||
)
|
||||
|
||||
# attention computation end
|
||||
@@ -260,6 +262,7 @@ class MMSingleStreamBlock(nn.Module):
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
mask_strategy=None,
|
||||
selected_attn_processor=None,
|
||||
) -> torch.Tensor:
|
||||
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
||||
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
|
||||
@@ -296,6 +299,7 @@ class MMSingleStreamBlock(nn.Module):
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
mask_strategy=mask_strategy,
|
||||
selected_attn_processor=selected_attn_processor,
|
||||
)
|
||||
|
||||
# attention computation end
|
||||
@@ -520,11 +524,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
guidance=None,
|
||||
selected_attn_processor=None,
|
||||
) -> 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
|
||||
@@ -568,7 +571,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index], selected_attn_processor]
|
||||
img, txt = block(*double_block_args)
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
x = torch.cat((img, txt), 1)
|
||||
@@ -583,6 +586,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
mask_strategy[index + len(self.double_blocks)],
|
||||
selected_attn_processor,
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
if output_features and _ % output_features_stride == 0:
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
import torch
|
||||
import time
|
||||
from einops import rearrange
|
||||
|
||||
def tile(x, sp_size):
|
||||
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=12 // sp_size, h=48, w=80)
|
||||
return rearrange(x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
n_t=2,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
|
||||
q = torch.load("query.pt")
|
||||
k = torch.load("key.pt")
|
||||
v = torch.load("value.pt")
|
||||
text_mask = torch.load("text_mask.pt")
|
||||
|
||||
query, encoder_query = q.split_with_sizes((q.shape[1] - 256, 256), dim=1)
|
||||
key, encoder_key = k.split_with_sizes((k.shape[1] - 256, 256), dim=1)
|
||||
value, encoder_value = v.split_with_sizes((v.shape[1] - 256, 256), dim=1)
|
||||
|
||||
q = torch.cat([tile(query, 1), encoder_query], dim=1).transpose(1, 2)
|
||||
k = torch.cat([tile(key, 1), encoder_key], dim=1).transpose(1, 2)
|
||||
v = torch.cat([tile(value, 1), encoder_value], dim=1).transpose(1, 2)
|
||||
|
||||
selected_strategies = [(2, 6, 1), (1, 6, 10), (2, 3, 3), (2, 6, 10), (2, 1, 10), (2, 3, 5)]
|
||||
text_length = text_mask.sum()
|
||||
selected_attn_processor = []
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
from functools import partial
|
||||
from csrc.sliding_tile_attention.test.sba import get_sliding_block_attention_mask
|
||||
|
||||
for ms in selected_strategies:
|
||||
mask = get_sliding_block_attention_mask(ms, (6, 8, 8), (12, 48, 80), text_length, "cuda")
|
||||
attn_processor = torch.compile(partial(flex_attention, block_mask=mask), mode="max-autotune-no-cudagraphs")
|
||||
selected_attn_processor.append(attn_processor)
|
||||
|
||||
warmup_time = 1
|
||||
|
||||
print(q.shape)
|
||||
|
||||
for processor in selected_attn_processor:
|
||||
processor(q, k, v)
|
||||
|
||||
for processor in selected_attn_processor:
|
||||
torch.cuda.synchronize()
|
||||
start_time = time.time()
|
||||
processor(q, k, v)
|
||||
torch.cuda.synchronize()
|
||||
actuall_time = time.time()-start_time
|
||||
strategy = selected_strategies[selected_attn_processor.index(processor)]
|
||||
t, h, w = strategy
|
||||
print(f"for startegy {selected_strategies[selected_attn_processor.index(processor)]}")
|
||||
print(f"theortical speed up is {(2*6*10)/(t*h*w)}")
|
||||
print(f"actual speed up is {0.15808820724487305/actuall_time}")
|
||||
|
||||
@@ -1,17 +1,178 @@
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import json
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.models.hunyuan.modules.modulate_layers import modulate
|
||||
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
|
||||
from typing import Any, Dict, Optional, Union
|
||||
def teacache_forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
mask_strategy=None,
|
||||
output_features=False,
|
||||
output_features_stride=8,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
guidance=None,
|
||||
selected_attn_processor=None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
|
||||
img = x = hidden_states
|
||||
text_mask = encoder_attention_mask
|
||||
t = timestep
|
||||
txt = encoder_hidden_states[:, 1:]
|
||||
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
|
||||
_, _, ot, oh, ow = x.shape # codespell:ignore
|
||||
tt, th, tw = (
|
||||
ot // self.patch_size[0], # codespell:ignore
|
||||
oh // self.patch_size[1], # codespell:ignore
|
||||
ow // self.patch_size[2], # codespell:ignore
|
||||
)
|
||||
original_tt = nccl_info.sp_size * tt
|
||||
freqs_cos, freqs_sin = self.get_rotary_pos_embed((original_tt, th, tw))
|
||||
# Prepare modulation vectors.
|
||||
vec = self.time_in(t)
|
||||
|
||||
# text modulation
|
||||
vec = vec + self.vector_in(text_states_2)
|
||||
|
||||
# guidance modulation
|
||||
if self.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
|
||||
# our timestep_embedding is merged into guidance_in(TimestepEmbedder)
|
||||
vec = vec + self.guidance_in(guidance)
|
||||
|
||||
# Embed image and text.
|
||||
img = self.img_in(img)
|
||||
if self.text_projection == "linear":
|
||||
txt = self.txt_in(txt)
|
||||
elif self.text_projection == "single_refiner":
|
||||
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
|
||||
|
||||
txt_seq_len = txt.shape[1]
|
||||
img_seq_len = img.shape[1]
|
||||
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
if self.enable_teacache:
|
||||
inp = img.clone()
|
||||
vec_ = vec.clone()
|
||||
(
|
||||
img_mod1_shift,
|
||||
img_mod1_scale,
|
||||
img_mod1_gate,
|
||||
img_mod2_shift,
|
||||
img_mod2_scale,
|
||||
img_mod2_gate,
|
||||
) = self.double_blocks[0].img_mod(vec_).chunk(6, dim=-1)
|
||||
normed_inp = self.double_blocks[0].img_norm1(inp)
|
||||
modulated_inp = modulate(normed_inp, shift=img_mod1_shift, scale=img_mod1_scale).to("cpu")
|
||||
del inp, vec_, img_mod1_shift, img_mod1_scale, normed_inp
|
||||
|
||||
if self.cnt == 0 or self.cnt == self.num_steps - 1:
|
||||
should_calc = True
|
||||
self.accumulated_rel_l1_distance = 0
|
||||
else:
|
||||
coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02]
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
self.accumulated_rel_l1_distance += rescale_func(
|
||||
((modulated_inp - self.previous_modulated_input).abs().mean() /
|
||||
self.previous_modulated_input.abs().mean()).cpu().item())
|
||||
if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
|
||||
should_calc = False
|
||||
else:
|
||||
should_calc = True
|
||||
self.accumulated_rel_l1_distance = 0
|
||||
self.previous_modulated_input = modulated_inp
|
||||
self.cnt += 1
|
||||
if self.cnt == self.num_steps:
|
||||
self.cnt = 0
|
||||
if self.enable_teacache:
|
||||
if not should_calc:
|
||||
img += self.previous_residual.to(img.device)
|
||||
self.previous_residual = self.previous_residual.to(img.device)
|
||||
else:
|
||||
ori_img = img.clone().to("cpu")
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index], selected_attn_processor]
|
||||
img, txt = block(*double_block_args)
|
||||
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
x = torch.cat((img, txt), 1)
|
||||
if output_features:
|
||||
features_list = []
|
||||
if len(self.single_blocks) > 0:
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
mask_strategy[index + len(self.double_blocks)],
|
||||
selected_attn_processor,
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
if output_features and _ % output_features_stride == 0:
|
||||
features_list.append(x[:, :img_seq_len, ...])
|
||||
|
||||
img = x[:, :img_seq_len, ...]
|
||||
self.previous_residual = (img.clone().to("cpu") - ori_img).to("cpu")
|
||||
del ori_img
|
||||
else:
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index], selected_attn_processor]
|
||||
img, txt = block(*double_block_args)
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
x = torch.cat((img, txt), 1)
|
||||
if output_features:
|
||||
features_list = []
|
||||
if len(self.single_blocks) > 0:
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
mask_strategy[index + len(self.double_blocks)],
|
||||
selected_attn_processor,
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
if output_features and _ % output_features_stride == 0:
|
||||
features_list.append(x[:, :img_seq_len, ...])
|
||||
|
||||
img = x[:, :img_seq_len, ...]
|
||||
|
||||
# ---------------------------- Final layer ------------------------------
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
|
||||
img = self.unpatchify(img, tt, th, tw)
|
||||
assert not return_dict, "return_dict is not supported."
|
||||
if output_features:
|
||||
features_list = torch.stack(features_list, dim=0)
|
||||
else:
|
||||
features_list = None
|
||||
|
||||
return (img, features_list)
|
||||
|
||||
def initialize_distributed():
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
@@ -28,6 +189,12 @@ def main(args):
|
||||
|
||||
print(args)
|
||||
models_root_path = Path(args.model_path)
|
||||
|
||||
if os.path.exists(args.mask_strategy_file_path):
|
||||
with open(args.mask_strategy_file_path, 'r') as f:
|
||||
mask_strategy = json.load(f)
|
||||
else:
|
||||
mask_strategy = None
|
||||
if not models_root_path.exists():
|
||||
raise ValueError(f"`models_root` not exists: {models_root_path}")
|
||||
|
||||
@@ -40,6 +207,15 @@ def main(args):
|
||||
|
||||
# Get the updated args
|
||||
args = hunyuan_video_sampler.args
|
||||
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.enable_teacache = args.enable_teacache
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.cnt = 0
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.num_steps = args.num_inference_steps
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.rel_l1_thresh = args.rel_l1_thresh # 0.1 for 1.6x speedup, 0.15 for 2.1x speedup
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.accumulated_rel_l1_distance = 0
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.previous_modulated_input = None
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.previous_residual = None
|
||||
hunyuan_video_sampler.pipeline.transformer.__class__.forward = teacache_forward
|
||||
|
||||
if args.prompt.endswith('.txt'):
|
||||
with open(args.prompt) as f:
|
||||
@@ -61,6 +237,7 @@ def main(args):
|
||||
flow_shift=args.flow_shift,
|
||||
batch_size=args.batch_size,
|
||||
embedded_guidance_scale=args.embedded_cfg_scale,
|
||||
mask_strategy=mask_strategy,
|
||||
)
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
outputs = []
|
||||
@@ -162,6 +339,7 @@ if __name__ == "__main__":
|
||||
)
|
||||
|
||||
# Model parameters
|
||||
parser.add_argument("--use-fp8", action='store_true')
|
||||
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
|
||||
parser.add_argument("--latent-channels", type=int, default=16)
|
||||
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
|
||||
@@ -202,7 +380,19 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--text-states-dim-2", type=int, default=768)
|
||||
parser.add_argument("--tokenizer-2", type=str, default="clipL")
|
||||
parser.add_argument("--text-len-2", type=int, default=77)
|
||||
|
||||
parser.add_argument("--vae_tiling", action='store_true')
|
||||
parser.add_argument("--mask_strategy_file_path", type=str, default="assets/mask_strategy.json")
|
||||
parser.add_argument(
|
||||
"--rel_l1_thresh",
|
||||
type=float,
|
||||
default=0.15,
|
||||
help="0.1 for 1.6x speedup, 0.15 for 2.1x speedup",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable_teacache",
|
||||
action="store_true",
|
||||
help="Use teacache for speeding up inference",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
# process for vae sequence parallel
|
||||
if args.vae_sp and not args.vae_tiling:
|
||||
|
||||
+2
-5
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "fastvideo"
|
||||
version = "0.0.1.dev1"
|
||||
version = "0.0.1"
|
||||
description = "FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.8"
|
||||
@@ -39,10 +39,7 @@ dependencies = [
|
||||
"gpustat", "watch",
|
||||
|
||||
# Kernel & Packaging
|
||||
"wheel",
|
||||
|
||||
# Sliding Tile Atteniton Kernel
|
||||
"st_attn>=0.0.1"
|
||||
"wheel"
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -1,32 +1,33 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=4
|
||||
export MODEL_BASE=data/FastHunyuan
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_hunyuan.py \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_inference_steps 6 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 17 \
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 1024 \
|
||||
--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
|
||||
# num_gpus=4
|
||||
# export MODEL_BASE=data/FastHunyuan
|
||||
# torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
# fastvideo/sample/sample_t2v_hunyuan.py \
|
||||
# --height 720 \
|
||||
# --width 1280 \
|
||||
# --num_frames 125 \
|
||||
# --num_inference_steps 6 \
|
||||
# --guidance_scale 1 \
|
||||
# --embedded_cfg_scale 6 \
|
||||
# --flow_shift 17 \
|
||||
# --flow-reverse \
|
||||
# --prompt ./assets/prompt.txt \
|
||||
# --seed 1024 \
|
||||
# --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
|
||||
|
||||
# Inference original Hunyuan
|
||||
num_gpus=4
|
||||
num_gpus=1
|
||||
export MODEL_BASE=data/hunyuan
|
||||
mask_strategy_file_path=assets/test_mask_strategy_hunyuan.json
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_hunyuan.py \
|
||||
--height 720 \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_frames 45 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
@@ -34,7 +35,13 @@ torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/hunyuan/vae_sp/ \
|
||||
--output_path outputs_video/hunyuan/flex_attention_235/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
|
||||
--vae-sp
|
||||
--mask_strategy_file_path $mask_strategy_file_path \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states_fp8.pt \
|
||||
--vae-sp \
|
||||
--use-fp8 \
|
||||
--use-cpu-offload \
|
||||
--vae_tiling \
|
||||
--enable_teacache \
|
||||
--rel_l1_thresh 0.15 \
|
||||
|
||||
@@ -1,39 +1,42 @@
|
||||
#!/bin/bash
|
||||
qshape x strategy_size
|
||||
|
||||
|
||||
|
||||
# Inference with STA + Teacache
|
||||
num_gpus=1 # currently it should be a factor of 30
|
||||
mask_strategy_file_path=assets/mask_strategy_hunyuan.json
|
||||
export MODEL_BASE=data/hunyuan
|
||||
rel_l1_thresh=0.15
|
||||
CUDA_VISIBLE_DEVICES=0,1,2,3 torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29603 \
|
||||
fastvideo/sample/sample_t2v_hunyuan_STA.py \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
--num_frames 117 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 7 \
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 12345 \
|
||||
--output_path outputs_video/hunyuan_STA/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--mask_strategy_file_path $mask_strategy_file_path \
|
||||
--rel_l1_thresh $rel_l1_thresh \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
|
||||
--vae-sp \
|
||||
--enable_teacache
|
||||
# num_gpus=1 # currently it should be a factor of 30
|
||||
# mask_strategy_file_path=assets/mask_strategy_hunyuan.json
|
||||
# export MODEL_BASE=data/hunyuan
|
||||
# rel_l1_thresh=0.15
|
||||
# CUDA_VISIBLE_DEVICES=0,1,2,3 torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29603 \
|
||||
# fastvideo/sample/sample_t2v_hunyuan_STA.py \
|
||||
# --height 768 \
|
||||
# --width 1280 \
|
||||
# --num_frames 117 \
|
||||
# --num_inference_steps 50 \
|
||||
# --guidance_scale 1 \
|
||||
# --embedded_cfg_scale 6 \
|
||||
# --flow_shift 7 \
|
||||
# --flow-reverse \
|
||||
# --prompt ./assets/prompt.txt \
|
||||
# --seed 12345 \
|
||||
# --output_path outputs_video/hunyuan_STA/ \
|
||||
# --model_path $MODEL_BASE \
|
||||
# --mask_strategy_file_path $mask_strategy_file_path \
|
||||
# --rel_l1_thresh $rel_l1_thresh \
|
||||
# --dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
|
||||
# --vae-sp \
|
||||
# --enable_teacache
|
||||
|
||||
# Inference with STA only
|
||||
num_gpus=1
|
||||
mask_strategy_file_path=assets/mask_strategy_hunyuan.json
|
||||
mask_strategy_file_path=assets/test_mask_strategy_hunyuan.json
|
||||
export MODEL_BASE=data/hunyuan
|
||||
CUDA_VISIBLE_DEVICES=1 torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29603 \
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29603 \
|
||||
fastvideo/sample/sample_t2v_hunyuan_STA.py \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
--num_frames 117 \
|
||||
--num_frames 45 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
@@ -41,9 +44,11 @@ CUDA_VISIBLE_DEVICES=1 torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_p
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 12345 \
|
||||
--output_path outputs_video/hunyuan_STA/ \
|
||||
--output_path outputs_video/hunyuan_STA/test_4090/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--mask_strategy_file_path $mask_strategy_file_path \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states_fp8.pt \
|
||||
--vae-sp \
|
||||
--enable_torch_compile
|
||||
--use-fp8 \
|
||||
--use-cpu-offload \
|
||||
--vae_tiling
|
||||
@@ -1,6 +1,6 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=4
|
||||
num_gpus=1
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_hunyuan_hf.py \
|
||||
--model_path ~/data/hunyuan_diffusers/ \
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
recursive-include tk *
|
||||
include config.py
|
||||
@@ -1,252 +0,0 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,58 @@
|
||||
import torch
|
||||
import thunderkittens_cuda as tk
|
||||
import time
|
||||
|
||||
def test_attention_kernel():
|
||||
# Create small test tensors
|
||||
batch_size = 1
|
||||
seq_len = 1 # Start small
|
||||
n_heads = 4
|
||||
head_dim = 64 # Must match your kernel's supported dimensions (64 or 128)
|
||||
|
||||
# Create random inputs
|
||||
q = torch.randn(batch_size, seq_len, n_heads, head_dim,
|
||||
device='cuda', dtype=torch.bfloat16)
|
||||
k = torch.randn_like(q)
|
||||
v = torch.randn_like(q)
|
||||
|
||||
print("Input shapes:", q.shape)
|
||||
print("Input tensors created")
|
||||
|
||||
o = torch.empty_like(q)
|
||||
try:
|
||||
# Add timing
|
||||
start = time.time()
|
||||
print("Calling thunderkittens kernel...")
|
||||
_ = tk.attention_fwd_4090(q, k, v, o, False)
|
||||
print("yes")
|
||||
torch.cuda.synchronize()
|
||||
end = time.time()
|
||||
print(f"Kernel completed in {(end-start)*1000:.2f} ms")
|
||||
print("Output shape:", o.shape)
|
||||
|
||||
# Run again with different parameters
|
||||
_ = tk.attention_fwd_4090(q, k, v, o, True) # Try with causal=True
|
||||
print("Causal version completed")
|
||||
|
||||
# Validate results (compare with PyTorch implementation)
|
||||
print("Computing reference result...")
|
||||
q_scaled = q * (1.0 / torch.sqrt(torch.tensor(head_dim, dtype=torch.float)))
|
||||
attn = torch.matmul(q_scaled, k.transpose(-1, -2))
|
||||
attn = torch.softmax(attn, dim=-1)
|
||||
ref_out = torch.matmul(attn, v)
|
||||
|
||||
# Check if results are reasonable
|
||||
error = (o - ref_out).abs().mean().item()
|
||||
print(f"Mean absolute error: {error}")
|
||||
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"Error calling kernel: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return False
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Testing attention kernel...")
|
||||
success = test_attention_kernel()
|
||||
print(f"Test {'succeeded' if success else 'failed'}")
|
||||
@@ -0,0 +1,297 @@
|
||||
import torch
|
||||
import time
|
||||
import math
|
||||
from einops import rearrange
|
||||
from functools import partial, lru_cache
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
from csrc.sliding_tile_attention.test.sba import get_sliding_block_attention_mask
|
||||
import tk_4090_cuda as tk
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
import torch.nn.functional as F
|
||||
# Import benchmark utilities from flash_attn
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
torch._dynamo.config.cache_size_limit = 64
|
||||
|
||||
@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_block_attention_mask(
|
||||
adjusted_strategy,
|
||||
tile_size,
|
||||
image_size,
|
||||
text_length,
|
||||
device
|
||||
)
|
||||
|
||||
# Compile the flex attention with the mask
|
||||
compiled_flex_attn = torch.compile(
|
||||
partial(flex_attention, block_mask=mask),
|
||||
dynamic=False
|
||||
)
|
||||
|
||||
return compiled_flex_attn
|
||||
|
||||
def 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 tile(x, sp_size):
|
||||
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=12 // sp_size, h=48, w=80)
|
||||
return rearrange(x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
n_t=2,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
|
||||
def flops_attention_forward(batch, seqlen, headdim, nheads):
|
||||
"""
|
||||
Calculate FLOPs for attention forward pass operation.
|
||||
"""
|
||||
return 4 * batch * seqlen**2 * nheads * headdim
|
||||
|
||||
def efficiency(flop, time):
|
||||
"""
|
||||
Calculate efficiency in TFLOPs/s
|
||||
"""
|
||||
return (flop / time / 10**12) if not math.isnan(time) and time > 0 else 0.
|
||||
|
||||
def benchmark_tk_attention(func, q, k, v, o, text_length, repeats=100, desc=""):
|
||||
time_f, m = benchmark_forward(func, q, k, v, o, text_length, repeats=repeats, desc=desc, verbose=False)
|
||||
return m.mean
|
||||
|
||||
def benchmark_baseline_attention(func, qkv, attn_mask, causal, dropout_p, softmax_sacle, repeats=100, desc=""):
|
||||
time_f, m = benchmark_forward(func, qkv, attn_mask, causal, dropout_p, softmax_sacle, repeats=repeats, desc=desc, verbose=False)
|
||||
return m.mean
|
||||
|
||||
def benchmark_compiled_attention(func, q, k, v, strategy, tile_size, image_size, text_length, repeats=100, desc=""):
|
||||
time_f, m = benchmark_forward(func, q, k, v, strategy, tile_size, image_size, text_length, repeats=repeats, desc=desc, verbose=False)
|
||||
return m.mean
|
||||
|
||||
|
||||
def benchmark_attention(func, q, k, v, repeats=100, desc=""):
|
||||
"""
|
||||
Benchmark forward pass of attention function
|
||||
"""
|
||||
time_f, m = benchmark_forward(func, q, k, v, repeats=repeats, desc=desc, verbose=False)
|
||||
return m.mean
|
||||
|
||||
def main():
|
||||
debug = False
|
||||
|
||||
# Load tensors
|
||||
print("Loading tensors...")
|
||||
if debug:
|
||||
debug_seq_len = 46080
|
||||
debug_batch_size = 1
|
||||
debug_nheads = 24
|
||||
debug_headdim = 128
|
||||
q = torch.randn((debug_batch_size, debug_seq_len+256, debug_nheads, debug_headdim), device="cuda", dtype=torch.bfloat16)
|
||||
k = torch.randn((debug_batch_size, debug_seq_len+256, debug_nheads, debug_headdim), device="cuda", dtype=torch.bfloat16)
|
||||
v = torch.randn((debug_batch_size, debug_seq_len+256, debug_nheads, debug_headdim), device="cuda", dtype=torch.bfloat16)
|
||||
else:
|
||||
q = torch.load("query.pt").to("cuda")
|
||||
k = torch.load("key.pt").to("cuda")
|
||||
v = torch.load("value.pt").to("cuda")
|
||||
text_mask = torch.load("text_mask.pt")
|
||||
text_length = text_mask.sum()
|
||||
|
||||
batch_size = q.shape[0]
|
||||
nheads = q.shape[2]
|
||||
seqlen = q.shape[1]
|
||||
headdim = q.shape[3]
|
||||
warmup = 50
|
||||
repeats = 100 # Number of repeats for reliable timing
|
||||
|
||||
# Baseline measurement (standard flex attention without masking)
|
||||
baseline_attn = flash_attn_no_pad
|
||||
|
||||
# Perform warmup for baseline
|
||||
print("\nWarming up baseline...")
|
||||
if debug:
|
||||
attn_mask = F.pad(text_mask, (debug_seq_len, 0), value=True)
|
||||
else:
|
||||
attn_mask = F.pad(text_mask, (seqlen - 256, 0), value=True)
|
||||
for _ in range(warmup):
|
||||
flash_attn_no_pad(torch.stack([q, k, v], dim=2), attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
|
||||
# Benchmark baseline
|
||||
print("\nBenchmarking baseline flex attention...")
|
||||
torch.cuda.synchronize()
|
||||
baseline_time = benchmark_baseline_attention(baseline_attn, torch.stack([q, k, v], dim=2), attn_mask, causal=False, dropout_p=0.0, softmax_sacle=None, repeats=repeats, desc="Baseline")
|
||||
|
||||
# benchmark tk_4090_cuda
|
||||
# tk.attention_fwd_4090(query, key, value, result, text_length)
|
||||
print("Benchmarking tk_4090_cuda...")
|
||||
result = torch.empty_like(q)
|
||||
for _ in range(warmup):
|
||||
tk.attention_fwd_4090(q, k, v, result, text_length)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
tk_time = benchmark_tk_attention(tk.attention_fwd_4090, q, k, v, result, text_length, repeats=repeats, desc="tk_4090_cuda")
|
||||
|
||||
|
||||
# Prepare tensors
|
||||
strategies = []
|
||||
if not debug:
|
||||
query, encoder_query = q.split_with_sizes((q.shape[1] - 256, 256), dim=1)
|
||||
key, encoder_key = k.split_with_sizes((k.shape[1] - 256, 256), dim=1)
|
||||
value, encoder_value = v.split_with_sizes((v.shape[1] - 256, 256), dim=1)
|
||||
|
||||
q = torch.cat([tile(query, 1), encoder_query], dim=1).transpose(1, 2)
|
||||
k = torch.cat([tile(key, 1), encoder_key], dim=1).transpose(1, 2)
|
||||
v = torch.cat([tile(value, 1), encoder_value], dim=1).transpose(1, 2)
|
||||
|
||||
# Strategy definitions
|
||||
strategies = [
|
||||
(2, 6, 1), # (t, h, w)
|
||||
(1, 6, 10),
|
||||
(2, 3, 3),
|
||||
(2, 6, 10),
|
||||
(2, 1, 10),
|
||||
(2, 3, 5)
|
||||
]
|
||||
|
||||
# Strategy names for better reporting
|
||||
strategy_names = [f"Strategy_{t}_{h}_{w}" for t, h, w in strategies]
|
||||
|
||||
print(f"Text length: {text_length}")
|
||||
|
||||
# flex full baseline
|
||||
print("Benchmarking flex full baseline...")
|
||||
flex_baseline = torch.compile(
|
||||
partial(flex_attention, block_mask=None),
|
||||
)
|
||||
|
||||
for _ in range(warmup):
|
||||
flex_baseline(q, k, v)
|
||||
torch.cuda.synchronize()
|
||||
flex_baseline_time = benchmark_attention(flex_baseline, q, k, v, repeats=repeats, desc="FA2 Baseline")
|
||||
|
||||
|
||||
# Create attention processors for each strategy
|
||||
attn_processors = []
|
||||
for strategy in strategies:
|
||||
mask = get_sliding_block_attention_mask(strategy, (6, 8, 8), (12, 48, 80), text_length, "cuda")
|
||||
attn_processor = torch.compile(
|
||||
partial(flex_attention, block_mask=mask),
|
||||
)
|
||||
attn_processors.append(attn_processor)
|
||||
|
||||
# Setup benchmarking parameters
|
||||
|
||||
print(f"Tensor shapes - Q/K/V: {q.shape}")
|
||||
print(f"Benchmarking with: batch_size={batch_size}, nheads={nheads}, seqlen={seqlen}, headdim={headdim}")
|
||||
|
||||
# Calculate theoretical FLOPS for forward pass
|
||||
flops = flops_attention_forward(batch_size, seqlen, headdim, nheads)
|
||||
print(f"Theoretical FLOPs: {flops / 10**12:.2f} TFLOPs")
|
||||
|
||||
tk_efficiency = efficiency(flops, tk_time)
|
||||
print(f"tk_4090_cuda forward: {tk_time:.6f}s, {tk_efficiency:.2f} TFLOPs/s")
|
||||
|
||||
baseline_efficiency = efficiency(flops, baseline_time)
|
||||
|
||||
print(f"Baseline forward: {baseline_time:.6f}s, {baseline_efficiency:.2f} TFLOPs/s")
|
||||
|
||||
# Create result containers
|
||||
times = {}
|
||||
speeds = {}
|
||||
speedups = {}
|
||||
|
||||
# Benchmark each attention strategy
|
||||
print("\nBenchmarking attention strategies...")
|
||||
if not debug:
|
||||
for i, (processor, strategy, name) in enumerate(zip(attn_processors, strategies, strategy_names)):
|
||||
t, h, w = strategy
|
||||
|
||||
# Warmup
|
||||
print(f"\nWarming up {name}...")
|
||||
for _ in range(warmup):
|
||||
processor(q, k, v)
|
||||
sliding_tile_attention(q, k, v, strategy, (6, 8, 8), (12, 48, 80), text_length)
|
||||
|
||||
# Benchmark
|
||||
print(f"Benchmarking {name} (t={t}, h={h}, w={w})...")
|
||||
torch.cuda.synchronize()
|
||||
fwd_time = benchmark_attention(processor, q, k, v, repeats=repeats, desc=name)
|
||||
print(f"processor {name} (t={t}, h={h}, w={w}): {fwd_time:.6f}s")
|
||||
|
||||
compiled_fwd_time = benchmark_compiled_attention(sliding_tile_attention, q, k, v, strategy, (6, 8, 8), (12, 48, 80), text_length, repeats=repeats, desc=name)
|
||||
|
||||
|
||||
print(f"Compiled {name} (t={t}, h={h}, w={w}): {compiled_fwd_time:.6f}s")
|
||||
|
||||
# Save results
|
||||
times[name] = fwd_time
|
||||
speeds[name] = efficiency(flops, fwd_time)
|
||||
|
||||
# Calculate theoretical and actual speedup
|
||||
theoretical_speedup = (2*6*10)/(t*h*w)
|
||||
actual_speedup = baseline_time / fwd_time
|
||||
speedups[name] = (theoretical_speedup, actual_speedup)
|
||||
|
||||
# Print results
|
||||
print(f"Strategy: {name} (t={t}, h={h}, w={w})")
|
||||
print(f"Forward: {fwd_time:.6f}s, {speeds[name]:.2f} TFLOPs/s")
|
||||
print(f"Theoretical speedup: {theoretical_speedup:.2f}x")
|
||||
print(f"Actual speedup: {actual_speedup:.2f}x")
|
||||
print(f"Efficiency ratio (actual/theoretical): {actual_speedup/theoretical_speedup:.2f}")
|
||||
|
||||
# Print summary table
|
||||
print("\n" + "="*80)
|
||||
print("SUMMARY OF RESULTS")
|
||||
print("="*80)
|
||||
print(f"{'Strategy':<15} {'Time (ms)':<15} {'TFLOPs/s':<10} {'Actual Speedup':<15} {'Theoretical':<15} {'Efficiency':<10}")
|
||||
print("-"*80)
|
||||
|
||||
# Baseline row
|
||||
print(f"{'FA2 Baseline':<15} {baseline_time*1000:13.2f} {baseline_efficiency:10.2f} {1.00:15.2f} {1.00:15.2f} {1.00:10.2f}")
|
||||
if not debug:
|
||||
print(f"{'FA2 Flex':<15} {flex_baseline_time*1000:13.2f} {efficiency(flops, flex_baseline_time):10.2f} {baseline_time/flex_baseline_time:15.2f} {1.00:15.2f} {baseline_time/flex_baseline_time:10.2f}")
|
||||
print(f"{'Tk_4090_cuda':<15} {tk_time*1000:13.2f} {tk_efficiency:10.2f} {baseline_time/tk_time:15.2f} {1.00:15.2f} {baseline_time/tk_time:10.2f}")
|
||||
|
||||
# Strategy rows
|
||||
if not debug:
|
||||
for name in strategy_names:
|
||||
theoretical, actual = speedups[name]
|
||||
print(f"{name:<15} {times[name]*1000:13.2f} {speeds[name]:10.2f} {actual:15.2f} {theoretical:15.2f} {actual/theoretical:10.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user