Compare commits

..
10 Commits
26 changed files with 454 additions and 622 deletions
-70
View File
@@ -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/
-221
View File
@@ -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/
-3
View File
@@ -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
View File
@@ -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
+26
View File
@@ -45,6 +45,32 @@ The code is tested on Python 3.10.0, CUDA 12.4 and H100.
```
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
## 🎯 STA mask search pipeline
### Overview
The STA mask search pipeline consists of three sequential steps:
1. **Searching**: Choose sparse attention mask candidates and do searching
2. **Tuning**: Use L2 loss to determine optimal mask strategy
3. **Inference**: Apply selected strategy for fast video generation
```bash
sh scripts/inference/inference_hunyuan.sh # Inference stepvideo with STA
```
The only thing you need to do is to specify ```--STA_mode``` with original hunyuan inference script.
#### Step 1: Searching
Run with ```--STA_mode STA_searching```, and this step generates a folder containing mask search results in JSON format for each prompt.
#### Step 2: Tuning
Run with ```--STA_mode STA_tuning```. During this step, the system will:
1. Reads all JSON files from the search results folder
2. Averages L2 distances across different masks to determine the optimal mask strategy per attention head. (First 12-15 steps will be full mask to get better quality)
3. Generates accelerated videos for evaluation
4. Saves the best strategy to a single json file
#### Step 3: Inference
After determining the optimal strategy, run with ```--STA_mode STA_inference```. This step reads the strategy file and runs inference with the optimized settings.
#### Configuration
You can modify various STA configuration parameters in:
```fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py```
## 🚀 Inference
### Inference StepVideo with Sliding Tile Attention
First, download the model:
@@ -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})
+1 -1
View File
@@ -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
@@ -550,7 +550,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
enable_vae_sp: bool = False,
n_tokens: Optional[int] = None,
embedded_guidance_scale: Optional[float] = None,
mask_strategy: Optional[Dict[str, list]] = None,
STA_mode: Optional[str] = None,
**kwargs,
):
r"""
@@ -782,6 +782,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
generator,
latents,
)
img_size = latents.shape[-3:]
img_size = (img_size[0], img_size[1] // 2, img_size[2] // 2)
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
if get_sequence_parallel_state():
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
@@ -801,26 +804,36 @@ class HunyuanVideoPipeline(DiffusionPipeline):
vae_dtype = PRECISION_TO_TYPE[self.args.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32) and not self.args.disable_autocast
# STA
from fastvideo.utils.STA_configuration import configure_sta
mask_search_final_result = []
sparse_mask_candidates = ["1,6,10", "3,3,5", "5,1,10", "5,3,3", "5,6,1"]
full_mask = ["5,6,10"]
STA_param = None
if STA_mode == 'STA_searching':
STA_param = configure_sta(
mode='STA_searching',
mask_candidates=sparse_mask_candidates +
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
)
elif STA_mode == 'STA_tuning':
STA_param = configure_sta(
mode='STA_tuning',
mask_search_files_path='output/mask_search_result/',
mask_candidates=sparse_mask_candidates,
skip_time_steps=15, # Use full attention for first 15 steps
save_dir='output/mask_strategy' # Custom save directory
)
elif STA_mode == 'STA_inference':
STA_param = configure_sta(mode='STA_inference', load_path='output/mask_strategy/mask_strategy.json')
# 7. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, l, h = map(int, key.split('_'))
result[t][l][h] = value
return result
mask_strategy = dict_to_3d_list(mask_strategy)
# if is_progress_bar:
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
# expand the latents if we are doing classifier free guidance
latent_model_input = (torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents)
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
@@ -841,15 +854,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
value=0,
).unsqueeze(1)
encoder_hidden_states = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
noise_pred, _, mask_search_result = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
latent_model_input,
encoder_hidden_states,
t_expand,
prompt_mask,
mask_strategy=mask_strategy[i],
STA_param=STA_param[i],
guidance=guidance_expand,
return_dict=False,
)[0]
)
mask_search_final_result.append(mask_search_result)
# perform guidance
if self.do_classifier_free_guidance:
@@ -888,6 +902,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if get_sequence_parallel_state():
latents = all_gather(latents, dim=2)
if STA_mode == 'STA_searching':
from fastvideo.utils.STA_configuration import save_mask_search_results
save_mask_search_results(mask_search_final_result,
prompt=prompt,
mask_strategies=sparse_mask_candidates,
output_dir='output/mask_search_result_test/')
if not output_type == "latent":
expand_temporal_dim = False
if len(latents.shape) == 4:
+2 -2
View File
@@ -338,7 +338,7 @@ class HunyuanVideoSampler(Inference):
embedded_guidance_scale=None,
batch_size=1,
num_videos_per_prompt=1,
mask_strategy=None,
STA_mode=None,
**kwargs,
):
"""
@@ -471,7 +471,7 @@ class HunyuanVideoSampler(Inference):
vae_ver=self.args.vae,
enable_tiling=self.args.vae_tiling,
enable_vae_sp=self.args.vae_sp,
mask_strategy=mask_strategy,
STA_mode=STA_mode,
)[0]
out_dict["samples"] = samples
out_dict["prompts"] = prompt
+37 -10
View File
@@ -58,7 +58,7 @@ def untile(x, sp_size):
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)
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, STA_param=None):
query, encoder_query = q
key, encoder_key = k
value, encoder_value = v
@@ -81,18 +81,45 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
if mask_strategy[0] is not None:
loss_result = None
if STA_param[0] is not None:
query = torch.cat([tile(query, nccl_info.sp_size), encoder_query], dim=1).transpose(1, 2)
key = torch.cat([tile(key, nccl_info.sp_size), encoder_key], dim=1).transpose(1, 2)
value = torch.cat([tile(value, nccl_info.sp_size), encoder_value], dim=1).transpose(1, 2)
head_num = query.size(1)
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)]
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
if len(STA_param) < 24: # searching mode; thus do not use more than 24 mask candidates
sparse_attn_hidden_states_all = []
full_mask_window = STA_param[-1]
for window_size in STA_param[:-1]:
hidden_states = sliding_tile_attention(query, key, value, [window_size] * head_num,
text_length).transpose(1, 2)
sparse_attn_hidden_states_all.append(hidden_states)
hidden_states = sliding_tile_attention(query, key, value, [full_mask_window] * head_num,
text_length).transpose(1, 2) # torch.Size([1, 115456, 24, 128])
attn_L2_loss = []
attn_L1_loss = []
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
# L2 loss
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() - hidden_states.float())**2,
dim=[0, 1, 3]).cpu().numpy()
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
attn_L2_loss.append(attn_L2_loss_)
# L1 loss
attn_L1_loss_ = torch.mean(torch.abs(sparse_attn_hidden_states.float() - hidden_states.float()),
dim=[0, 1, 3]).cpu().numpy()
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
attn_L1_loss.append(attn_L1_loss_)
loss_result = [attn_L2_loss, attn_L1_loss]
else:
current_rank = nccl_info.rank_within_group
start_head = current_rank * head_num
windows = [STA_param[head_idx + start_head] for head_idx in range(head_num)]
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)
@@ -106,7 +133,7 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes((sequence_length, encoder_sequence_length),
dim=1)
if mask_strategy[0] is not None:
if STA_param[0] is not None:
hidden_states = untile(hidden_states, nccl_info.sp_size)
if get_sequence_parallel_state():
@@ -121,4 +148,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, loss_result
+33 -17
View File
@@ -109,7 +109,7 @@ class MMDoubleStreamBlock(nn.Module):
vec: torch.Tensor,
freqs_cis: tuple = None,
text_mask: torch.Tensor = None,
mask_strategy=None,
STA_param=None,
) -> Tuple[torch.Tensor, torch.Tensor]:
(
img_mod1_shift,
@@ -163,16 +163,23 @@ class MMDoubleStreamBlock(nn.Module):
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
attn = parallel_attention(
attn, loss_result = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
(img_v, txt_v),
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
mask_strategy=mask_strategy,
STA_param=STA_param,
)
if loss_result is not None:
layer_loss_save = {
"L2_loss": loss_result[0],
"L1_loss": loss_result[1],
}
else:
layer_loss_save = None
# attention computation end
img_attn, txt_attn = attn[:, :img.shape[1]], attn[:, img.shape[1]:]
@@ -190,7 +197,7 @@ class MMDoubleStreamBlock(nn.Module):
self.txt_mlp(modulate(self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale)),
gate=txt_mod2_gate,
)
return img, txt
return img, txt, layer_loss_save
class MMSingleStreamBlock(nn.Module):
@@ -259,7 +266,7 @@ class MMSingleStreamBlock(nn.Module):
txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
text_mask: torch.Tensor = None,
mask_strategy=None,
STA_param=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)
@@ -288,21 +295,28 @@ class MMSingleStreamBlock(nn.Module):
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
attn = parallel_attention(
attn, loss_result = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
(img_v, txt_v),
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
mask_strategy=mask_strategy,
STA_param=STA_param,
)
if loss_result is not None:
layer_loss_save = {
"L2_loss": loss_result[0],
"L1_loss": loss_result[1],
}
else:
layer_loss_save = None
# attention computation end
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + apply_gate(output, gate=mod_gate)
return x + apply_gate(output, gate=mod_gate), layer_loss_save
class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
@@ -514,7 +528,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
mask_strategy=None,
STA_param=None,
output_features=False,
output_features_stride=8,
attention_kwargs: Optional[Dict[str, Any]] = None,
@@ -523,8 +537,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))]
if STA_param is None:
STA_param = [[None] * len(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
@@ -566,10 +580,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
mask_search_result_save = []
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
img, txt = block(*double_block_args)
double_block_args = [img, txt, vec, freqs_cis, text_mask, STA_param[index]]
img, txt, layer_loss_save = block(*double_block_args)
mask_search_result_save.append(layer_loss_save)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if output_features:
@@ -582,9 +597,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
txt_seq_len,
(freqs_cos, freqs_sin),
text_mask,
mask_strategy[index + len(self.double_blocks)],
STA_param[index + len(self.double_blocks)],
]
x = block(*single_block_args)
x, layer_loss_save = block(*single_block_args)
mask_search_result_save.append(layer_loss_save)
if output_features and _ % output_features_stride == 0:
features_list.append(x[:, :img_seq_len, ...])
@@ -599,7 +615,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
features_list = torch.stack(features_list, dim=0)
else:
features_list = None
return (img, features_list)
return (img, features_list, mask_search_result_save)
def unpatchify(self, x, t, h, w):
"""
+5
View File
@@ -61,6 +61,7 @@ def main(args):
flow_shift=args.flow_shift,
batch_size=args.batch_size,
embedded_guidance_scale=args.embedded_cfg_scale,
STA_mode=args.STA_mode,
)
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
outputs = []
@@ -202,6 +203,10 @@ 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("--STA_mode",
type=str,
default="STA_inference",
help="STA_modes should be one of ['STA_searching', 'STA_tuning', 'STA_inference']")
args = parser.parse_args()
# process for vae sequence parallel
+284
View File
@@ -0,0 +1,284 @@
import json
import os
from collections import defaultdict
import numpy as np
def configure_sta(mode='STA_searching', **kwargs):
"""
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
Parameters:
----------
mode : str
The STA mode to use. Options are:
- 'STA_searching': Generate a set of mask candidates for initial search
- 'STA_tuning': Select best mask strategy based on previously saved results
- 'STA_inference': Load and use a previously tuned mask strategy
**kwargs : dict
Mode-specific parameters:
For 'STA_searching':
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
For 'STA_tuning':
- mask_search_files_path: str, required, path to mask search results
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
- skip_time_steps: int, optional, number of time steps to use full attention (default 15)
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
For 'STA_inference':
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
Returns:
-------
list
The configured STA parameter (STA_param) for the specified mode
"""
valid_modes = ['STA_searching', 'STA_tuning', 'STA_inference']
if mode not in valid_modes:
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
if mode == 'STA_searching':
# Get parameters with defaults
mask_candidates = kwargs.get('mask_candidates', ["1,6,10", "3,3,5", "5,1,10", "5,3,3", "5,6,1", "5,6,10"])
mask_selected = kwargs.get('mask_selected', list(range(len(mask_candidates))))
# Parse selected masks
selected_masks = []
for index in mask_selected:
mask = mask_candidates[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Create 3D mask structure with fixed dimensions (t=50, l=60)
masks_3d = []
for i in range(50): # Fixed t dimension = 50
row = []
for j in range(60): # Fixed l dimension = 60
row.append(selected_masks) # Add all masks at each position
masks_3d.append(row)
return masks_3d
elif mode == 'STA_tuning':
# Get required parameters
mask_search_files_path = kwargs.get('mask_search_files_path')
if not mask_search_files_path:
raise ValueError("mask_search_files_path is required for STA_tuning mode")
# Get optional parameters with defaults
mask_candidates = kwargs.get('mask_candidates', ["1,6,10", "3,3,5", "5,1,10", "5,3,3", "5,6,1"])
mask_selected = kwargs.get('mask_selected', list(range(len(mask_candidates))))
skip_time_steps = kwargs.get('skip_time_steps', 15)
save_dir = kwargs.get('save_dir', "output/mask_strategy")
# Parse selected masks
selected_masks = []
for index in mask_selected:
mask = mask_candidates[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Read JSON results
results = read_specific_json_files(mask_search_files_path)
averaged_results = average_head_losses(results, selected_masks)
# Add full attention mask for specific cases
full_attention_mask = kwargs.get('full_attention_mask', [5, 6, 10])
selected_masks.append(full_attention_mask)
# Select best mask strategy
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(averaged_results, selected_masks,
skip_time_steps)
# Save mask strategy
os.makedirs(save_dir, exist_ok=True)
file_path = os.path.join(save_dir, 'mask_strategy.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = 50 * 60 * 24 # Fixed dimensions
for strategy, count in strategy_counts.items():
print(f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)")
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy)
return mask_strategy_3d
else: # STA_inference
# Get parameters with defaults
load_path = kwargs.get('load_path', os.path.join("mask_candidates", 'mask_strategy.json'))
# Load previously saved mask strategy
with open(load_path, 'r') as f:
mask_strategy = json.load(f)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy)
return mask_strategy_3d
# Helper functions
def read_specific_json_files(folder_path):
"""Read and parse JSON files containing mask search results."""
json_contents = []
# List files only in the current directory (no walk)
files = os.listdir(folder_path)
# Filter files
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
print(f"Found {len(matching_files)} matching files: {matching_files}")
for file_name in matching_files:
file_path = os.path.join(folder_path, file_name)
with open(file_path, 'r') as file:
data = json.load(file)
json_contents.append(data)
return json_contents
def average_head_losses(results, selected_masks):
"""Average losses across all prompts for each mask strategy."""
# Initialize a dictionary to store the averaged results
averaged_losses = {}
loss_type = 'L2_loss'
# Get all loss types (e.g., 'L2_loss')
averaged_losses[loss_type] = {}
for mask in selected_masks:
mask_str = str(mask)
data_shape = np.array(results[0][loss_type][mask_str]).shape
accumulated_data = np.zeros(data_shape)
# Sum across all prompts
for prompt_result in results:
accumulated_data += np.array(prompt_result[loss_type][mask_str])
# Average by dividing by number of prompts
averaged_data = accumulated_data / len(results)
averaged_losses[loss_type][mask_str] = averaged_data
return averaged_losses
def select_best_mask_strategy(averaged_results, selected_masks, skip_time_steps=15):
"""Select the best mask strategy for each head based on loss minimization."""
best_mask_strategy = {}
loss_type = 'L2_loss'
# Get the shape of time steps and layers
time_steps = len(averaged_results[loss_type][str(selected_masks[0])])
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
# Counter for sparsity calculation
total_tokens = 0 # total number of masked tokens
total_length = 0 # total sequence length
strategy_counts = {str(strategy): 0 for strategy in selected_masks}
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
for t in range(time_steps):
for l in range(layers):
for h in range(24):
if t < skip_time_steps: # First steps use full attention
strategy = full_attn_strategy
else:
# Get losses for this head across all strategies
head_losses = []
for strategy in selected_masks[:-1]: # Exclude full attention
head_losses.append(averaged_results[loss_type][str(strategy)][t][l][h])
# Find which strategy gives minimum loss
best_strategy_idx = np.argmin(head_losses)
strategy = selected_masks[best_strategy_idx]
best_mask_strategy[f'{t}_{l}_{h}'] = strategy
# Calculate sparsity
nums = strategy # strategy is already a list of numbers
total_tokens += nums[0] * nums[1] * nums[2] # masked tokens for chosen strategy
total_length += 300 # total length always 5*6*10=300
# Count strategy usage
strategy_counts[str(strategy)] += 1
overall_sparsity = 1 - total_tokens / total_length
return best_mask_strategy, overall_sparsity, strategy_counts
def dict_to_3d_list(mask_strategy):
"""Convert a mask strategy dictionary to a 3D list structure with fixed dimensions."""
# Fixed dimensions for t, l, h (50, 60, 24)
result = [[[None for _ in range(24)] for _ in range(60)] for _ in range(50)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, l, h = map(int, key.split('_'))
result[t][l][h] = value
return result
def save_mask_search_results(mask_search_final_result,
prompt,
mask_strategies,
output_dir='output/mask_search_result/'):
if not mask_search_final_result:
print("No mask search results to save")
return None
# Create result dictionary with defaultdict for nested lists
mask_search_dict = {"L2_loss": defaultdict(list), "L1_loss": defaultdict(list)}
mask_selected = list(range(len(mask_strategies)))
selected_masks = []
for index in mask_selected:
mask = mask_strategies[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Process each mask strategy
for i, mask_strategy in enumerate(selected_masks):
mask_strategy = str(mask_strategy)
# Process L2 loss
step_results = []
for step_data in mask_search_final_result:
layer_losses = [layer_data["L2_loss"][i] for layer_data in step_data]
step_results.append(layer_losses)
mask_search_dict["L2_loss"][mask_strategy] = step_results
step_results = []
for step_data in mask_search_final_result:
layer_losses = [layer_data["L1_loss"][i] for layer_data in step_data]
step_results.append(layer_losses)
mask_search_dict["L1_loss"][mask_strategy] = step_results
# Create the output directory if it doesn't exist
os.makedirs(output_dir, exist_ok=True)
# Create a filename based on the first 20 characters of the prompt
filename = prompt[0][:20].replace(" ", "_")
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
# Save the results to a JSON file
with open(filepath, 'w') as f:
json.dump(mask_search_dict, f, indent=4)
print(f"Successfully saved mask research results to {filepath}")
return filepath
+3 -6
View File
@@ -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"
]
@@ -69,7 +66,7 @@ skip ="./data,./wandb,./csrc/sliding_tile_attention/tk"
"fastvideo/models/stepvideo/__init__.py" = ["F403"]
"fastvideo/models/stepvideo/utils/__init__.py" = ["F403"]
# Ignore all files that end in `_test.py`.
"fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py" = ["E741"]
"fastvideo/utils/STA_configuration.py" = ["E741"]
[tool.yapf]
+21
View File
@@ -38,3 +38,24 @@ torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
--model_path $MODEL_BASE \
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
--vae-sp
# Mask search for STA
num_gpus=1
export MODEL_BASE=data/hunyuan
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/sample/sample_t2v_hunyuan.py \
--height 768 \
--width 1280 \
--num_frames 117 \
--num_inference_steps 1 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 7 \
--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 \
--STA_mode STA_searching \
-2
View File
@@ -1,2 +0,0 @@
recursive-include tk *
include config.py
-252
View File
@@ -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()