Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d9d4bb5392 | ||
|
|
1d9a593ba9 | ||
|
|
1d9364c29c | ||
|
|
68d012210c | ||
|
|
89ed16efa4 | ||
|
|
b2a581c45d | ||
|
|
b267d0d041 | ||
|
|
782dae2739 | ||
|
|
722c47932f | ||
|
|
377d1607ba |
@@ -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
|
||||
|
||||
@@ -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})
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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 \
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user