Compare commits

...
53 Commits
Author SHA1 Message Date
Peiyuan Zhang c83432bdea syn 2025-03-28 16:46:39 +00:00
Peiyuan Zhang e90d7595b1 syn 2025-03-25 00:19:51 +00:00
Peiyuan Zhang 83db422e9c add lora convert 2025-03-25 00:04:10 +00:00
Zhang Peiyuan dc2a4514e8 Streamline pipeline (#271) 2025-03-16 19:23:11 -07:00
SolitaryThinker dd10588fb7 cleanup 2025-03-16 20:18:03 +00:00
SolitaryThinker 8140cb269f cleanup 2025-03-16 19:57:09 +00:00
SolitaryThinker bb6ae368f2 moved loader/ into models/ 2025-03-16 19:40:16 +00:00
SolitaryThinker e14d384a71 model/loader.py -> component_loader.py 2025-03-16 19:27:45 +00:00
William Lin 6c528468b5 Merge pull request #1 from SolitaryThinker/wei
Add wan dit
2025-03-16 12:22:27 -07:00
Peiyuan Zhang ae5ed0c7e0 update 2025-03-16 04:24:51 +00:00
Peiyuan Zhang c37535ab0b magic line 2025-03-16 03:33:42 +00:00
Peiyuan Zhang 4e056b92cd Merge branch 'rebased-refactor' of https://github.com/SolitaryThinker/FastVideo into rebased-refactor 2025-03-16 03:05:04 +00:00
Peiyuan Zhang 83348c6ed0 update 2025-03-16 03:05:01 +00:00
SolitaryThinker 691f9d1064 vae update 2025-03-16 01:37:58 +00:00
SolitaryThinker 796eaf809f add toggle flags for v0 pipeline components 2025-03-16 00:51:52 +00:00
SolitaryThinker f51e9d486b running, correctness isues 2025-03-15 23:38:37 +00:00
Peiyuan Zhang 759f243cf3 fix attn 2025-03-15 21:13:38 +00:00
SolitaryThinker fb44fbaa1c debugging encoders 2025-03-15 19:30:12 +00:00
Zhou, Wei e976583cf4 Add wan dit 2025-03-14 20:00:01 -04:00
SolitaryThinker d2db0d475b fix import paths 2025-03-14 23:13:34 +00:00
SolitaryThinker d5ac1e9bee remove unused attention 2025-03-14 19:31:38 +00:00
SolitaryThinker 1976b23121 move v1's v0 code into v1/v0_reference_src 2025-03-14 19:29:29 +00:00
SolitaryThinker eafeea4a3f revert pyproject.toml 2025-03-14 19:19:35 +00:00
SolitaryThinker bf1fc27989 remove unneeded file 2025-03-14 19:17:53 +00:00
8d99ec3c85 V1 (#257)
Signed-off-by: <>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
Co-authored-by: Ubuntu <ubuntu@awesome-gpu-name-8-inst-2tbsnfodvpomxv4tukw2dkfgyvz.c.nv-brev-20240723.internal>
Co-authored-by: Ubuntu <ubuntu@awesome-gpu-name-9-inst-2tpydiudxfu1jg9xvpflm7oexie.c.nv-brev-20240723.internal>
2025-03-14 19:14:54 +00:00
William Lin b631546e18 move refactor to fastvideo/v1 (#265) 2025-03-14 19:14:54 +00:00
Zhang Peiyuan b2def4b57c DiT done and plub in pipeline (#252) 2025-03-14 19:14:54 +00:00
William LinandPeiyuan Zhang 23fd3ed3c7 [Do not merge] V1 encoders and model loading (#261)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-03-14 19:14:54 +00:00
Zhang PeiyuanandWill Lin 6e8b11c137 v1 staging architecture
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-03-14 19:14:54 +00:00
Zhang Peiyuan fc5a4bc236 Refactor py (#246) 2025-03-14 19:14:54 +00:00
Zhang Peiyuan 0f4c8d1360 [Refactor] Add Hunyuan DiT Modeling (#241) 2025-03-14 19:14:51 +00:00
William Lin 5252d50b25 Initial clip encoder and cli args organization (#232) 2025-03-14 19:07:58 +00:00
William Lin ac07e436bb add sp comm (#231) 2025-03-14 19:07:58 +00:00
William Lin 42f902cf23 Initial set of common files and layers from vLLM (#226) 2025-03-14 19:07:58 +00:00
You Zhou 8a77cf22c9 Establish cicd workflow to build and publish FastVideo and STA Kernel (#227) 2025-03-11 20:27:36 -07:00
Yongqi ChenandPeiyuan Zhang d869d90d12 fix training mask strategy issue (#248)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-03-05 20:00:16 -08:00
Zhang Peiyuan 554ee17de5 [BUG] update cfg bug? (#223) 2025-02-27 16:02:44 -08:00
Yongqi ChenandPeiyuan Zhang 0be4fc62c9 fix train/distill issue (#215)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-25 08:11:17 -08:00
Yongqi ChenandPeiyuan Zhang 1e08893546 Added multi-GPU support for Hunyuan STA (#211)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-21 14:16:28 -08:00
Zhang Peiyuan 09ab452610 Update STA README.md (#206) 2025-02-20 22:26:26 -08:00
Yongqi ChenandPeiyuan Zhang e768b5ec5b Update readme (#202)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-20 13:16:25 -08:00
Zhang Peiyuan 59ec42f40e [FIX] Make STA optinal (#204) 2025-02-20 13:09:50 -08:00
rlsu9 5ae5b247b3 [FIX] fix isort format (#203) 2025-02-20 12:20:20 -08:00
ead6c62be4 [Feat] Add STA for StepVideo (#200)
Co-authored-by: rlsu9 <r3su@ucsd.edu>
Co-authored-by: BrianChen1129 <yongqich@umich.edu>
2025-02-20 11:33:58 -08:00
Yongqi ChenandPeiyuan Zhang 6805eaa06c [bug]: fix ori hunyuan inference issue (#199)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-19 14:18:15 -08:00
Zhang Peiyuan c39a15551c Update typo (#198) 2025-02-18 19:34:45 -08:00
Zhang Peiyuan e6dda263b0 Update Cite (#195) 2025-02-18 21:01:46 -05:00
Zhang Peiyuan f9482d113c update env (#194) 2025-02-18 20:45:08 -05:00
rlsu9 a3ec969397 [feat]: fix readme demo and add video to readme (#191) 2025-02-18 17:46:32 -05:00
Yongqi ChenandPeiyuan Zhang 76a12cc8a1 Infer sta tea with torch.compile (#190)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-18 11:29:36 -08:00
Yongqi ChenandPeiyuan Zhang ac490399c6 fix kernel issue (#185)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-16 21:35:56 -08:00
Yongqi ChenandPeiyuan Zhang 9ea39cee57 Add STA and teacache forward (#184)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-15 16:22:01 -08:00
Zhang Peiyuanandrlsu9 52e6e612a2 add sliding tile attn (#182)
Co-authored-by: rlsu9 <r3su@ucsd.edu>
2025-02-15 15:44:34 -08:00
195 changed files with 962550 additions and 3688 deletions
+70
View File
@@ -0,0 +1,70 @@
name: Publish FastVideo to PyPI on Version Change
on:
push:
branches:
- main
paths:
- 'pyproject.toml' # Trigger when pyproject.toml changes
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v3
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
# Get current commit's version
NEW_VERSION=$(grep -oP 'version\s*=\s*"\K[^"]+' pyproject.toml)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP 'version\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build-publish-main:
needs: check-version-change
if: needs.check-version-change.outputs.version-changed == 'true'
runs-on: ubuntu-latest
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- name: Checkout code
uses: actions/checkout@v3
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.10'
- name: Install build dependencies
run: |
python -m pip install --upgrade pip
pip install build twine wheel
- name: Build package
run: |
python -m build
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: dist/
+221
View File
@@ -0,0 +1,221 @@
name: Publish Sliding Tile Attention Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v3
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
cd csrc/sliding_tile_attention
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build_wheels:
name: Build Wheel
needs: check-version-change
if: needs.check-version-change.outputs.version-changed == 'true'
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12', '3.13']
torch-version: ['2.5.1', '2.6.0']
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Install CUDA ${{ matrix.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.cuda-version }}
linux-local-args: '["--toolkit"]'
method: 'network'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-${{ matrix.cuda-version }}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch ${{ matrix.torch-version }}+cu${{ matrix.cuda-version }}
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==${{ matrix.torch-version }} --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build wheel
run: |
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/sliding_tile_attention
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
# Get the correct version format
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
# Rename with version information
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
- name: Upload wheel artifact
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/sliding_tile_attention/dist/*.whl
retention-days: 90
publish_package:
name: Publish package
needs: [build_wheels]
if: needs.check-version-change.outputs.version-changed == 'true'
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Install CUDA 12.4.1
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: 12.4.1
linux-local-args: '["--toolkit"]'
method: 'network'
sub-packages: '["nvcc"]'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-12.4.1
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch 2.5.1+cu12.4.1
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build source distribution
run: |
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/sliding_tile_attention/dist/
+4 -1
View File
@@ -23,8 +23,11 @@ jobs:
python -m pip install --upgrade pip setuptools wheel
pip install torch
pip install packaging ninja
# remove st-attn dependency because no cuda environment
sed -i '/st_attn/d' pyproject.toml
pip install -e .
pip install pytest
- name: Run Pytest
run: |
pytest
pytest --ignore csrc/sliding_tile_attention/test
+8 -25
View File
@@ -1,4 +1,3 @@
ucf101_stride4x4x4
__pycache__
*.mp4
.ipynb_checkpoints
@@ -8,10 +7,8 @@ results/
build/
fastvideo.egg-info/
wandb/
.idea
*.ipynb
*.jpg
*.mp3
*.safetensors
*.mp4
*.png
@@ -20,28 +17,6 @@ wandb/
*.pt
cache_dir/
wandb/
sample_video*
sample_image*
512*
720*
1024*
debug*
private*
caption*
*deepspeed*
revised*
129f*
all*
read*
YSH*
*pick*
*ysh*
hw*
257f*
513f*
taming*
221hw*
65x512x512
runs/
samples/
*validation/
@@ -51,3 +26,11 @@ outputs_video
sbatch.sh
*.out
env
dist/
*.o
**/build/
**.egg-info
**.pyc
**.egg
**.txt
**.json
+3
View File
@@ -0,0 +1,3 @@
[submodule "csrc/sliding_tile_attention/tk"]
path = csrc/sliding_tile_attention/tk
url = https://github.com/HazyResearch/ThunderKittens.git
+1 -15
View File
@@ -184,18 +184,4 @@
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [2023] Lightning AI
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
identification within third-party archives.
+60 -17
View File
@@ -4,16 +4,22 @@
FastVideo is a lightweight framework for accelerating large video diffusion models.
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
<p align="center">
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🎮 <a href="https://discord.gg/REBzDQTWWt" target="_blank"> Discord </a> | 🕹️ <a href="https://replicate.com/lucataco/fast-hunyuan-video" target="_blank"> Replicate </a>
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> Slack </a>
</p>
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
FastVideo currently offers: (with more to come)
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
@@ -22,33 +28,46 @@ FastVideo currently offers: (with more to come)
Dev in progress and highly experimental.
## 🎥 More Demos
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
## Change Log
- ```2025/02/20```: FastVideo now supports STA on [StepVideo](https://github.com/stepfun-ai/Step-Video-T2V) with 3.4X speedup!
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
- ```2024/12/17```: `FastVideo` v1.0 is released.
## 🔧 Installation
The code is tested on Python 3.10.0, CUDA 12.1 and H100.
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
```
./env_setup.sh fastvideo
```
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.
## 🚀 Inference
### Inference StepVideo with Sliding Tile Attention
First, download the model:
```
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
```
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
```bash
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
```
### Inference HunyuanVideo with Sliding Tile Attention
First, download the model:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
```
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
```bash
sh scripts/inference/inference_hunyuan_STA.sh
```
### Video Demos using STA + Teacache
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
### Inference FastHunyuan on single RTX4090
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
@@ -176,7 +195,7 @@ For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` opt
## 🤝 Contributing
We welcome all contributions. Please run `bash format.sh` before submitting a pull request.
We welcome all contributions. Please run `bash format.sh --all` before submitting a pull request.
## 🔧 Testing
Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence parallel pipelines. We recommend adding corresponding test cases in the `test` folder to support your contribution.
@@ -185,3 +204,27 @@ Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence p
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
We thank MBZUAI and Anyscale for their support throughout this project.
## Citation
If you use FastVideo for your research, please cite our paper:
```bibtex
@misc{zhang2025fastvideogenerationsliding,
title={Fast Video Generation with Sliding Tile Attention},
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
year={2025},
eprint={2502.04507},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.04507},
}
@misc{ding2025efficientvditefficientvideodiffusion,
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
year={2025},
eprint={2502.06155},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.06155},
}
```
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 751 KiB

+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config.py
+68
View File
@@ -0,0 +1,68 @@
# Sliding Tile Atteniton Kernel
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
Install STA:
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
git submodule update --init --recursive
python setup.py install
```
## Usage
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
# a tile is a cube of size (6, 8, 8)
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
# text_length: int ranging from 0 to 256
# If your attention contains text token (Hunyuan)
out = sliding_tile_attention(q, k, v, window_size, text_length)
# If your attention does not contain text token (StepVideo)
out = sliding_tile_attention(q, k, v, window_size, 0, False)
```
## Test
```bash
python test/test_sta.py
```
## How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
## Why is STA Fast?
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
STA removes mixed blocks.
<div align="center">
<img src=../../assets/sliding_tile_attn_map.png width="80%"/>
</div>
## Acknowledgement
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
+15
View File
@@ -0,0 +1,15 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'attn': {
'source_files': {
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['attn']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
+76
View File
@@ -0,0 +1,76 @@
import os
import subprocess
from config import kernels, sources, target
from setuptools import find_packages, setup
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',
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
torch_include = subprocess.check_output([
'python', '-c',
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip()
print('st_attn root:', tk_root)
print('Python include:', python_include)
print('Torch include directories:', torch_include)
# CUDA flags
cuda_flags = [
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
] + torch_include.split()
cpp_flags = ['-std=c++20', '-O3']
if target == 'h100':
cuda_flags.append('-DKITTENS_HOPPER')
cuda_flags.append('-arch=sm_90a')
else:
raise ValueError(f'Target {target} not supported')
source_files = ['st_attn.cpp']
for k in kernels:
if target not in sources[k]['source_files']:
raise KeyError(f'Target {target} not found in source files for kernel {k}')
if isinstance(sources[k]['source_files'][target], list):
source_files.extend(sources[k]['source_files'][target])
else:
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,
packages=find_packages(),
ext_modules=[
CUDAExtension('st_attn_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
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"])
+24
View File
@@ -0,0 +1,24 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
@@ -0,0 +1,35 @@
import math
import torch
from st_attn_cuda import sta_fwd
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True):
seq_length = q_all.shape[2]
if has_text:
assert q_all.shape[
2] == 115456, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
target_size = math.ceil(seq_length / 384) * 384
pad_size = target_size - seq_length
if pad_size > 0:
q_all = torch.cat([q_all, q_all[:, :, -pad_size:]], dim=2)
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
else:
assert q_all.shape[2] == 82944
hidden_states = torch.empty_like(q_all)
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
for batch in range(q_all.shape[0]):
q_head, k_head, v_head, o_head = (q_all[batch:batch + 1, head_index:head_index + 1],
k_all[batch:batch + 1,
head_index:head_index + 1], v_all[batch:batch + 1,
head_index:head_index + 1],
hidden_states[batch:batch + 1, head_index:head_index + 1])
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True)
return hidden_states[:, :, :seq_length]
@@ -0,0 +1,687 @@
// # Define TORCH_COMPILE macro
#include "kittens.cuh"
#include <cooperative_groups.h>
#include <iostream>
#include <stdio.h>
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
#define ABS(x) ((x) < 0 ? -(x) : (x))
constexpr int CONSUMER_WARPGROUPS = (3);
constexpr int PRODUCER_WARPGROUPS = (1);
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
using namespace kittens;
namespace cg = cooperative_groups;
template<int D> struct fwd_attend_ker_tile_dims {};
template<> struct fwd_attend_ker_tile_dims<64> {
constexpr static int tile_width = (64);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (4);
};
template<> struct fwd_attend_ker_tile_dims<128> {
constexpr static int tile_width = (128);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (2);
};
template<int D> struct fwd_globals {
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
q_gl q;
k_gl k;
v_gl v;
l_gl l;
o_gl o;
const int N;
const int text_L;
const int hr;
};
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
extern __shared__ int __shm[];
tma_swizzle_allocator al((int*)&__shm[0]);
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
using K = fwd_attend_ker_tile_dims<D>;
using q_tile = st_bf<K::qo_height, K::tile_width>;
using k_tile = st_bf<K::kv_height, K::tile_width>;
using v_tile = st_bf<K::kv_height, K::tile_width>;
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
using o_tile = st_bf<K::qo_height, K::tile_width>;
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
int img_kv_blocks;
int kv_blocks = g.N / (K::kv_height);
if constexpr (text_kv) {
img_kv_blocks = kv_blocks - 3;
} else {
img_kv_blocks = kv_blocks;
}
int kv_head_idx = blockIdx.y / g.hr;
int seq_idx;
if constexpr (text_q) {
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
} else {
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
}
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
if (threadIdx.x == 0) {
init_semaphore(qsmem_semaphore, 0, 1);
for(int j = 0; j < K::stages; j++) {
init_semaphore(k_smem_arrived[j], 0, 1);
init_semaphore(v_smem_arrived[j], 0, 1);
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
}
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
}
if constexpr (text_q){
for (int j = 0; j < K::stages - 1; j++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
int count = 0;
int j = 0;
while (count < K::stages - 1) {
int kt = j / 3 / (CH * CW);
int kh = (j / 3) % (CH * CW) / CW;
int kw = (j / 3) % CW;
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
if (mask){
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
count += 1;
}
j += 1;
}
}
}
__syncthreads();
int pipe_idx = K::stages - 1;
if(warpgroupid == NUM_WARPGROUPS-1) {
warpgroup::decrease_registers<32>();
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
}
else { kv_iters = kv_blocks-2;}
if(warpid == NUM_WORKERS-4) {
if constexpr (text_q){
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
int k_t_min = CLAMP(qt-DT, 0, CT-1);
int k_t_max = CLAMP(qt+DT, 0, CT-1);
int k_h_min = CLAMP(qh-DH, 0, CH-1);
int k_h_max = CLAMP(qh+DH, 0, CH-1);
int k_w_min = CLAMP(qw-DW, 0, CW-1);
int k_w_max = CLAMP(qw+DW, 0, CW-1);
int count = 0;
for (int kt = k_t_min; kt <= k_t_max; kt++) {
for (int kh = k_h_min; kh <= k_h_max; kh++) {
for (int kw = k_w_min; kw <= k_w_max; kw++) {
for (int j = 0; j <= 2; j++){
if (count >= K::stages - 1) {
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
} else {
count += 1;
}
}
}
}
}
// for text
for (int index = img_kv_blocks; index < kv_blocks; index++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
}
}
}
}
else {
warpgroup::increase_registers<160>();
rt_fl<16, K::kv_height> att_block;
rt_bf<16, K::kv_height> att_block_mma;
rt_fl<16, K::tile_width> o_reg;
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
neg_infty(max_vec);
zero(norm_vec);
zero(o_reg);
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
kv_iters = (kv_iters/8);
}
else if constexpr (text_q){
// the last three kv blocks are for text, we process them separately
kv_iters = img_kv_blocks - 1;
} else {
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
// the last three kv blocks are for text, we process them separately
if constexpr(text_kv) {
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
// apply non-pad mask
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
}
div_row(o_reg, o_reg, norm_vec);
warpgroup::store(o_smem[warpgroupid], o_reg);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
}
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
log(norm_vec, norm_vec);
add(norm_vec, norm_vec, max_vec_scaled);
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
else { mul(norm_vec, norm_vec, -11.313708499f); }
warpgroup::store(l_smem[warpgroupid], norm_vec);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
}
tma::store_async_wait();
}
}
#include "pyutils/torch_helpers.cuh"
#include <ATen/cuda/CUDAContext.h>
#include <iostream>
torch::Tensor
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
CHECK_INPUT(v);
auto batch = q.size(0);
auto seq_len = q.size(2);
auto head_dim = q.size(3);
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
auto hr = qo_heads / kv_heads;
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<128>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4-2), qo_heads, batch);
dim3 grid_text(2, qo_heads, batch);
if (!process_text) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 1, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true,1, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==5 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==5 && kernel_h_size == 3 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 0, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true,2, 0, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
} else {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
}
} else {
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
}
return o;
cudaDeviceSynchronize();
}
+151
View File
@@ -0,0 +1,151 @@
import os
from collections import defaultdict
import matplotlib.pyplot as plt
import numpy as np
import torch
from st_attn import sliding_tile_attention
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
assert mode in ["fwd", "bwd", "fwd_bwd"]
f = 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
def efficiency(flop, time):
flop = flop / 1e12
time = time / 1e6
return flop / time
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
# Prepare for timing forward pass
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
torch.cuda.empty_cache()
torch.cuda.synchronize()
# Warmup for forward pass
for _ in range(10):
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
# Time the forward pass
for i in range(10):
start_events_fwd[i].record()
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
end_events_fwd[i].record()
torch.cuda.synchronize()
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
time_us_fwd = np.mean(times_fwd) * 1000
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
print("-" * 60)
# torch.cuda.empty_cache()
# torch.cuda.synchronize()
# # Prepare for timing backward pass
# start_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# end_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# # Warmup for backward pass
# for _ in range(10):
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# # Time the backward pass
# for i in range(10):
# start_events_bwd[i].record()
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# end_events_bwd[i].record()
# torch.cuda.synchronize()
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
# time_us_bwd = np.mean(times_bwd) * 1000
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
print("=" * 60)
torch.cuda.empty_cache()
torch.cuda.synchronize()
return results
def plot_results(results):
os.makedirs('benchmark_results', exist_ok=True)
for mode in ['fwd', 'bwd']:
for (D, causal), values in results[mode].items():
seq_lens = [x[0] for x in values]
tflops = [x[1] for x in values]
plt.figure(figsize=(10, 6))
bars = plt.bar(range(len(seq_lens)), tflops, tick_label=seq_lens)
plt.xlabel('Sequence Length')
plt.ylabel('TFLOPS')
plt.title(f'{mode.upper()} Pass - Head Dim: {D}, Causal: {causal}')
plt.grid(True)
# Adding the numerical y value on top of each bar
for bar in bars:
yval = bar.get_height()
plt.text(bar.get_x() + bar.get_width() / 2, yval, round(yval, 2), ha='center', va='bottom')
filename = f'benchmark_results/{mode}_D{D}_causal{causal}.png'
plt.savefig(filename)
plt.close()
# Example list of configurations to test
configurations = [
(2, 24, 82944, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
# (16, 16, 768*8, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768, 128, True),
# (16, 16, 768*2, 128, True),
# (16, 16, 768*4, 128, True),
# (16, 16, 768*8, 128, True),
# (16, 16, 768*16, 128, True),
# (16, 32, 768, 64, False),
# (16, 32, 768*2, 64, False),
# (16, 32, 768*4, 64, False),
# (16, 32, 768*8, 64, False),
# (16, 32, 768*16, 64, False),
# (16, 32, 768, 64, True),
# (16, 32, 768*2, 64, True),
# (16, 32, 768*4, 64, True),
# (16, 32, 768*8, 64, True),
# (16, 32, 768*16, 64, True),
]
results = benchmark_attention(configurations)
# plot_results(results)
@@ -0,0 +1,71 @@
from typing import Tuple
import torch
from torch import BoolTensor, IntTensor
from torch.nn.attention.flex_attention import create_block_mask
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
torch._inductor.config.realize_opcount_threshold = 100
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
"""Generates a 3D NATTEN attention mask with a given kernel size.
Args:
canvas_t: The time dimension of the canvas.
canvas_h: The height of the canvas.
canvas_w: The width of the canvas.
kernel_t: The time dimension of the kernel.
kernel_h: The height of the kernel.
kernel_w: The width of the kernel.
"""
canvas_t, canvas_h, canvas_w = canvas_twh
kernel_t, kernel_h, kernel_w = kernel_twh
tile_t_size, tile_h_size, tile_w_size = tile_twh
total_tile_size = tile_t_size * tile_h_size * tile_w_size
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
img_seq_len = canvas_t * canvas_h * canvas_w
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
tile_id = idx // total_tile_size
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
tile_w = tile_id % canvas_tile_w
return tile_t, tile_h, tile_w
def sta_mask_3d(
b: IntTensor,
h: IntTensor,
q_idx: IntTensor,
kv_idx: IntTensor,
) -> BoolTensor:
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
# kernel nominally attempts to center itself on the query, but kernel center
# is clamped to a fixed distance (kernel half-length) from the canvas edge
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
return sta_mask_3d
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
img_seq_len = img_size[0] * img_size[1] * img_size[2]
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
mask = create_block_mask(image_mask,
B=None,
H=None,
Q_LEN=img_seq_len + text_max_len,
KV_LEN=img_seq_len + text_max_len,
device=device,
_compile=True)
return mask
@@ -0,0 +1,96 @@
import torch
from flex_sta_ref import get_sliding_tile_attention_mask
from st_attn import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
return o
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mode='all'):
results = {
'TK vs FLEX': {
'sum_diff': 0,
'sum_abs': 0,
'max_diff': 0
},
}
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
from tqdm import tqdm
for kernel_size in tqdm(kernel_size_ls):
for _ in range(num_iterations):
torch.manual_seed(0)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
pt_o = flex_test(Q, K, V, kernel_size)
diff = pt_o - tk_o
abs_diff = torch.abs(diff)
results['TK vs FLEX']['sum_diff'] += torch.sum(abs_diff).item()
results['TK vs FLEX']['max_diff'] = max(results['TK vs FLEX']['max_diff'], torch.max(abs_diff).item())
torch.cuda.empty_cache()
print("kernel_size", kernel_size)
print("max_diff", torch.max(abs_diff).item())
print(
"avg_diff",
torch.sum(abs_diff).item() / (b * h * n * d *
(1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)))
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else
3 if error_mode == 'backward' else 4) * len(kernel_size_ls)
for name, data in results.items():
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
def generate_error_graphs(b, h, d, causal, mean, std, error_mode='all'):
seq_lengths = [82944]
tk_avg_errors, tk_max_errors = [], []
for n in tqdm(seq_lengths, desc="Generating error data"):
results = check_correctness(b, h, n, d, causal, mean, std, error_mode=error_mode)
tk_avg_errors.append(results['TK vs FLEX']['avg_diff'])
tk_max_errors.append(results['TK vs FLEX']['max_diff'])
# Example usage
b, h, d = 2, 24, 128
causal = False
mean = 1e-1
std = 10
for mode in ['output']:
generate_error_graphs(b, h, d, causal, mean, std, error_mode=mode)
print("Error graphs generated and saved for all modes.")
+7 -22
View File
@@ -23,9 +23,7 @@ def init_args():
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--scheduler_type",
type=str,
default="pcm_linear_quadratic")
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
parser.add_argument("--shift", type=float, default=8.0)
parser.add_argument("--num_euler_timesteps", type=int, default=50)
@@ -50,15 +48,11 @@ def load_model(args):
)
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(
args.transformer_path)
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/")
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(args.model_path,
transformer=transformer,
scheduler=scheduler)
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
# pipe.to(device)
# if args.cpu_offload:
@@ -137,11 +131,7 @@ with gr.Blocks() as demo:
step=32,
value=args.height,
)
width = gr.Slider(label="Width",
minimum=256,
maximum=1024,
step=32,
value=args.width)
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
with gr.Row():
num_frames = gr.Slider(
@@ -164,8 +154,7 @@ with gr.Blocks() as demo:
)
with gr.Row():
use_negative_prompt = gr.Checkbox(label="Use negative prompt",
value=False)
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
@@ -173,11 +162,7 @@ with gr.Blocks() as demo:
visible=False,
)
seed = gr.Slider(label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=args.seed)
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
+15
View File
@@ -0,0 +1,15 @@
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
+3 -1
View File
@@ -1,12 +1,14 @@
#!/bin/bash
# install torch
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu124
# install FA2 and diffusers
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
pip install -r requirements-lint.txt
pip install -r requirements.txt
# install fastvideo
pip install -e .
@@ -27,8 +27,7 @@ class T5dataset(Dataset):
self.vae_debug = vae_debug
with open(self.json_path, "r") as f:
train_dataset = json.load(f)
self.train_dataset = sorted(train_dataset,
key=lambda x: x["latent_path"])
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
def __getitem__(self, idx):
caption = self.train_dataset[idx]["caption"]
@@ -36,17 +35,13 @@ class T5dataset(Dataset):
length = self.train_dataset[idx]["length"]
if self.vae_debug:
latents = torch.load(
os.path.join(args.output_dir, "latent",
self.train_dataset[idx]["latent_path"]),
os.path.join(args.output_dir, "latent", self.train_dataset[idx]["latent_path"]),
map_location="cpu",
)
else:
latents = []
return dict(caption=caption,
latents=latents,
filename=filename,
length=length)
return dict(caption=caption, latents=latents, filename=filename, length=length)
def __len__(self):
return len(self.train_dataset)
@@ -60,31 +55,21 @@ def main(args):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
videoprocessor = VideoProcessor(vae_scale_factor=8)
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"),
exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
latents_json_path = os.path.join(args.output_dir,
"videos2caption_temp.json")
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
train_dataset = T5dataset(latents_json_path, args.vae_debug)
text_encoder = load_text_encoder(args.model_type,
args.model_path,
device=device)
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
vae.enable_tiling()
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
@@ -96,26 +81,19 @@ def main(args):
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt=data["caption"], )
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(prompt=data["caption"], )
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device),
return_dict=False)[0]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.postprocess_video(video)
for idx, video_name in enumerate(data["filename"]):
prompt_embed_path = os.path.join(args.output_dir,
"prompt_embed",
video_name + ".pt")
video_path = os.path.join(args.output_dir, "video",
video_name + ".mp4")
prompt_attention_mask_path = os.path.join(
args.output_dir, "prompt_attention_mask",
video_name + ".pt")
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
video_name + ".pt")
# save latent
torch.save(prompt_embeds[idx], prompt_embed_path)
torch.save(prompt_attention_mask[idx],
prompt_attention_mask_path)
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
print(f"sample {video_name} saved")
if args.vae_debug:
export_to_video(video[idx], video_path, fps=fps)
@@ -133,8 +111,7 @@ def main(args):
if local_rank == 0:
# os.remove(latents_json_path)
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption.json"),
"w") as f:
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
@@ -148,8 +125,7 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=1,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
@@ -157,16 +133,13 @@ if __name__ == "__main__":
default=1,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--text_encoder_name",
type=str,
default="google/t5-v1_1-xxl")
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument("--vae_debug", action="store_true")
args = parser.parse_args()
@@ -20,10 +20,7 @@ def main(args):
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size, "local rank", local_rank)
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
@@ -31,14 +28,10 @@ def main(args):
num_workers=args.dataloader_num_workers,
)
encoder_device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu")
encoder_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
@@ -48,12 +41,10 @@ def main(args):
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
latents = vae.encode(data["pixel_values"].to(
encoder_device))["latent_dist"].sample()
latents = vae.encode(data["pixel_values"].to(encoder_device))["latent_dist"].sample()
for idx, video_path in enumerate(data["path"]):
video_name = os.path.basename(video_path).split(".")[0]
latent_path = os.path.join(args.output_dir, "latent",
video_name + ".pt")
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
torch.save(latents[idx].to(torch.bfloat16), latent_path)
item = {}
item["length"] = latents[idx].shape[1]
@@ -67,8 +58,7 @@ def main(args):
dist.all_gather_object(gathered_data, local_data)
if local_rank == 0:
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption_temp.json"),
"w") as f:
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
@@ -83,8 +73,7 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=1,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
@@ -92,15 +81,10 @@ if __name__ == "__main__":
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range",
type=int,
default=2.0)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
@@ -110,25 +94,21 @@ if __name__ == "__main__":
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name",
type=str,
default="google/t5-v1_1-xxl")
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
@@ -18,14 +18,9 @@ def main(args):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
text_encoder = load_text_encoder(args.model_type,
args.model_path,
device=device)
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
# output_dir/validation/prompt_attention_mask
# output_dir/validation/prompt_embed
@@ -34,8 +29,7 @@ def main(args):
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
exist_ok=True,
)
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"),
exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True)
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
lines = file.readlines()
@@ -43,12 +37,9 @@ def main(args):
for prompt in prompts:
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt)
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(prompt)
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(args.output_dir, "validation",
"prompt_embed",
f"{file_name}.pt")
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
prompt_attention_mask_path = os.path.join(
args.output_dir,
"validation",
@@ -56,8 +47,7 @@ def main(args):
f"{file_name}.pt",
)
torch.save(prompt_embeds[0], prompt_embed_path)
torch.save(prompt_attention_mask[0],
prompt_attention_mask_path)
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
print(f"sample {file_name} saved")
@@ -71,8 +61,7 @@ if __name__ == "__main__":
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
help="The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
+5 -12
View File
@@ -3,16 +3,14 @@ from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.dataset.t2v_datasets import T2V_dataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
def getdataset(args):
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
CenterCropResizeVideo((args.max_height, args.max_width),
top_crop=True),
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
]
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
@@ -27,8 +25,7 @@ def getdataset(args):
norm_fun,
])
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name,
cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(
args,
@@ -66,8 +63,7 @@ if __name__ == "__main__":
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data":
"/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
@@ -84,10 +80,7 @@ if __name__ == "__main__":
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [
i["cap"] if isinstance(i["cap"], list) else [i["cap"]]
for i in image_data
]
caps = [i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
+7 -19
View File
@@ -20,10 +20,8 @@ class LatentDataset(Dataset):
self.datase_dir_path = os.path.dirname(json_path)
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path,
"prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
"prompt_attention_mask")
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
with open(self.json_path, "r") as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
@@ -33,16 +31,12 @@ class LatentDataset(Dataset):
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item["length"] if "length" in data_item else 1
for data_item in self.data_anno
]
self.lengths = [data_item["length"] if "length" in data_item else 1 for data_item in self.data_anno]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
prompt_attention_mask_file = self.data_anno[idx][
"prompt_attention_mask"]
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
# load
latent = torch.load(
os.path.join(self.latent_dir, latent_file),
@@ -60,8 +54,7 @@ class LatentDataset(Dataset):
weights_only=True,
)
prompt_attention_mask = torch.load(
os.path.join(self.prompt_attention_mask_dir,
prompt_attention_mask_file),
os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file),
map_location="cpu",
weights_only=True,
)
@@ -111,13 +104,8 @@ def latent_collate_function(batch):
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
num_latent_t=28)
dataloader = torch.utils.data.DataLoader(
dataset,
batch_size=2,
shuffle=False,
collate_fn=latent_collate_function)
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
+21 -46
View File
@@ -46,8 +46,7 @@ class DataSetProg(metaclass=SingletonMeta):
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(
math.ceil(len(self.elements) / float(self.num_workers)))
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
@@ -58,9 +57,7 @@ class DataSetProg(metaclass=SingletonMeta):
else:
worker_id = work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
len(self.worker_elements[worker_id])]
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
@@ -68,10 +65,7 @@ class DataSetProg(metaclass=SingletonMeta):
dataset_prog = DataSetProg()
def filter_resolution(h,
w,
max_h_div_w_ratio=17 / 16,
min_h_div_w_ratio=8 / 16):
def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
@@ -79,8 +73,7 @@ def filter_resolution(h,
class T2V_dataset(Dataset):
def __init__(self, args, transform, temporal_sample, tokenizer,
transform_topcrop):
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
@@ -109,8 +102,7 @@ class T2V_dataset(Dataset):
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
n_elements)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
@@ -137,8 +129,7 @@ class T2V_dataset(Dataset):
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
torchvision_video, _, metadata = torchvision.io.read_video(video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
@@ -178,8 +169,7 @@ class T2V_dataset(Dataset):
)
def get_image(self, idx):
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
@@ -188,15 +178,13 @@ class T2V_dataset(Dataset):
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images"
in image_data["path"] else self.transform(image)
image = (self.transform_topcrop(image) if "human_images" in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else
[image_data["cap"]])
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
@@ -250,12 +238,10 @@ class T2V_dataset(Dataset):
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None
or resolution.get("width", None) is None):
if (resolution.get("height", None) is None or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"][
"width"]
height, width = i["resolution"]["height"], i["resolution"]["width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
@@ -273,34 +259,29 @@ class T2V_dataset(Dataset):
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps * self.speed_factor
): # too long video is not suitable for this training stage (self.num_frames)
self.num_frames / self.train_fps *
self.speed_factor): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"],
frame_interval).astype(int)
frame_indices = np.arange(start_frame_idx, i["num_frames"], frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio):
if (len(frame_indices) < self.num_frames and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(
len(frame_indices))
begin_index, end_index = self.temporal_sample(len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
i["sample_num_frames"] = len(i["sample_frame_index"]) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
@@ -309,32 +290,26 @@ class T2V_dataset(Dataset):
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}")
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices):
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1,
2) # (T, H, W, C) -> (T C H W)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data):
cap_lists = []
with open(data, "r") as f:
folder_anno = [
i.strip().split(",") for i in f.readlines()
if len(i.strip()) > 0
]
folder_anno = [i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno, "r") as f:
+21 -58
View File
@@ -21,19 +21,15 @@ def center_crop_arr(pil_image, image_size):
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
"""
while min(*pil_image.size) >= 2 * image_size:
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size),
resample=Image.BOX)
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
scale = image_size / min(*pil_image.size)
pil_image = pil_image.resize(tuple(
round(x * scale) for x in pil_image.size),
resample=Image.BICUBIC)
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
arr = np.array(pil_image)
crop_y = (arr.shape[0] - image_size) // 2
crop_x = (arr.shape[1] - image_size) // 2
return Image.fromarray(arr[crop_y:crop_y + image_size,
crop_x:crop_x + image_size])
return Image.fromarray(arr[crop_y:crop_y + image_size, crop_x:crop_x + image_size])
def crop(clip, i, j, h, w):
@@ -48,9 +44,7 @@ def crop(clip, i, j, h, w):
def resize(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
return torch.nn.functional.interpolate(
clip,
size=target_size,
@@ -62,9 +56,7 @@ def resize(clip, target_size, interpolation_mode):
def resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
@@ -174,8 +166,7 @@ def normalize_video(clip):
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError("clip tensor should have data type uint8. Got %s" %
str(clip.dtype))
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
@@ -236,9 +227,7 @@ class RandomCropVideo:
th, tw = self.size
if h < th or w < tw:
raise ValueError(
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
)
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
if w == tw and h == th:
return 0, 0, h, w
@@ -312,9 +301,7 @@ class LongSideResizeVideo:
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(clip,
target_size=(h, w),
interpolation_mode=self.interpolation_mode)
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
return resize_clip
def __repr__(self) -> str:
@@ -334,8 +321,7 @@ class CenterCropResizeVideo:
interpolation_mode="bilinear",
):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
@@ -349,10 +335,7 @@ class CenterCropResizeVideo:
size is (T, C, crop_size, crop_size)
"""
# clip_center_crop = center_crop_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(clip,
self.size[0],
self.size[1],
top_crop=self.top_crop)
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
# import ipdb;ipdb.set_trace()
clip_center_crop_resize = resize(
clip_center_crop,
@@ -378,9 +361,7 @@ class UCFCenterCropVideo:
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -395,9 +376,7 @@ class UCFCenterCropVideo:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_resize = resize_scale(clip=clip,
target_size=self.size,
interpolation_mode=self.interpolation_mode)
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
@@ -417,9 +396,7 @@ class KineticsRandomCropResizeVideo:
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -428,8 +405,7 @@ class KineticsRandomCropResizeVideo:
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size,
self.interpolation_mode)
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
return clip_resize
@@ -442,9 +418,7 @@ class CenterCropVideo:
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -571,8 +545,7 @@ class DynamicSampleDuration(object):
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(
range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t_list = list(range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
@@ -587,18 +560,14 @@ if __name__ == "__main__":
from torchvision import transforms
from torchvision.utils import save_image
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi",
pts_unit="sec",
output_format="TCHW")
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW")
trans = transforms.Compose([
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5],
std=[0.5, 0.5, 0.5],
inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
target_video_len = 32
@@ -613,10 +582,7 @@ if __name__ == "__main__":
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(start_frame_ind,
end_frame_ind - 1,
target_video_len,
dtype=int)
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
@@ -627,14 +593,11 @@ if __name__ == "__main__":
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) *
255).to(dtype=torch.uint8)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi",
select_vframes_trans_int.permute(0, 2, 3, 1),
fps=8)
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
for i in range(target_video_len):
save_image(
+100 -215
View File
@@ -21,23 +21,17 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint,
save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing,
get_dit_fsdp_kwargs)
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_dit_fsdp_kwargs)
from fastvideo.utils.load import load_transformer
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
get_sequence_parallel_state,
initialize_sequence_parallel_state
)
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
@@ -59,18 +53,14 @@ def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
gradient_accumulation_steps)
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
gradient_accumulation_steps)
absolute_mean = torch.mean(
torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(
torch.abs(model_pred)) / gradient_accumulation_steps
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) / gradient_accumulation_steps)
absolute_mean = torch.mean(torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(
largest_singular_value).item()
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
@@ -118,23 +108,18 @@ def distill_one_step(
model_input = normalize_dit_input(model_type, latents)
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
index = torch.randint(0, num_euler_timesteps, (bsz, ), device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index, model_input.shape)
timesteps = (sigmas *
noise_scheduler.config.num_train_timesteps).view(-1)
timesteps = (sigmas * noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
timesteps_prev = (sigmas_prev * noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
@@ -146,15 +131,13 @@ def distill_one_step(
"return_dict": False,
}
if hunyuan_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
teacher_kwargs["guidance"] = torch.tensor([1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
model_pred = transformer(**teacher_kwargs)[0]
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
model_pred, end_index = solver.euler_style_multiphase_pred(noisy_model_input, model_pred, index, multiphase)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
@@ -177,10 +160,8 @@ def distill_one_step(
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = cond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output,
index)
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
@@ -202,33 +183,26 @@ def distill_one_step(
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
huber_c = 0.001
# loss = loss.mean()
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
loss = (torch.mean(torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) - huber_c) /
gradient_accumulation_steps)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float()**2)) *
pred_decay_weight / gradient_accumulation_steps)
pred_decay_loss = (torch.mean(torch.sqrt(model_pred.float()**2)) * pred_decay_weight /
gradient_accumulation_steps)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (torch.mean(model_pred.float()**2) *
pred_decay_weight /
gradient_accumulation_steps)
pred_decay_loss = (torch.mean(model_pred.float()**2) * pred_decay_weight / gradient_accumulation_steps)
loss += pred_decay_loss
else:
assert NotImplementedError(
"pred_decay_type is not implemented")
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(model_pred.detach().float(), model_pred_norm,
gradient_accumulation_steps)
get_norm(model_pred.detach().float(), model_pred_norm, gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
@@ -238,12 +212,9 @@ def distill_one_step(
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(ema_transformer.parameters(),
transformer.parameters()):
for p_averaged, p_model in zip(ema_transformer.parameters(), transformer.parameters()):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(),
1 - ema_decay))
p_averaged.copy_(torch.lerp(p_averaged.detach(), p_model.detach(), 1 - ema_decay))
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
optimizer.step()
@@ -306,11 +277,8 @@ def main(args):
transformer.add_adapter(transformer_lora_config)
main_print(
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
transformer,
args.fsdp_sharding_startegy,
@@ -322,12 +290,9 @@ def main(args):
if args.use_lora:
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = [
"to_k", "to_q", "to_v", "to_out.0"
]
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
transformer)
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
@@ -345,13 +310,10 @@ def main(args):
main_print("--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules, args.selective_checkpointing)
if args.use_ema:
apply_fsdp_checkpointing(ema_transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(ema_transformer, no_split_modules, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
teacher_transformer.requires_grad_(False)
@@ -359,8 +321,7 @@ def main(args):
ema_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
linear_steps = int(noise_scheduler.config.num_train_timesteps *
args.linear_range)
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
sigmas = linear_quadratic_schedule(
noise_scheduler.config.num_train_timesteps,
args.linear_quadratic_threshold,
@@ -376,8 +337,7 @@ def main(args):
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
@@ -389,8 +349,8 @@ def main(args):
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer)
transformer, optimizer, init_steps = resume_lora_optimizer(transformer, args.resume_from_lora_checkpoint,
optimizer)
main_print(f"optimizer: {optimizer}")
# todo add lr scheduler
@@ -404,8 +364,7 @@ def main(args):
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
@@ -429,42 +388,33 @@ def main(args):
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
total_batch_size = (world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(
f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
assert NotImplementedError("resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
@@ -546,37 +496,26 @@ def main(args):
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"pred_fro_norm": pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value": pred_norm["largest singular value"],
"pred_absolute_mean": pred_norm["absolute mean"],
"pred_absolute_max": pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
save_checkpoint(ema_transformer, rank, args.output_dir, step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
@@ -610,11 +549,9 @@ def main(args):
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
@@ -623,10 +560,7 @@ def main(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
parser.add_argument("--model_type", type=str, default="mochi", help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
@@ -637,8 +571,7 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
@@ -646,10 +579,7 @@ if __name__ == "__main__":
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
@@ -671,16 +601,12 @@ if __name__ == "__main__":
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
@@ -692,37 +618,31 @@ if __name__ == "__main__":
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
help=("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
help=("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
@@ -731,29 +651,25 @@ if __name__ == "__main__":
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
help="Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help=
"Initial learning rate (after the potential warmup period) to use.",
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
@@ -761,47 +677,36 @@ if __name__ == "__main__":
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
help="Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
@@ -815,14 +720,8 @@ if __name__ == "__main__":
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
@@ -830,9 +729,8 @@ if __name__ == "__main__":
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
help=('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
@@ -852,15 +750,9 @@ if __name__ == "__main__":
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
parser.add_argument("--distill_cfg", type=float, default=3.0, help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument("--scheduler_type", type=str, default="pcm", help="The scheduler type to use.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
@@ -873,16 +765,9 @@ if __name__ == "__main__":
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--use_ema",
action="store_true",
help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--weight_decay", type=float, default=0.001, help="Weight decay to apply.")
parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
+4 -10
View File
@@ -12,16 +12,12 @@ class DiscriminatorHead(nn.Module):
self.conv1 = nn.Sequential(
nn.Conv2d(input_channel, inner_channel, 1, 1, 0),
nn.GroupNorm(32, inner_channel),
nn.LeakyReLU(
inplace=True
), # use LeakyReLu instead of GELU shown in the paper to save memory
nn.LeakyReLU(inplace=True), # use LeakyReLu instead of GELU shown in the paper to save memory
)
self.conv2 = nn.Sequential(
nn.Conv2d(inner_channel, inner_channel, 1, 1, 0),
nn.GroupNorm(32, inner_channel),
nn.LeakyReLU(
inplace=True
), # use LeakyReLu instead of GELU shown in the paper to save memory
nn.LeakyReLU(inplace=True), # use LeakyReLu instead of GELU shown in the paper to save memory
)
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
@@ -53,10 +49,8 @@ class Discriminator(nn.Module):
self.num_h_per_head = num_h_per_head
self.head_num = len(adapter_channel_dims)
self.heads = nn.ModuleList([
nn.ModuleList([
DiscriminatorHead(adapter_channel)
for _ in range(self.num_h_per_head)
]) for adapter_channel in adapter_channel_dims
nn.ModuleList([DiscriminatorHead(adapter_channel) for _ in range(self.num_h_per_head)])
for adapter_channel in adapter_channel_dims
])
def forward(self, features):
+24 -54
View File
@@ -39,28 +39,21 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
):
if linear_quadratic:
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(num_train_timesteps,
linear_quadratic_threshold,
linear_steps)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
timesteps = np.linspace(1,
num_train_timesteps,
num_train_timesteps,
dtype=np.float32)[::-1].copy()
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.euler_timesteps = (np.arange(1, pcm_timesteps + 1) *
(num_train_timesteps //
pcm_timesteps)).round().astype(np.int64) - 1
(num_train_timesteps // pcm_timesteps)).round().astype(np.int64) - 1
self.sigmas = sigmas.numpy()[::-1][self.euler_timesteps]
self.sigmas = torch.from_numpy((self.sigmas[::-1].copy()))
self.timesteps = self.sigmas * num_train_timesteps
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@@ -119,9 +112,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def set_timesteps(self,
num_inference_steps: int,
device: Union[str, torch.device] = None):
def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.device] = None):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -132,19 +123,14 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
"""
self.num_inference_steps = num_inference_steps
inference_indices = np.linspace(0,
self.config.pcm_timesteps,
num=num_inference_steps,
endpoint=False)
inference_indices = np.linspace(0, self.config.pcm_timesteps, num=num_inference_steps, endpoint=False)
inference_indices = np.floor(inference_indices).astype(np.int64)
inference_indices = torch.from_numpy(inference_indices).long()
self.sigmas_ = self.sigmas[inference_indices]
timesteps = self.sigmas_ * self.config.num_train_timesteps
self.timesteps = timesteps.to(device=device)
self.sigmas_ = torch.cat(
[self.sigmas_,
torch.zeros(1, device=self.sigmas_.device)])
self.sigmas_ = torch.cat([self.sigmas_, torch.zeros(1, device=self.sigmas_.device)])
self._step_index = None
self._begin_index = None
@@ -208,10 +194,9 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)):
raise ValueError((
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
self._init_step_index(timestep)
@@ -241,18 +226,14 @@ class EulerSolver:
def __init__(self, sigmas, timesteps=1000, euler_timesteps=50):
self.step_ratio = timesteps // euler_timesteps
self.euler_timesteps = (np.arange(1, euler_timesteps + 1) *
self.step_ratio).round().astype(np.int64) - 1
self.euler_timesteps_prev = np.asarray(
[0] + self.euler_timesteps[:-1].tolist())
self.euler_timesteps = (np.arange(1, euler_timesteps + 1) * self.step_ratio).round().astype(np.int64) - 1
self.euler_timesteps_prev = np.asarray([0] + self.euler_timesteps[:-1].tolist())
self.sigmas = sigmas[self.euler_timesteps]
self.sigmas_prev = np.asarray(
[sigmas[0]] + sigmas[self.euler_timesteps[:-1]].tolist()
) # either use sigma0 or 0
self.sigmas_prev = np.asarray([sigmas[0]] +
sigmas[self.euler_timesteps[:-1]].tolist()) # either use sigma0 or 0
self.euler_timesteps = torch.from_numpy(self.euler_timesteps).long()
self.euler_timesteps_prev = torch.from_numpy(
self.euler_timesteps_prev).long()
self.euler_timesteps_prev = torch.from_numpy(self.euler_timesteps_prev).long()
self.sigmas = torch.from_numpy(self.sigmas)
self.sigmas_prev = torch.from_numpy(self.sigmas_prev)
@@ -265,10 +246,8 @@ class EulerSolver:
return self
def euler_step(self, sample, model_pred, timestep_index):
sigma = extract_into_tensor(self.sigmas, timestep_index,
model_pred.shape)
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index,
model_pred.shape)
sigma = extract_into_tensor(self.sigmas, timestep_index, model_pred.shape)
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index, model_pred.shape)
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev
@@ -280,29 +259,20 @@ class EulerSolver:
multiphase,
is_target=False,
):
inference_indices = np.linspace(0,
len(self.euler_timesteps),
num=multiphase,
endpoint=False)
inference_indices = np.linspace(0, len(self.euler_timesteps), num=multiphase, endpoint=False)
inference_indices = np.floor(inference_indices).astype(np.int64)
inference_indices = (torch.from_numpy(inference_indices).long().to(
self.euler_timesteps.device))
expanded_timestep_index = timestep_index.unsqueeze(1).expand(
-1, inference_indices.size(0))
inference_indices = (torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device))
expanded_timestep_index = timestep_index.unsqueeze(1).expand(-1, inference_indices.size(0))
valid_indices_mask = expanded_timestep_index >= inference_indices
last_valid_index = valid_indices_mask.flip(dims=[1]).long().argmax(
dim=1)
last_valid_index = valid_indices_mask.flip(dims=[1]).long().argmax(dim=1)
last_valid_index = inference_indices.size(0) - 1 - last_valid_index
timestep_index_end = inference_indices[last_valid_index]
if is_target:
sigma = extract_into_tensor(self.sigmas_prev, timestep_index,
sample.shape)
sigma = extract_into_tensor(self.sigmas_prev, timestep_index, sample.shape)
else:
sigma = extract_into_tensor(self.sigmas, timestep_index,
sample.shape)
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index_end,
sample.shape)
sigma = extract_into_tensor(self.sigmas, timestep_index, sample.shape)
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index_end, sample.shape)
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev, timestep_index_end
+81 -178
View File
@@ -20,27 +20,20 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (
resume_lora_optimizer, resume_training_generator_discriminator,
save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing,
get_discriminator_fsdp_kwargs,
get_dit_fsdp_kwargs)
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_discriminator_fsdp_kwargs, get_dit_fsdp_kwargs)
from fastvideo.utils.load import load_transformer
from fastvideo.utils.logging_ import main_print
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
get_sequence_parallel_state,
initialize_sequence_parallel_state
)
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
@@ -83,9 +76,8 @@ def gan_d_loss(
fake_outputs = discriminator(fake_features)
real_outputs = discriminator(real_features)
for fake_output, real_output in zip(fake_outputs, real_outputs):
loss += (torch.mean(weight * torch.relu(fake_output.float() + 1)) +
torch.mean(weight * torch.relu(1 - real_output.float()))) / (
discriminator.head_num * discriminator.num_h_per_head)
loss += (torch.mean(weight * torch.relu(fake_output.float() + 1)) + torch.mean(
weight * torch.relu(1 - real_output.float()))) / (discriminator.head_num * discriminator.num_h_per_head)
return loss
@@ -111,8 +103,8 @@ def gan_g_loss(
)[1]
fake_outputs = discriminator(features, )
for fake_output in fake_outputs:
loss += torch.mean(weight * torch.relu(1 - fake_output.float())) / (
discriminator.head_num * discriminator.num_h_per_head)
loss += torch.mean(
weight * torch.relu(1 - fake_output.float())) / (discriminator.head_num * discriminator.num_h_per_head)
return loss
@@ -151,22 +143,18 @@ def distill_one_step_adv(
model_input = normalize_dit_input(model_type, latents)
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
index = torch.randint(0, num_euler_timesteps, (bsz, ), device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index, model_input.shape)
timesteps = (sigmas * noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
timesteps_prev = (sigmas_prev * noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
@@ -180,8 +168,7 @@ def distill_one_step_adv(
)[0]
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
model_pred, end_index = solver.euler_style_multiphase_pred(noisy_model_input, model_pred, index, multiphase)
# # simplified flow matching aka 0-rectified flow matching loss
# # target = model_input - noise
@@ -196,12 +183,9 @@ def distill_one_step_adv(
device=end_index.device,
)
sigmas_end = extract_into_tensor(solver.sigmas_prev, end_index,
model_input.shape)
sigmas_adv = extract_into_tensor(solver.sigmas_prev, adv_index,
model_input.shape)
timesteps_adv = (sigmas_adv *
noise_scheduler.config.num_train_timesteps).view(-1)
sigmas_end = extract_into_tensor(solver.sigmas_prev, end_index, model_input.shape)
sigmas_adv = extract_into_tensor(solver.sigmas_prev, adv_index, model_input.shape)
timesteps_adv = (sigmas_adv * noise_scheduler.config.num_train_timesteps).view(-1)
with torch.no_grad():
w = distill_cfg
@@ -225,8 +209,7 @@ def distill_one_step_adv(
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = cond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
teacher_output = 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
@@ -240,20 +223,14 @@ def distill_one_step_adv(
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
real_adv = ((1 - sigmas_adv) * target +
(sigmas_adv - sigmas_end) * torch.randn_like(target)) / (
1 - sigmas_end)
real_adv = ((1 - sigmas_adv) * target + (sigmas_adv - sigmas_end) * torch.randn_like(target)) / (1 - sigmas_end)
fake_adv = ((1 - sigmas_adv) * model_pred +
(sigmas_adv - sigmas_end) * torch.randn_like(model_pred)) / (
1 - sigmas_end)
(sigmas_adv - sigmas_end) * torch.randn_like(model_pred)) / (1 - sigmas_end)
huber_c = 0.001
g_loss = torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c)
g_loss = torch.mean(torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) - huber_c)
discriminator.requires_grad_(False)
with torch.autocast("cuda", dtype=torch.bfloat16):
g_gan_loss = adv_weight * gan_g_loss(
@@ -358,9 +335,7 @@ def main(args):
main_print(
f" Total discriminator parameters = {sum(p.numel() for p in discriminator.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
transformer,
args.fsdp_sharding_startegy,
@@ -368,18 +343,14 @@ def main(args):
args.use_cpu_offload,
args.master_weight_type,
)
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs(
args.master_weight_type)
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs(args.master_weight_type)
if args.use_lora:
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = [
"to_k", "to_q", "to_v", "to_out.0"
]
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
transformer)
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
@@ -396,18 +367,14 @@ def main(args):
main_print("--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
teacher_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
sigmas = linear_quadratic_schedule(
noise_scheduler.config.num_train_timesteps,
args.linear_quadratic_threshold)
sigmas = linear_quadratic_schedule(noise_scheduler.config.num_train_timesteps, args.linear_quadratic_threshold)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
sigmas = noise_scheduler.sigmas
@@ -418,8 +385,7 @@ def main(args):
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
@@ -439,8 +405,8 @@ def main(args):
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer)
transformer, optimizer, init_steps = resume_lora_optimizer(transformer, args.resume_from_lora_checkpoint,
optimizer)
elif args.resume_from_checkpoint:
(
transformer,
@@ -469,8 +435,7 @@ def main(args):
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
@@ -494,37 +459,29 @@ def main(args):
)
assert args.gradient_accumulation_steps == 1
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
total_batch_size = (world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(
f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
progress_bar = tqdm(
range(0, args.max_train_steps),
@@ -619,8 +576,7 @@ def main(args):
main_print(f"--> saving checkpoint at step {step}")
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
else:
# Your existing checkpoint saving code
# TODO
@@ -652,11 +608,9 @@ def main(args):
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
@@ -665,10 +619,7 @@ def main(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
parser.add_argument("--model_type", type=str, default="mochi", help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
@@ -678,8 +629,7 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
@@ -687,10 +637,7 @@ if __name__ == "__main__":
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
@@ -709,16 +656,12 @@ if __name__ == "__main__":
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
@@ -730,10 +673,9 @@ if __name__ == "__main__":
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
help=("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--shift", type=float, default=1.0)
@@ -741,27 +683,22 @@ if __name__ == "__main__":
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
help=("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
@@ -770,29 +707,25 @@ if __name__ == "__main__":
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help=
"Initial learning rate (after the potential warmup period) to use.",
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--discriminator_learning_rate",
type=float,
default=1e-5,
help=
"Initial learning rate (after the potential warmup period) to use.",
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
@@ -800,47 +733,36 @@ if __name__ == "__main__":
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
help="Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
@@ -854,24 +776,15 @@ if __name__ == "__main__":
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
help="Number of updates steps to accumulate before performing a backward/update pass.",
)
# lr_scheduler
@@ -879,9 +792,8 @@ if __name__ == "__main__":
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
help=('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
@@ -901,15 +813,9 @@ if __name__ == "__main__":
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
parser.add_argument("--distill_cfg", type=float, default=3.0, help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument("--scheduler_type", type=str, default="pcm", help="The scheduler type to use.")
parser.add_argument(
"--adv_weight",
type=float,
@@ -928,10 +834,7 @@ if __name__ == "__main__":
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--weight_decay", type=float, default=0.001, help="Weight decay to apply.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
+4 -13
View File
@@ -3,23 +3,15 @@ from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
def flash_attn_no_pad(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None):
def flash_attn_no_pad(qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None):
# adapted from https://github.com/Dao-AILab/flash-attention/blob/13403e81157ba37ca525890f2f0f2137edf75311/flash_attn/flash_attention.py#L27
batch_size = qkv.shape[0]
seqlen = qkv.shape[1]
nheads = qkv.shape[-2]
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
x, key_padding_mask)
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad,
"nnz (three h d) -> nnz three h d",
three=3,
h=nheads)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads)
output_unpad = flash_attn_varlen_qkvpacked_func(
x_unpad,
cu_seqlens,
@@ -29,8 +21,7 @@ def flash_attn_no_pad(qkv,
causal=causal,
)
output = rearrange(
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices,
batch_size, seqlen),
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size, seqlen),
"b s (h d) -> b s h d",
h=nheads,
)
@@ -32,14 +32,13 @@ from diffusers.models import AutoencoderKL
from diffusers.models.lora import adjust_lora_scale_text_encoder
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.schedulers import KarrasDiffusionSchedulers
from diffusers.utils import (USE_PEFT_BACKEND, BaseOutput, deprecate, logging,
replace_example_docstring, scale_lora_layers)
from diffusers.utils import (USE_PEFT_BACKEND, BaseOutput, deprecate, logging, replace_example_docstring,
scale_lora_layers)
from diffusers.utils.torch_utils import randn_tensor
from einops import rearrange
from fastvideo.utils.communications import all_gather
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from ...constants import PRECISION_TO_TYPE
from ...modules import HYVideoDiffusionTransformer
@@ -56,14 +55,12 @@ def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
"""
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)),
keepdim=True)
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
# rescale the results from guidance (fixes overexposure)
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
noise_cfg = (guidance_rescale * noise_pred_rescaled +
(1 - guidance_rescale) * noise_cfg)
noise_cfg = (guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg)
return noise_cfg
@@ -99,28 +96,22 @@ def retrieve_timesteps(
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError(
"Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
)
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
if timesteps is not None:
accepts_timesteps = "timesteps" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys())
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
f" timestep schedules. Please check whether you are using the correct scheduler.")
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys())
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
f" sigmas schedules. Please check whether you are using the correct scheduler.")
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
@@ -158,9 +149,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
_optional_components = ["text_encoder_2"]
_exclude_from_cpu_offload = ["transformer"]
_callback_tensor_inputs = [
"latents", "prompt_embeds", "negative_prompt_embeds"
]
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(
self,
@@ -184,8 +173,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
self.args = args
# ==========================================================================================
if (hasattr(scheduler.config, "steps_offset")
and scheduler.config.steps_offset != 1):
if (hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1):
deprecation_message = (
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
@@ -193,27 +181,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
" file")
deprecate("steps_offset!=1",
"1.0.0",
deprecation_message,
standard_warn=False)
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
new_config = dict(scheduler.config)
new_config["steps_offset"] = 1
scheduler._internal_dict = FrozenDict(new_config)
if (hasattr(scheduler.config, "clip_sample")
and scheduler.config.clip_sample is True):
if (hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True):
deprecation_message = (
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
)
deprecate("clip_sample not set",
"1.0.0",
deprecation_message,
standard_warn=False)
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file")
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
new_config = dict(scheduler.config)
new_config["clip_sample"] = False
scheduler._internal_dict = FrozenDict(new_config)
@@ -225,10 +205,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
scheduler=scheduler,
text_encoder_2=text_encoder_2,
)
self.vae_scale_factor = 2**(len(self.vae.config.block_out_channels) -
1)
self.image_processor = VaeImageProcessor(
vae_scale_factor=self.vae_scale_factor)
self.vae_scale_factor = 2**(len(self.vae.config.block_out_channels) - 1)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
def encode_prompt(
self,
@@ -296,14 +274,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if prompt_embeds is None:
# textual inversion: process multi-vector tokens if necessary
if isinstance(self, TextualInversionLoaderMixin):
prompt = self.maybe_convert_prompt(prompt,
text_encoder.tokenizer)
prompt = self.maybe_convert_prompt(prompt, text_encoder.tokenizer)
text_inputs = text_encoder.text2tokens(prompt, data_type=data_type)
if clip_skip is None:
prompt_outputs = text_encoder.encode(text_inputs,
data_type=data_type,
device=device)
prompt_outputs = text_encoder.encode(text_inputs, data_type=data_type, device=device)
prompt_embeds = prompt_outputs.hidden_state
else:
prompt_outputs = text_encoder.encode(
@@ -315,23 +290,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
# Access the `hidden_states` first, that contains a tuple of
# all the hidden states from the encoder layers. Then index into
# the tuple to access the hidden states from the desired layer.
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip +
1)]
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip + 1)]
# We also need to apply the final LayerNorm here to not mess with the
# representations. The `last_hidden_states` that we typically use for
# obtaining the final prompt representations passes through the LayerNorm
# layer.
prompt_embeds = text_encoder.model.text_model.final_layer_norm(
prompt_embeds)
prompt_embeds = text_encoder.model.text_model.final_layer_norm(prompt_embeds)
attention_mask = prompt_outputs.attention_mask
if attention_mask is not None:
attention_mask = attention_mask.to(device)
bs_embed, seq_len = attention_mask.shape
attention_mask = attention_mask.repeat(1,
num_videos_per_prompt)
attention_mask = attention_mask.view(
bs_embed * num_videos_per_prompt, seq_len)
attention_mask = attention_mask.repeat(1, num_videos_per_prompt)
attention_mask = attention_mask.view(bs_embed * num_videos_per_prompt, seq_len)
if text_encoder is not None:
prompt_embeds_dtype = text_encoder.dtype
@@ -340,21 +311,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
prompt_embeds_dtype = prompt_embeds.dtype
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype,
device=device)
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
if prompt_embeds.ndim == 2:
bs_embed, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, -1)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
else:
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, seq_len, -1)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, seq_len, -1)
return (
prompt_embeds,
@@ -365,10 +333,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
def decode_latents(self, latents, enable_tiling=True):
deprecation_message = "The decode_latents method is deprecated and will be removed in 1.0.0. Please use VaeImageProcessor.postprocess(...) instead"
deprecate("decode_latents",
"1.0.0",
deprecation_message,
standard_warn=False)
deprecate("decode_latents", "1.0.0", deprecation_message, standard_warn=False)
latents = 1 / self.vae.config.scaling_factor * latents
if enable_tiling:
@@ -409,30 +374,21 @@ class HunyuanVideoPipeline(DiffusionPipeline):
vae_ver="88-4c-sd",
):
if height % 8 != 0 or width % 8 != 0:
raise ValueError(
f"`height` and `width` have to be divisible by 8 but are {height} and {width}."
)
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
if video_length is not None:
if "884" in vae_ver:
if video_length != 1 and (video_length - 1) % 4 != 0:
raise ValueError(
f"`video_length` has to be 1 or a multiple of 4 but is {video_length}."
)
raise ValueError(f"`video_length` has to be 1 or a multiple of 4 but is {video_length}.")
elif "888" in vae_ver:
if video_length != 1 and (video_length - 1) % 8 != 0:
raise ValueError(
f"`video_length` has to be 1 or a multiple of 8 but is {video_length}."
)
raise ValueError(f"`video_length` has to be 1 or a multiple of 8 but is {video_length}.")
if callback_steps is not None and (not isinstance(callback_steps, int)
or callback_steps <= 0):
raise ValueError(
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
f" {type(callback_steps)}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs):
if callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0):
raise ValueError(f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
f" {type(callback_steps)}.")
if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
@@ -443,19 +399,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
" only forward one of the two.")
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str)
and not isinstance(prompt, list)):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
)
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined.")
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
raise ValueError(f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two.")
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
@@ -486,14 +436,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
f" size of {batch_size}. Make sure the batch size matches the length of the generators.")
if latents is None:
latents = randn_tensor(shape,
generator=generator,
device=device,
dtype=dtype)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
@@ -585,8 +531,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
negative_prompt: Optional[Union[str, List[str]]] = None,
num_videos_per_prompt: Optional[int] = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator,
List[torch.Generator]]] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
@@ -597,8 +542,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
guidance_rescale: float = 0.0,
clip_skip: Optional[int] = None,
callback_on_step_end: Optional[Union[Callable[[int, int, Dict],
None], PipelineCallback,
callback_on_step_end: Optional[Union[Callable[[int, int, Dict], None], PipelineCallback,
MultiPipelineCallbacks, ]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
vae_ver: str = "88-4c-sd",
@@ -606,6 +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,
**kwargs,
):
r"""
@@ -706,8 +651,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
"Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
)
if isinstance(callback_on_step_end,
(PipelineCallback, MultiPipelineCallbacks)):
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
# 0. Default height and width to unet
@@ -743,8 +687,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
batch_size = prompt_embeds.shape[0]
device = (torch.device(f"cuda:{dist.get_rank()}")
if dist.is_initialized() else self._execution_device)
device = (torch.device(f"cuda:{dist.get_rank()}") if dist.is_initialized() else self._execution_device)
# 3. Encode input prompt
lora_scale = (self.cross_attention_kwargs.get("scale", None)
@@ -804,15 +747,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if prompt_mask is not None:
prompt_mask = torch.cat([negative_prompt_mask, prompt_mask])
if prompt_embeds_2 is not None:
prompt_embeds_2 = torch.cat(
[negative_prompt_embeds_2, prompt_embeds_2])
prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2])
if prompt_mask_2 is not None:
prompt_mask_2 = torch.cat(
[negative_prompt_mask_2, prompt_mask_2])
prompt_mask_2 = torch.cat([negative_prompt_mask_2, prompt_mask_2])
# 4. Prepare timesteps
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.set_timesteps, {"n_tokens": n_tokens})
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(self.scheduler.set_timesteps,
{"n_tokens": n_tokens})
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -841,12 +782,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
generator,
latents,
)
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()
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
@@ -859,17 +797,24 @@ class HunyuanVideoPipeline(DiffusionPipeline):
)
target_dtype = PRECISION_TO_TYPE[self.args.precision]
autocast_enabled = (target_dtype !=
torch.float32) and not self.args.disable_autocast
autocast_enabled = (target_dtype != torch.float32) and not self.args.disable_autocast
vae_dtype = PRECISION_TO_TYPE[self.args.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not self.args.disable_autocast
vae_autocast_enabled = (vae_dtype != torch.float32) and not self.args.disable_autocast
# 7. Denoising loop
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
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):
@@ -877,38 +822,31 @@ class HunyuanVideoPipeline(DiffusionPipeline):
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)
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)
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (torch.tensor(
[embedded_guidance_scale] * latent_model_input.shape[0],
dtype=torch.float32,
device=device,
).to(target_dtype) * 1000.0 if embedded_guidance_scale
is not None else None)
).to(target_dtype) * 1000.0 if embedded_guidance_scale is not None else None)
# predict the noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
# concat prompt_embeds_2 and prompt_embeds. Mismatch fill with zeros
if prompt_embeds_2.shape[-1] != prompt_embeds.shape[-1]:
prompt_embeds_2 = F.pad(
prompt_embeds_2,
(0, prompt_embeds.shape[2] -
prompt_embeds_2.shape[1]),
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
value=0,
).unsqueeze(1)
encoder_hidden_states = torch.cat(
[prompt_embeds_2, prompt_embeds], dim=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)
latent_model_input, # [2, 16, 33, 24, 42]
latent_model_input,
encoder_hidden_states,
t_expand, # [2]
prompt_mask, # [2, 256]fpdb
t_expand,
prompt_mask,
mask_strategy=mask_strategy[i],
guidance=guidance_expand,
return_dict=False,
)[0]
@@ -916,8 +854,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
# perform guidance
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (
noise_pred_text - noise_pred_uncond)
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
if self.do_classifier_free_guidance and self.guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
@@ -928,29 +865,20 @@ class HunyuanVideoPipeline(DiffusionPipeline):
)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred,
t,
latents,
**extra_step_kwargs,
return_dict=False)[0]
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(
self, i, t, callback_kwargs)
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop(
"prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop(
"negative_prompt_embeds", negative_prompt_embeds)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
if progress_bar is not None:
progress_bar.update()
if callback is not None and i % callback_steps == 0:
@@ -970,26 +898,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
pass
else:
raise ValueError(
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}."
)
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}.")
if (hasattr(self.vae.config, "shift_factor")
and self.vae.config.shift_factor):
latents = (latents / self.vae.config.scaling_factor +
self.vae.config.shift_factor)
if (hasattr(self.vae.config, "shift_factor") and self.vae.config.shift_factor):
latents = (latents / self.vae.config.scaling_factor + self.vae.config.shift_factor)
else:
latents = latents / self.vae.config.scaling_factor
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
if enable_tiling:
self.vae.enable_tiling()
if enable_vae_sp:
self.vae.enable_parallel()
image = self.vae.decode(latents,
return_dict=False,
generator=generator)[0]
image = self.vae.decode(latents, return_dict=False, generator=generator)[0]
if expand_temporal_dim or image.shape[2] == 1:
image = image.squeeze(2)
@@ -80,17 +80,14 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] *
num_train_timesteps).to(dtype=torch.float32)
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
self._step_index = None
self._begin_index = None
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
@property
def step_index(self):
@@ -146,8 +143,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
sigmas = 1 - sigmas
self.sigmas = sigmas
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
dtype=torch.float32, device=device)
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(dtype=torch.float32, device=device)
# Reset step index
self._step_index = None
@@ -174,9 +170,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
else:
self._step_index = self._begin_index
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
@@ -216,10 +210,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)):
raise ValueError((
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
self._init_step_index(timestep)
@@ -232,9 +225,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
# upon completion increase step index by one
self._step_index += 1
+22 -57
View File
@@ -7,8 +7,7 @@ from .modules.models import HUNYUAN_VIDEO_CONFIG
def parse_args(namespace=None):
parser = argparse.ArgumentParser(
description="HunyuanVideo inference script")
parser = argparse.ArgumentParser(description="HunyuanVideo inference script")
parser = add_network_args(parser)
parser = add_extra_models_args(parser)
@@ -36,8 +35,7 @@ def add_network_args(parser: argparse.ArgumentParser):
"--latent-channels",
type=str,
default=16,
help=
"Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, "
help="Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, "
"it still needs to match the latent channels of the VAE model.",
)
group.add_argument(
@@ -45,22 +43,16 @@ def add_network_args(parser: argparse.ArgumentParser):
type=str,
default="bf16",
choices=PRECISIONS,
help=
"Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.",
help="Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.",
)
# RoPE
group.add_argument("--rope-theta",
type=int,
default=256,
help="Theta used in RoPE.")
group.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
return parser
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(
title="Extra models args, including vae, text encoders and tokenizers)"
)
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
# - VAE
group.add_argument(
@@ -104,10 +96,7 @@ def add_extra_models_args(parser: argparse.ArgumentParser):
default=4096,
help="Dimension of the text encoder hidden states.",
)
group.add_argument("--text-len",
type=int,
default=256,
help="Maximum length of the text input.")
group.add_argument("--text-len", type=int, default=256, help="Maximum length of the text input.")
group.add_argument(
"--tokenizer",
type=str,
@@ -138,8 +127,7 @@ def add_extra_models_args(parser: argparse.ArgumentParser):
group.add_argument(
"--apply-final-norm",
action="store_true",
help=
"Apply final normalization to the used text encoder hidden states.",
help="Apply final normalization to the used text encoder hidden states.",
)
# - CLIP
@@ -232,16 +220,13 @@ def add_inference_args(parser: argparse.ArgumentParser):
"--model-base",
type=str,
default="ckpts",
help=
"Root path of all the models, including t2v models and extra models.",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--dit-weight",
type=str,
default=
"ckpts/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
help=
"Path to the HunyuanVideo model. If None, search the model in the args.model_root."
default="ckpts/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
help="Path to the HunyuanVideo model. If None, search the model in the args.model_root."
"1. If it is a file, load the model directly."
"2. If it is a directory, search the model in the directory. Support two types of models: "
"1) named `pytorch_model_*.pt`"
@@ -252,15 +237,13 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=str,
default="540p",
choices=["540p", "720p"],
help=
"Root path of all the models, including t2v models and extra models.",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--load-key",
type=str,
default="module",
help=
"Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
group.add_argument(
"--use-cpu-offload",
@@ -284,8 +267,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
group.add_argument(
"--disable-autocast",
action="store_true",
help=
"Disable autocast for denoising loop and vae decoding in pipeline sampling.",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
group.add_argument(
"--save-path",
@@ -317,8 +299,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=int,
nargs="+",
default=(720, 1280),
help=
"Video size for training. If a single value is provided, it will be used for both height "
help="Video size for training. If a single value is provided, it will be used for both height "
"and width. If two values are provided, they will be used for height and width "
"respectively.",
)
@@ -326,8 +307,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
"--video-length",
type=int,
default=129,
help=
"How many frames to sample from a video. if using 3d vae, the number should be 4n+1",
help="How many frames to sample from a video. if using 3d vae, the number should be 4n+1",
)
# --- prompt ---
group.add_argument(
@@ -341,26 +321,16 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=str,
default="auto",
choices=["file", "random", "fixed", "auto"],
help=
"Seed type for evaluation. If file, use the seed from the CSV file. If random, generate a "
help="Seed type for evaluation. If file, use the seed from the CSV file. If random, generate a "
"random seed. If fixed, use the fixed seed given by `--seed`. If auto, `csv` will use the "
"seed column if available, otherwise use the fixed `seed` value. `prompt` will use the "
"fixed `seed` value.",
)
group.add_argument("--seed",
type=int,
default=None,
help="Seed for evaluation.")
group.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument("--neg-prompt",
type=str,
default=None,
help="Negative prompt for sampling.")
group.add_argument("--cfg-scale",
type=float,
default=1.0,
help="Classifier free guidance scale.")
group.add_argument("--neg-prompt", type=str, default=None, help="Negative prompt for sampling.")
group.add_argument("--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale.")
group.add_argument(
"--embedded-cfg-scale",
type=float,
@@ -371,8 +341,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
group.add_argument(
"--reproduce",
action="store_true",
help=
"Enable reproducibility by setting random seeds and deterministic algorithms.",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
return parser
@@ -402,14 +371,10 @@ def sanity_check_args(args):
# VAE channels
vae_pattern = r"\d{2,3}-\d{1,2}c-\w+"
if not re.match(vae_pattern, args.vae):
raise ValueError(
f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'."
)
raise ValueError(f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'.")
vae_channels = int(args.vae.split("-")[1][:-1])
if args.latent_channels is None:
args.latent_channels = vae_channels
if vae_channels != args.latent_channels:
raise ValueError(
f"Latent channels ({args.latent_channels}) must match the VAE channels ({vae_channels})."
)
raise ValueError(f"Latent channels ({args.latent_channels}) must match the VAE channels ({vae_channels}).")
return args
+46 -98
View File
@@ -7,12 +7,9 @@ import torch
from loguru import logger
from safetensors.torch import load_file as safetensors_load_file
from fastvideo.models.hunyuan.constants import (NEGATIVE_PROMPT,
PRECISION_TO_TYPE,
PROMPT_TEMPLATE)
from fastvideo.models.hunyuan.constants import NEGATIVE_PROMPT, PRECISION_TO_TYPE, PROMPT_TEMPLATE
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
from fastvideo.models.hunyuan.diffusion.schedulers import \
FlowMatchDiscreteScheduler
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
from fastvideo.models.hunyuan.modules import load_model
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.models.hunyuan.utils.data_utils import align_to
@@ -47,17 +44,12 @@ class Inference(object):
self.use_cpu_offload = use_cpu_offload
self.args = args
self.device = (device if device is not None else
"cuda" if torch.cuda.is_available() else "cpu")
self.device = (device if device is not None else "cuda" if torch.cuda.is_available() else "cpu")
self.logger = logger
self.parallel_args = parallel_args
@classmethod
def from_pretrained(cls,
pretrained_model_path,
args,
device=None,
**kwargs):
def from_pretrained(cls, pretrained_model_path, args, device=None, **kwargs):
"""
Initialize the Inference pipeline.
@@ -67,8 +59,7 @@ class Inference(object):
device (int): The device for inference. Default is 0.
"""
# ========================================================================
logger.info(
f"Got text-to-video model root path: {pretrained_model_path}")
logger.info(f"Got text-to-video model root path: {pretrained_model_path}")
# ==================== Initialize Distributed Environment ================
if nccl_info.sp_size > 1:
@@ -85,10 +76,7 @@ class Inference(object):
# =========================== Build main model ===========================
logger.info("Building model...")
factor_kwargs = {
"device": device,
"dtype": PRECISION_TO_TYPE[args.precision]
}
factor_kwargs = {"device": device, "dtype": PRECISION_TO_TYPE[args.precision]}
in_channels = args.latent_channels
out_channels = args.latent_channels
@@ -100,6 +88,8 @@ class Inference(object):
)
model = model.to(device)
model = Inference.load_state_dict(args, model, pretrained_model_path)
if args.enable_torch_compile:
model = torch.compile(model)
model.eval()
# ============================= Build extra models ========================
@@ -114,23 +104,19 @@ class Inference(object):
# Text encoder
if args.prompt_template_video is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get(
"crop_start", 0)
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get("crop_start", 0)
elif args.prompt_template is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template].get(
"crop_start", 0)
crop_start = PROMPT_TEMPLATE[args.prompt_template].get("crop_start", 0)
else:
crop_start = 0
max_length = args.text_len + crop_start
# prompt_template
prompt_template = (PROMPT_TEMPLATE[args.prompt_template]
if args.prompt_template is not None else None)
prompt_template = (PROMPT_TEMPLATE[args.prompt_template] if args.prompt_template is not None else None)
# prompt_template_video
prompt_template_video = (PROMPT_TEMPLATE[args.prompt_template_video]
if args.prompt_template_video is not None else
None)
if args.prompt_template_video is not None else None)
text_encoder = TextEncoder(
text_encoder_type=args.text_encoder,
@@ -184,23 +170,17 @@ class Inference(object):
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
bare_model = True
elif any(str(f).endswith("_model_states.pt") for f in files):
files = [
f for f in files if str(f).endswith("_model_states.pt")
]
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
logger.warning(f"Multiple model weights found in {dit_weight}, using {model_path}")
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
raise ValueError(f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file.")
else:
if dit_weight.is_dir():
files = list(dit_weight.glob("*.pt"))
@@ -210,23 +190,17 @@ class Inference(object):
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
bare_model = True
elif any(str(f).endswith("_model_states.pt") for f in files):
files = [
f for f in files if str(f).endswith("_model_states.pt")
]
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
logger.warning(f"Multiple model weights found in {dit_weight}, using {model_path}")
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
raise ValueError(f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file.")
elif dit_weight.is_file():
model_path = dit_weight
bare_model = "unknown"
@@ -241,21 +215,18 @@ class Inference(object):
state_dict = safetensors_load_file(model_path)
elif model_path.suffix == ".pt":
# Use torch for .pt files
state_dict = torch.load(model_path,
map_location=lambda storage, loc: storage)
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage)
else:
raise ValueError(f"Unsupported file format: {model_path}")
if bare_model == "unknown" and ("ema" in state_dict
or "module" in state_dict):
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
bare_model = False
if bare_model is False:
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}.")
raise KeyError(f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}.")
model.load_state_dict(state_dict, strict=True)
return model
@@ -264,13 +235,11 @@ class Inference(object):
if isinstance(size, int):
size = [size]
if not isinstance(size, (list, tuple)):
raise ValueError(
f"Size must be an integer or (height, width), got {size}.")
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
if len(size) == 1:
size = [size[0], size[0]]
if len(size) != 2:
raise ValueError(
f"Size must be an integer or (height, width), got {size}.")
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
return size
@@ -369,6 +338,7 @@ class HunyuanVideoSampler(Inference):
embedded_guidance_scale=None,
batch_size=1,
num_videos_per_prompt=1,
mask_strategy=None,
**kwargs,
):
"""
@@ -395,36 +365,22 @@ class HunyuanVideoSampler(Inference):
if isinstance(seed, torch.Tensor):
seed = seed.tolist()
if seed is None:
seeds = [
random.randint(0, 1_000_000)
for _ in range(batch_size * num_videos_per_prompt)
]
seeds = [random.randint(0, 1_000_000) for _ in range(batch_size * num_videos_per_prompt)]
elif isinstance(seed, int):
seeds = [
seed + i for _ in range(batch_size)
for i in range(num_videos_per_prompt)
]
seeds = [seed + i for _ in range(batch_size) for i in range(num_videos_per_prompt)]
elif isinstance(seed, (list, tuple)):
if len(seed) == batch_size:
seeds = [
int(seed[i]) + j for i in range(batch_size)
for j in range(num_videos_per_prompt)
]
seeds = [int(seed[i]) + j for i in range(batch_size) for j in range(num_videos_per_prompt)]
elif len(seed) == batch_size * num_videos_per_prompt:
seeds = [int(s) for s in seed]
else:
raise ValueError(
f"Length of seed must be equal to number of prompt(batch_size) or "
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}."
)
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}.")
else:
raise ValueError(
f"Seed must be an integer, a list of integers, or None, got {seed}."
)
raise ValueError(f"Seed must be an integer, a list of integers, or None, got {seed}.")
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
generator = [
torch.Generator("cpu").manual_seed(seed) for seed in seeds
]
generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
out_dict["seeds"] = seeds
# ========================================================================
@@ -435,13 +391,9 @@ class HunyuanVideoSampler(Inference):
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
)
if (video_length - 1) % 4 != 0:
raise ValueError(
f"`video_length-1` must be a multiple of 4, got {video_length}"
)
raise ValueError(f"`video_length-1` must be a multiple of 4, got {video_length}")
logger.info(
f"Input (height, width, video_length) = ({height}, {width}, {video_length})"
)
logger.info(f"Input (height, width, video_length) = ({height}, {width}, {video_length})")
target_height = align_to(height, 16)
target_width = align_to(width, 16)
@@ -453,17 +405,14 @@ class HunyuanVideoSampler(Inference):
# Arguments: prompt, new_prompt, negative_prompt
# ========================================================================
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
prompt = [prompt.strip()]
# negative prompt
if negative_prompt is None or negative_prompt == "":
negative_prompt = self.default_negative_prompt
if not isinstance(negative_prompt, str):
raise TypeError(
f"`negative_prompt` must be a string, but got {type(negative_prompt)}"
)
raise TypeError(f"`negative_prompt` must be a string, but got {type(negative_prompt)}")
negative_prompt = [negative_prompt.strip()]
# ========================================================================
@@ -477,11 +426,9 @@ class HunyuanVideoSampler(Inference):
self.pipeline.scheduler = scheduler
if "884" in self.args.vae:
latents_size = [(video_length - 1) // 4 + 1, height // 8,
width // 8]
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
elif "888" in self.args.vae:
latents_size = [(video_length - 1) // 8 + 1, height // 8,
width // 8]
latents_size = [(video_length - 1) // 8 + 1, height // 8, width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# ========================================================================
@@ -524,6 +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,
)[0]
out_dict["samples"] = samples
out_dict["prompts"] = prompt
+65 -31
View File
@@ -1,10 +1,16 @@
import torch
import torch.nn.functional as F
from einops import rearrange
try:
from st_attn import sliding_tile_attention
except ImportError:
print("Could not load Sliding Tile Attention.")
sliding_tile_attention = None
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
def attention(
@@ -21,23 +27,43 @@ def attention(
if attn_mask is not None and attn_mask.dtype != torch.bool:
attn_mask = attn_mask.bool()
x = flash_attn_no_pad(qkv,
attn_mask,
causal=causal,
dropout_p=drop_rate,
softmax_scale=None)
x = flash_attn_no_pad(qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None)
b, s, a, d = x.shape
out = x.reshape(b, s, -1)
return out
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
# 1GPU torch.Size([1, 11264, 24, 128]) tensor([ 0, 11275, 11520], device='cuda:0', dtype=torch.int32)
# 2GPU torch.Size([1, 5632, 24, 128]) tensor([ 0, 5643, 5888], device='cuda:0', dtype=torch.int32)
def tile(x, sp_size):
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
return rearrange(x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
n_t=5,
n_h=6,
n_w=10,
ts_t=6,
ts_h=8,
ts_w=8)
def untile(x, sp_size):
x = rearrange(x,
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
n_t=5,
n_h=6,
n_w=10,
ts_t=6,
ts_h=8,
ts_w=8)
return rearrange(x, "b (t sp h w) head d -> b (sp t h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=None):
query, encoder_query = q
key, encoder_key = k
value, encoder_value = v
text_length = text_mask.sum()
if get_sequence_parallel_state():
# batch_size, seq_len, attn_heads, head_dim
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
@@ -46,8 +72,7 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(
dim, nccl_info.rank_within_group * local_heads, local_heads)
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
encoder_query = shrink_head(encoder_query, dim=2)
encoder_key = shrink_head(encoder_key, dim=2)
@@ -57,28 +82,37 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# Hint: please check encoder_query.shape
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
value = torch.cat([value, encoder_value], dim=1)
# B, S, 3, H, D
qkv = torch.stack([query, key, value], dim=2)
if mask_strategy[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)
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
hidden_states = flash_attn_no_pad(qkv,
attn_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None)
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)
else:
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
value = torch.cat([value, encoder_value], dim=1)
# B, S, 3, H, D
qkv = torch.stack([query, key, value], dim=2)
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes((sequence_length, encoder_sequence_length),
dim=1)
if mask_strategy[0] is not None:
hidden_states = untile(hidden_states, nccl_info.sp_size)
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1)
if get_sequence_parallel_state():
hidden_states = all_to_all_4D(hidden_states,
scatter_dim=1,
gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states,
dim=2).contiguous()
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
hidden_states = hidden_states.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
@@ -45,8 +45,7 @@ class PatchEmbed(nn.Module):
bias=bias,
**factory_kwargs,
)
nn.init.xavier_uniform_(
self.proj.weight.view(self.proj.weight.size(0), -1))
nn.init.xavier_uniform_(self.proj.weight.view(self.proj.weight.size(0), -1))
if bias:
nn.init.zeros_(self.proj.bias)
@@ -67,12 +66,7 @@ class TextProjection(nn.Module):
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
"""
def __init__(self,
in_channels,
hidden_size,
act_layer,
dtype=None,
device=None):
def __init__(self, in_channels, hidden_size, act_layer, dtype=None, device=None):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.linear_1 = nn.Linear(
@@ -111,14 +105,12 @@ def timestep_embedding(t, dim, max_period=10000):
.. ref_link: https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
"""
half = dim // 2
freqs = torch.exp(-math.log(max_period) *
torch.arange(start=0, end=half, dtype=torch.float32) /
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) /
half).to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat(
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
@@ -145,10 +137,7 @@ class TimestepEmbedder(nn.Module):
out_size = hidden_size
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size,
hidden_size,
bias=True,
**factory_kwargs),
nn.Linear(frequency_embedding_size, hidden_size, bias=True, **factory_kwargs),
act_layer(),
nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),
)
@@ -156,8 +145,6 @@ class TimestepEmbedder(nn.Module):
nn.init.normal_(self.mlp[2].weight, std=0.02)
def forward(self, t):
t_freq = timestep_embedding(t, self.frequency_embedding_size,
self.max_period).type(
self.mlp[0].weight.dtype)
t_freq = timestep_embedding(t, self.frequency_embedding_size, self.max_period).type(self.mlp[0].weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb
+9 -35
View File
@@ -32,21 +32,13 @@ class MLP(nn.Module):
hidden_channels = hidden_channels or in_channels
bias = to_2tuple(bias)
drop_probs = to_2tuple(drop)
linear_layer = partial(nn.Conv2d,
kernel_size=1) if use_conv else nn.Linear
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
self.fc1 = linear_layer(in_channels,
hidden_channels,
bias=bias[0],
**factory_kwargs)
self.fc1 = linear_layer(in_channels, hidden_channels, bias=bias[0], **factory_kwargs)
self.act = act_layer()
self.drop1 = nn.Dropout(drop_probs[0])
self.norm = (norm_layer(hidden_channels, **factory_kwargs)
if norm_layer is not None else nn.Identity())
self.fc2 = linear_layer(hidden_channels,
out_features,
bias=bias[1],
**factory_kwargs)
self.norm = (norm_layer(hidden_channels, **factory_kwargs) if norm_layer is not None else nn.Identity())
self.fc2 = linear_layer(hidden_channels, out_features, bias=bias[1], **factory_kwargs)
self.drop2 = nn.Dropout(drop_probs[1])
def forward(self, x):
@@ -66,15 +58,9 @@ class MLPEmbedder(nn.Module):
def __init__(self, in_dim: int, hidden_dim: int, device=None, dtype=None):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.in_layer = nn.Linear(in_dim,
hidden_dim,
bias=True,
**factory_kwargs)
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True, **factory_kwargs)
self.silu = nn.SiLU()
self.out_layer = nn.Linear(hidden_dim,
hidden_dim,
bias=True,
**factory_kwargs)
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True, **factory_kwargs)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.out_layer(self.silu(self.in_layer(x)))
@@ -83,21 +69,12 @@ class MLPEmbedder(nn.Module):
class FinalLayer(nn.Module):
"""The final layer of DiT."""
def __init__(self,
hidden_size,
patch_size,
out_channels,
act_layer,
device=None,
dtype=None):
def __init__(self, hidden_size, patch_size, out_channels, act_layer, device=None, dtype=None):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
# Just use LayerNorm for the final layer
self.norm_final = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
if isinstance(patch_size, int):
self.linear = nn.Linear(
hidden_size,
@@ -117,10 +94,7 @@ class FinalLayer(nn.Module):
# Here we don't distinguish between the modulate types. Just use the simple one.
self.adaLN_modulation = nn.Sequential(
act_layer(),
nn.Linear(hidden_size,
2 * hidden_size,
bias=True,
**factory_kwargs),
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
)
# Zero-initialize the modulation
nn.init.zeros_(self.adaLN_modulation[1].weight)
+74 -158
View File
@@ -6,8 +6,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models import ModelMixin
from einops import rearrange
from fastvideo.models.hunyuan.modules.posemb_layers import \
get_nd_rotary_pos_embed
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
from fastvideo.utils.parallel_states import nccl_info
from .activation_layers import get_activation_layer
@@ -53,31 +52,17 @@ class MMDoubleStreamBlock(nn.Module):
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.img_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.img_attn_qkv = nn.Linear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
**factory_kwargs)
self.img_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.img_attn_q_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.img_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.img_attn_k_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.img_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.img_attn_proj = nn.Linear(hidden_size,
hidden_size,
bias=qkv_bias,
**factory_kwargs)
self.img_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
self.img_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.img_mlp = MLP(
hidden_size,
mlp_hidden_dim,
@@ -92,30 +77,16 @@ class MMDoubleStreamBlock(nn.Module):
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.txt_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.txt_attn_qkv = nn.Linear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
**factory_kwargs)
self.txt_attn_q_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.txt_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
self.txt_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.txt_attn_k_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.txt_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.txt_attn_proj = nn.Linear(hidden_size,
hidden_size,
bias=qkv_bias,
**factory_kwargs)
self.txt_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
self.txt_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.txt_mlp = MLP(
hidden_size,
mlp_hidden_dim,
@@ -138,6 +109,7 @@ class MMDoubleStreamBlock(nn.Module):
vec: torch.Tensor,
freqs_cis: tuple = None,
text_mask: torch.Tensor = None,
mask_strategy=None,
) -> Tuple[torch.Tensor, torch.Tensor]:
(
img_mod1_shift,
@@ -158,14 +130,9 @@ class MMDoubleStreamBlock(nn.Module):
# Prepare image for attention.
img_modulated = self.img_norm1(img)
img_modulated = modulate(img_modulated,
shift=img_mod1_shift,
scale=img_mod1_scale)
img_modulated = modulate(img_modulated, shift=img_mod1_shift, scale=img_mod1_scale)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(img_qkv,
"B L (K H D) -> K B L H D",
K=3,
H=self.heads_num)
img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
@@ -175,34 +142,23 @@ class MMDoubleStreamBlock(nn.Module):
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(
dim, nccl_info.rank_within_group * local_heads,
local_heads)
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
freqs_cis = (
shrink_head(freqs_cis[0], dim=0),
shrink_head(freqs_cis[1], dim=0),
)
img_qq, img_kk = apply_rotary_emb(img_q,
img_k,
freqs_cis,
head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), 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
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(txt_modulated,
shift=txt_mod1_shift,
scale=txt_mod1_scale)
txt_modulated = modulate(txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(txt_qkv,
"B L (K H D) -> K B L H D",
K=3,
H=self.heads_num)
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
@@ -214,6 +170,7 @@ class MMDoubleStreamBlock(nn.Module):
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
mask_strategy=mask_strategy,
)
# attention computation end
@@ -221,27 +178,18 @@ class MMDoubleStreamBlock(nn.Module):
img_attn, txt_attn = attn[:, :img.shape[1]], attn[:, img.shape[1]:]
# Calculate the img blocks.
img = img + apply_gate(self.img_attn_proj(img_attn),
gate=img_mod1_gate)
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(self.img_norm2(img),
shift=img_mod2_shift,
scale=img_mod2_scale)),
self.img_mlp(modulate(self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale)),
gate=img_mod2_gate,
)
# Calculate the txt blocks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn),
gate=txt_mod1_gate)
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
txt = txt + apply_gate(
self.txt_mlp(
modulate(self.txt_norm2(txt),
shift=txt_mod2_shift,
scale=txt_mod2_scale)),
self.txt_mlp(modulate(self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale)),
gate=txt_mod2_gate,
)
return img, txt
@@ -277,24 +225,17 @@ class MMSingleStreamBlock(nn.Module):
self.scale = qk_scale or head_dim**-0.5
# qkv and mlp_in
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + mlp_hidden_dim,
**factory_kwargs)
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs)
# proj and mlp_out
self.linear2 = nn.Linear(hidden_size + mlp_hidden_dim, hidden_size,
**factory_kwargs)
self.linear2 = nn.Linear(hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.q_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.k_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.pre_norm = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.mlp_act = get_activation_layer(mlp_act_type)()
self.modulation = ModulateDiT(
@@ -318,17 +259,13 @@ class MMSingleStreamBlock(nn.Module):
txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
text_mask: torch.Tensor = None,
mask_strategy=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)
qkv, mlp = torch.split(self.linear1(x_mod),
[3 * self.hidden_size, self.mlp_hidden_dim],
dim=-1)
qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
q, k, v = rearrange(qkv,
"B L (K H D) -> K B L H D",
K=3,
H=self.heads_num)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
@@ -336,22 +273,19 @@ class MMSingleStreamBlock(nn.Module):
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(
dim, nccl_info.rank_within_group * local_heads, local_heads)
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
freqs_cis = (shrink_head(freqs_cis[0],
dim=0), shrink_head(freqs_cis[1], dim=0))
freqs_cis = (
shrink_head(freqs_cis[0], dim=0),
shrink_head(freqs_cis[1], dim=0),
)
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_v, txt_v = v[:, :-txt_len, :, :], v[:, -txt_len:, :, :]
img_qq, img_kk = apply_rotary_emb(img_q,
img_k,
freqs_cis,
head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), 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(
@@ -361,6 +295,7 @@ class MMSingleStreamBlock(nn.Module):
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
mask_strategy=mask_strategy,
)
# attention computation end
@@ -463,19 +398,15 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.text_projection = text_projection
if hidden_size % heads_num != 0:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}"
)
raise ValueError(f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}")
pe_dim = hidden_size // heads_num
if sum(rope_dim_list) != pe_dim:
raise ValueError(
f"Got {rope_dim_list} but expected positional dim {pe_dim}")
raise ValueError(f"Got {rope_dim_list} but expected positional dim {pe_dim}")
self.hidden_size = hidden_size
self.heads_num = heads_num
# image projection
self.img_in = PatchEmbed(self.patch_size, self.in_channels,
self.hidden_size, **factory_kwargs)
self.img_in = PatchEmbed(self.patch_size, self.in_channels, self.hidden_size, **factory_kwargs)
# text projection
if self.text_projection == "linear":
@@ -494,21 +425,16 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
**factory_kwargs,
)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}")
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
# time modulation
self.time_in = TimestepEmbedder(self.hidden_size,
get_activation_layer("silu"),
**factory_kwargs)
self.time_in = TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
# text modulation
self.vector_in = MLPEmbedder(self.config.text_states_dim_2,
self.hidden_size, **factory_kwargs)
self.vector_in = MLPEmbedder(self.config.text_states_dim_2, self.hidden_size, **factory_kwargs)
# guidance modulation
self.guidance_in = (TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
self.guidance_in = (TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
if guidance_embed else None)
# double blocks
@@ -564,12 +490,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
head_dim = self.hidden_size // self.heads_num
rope_dim_list = self.rope_dim_list
if rope_dim_list is None:
rope_dim_list = [
head_dim // target_ndim for _ in range(target_ndim)
]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert (sum(rope_dim_list) == head_dim), "sum(rope_dim_list) should equal to head_dim of attention layer"
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
@@ -592,6 +514,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
mask_strategy=None,
output_features=False,
output_features_stride=8,
attention_kwargs: Optional[Dict[str, Any]] = None,
@@ -599,15 +522,14 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
guidance=None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance is None:
guidance = torch.tensor([6016.0],
device=hidden_states.device,
dtype=torch.bfloat16)
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
if mask_strategy is None:
mask_strategy = [[None] * self.heads_num for _ in range(len(self.double_blocks) + len(self.single_blocks))]
img = x = hidden_states
text_mask = encoder_attention_mask
t = timestep
txt = encoder_hidden_states[:, 1:]
text_states_2 = encoder_hidden_states[:, 0, :self.config.
text_states_dim_2]
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
_, _, ot, oh, ow = x.shape # codespell:ignore
tt, th, tw = (
ot // self.patch_size[0], # codespell:ignore
@@ -625,9 +547,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
# guidance modulation
if self.guidance_embed:
if guidance is None:
raise ValueError(
"Didn't get guidance strength for guidance distilled model."
)
raise ValueError("Didn't get guidance strength for guidance distilled model.")
# our timestep_embedding is merged into guidance_in(TimestepEmbedder)
vec = vec + self.guidance_in(guidance)
@@ -637,36 +557,33 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
if self.text_projection == "linear":
txt = self.txt_in(txt)
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t,
text_mask if self.use_attention_mask else None)
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}")
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
for _, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask]
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)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if output_features:
features_list = []
if len(self.single_blocks) > 0:
for _, block in enumerate(self.single_blocks):
for index, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
(freqs_cos, freqs_sin),
text_mask,
mask_strategy[index + len(self.double_blocks)],
]
x = block(*single_block_args)
if output_features and _ % output_features_stride == 0:
features_list.append(x[:, :img_seq_len, ...])
@@ -674,8 +591,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
img = x[:, :img_seq_len, ...]
# ---------------------------- Final layer ------------------------------
img = self.final_layer(img,
vec) # (N, T, patch_size ** 2 * out_channels)
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
img = self.unpatchify(img, tt, th, tw)
assert not return_dict, "return_dict is not supported."
@@ -704,18 +620,18 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
counts = {
"double":
sum([
sum(p.numel() for p in block.img_attn_qkv.parameters()) +
sum(p.numel() for p in block.img_attn_proj.parameters()) +
sum(p.numel() for p in block.img_mlp.parameters()) +
sum(p.numel() for p in block.txt_attn_qkv.parameters()) +
sum(p.numel() for p in block.txt_attn_proj.parameters()) +
sum(p.numel() for p in block.txt_mlp.parameters())
sum(p.numel()
for p in block.img_attn_qkv.parameters()) + sum(p.numel()
for p in block.img_attn_proj.parameters()) +
sum(p.numel() for p in block.img_mlp.parameters()) + sum(p.numel()
for p in block.txt_attn_qkv.parameters()) +
sum(p.numel() for p in block.txt_attn_proj.parameters()) + sum(p.numel()
for p in block.txt_mlp.parameters())
for block in self.double_blocks
]),
"single":
sum([
sum(p.numel() for p in block.linear1.parameters()) +
sum(p.numel() for p in block.linear2.parameters())
sum(p.numel() for p in block.linear1.parameters()) + sum(p.numel() for p in block.linear2.parameters())
for block in self.single_blocks
]),
"total":
@@ -18,10 +18,7 @@ class ModulateDiT(nn.Module):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.act = act_layer()
self.linear = nn.Linear(hidden_size,
factor * hidden_size,
bias=True,
**factory_kwargs)
self.linear = nn.Linear(hidden_size, factor * hidden_size, bias=True, **factory_kwargs)
# Zero-initialize the modulation
nn.init.zeros_(self.linear.weight)
nn.init.zeros_(self.linear.bias)
@@ -152,5 +149,4 @@ def get_norm_layer(norm_layer):
elif norm_layer == "rms":
return RMSNorm
else:
raise NotImplementedError(
f"Norm layer {norm_layer} is not implemented")
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
@@ -75,5 +75,4 @@ def get_norm_layer(norm_layer):
elif norm_layer == "rms":
return RMSNorm
else:
raise NotImplementedError(
f"Norm layer {norm_layer} is not implemented")
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
@@ -100,19 +100,13 @@ def reshape_for_broadcast(
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
else:
assert freqs_cis[0].shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [
d if i == 1 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
else:
# freqs_cis: values in complex space
@@ -121,25 +115,18 @@ def reshape_for_broadcast(
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
else:
assert freqs_cis.shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [
d if i == 1 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis.view(*shape)
def rotate_half(x):
x_real, x_imag = (x.float().reshape(*x.shape[:-1], -1,
2).unbind(-1)) # [B, S, H, D//2]
x_real, x_imag = (x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)) # [B, S, H, D//2]
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
@@ -177,15 +164,12 @@ def apply_rotary_emb(
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
else:
# view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex)
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1,
2)) # [B, S, H, D//2]
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(
xq.device) # [S, D//2] --> [1, S, 1, D//2]
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2)) # [B, S, H, D//2]
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(xq.device) # [S, D//2] --> [1, S, 1, D//2]
# (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin)
# view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real)
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1,
2)) # [B, S, H, D//2]
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2)) # [B, S, H, D//2]
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
return xq_out, xk_out
@@ -219,28 +203,21 @@ def get_nd_rotary_pos_embed(
pos_embed (torch.Tensor): [HW, D/2]
"""
grid = get_meshgrid_nd(start, *args,
dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
if isinstance(theta_rescale_factor, int) or isinstance(
theta_rescale_factor, float):
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor,
list) and len(theta_rescale_factor) == 1:
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(
rope_dim_list
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(
interpolation_factor, float):
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor,
list) and len(interpolation_factor) == 1:
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(
rope_dim_list
), "len(interpolation_factor) should equal to len(rope_dim_list)"
rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
@@ -300,8 +277,7 @@ def get_1d_rotary_pos_embed(
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor**(dim / (dim - 2))
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)
) # [D/2]
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) # [D/2]
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
if use_real:
@@ -309,6 +285,5 @@ def get_1d_rotary_pos_embed(
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
return freqs_cos, freqs_sin
else:
freqs_cis = torch.polar(torch.ones_like(freqs),
freqs) # complex64 # [S, D/2]
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
return freqs_cis
@@ -33,30 +33,16 @@ class IndividualTokenRefinerBlock(nn.Module):
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=True,
eps=1e-6,
**factory_kwargs)
self.self_attn_qkv = nn.Linear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
**factory_kwargs)
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.self_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.self_attn_q_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.self_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.self_attn_k_norm = (qk_norm_layer(
head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
self.self_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.self_attn_proj = nn.Linear(hidden_size,
hidden_size,
bias=qkv_bias,
**factory_kwargs)
self.self_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
self.norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=True,
eps=1e-6,
**factory_kwargs)
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs)
act_layer = get_activation_layer(act_type)
self.mlp = MLP(
in_channels=hidden_size,
@@ -68,10 +54,7 @@ class IndividualTokenRefinerBlock(nn.Module):
self.adaLN_modulation = nn.Sequential(
act_layer(),
nn.Linear(hidden_size,
2 * hidden_size,
bias=True,
**factory_kwargs),
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
)
# Zero-initialize the modulation
nn.init.zeros_(self.adaLN_modulation[1].weight)
@@ -80,18 +63,14 @@ class IndividualTokenRefinerBlock(nn.Module):
def forward(
self,
x: torch.Tensor,
c: torch.
Tensor, # timestep_aware_representations + context_aware_representations
c: torch.Tensor, # timestep_aware_representations + context_aware_representations
attn_mask: torch.Tensor = None,
):
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
norm_x = self.norm1(x)
qkv = self.self_attn_qkv(norm_x)
q, k, v = rearrange(qkv,
"B L (K H D) -> K B L H D",
K=3,
H=self.heads_num)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed
q = self.self_attn_q_norm(q).to(v)
k = self.self_attn_k_norm(k).to(v)
@@ -179,18 +158,13 @@ class SingleTokenRefiner(nn.Module):
self.attn_mode = attn_mode
assert self.attn_mode == "torch", "Only support 'torch' mode for token refiner."
self.input_embedder = nn.Linear(in_channels,
hidden_size,
bias=True,
**factory_kwargs)
self.input_embedder = nn.Linear(in_channels, hidden_size, bias=True, **factory_kwargs)
act_layer = get_activation_layer(act_type)
# Build timestep embedding layer
self.t_embedder = TimestepEmbedder(hidden_size, act_layer,
**factory_kwargs)
self.t_embedder = TimestepEmbedder(hidden_size, act_layer, **factory_kwargs)
# Build context embedding layer
self.c_embedder = TextProjection(in_channels, hidden_size, act_layer,
**factory_kwargs)
self.c_embedder = TextProjection(in_channels, hidden_size, act_layer, **factory_kwargs)
self.individual_token_refiner = IndividualTokenRefiner(
hidden_size=hidden_size,
@@ -217,10 +191,8 @@ class SingleTokenRefiner(nn.Module):
context_aware_representations = x.mean(dim=1)
else:
mask_float = mask.float().unsqueeze(-1) # [b, s1, 1]
context_aware_representations = (x * mask_float).sum(
dim=1) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(
context_aware_representations)
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(context_aware_representations)
c = timestep_aware_representations + context_aware_representations
x = self.input_embedder(x)
@@ -23,24 +23,20 @@ def load_text_encoder(
if text_encoder_path is None:
text_encoder_path = TEXT_ENCODER_PATH[text_encoder_type]
if logger is not None:
logger.info(
f"Loading text encoder model ({text_encoder_type}) from: {text_encoder_path}"
)
logger.info(f"Loading text encoder model ({text_encoder_type}) from: {text_encoder_path}")
if text_encoder_type == "clipL":
text_encoder = CLIPTextModel.from_pretrained(text_encoder_path)
text_encoder.final_layer_norm = text_encoder.text_model.final_layer_norm
elif text_encoder_type == "llm":
text_encoder = AutoModel.from_pretrained(text_encoder_path,
low_cpu_mem_usage=True)
text_encoder = AutoModel.from_pretrained(text_encoder_path, low_cpu_mem_usage=True)
text_encoder.final_layer_norm = text_encoder.norm
else:
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
# from_pretrained will ensure that the model is in eval mode.
if text_encoder_precision is not None:
text_encoder = text_encoder.to(
dtype=PRECISION_TO_TYPE[text_encoder_precision])
text_encoder = text_encoder.to(dtype=PRECISION_TO_TYPE[text_encoder_precision])
text_encoder.requires_grad_(False)
@@ -53,22 +49,16 @@ def load_text_encoder(
return text_encoder, text_encoder_path
def load_tokenizer(tokenizer_type,
tokenizer_path=None,
padding_side="right",
logger=None):
def load_tokenizer(tokenizer_type, tokenizer_path=None, padding_side="right", logger=None):
if tokenizer_path is None:
tokenizer_path = TOKENIZER_PATH[tokenizer_type]
if logger is not None:
logger.info(
f"Loading tokenizer ({tokenizer_type}) from: {tokenizer_path}")
logger.info(f"Loading tokenizer ({tokenizer_type}) from: {tokenizer_path}")
if tokenizer_type == "clipL":
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path,
max_length=77)
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path, max_length=77)
elif tokenizer_type == "llm":
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
padding_side=padding_side)
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path, padding_side=padding_side)
else:
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
@@ -125,16 +115,12 @@ class TextEncoder(nn.Module):
self.max_length = max_length
self.precision = text_encoder_precision
self.model_path = text_encoder_path
self.tokenizer_type = (tokenizer_type if tokenizer_type is not None
else text_encoder_type)
self.tokenizer_path = (tokenizer_path if tokenizer_path is not None
else text_encoder_path)
self.tokenizer_type = (tokenizer_type if tokenizer_type is not None else text_encoder_type)
self.tokenizer_path = (tokenizer_path if tokenizer_path is not None else text_encoder_path)
self.use_attention_mask = use_attention_mask
if prompt_template_video is not None:
assert (use_attention_mask is True
), "Attention mask is True required when training videos."
self.input_max_length = (input_max_length if input_max_length
is not None else max_length)
assert (use_attention_mask is True), "Attention mask is True required when training videos."
self.input_max_length = (input_max_length if input_max_length is not None else max_length)
self.prompt_template = prompt_template
self.prompt_template_video = prompt_template_video
self.hidden_state_skip_layer = hidden_state_skip_layer
@@ -144,10 +130,8 @@ class TextEncoder(nn.Module):
self.use_template = self.prompt_template is not None
if self.use_template:
assert (
isinstance(self.prompt_template, dict)
and "template" in self.prompt_template
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
assert (isinstance(self.prompt_template, dict) and "template" in self.prompt_template
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
assert "{}" in str(self.prompt_template["template"]), (
"`prompt_template['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template['template']}")
@@ -156,8 +140,7 @@ class TextEncoder(nn.Module):
if self.use_video_template:
if self.prompt_template_video is not None:
assert (
isinstance(self.prompt_template_video, dict)
and "template" in self.prompt_template_video
isinstance(self.prompt_template_video, dict) and "template" in self.prompt_template_video
), f"`prompt_template_video` must be a dictionary with a key 'template', got {self.prompt_template_video}"
assert "{}" in str(self.prompt_template_video["template"]), (
"`prompt_template_video['template']` must contain a placeholder `{}` for the input text, "
@@ -170,8 +153,7 @@ class TextEncoder(nn.Module):
elif "llm" in text_encoder_type or "glm" in text_encoder_type:
self.output_key = output_key or "last_hidden_state"
else:
raise ValueError(
f"Unsupported text encoder type: {text_encoder_type}")
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
self.model, self.model_path = load_text_encoder(
text_encoder_type=self.text_encoder_type,
@@ -226,10 +208,7 @@ class TextEncoder(nn.Module):
else:
raise ValueError(f"Unsupported data type: {data_type}")
if isinstance(text, (list, tuple)):
text = [
self.apply_text_to_template(one_text, prompt_template)
for one_text in text
]
text = [self.apply_text_to_template(one_text, prompt_template) for one_text in text]
if isinstance(text[0], list):
tokenize_input_type = "list"
elif isinstance(text, str):
@@ -262,8 +241,7 @@ class TextEncoder(nn.Module):
**kwargs,
)
else:
raise ValueError(
f"Unsupported tokenize_input_type: {tokenize_input_type}")
raise ValueError(f"Unsupported tokenize_input_type: {tokenize_input_type}")
def encode(
self,
@@ -291,27 +269,21 @@ class TextEncoder(nn.Module):
return_texts (bool): Whether to return the decoded texts. Defaults to False.
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask,
self.use_attention_mask)
hidden_state_skip_layer = use_default(hidden_state_skip_layer,
self.hidden_state_skip_layer)
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
hidden_state_skip_layer = use_default(hidden_state_skip_layer, self.hidden_state_skip_layer)
do_sample = use_default(do_sample, not self.reproduce)
attention_mask = (batch_encoding["attention_mask"].to(device)
if use_attention_mask else None)
attention_mask = (batch_encoding["attention_mask"].to(device) if use_attention_mask else None)
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
attention_mask=attention_mask,
output_hidden_states=output_hidden_states
or hidden_state_skip_layer is not None,
output_hidden_states=output_hidden_states or hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(
hidden_state_skip_layer + 1)]
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
# Real last hidden state already has layer norm applied. So here we only apply it
# for intermediate layers.
if hidden_state_skip_layer > 0 and self.apply_final_norm:
last_hidden_state = self.model.final_layer_norm(
last_hidden_state)
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
else:
last_hidden_state = outputs[self.output_key]
@@ -325,12 +297,10 @@ class TextEncoder(nn.Module):
raise ValueError(f"Unsupported data type: {data_type}")
if crop_start > 0:
last_hidden_state = last_hidden_state[:, crop_start:]
attention_mask = (attention_mask[:, crop_start:]
if use_attention_mask else None)
attention_mask = (attention_mask[:, crop_start:] if use_attention_mask else None)
if output_hidden_states:
return TextEncoderModelOutput(last_hidden_state, attention_mask,
outputs.hidden_states)
return TextEncoderModelOutput(last_hidden_state, attention_mask, outputs.hidden_states)
return TextEncoderModelOutput(last_hidden_state, attention_mask)
def forward(
+1 -5
View File
@@ -45,11 +45,7 @@ def safe_file(path):
return path
def save_videos_grid(videos: torch.Tensor,
path: str,
rescale=False,
n_rows=1,
fps=24):
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=1, fps=24):
"""save videos by video tensor
copy from https://github.com/guoyww/AnimateDiff/blob/e92bd5671ba62c0d774a32951453e328018b7c5b/animatediff/utils/util.py#L61
+2 -6
View File
@@ -31,8 +31,7 @@ def load_vae(
logger.info(f"Loading 3D VAE model ({vae_type}) from: {vae_path}")
config = AutoencoderKLCausal3D.load_config(vae_path)
if sample_size:
vae = AutoencoderKLCausal3D.from_config(config,
sample_size=sample_size)
vae = AutoencoderKLCausal3D.from_config(config, sample_size=sample_size)
else:
vae = AutoencoderKLCausal3D.from_config(config)
@@ -43,10 +42,7 @@ def load_vae(
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if any(k.startswith("vae.") for k in ckpt.keys()):
ckpt = {
k.replace("vae.", ""): v
for k, v in ckpt.items() if k.startswith("vae.")
}
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
vae.load_state_dict(ckpt)
spatial_compression_ratio = vae.config.spatial_compression_ratio
@@ -35,15 +35,13 @@ except ImportError:
from diffusers.loaders.single_file_model import (
FromOriginalModelMixin as FromOriginalVAEMixin, )
from diffusers.models.attention_processor import (
ADDED_KV_ATTENTION_PROCESSORS, CROSS_ATTENTION_PROCESSORS, Attention,
AttentionProcessor, AttnAddedKVProcessor, AttnProcessor)
from diffusers.models.attention_processor import (ADDED_KV_ATTENTION_PROCESSORS, CROSS_ATTENTION_PROCESSORS, Attention,
AttentionProcessor, AttnAddedKVProcessor, AttnProcessor)
from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils.accelerate_utils import apply_forward_hook
from .vae import (BaseOutput, DecoderCausal3D, DecoderOutput,
DiagonalGaussianDistribution, EncoderCausal3D)
from .vae import BaseOutput, DecoderCausal3D, DecoderOutput, DiagonalGaussianDistribution, EncoderCausal3D
@dataclass
@@ -113,12 +111,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
mid_block_add_attention=mid_block_add_attention,
)
self.quant_conv = nn.Conv3d(2 * latent_channels,
2 * latent_channels,
kernel_size=1)
self.post_quant_conv = nn.Conv3d(latent_channels,
latent_channels,
kernel_size=1)
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
self.use_slicing = False
self.use_spatial_tiling = False
@@ -130,11 +124,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
self.tile_latent_min_tsize = sample_tsize // time_compression_ratio
self.tile_sample_min_size = self.config.sample_size
sample_size = (self.config.sample_size[0] if isinstance(
self.config.sample_size,
(list, tuple)) else self.config.sample_size)
self.tile_latent_min_size = int(
sample_size / (2**(len(self.config.block_out_channels) - 1)))
sample_size = (self.config.sample_size[0] if isinstance(self.config.sample_size,
(list, tuple)) else self.config.sample_size)
self.tile_latent_min_size = int(sample_size / (2**(len(self.config.block_out_channels) - 1)))
self.tile_overlap_factor = 0.25
def _set_gradient_checkpointing(self, module, value=False):
@@ -207,12 +199,10 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
processors: Dict[str, AttentionProcessor],
):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor(
return_deprecated_lora=True)
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child,
processors)
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
return processors
@@ -244,21 +234,17 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
f" number of attention layers: {count}. Please make sure to pass {count} processor classes.")
def fn_recursive_attn_processor(name: str, module: torch.nn.Module,
processor):
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor, _remove_lora=_remove_lora)
else:
module.set_processor(processor.pop(f"{name}.processor"),
_remove_lora=_remove_lora)
module.set_processor(processor.pop(f"{name}.processor"), _remove_lora=_remove_lora)
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child,
processor)
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
@@ -268,11 +254,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
"""
Disables custom attention processors and sets the default attention implementation.
"""
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()):
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
processor = AttnAddedKVProcessor()
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()):
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
processor = AttnProcessor()
else:
raise ValueError(
@@ -282,11 +266,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
self.set_attn_processor(processor, _remove_lora=True)
@apply_forward_hook
def encode(
self,
x: torch.FloatTensor,
return_dict: bool = True
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
def encode(self,
x: torch.FloatTensor,
return_dict: bool = True) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
"""
Encode a batch of images/videos into latents.
@@ -304,9 +286,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
if self.use_temporal_tiling and x.shape[2] > self.tile_sample_min_tsize:
return self.temporal_tiled_encode(x, return_dict=return_dict)
if self.use_spatial_tiling and (
x.shape[-1] > self.tile_sample_min_size
or x.shape[-2] > self.tile_sample_min_size):
if self.use_spatial_tiling and (x.shape[-1] > self.tile_sample_min_size
or x.shape[-2] > self.tile_sample_min_size):
return self.spatial_tiled_encode(x, return_dict=return_dict)
if self.use_slicing and x.shape[0] > 1:
@@ -323,11 +304,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(
self,
z: torch.FloatTensor,
return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
assert len(z.shape) == 5, "The input tensor should have 5 dimensions."
if self.use_parallel:
@@ -336,9 +313,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
if self.use_temporal_tiling and z.shape[2] > self.tile_latent_min_tsize:
return self.temporal_tiled_decode(z, return_dict=return_dict)
if self.use_spatial_tiling and (
z.shape[-1] > self.tile_latent_min_size
or z.shape[-2] > self.tile_latent_min_size):
if self.use_spatial_tiling and (z.shape[-1] > self.tile_latent_min_size
or z.shape[-2] > self.tile_latent_min_size):
return self.spatial_tiled_decode(z, return_dict=return_dict)
z = self.post_quant_conv(z)
@@ -369,9 +345,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
"""
if self.use_slicing and z.shape[0] > 1:
decoded_slices = [
self._decode(z_slice).sample for z_slice in z.split(1)
]
decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
decoded = torch.cat(decoded_slices)
else:
decoded = self._decode(z).sample
@@ -381,28 +355,26 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
return DecoderOutput(sample=decoded)
def blend_v(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
for y in range(blend_extent):
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (
1 - y / blend_extent) + b[:, :, :, y, :] * (y / blend_extent)
b[:, :, :,
y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (y / blend_extent)
return b
def blend_h(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
for x in range(blend_extent):
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (
1 - x / blend_extent) + b[:, :, :, :, x] * (x / blend_extent)
b[:, :, :, :,
x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (x / blend_extent)
return b
def blend_t(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
for x in range(blend_extent):
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (
1 - x / blend_extent) + b[:, :, x, :, :] * (x / blend_extent)
b[:, :,
x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :,
x, :, :] * (x / blend_extent)
return b
def spatial_tiled_encode(
@@ -429,10 +401,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
overlap_size = int(self.tile_sample_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size *
self.tile_overlap_factor)
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split video into tiles and encode them separately.
@@ -440,8 +410,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
for i in range(0, x.shape[-2], overlap_size):
row = []
for j in range(0, x.shape[-1], overlap_size):
tile = x[:, :, :, i:i + self.tile_sample_min_size,
j:j + self.tile_sample_min_size, ]
tile = x[:, :, :, i:i + self.tile_sample_min_size, j:j + self.tile_sample_min_size, ]
tile = self.encoder(tile)
tile = self.quant_conv(tile)
row.append(tile)
@@ -471,8 +440,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
def spatial_tiled_decode(self,
z: torch.FloatTensor,
return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
r"""
Decode a batch of images/videos using a tiled decoder.
@@ -486,10 +454,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
overlap_size = int(self.tile_latent_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size *
self.tile_overlap_factor)
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
row_limit = self.tile_sample_min_size - blend_extent
# Split z into overlapping tiles and decode them separately.
@@ -498,8 +464,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
for i in range(0, z.shape[-2], overlap_size):
row = []
for j in range(0, z.shape[-1], overlap_size):
tile = z[:, :, :, i:i + self.tile_latent_min_size,
j:j + self.tile_latent_min_size, ]
tile = z[:, :, :, i:i + self.tile_latent_min_size, j:j + self.tile_latent_min_size, ]
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
row.append(decoded)
@@ -523,24 +488,19 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
return DecoderOutput(sample=dec)
def temporal_tiled_encode(self,
x: torch.FloatTensor,
return_dict: bool = True) -> AutoencoderKLOutput:
def temporal_tiled_encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
B, C, T, H, W = x.shape
overlap_size = int(self.tile_sample_min_tsize *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_tsize *
self.tile_overlap_factor)
overlap_size = int(self.tile_sample_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_latent_min_tsize - blend_extent
# Split the video into tiles and encode them separately.
row = []
for i in range(0, T, overlap_size):
tile = x[:, :, i:i + self.tile_sample_min_tsize + 1, :, :]
if self.use_spatial_tiling and (
tile.shape[-1] > self.tile_sample_min_size
or tile.shape[-2] > self.tile_sample_min_size):
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_sample_min_size
or tile.shape[-2] > self.tile_sample_min_size):
tile = self.spatial_tiled_encode(tile, return_moments=True)
else:
tile = self.encoder(tile)
@@ -566,25 +526,20 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
def temporal_tiled_decode(self,
z: torch.FloatTensor,
return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
# Split z into overlapping tiles and decode them separately.
B, C, T, H, W = z.shape
overlap_size = int(self.tile_latent_min_tsize *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_tsize *
self.tile_overlap_factor)
overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_sample_min_tsize - blend_extent
row = []
for i in range(0, T, overlap_size):
tile = z[:, :, i:i + self.tile_latent_min_tsize + 1, :, :]
if self.use_spatial_tiling and (
tile.shape[-1] > self.tile_latent_min_size
or tile.shape[-2] > self.tile_latent_min_size):
decoded = self.spatial_tiled_decode(tile,
return_dict=True).sample
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_latent_min_size
or tile.shape[-2] > self.tile_latent_min_size):
decoded = self.spatial_tiled_decode(tile, return_dict=True).sample
else:
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
@@ -605,22 +560,19 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
return DecoderOutput(sample=dec)
def _parallel_data_generator(self, gathered_results,
gathered_dim_metadata):
def _parallel_data_generator(self, gathered_results, gathered_dim_metadata):
global_idx = 0
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
_start_shape = 0
for shape in per_rank_metadata:
mul_shape = prod(shape)
yield (gathered_results[i, _start_shape:_start_shape +
mul_shape].reshape(shape), global_idx)
yield (gathered_results[i, _start_shape:_start_shape + mul_shape].reshape(shape), global_idx)
_start_shape += mul_shape
global_idx += 1
def parallel_tiled_decode(self,
z: torch.FloatTensor,
return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
"""
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
"""
@@ -628,16 +580,12 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
B, C, T, H, W = z.shape
# Calculate parameters
t_overlap_size = int(self.tile_latent_min_tsize *
(1 - self.tile_overlap_factor))
t_blend_extent = int(self.tile_sample_min_tsize *
self.tile_overlap_factor)
t_overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
t_blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_sample_min_tsize - t_blend_extent
s_overlap_size = int(self.tile_latent_min_size *
(1 - self.tile_overlap_factor))
s_blend_extent = int(self.tile_sample_min_size *
self.tile_overlap_factor)
s_overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
s_blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
s_row_limit = self.tile_sample_min_size - s_blend_extent
# Calculate tile dimensions
@@ -655,8 +603,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
local_results = []
local_dim_metadata = []
# Process assigned tiles
for local_idx, global_idx in enumerate(
range(start_tile_idx, end_tile_idx)):
for local_idx, global_idx in enumerate(range(start_tile_idx, end_tile_idx)):
# Convert flat index to 3D indices
t_idx = global_idx // total_spatial_tiles
spatial_idx = global_idx % total_spatial_tiles
@@ -670,8 +617,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
# Extract and process tile
tile = z[:, :, t_start:t_start + self.tile_latent_min_tsize + 1,
h_start:h_start + self.tile_latent_min_size,
w_start:w_start + self.tile_latent_min_size]
h_start:h_start + self.tile_latent_min_size, w_start:w_start + self.tile_latent_min_size]
# Process tile
tile = self.post_quant_conv(tile)
@@ -691,13 +637,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
del local_results
torch.cuda.empty_cache()
# first gather size to pad the results
local_size = torch.tensor([results.size(0)],
device=results.device,
dtype=torch.int64)
all_sizes = [
torch.zeros(1, device=results.device, dtype=torch.int64)
for _ in range(world_size)
]
local_size = torch.tensor([results.size(0)], device=results.device, dtype=torch.int64)
all_sizes = [torch.zeros(1, device=results.device, dtype=torch.int64) for _ in range(world_size)]
dist.all_gather(all_sizes, local_size)
max_size = max(size.item() for size in all_sizes)
padded_results = torch.zeros(max_size, device=results.device)
@@ -707,16 +648,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
# Gather all results
gathered_dim_metadata = [None] * world_size
gathered_results = torch.zeros_like(padded_results).repeat(
world_size, *[1] * len(padded_results.shape)
).contiguous(
) # use contiguous to make sure it won't copy data in the following operations
world_size, *[1] * len(padded_results.shape)).contiguous(
) # use contiguous to make sure it won't copy data in the following operations
dist.all_gather_into_tensor(gathered_results, padded_results)
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
# Process gathered results
data = [[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)]
for _ in range(num_t_tiles)]
for current_data, global_idx in self._parallel_data_generator(
gathered_results, gathered_dim_metadata):
data = [[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)] for _ in range(num_t_tiles)]
for current_data, global_idx in self._parallel_data_generator(gathered_results, gathered_dim_metadata):
t_idx = global_idx // total_spatial_tiles
spatial_idx = global_idx % total_spatial_tiles
h_idx = spatial_idx // num_w_tiles
@@ -726,11 +664,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
result_slices = []
last_slice_data = None
for i, tem_data in enumerate(data):
slice_data = self._merge_spatial_tiles(tem_data, s_blend_extent,
s_row_limit)
slice_data = self._merge_spatial_tiles(tem_data, s_blend_extent, s_row_limit)
if i > 0:
slice_data = self.blend_t(last_slice_data, slice_data,
t_blend_extent)
slice_data = self.blend_t(last_slice_data, slice_data, t_blend_extent)
result_slices.append(slice_data[:, :, :t_limit, :, :])
else:
result_slices.append(slice_data[:, :, :t_limit + 1, :, :])
@@ -748,8 +684,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
result_row = []
for j, tile in enumerate(row):
if i > 0:
tile = self.blend_v(spatial_rows[i - 1][j], tile,
blend_extent)
tile = self.blend_v(spatial_rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :, :row_limit, :row_limit])
@@ -806,9 +741,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
for _, attn_processor in self.attn_processors.items():
if "Added" in str(attn_processor.__class__.__name__):
raise ValueError(
"`fuse_qkv_projections()` is not supported for models having added KV projections."
)
raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.")
self.original_attn_processors = self.attn_processors
@@ -31,16 +31,9 @@ from torch import nn
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def prepare_causal_attention_mask(n_frame: int,
n_hw: int,
dtype,
device,
batch_size: int = None):
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len),
float("-inf"),
dtype=dtype,
device=device)
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, :(i_frame + 1) * n_hw] = 0
@@ -78,12 +71,7 @@ class CausalConv3d(nn.Module):
) # W, H, T
self.time_causal_padding = padding
self.conv = nn.Conv3d(chan_in,
chan_out,
kernel_size,
stride=stride,
dilation=dilation,
**kwargs)
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
@@ -135,10 +123,7 @@ class UpsampleCausal3D(nn.Module):
elif use_conv:
if kernel_size is None:
kernel_size = 3
conv = CausalConv3d(self.channels,
self.out_channels,
kernel_size=kernel_size,
bias=bias)
conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, bias=bias)
if name == "conv":
self.conv = conv
@@ -175,14 +160,10 @@ class UpsampleCausal3D(nn.Module):
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
if output_size is None:
if T > 1:
other_h = F.interpolate(other_h,
scale_factor=self.upsample_factor,
mode="nearest")
other_h = F.interpolate(other_h, scale_factor=self.upsample_factor, mode="nearest")
first_h = first_h.squeeze(2)
first_h = F.interpolate(first_h,
scale_factor=self.upsample_factor[1:],
mode="nearest")
first_h = F.interpolate(first_h, scale_factor=self.upsample_factor[1:], mode="nearest")
first_h = first_h.unsqueeze(2)
else:
raise NotImplementedError
@@ -260,15 +241,11 @@ class DownsampleCausal3D(nn.Module):
else:
self.conv = conv
def forward(self,
hidden_states: torch.FloatTensor,
scale: float = 1.0) -> torch.FloatTensor:
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
assert hidden_states.shape[1] == self.channels
if self.norm is not None:
hidden_states = self.norm(hidden_states.permute(0, 2, 3,
1)).permute(
0, 3, 1, 2)
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2)
assert hidden_states.shape[1] == self.channels
@@ -325,58 +302,36 @@ class ResnetBlockCausal3D(nn.Module):
groups_out = groups
if self.time_embedding_norm == "ada_group":
self.norm1 = AdaGroupNorm(temb_channels,
in_channels,
groups,
eps=eps)
self.norm1 = AdaGroupNorm(temb_channels, in_channels, groups, eps=eps)
elif self.time_embedding_norm == "spatial":
self.norm1 = SpatialNorm(in_channels, temb_channels)
else:
self.norm1 = torch.nn.GroupNorm(num_groups=groups,
num_channels=in_channels,
eps=eps,
affine=True)
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.conv1 = CausalConv3d(in_channels,
out_channels,
kernel_size=3,
stride=1)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
if temb_channels is not None:
if self.time_embedding_norm == "default":
self.time_emb_proj = linear_cls(temb_channels, out_channels)
elif self.time_embedding_norm == "scale_shift":
self.time_emb_proj = linear_cls(temb_channels,
2 * out_channels)
elif (self.time_embedding_norm == "ada_group"
or self.time_embedding_norm == "spatial"):
self.time_emb_proj = linear_cls(temb_channels, 2 * out_channels)
elif (self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial"):
self.time_emb_proj = None
else:
raise ValueError(
f"Unknown time_embedding_norm : {self.time_embedding_norm} "
)
raise ValueError(f"Unknown time_embedding_norm : {self.time_embedding_norm} ")
else:
self.time_emb_proj = None
if self.time_embedding_norm == "ada_group":
self.norm2 = AdaGroupNorm(temb_channels,
out_channels,
groups_out,
eps=eps)
self.norm2 = AdaGroupNorm(temb_channels, out_channels, groups_out, eps=eps)
elif self.time_embedding_norm == "spatial":
self.norm2 = SpatialNorm(out_channels, temb_channels)
else:
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out,
num_channels=out_channels,
eps=eps,
affine=True)
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
self.dropout = torch.nn.Dropout(dropout)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(out_channels,
conv_3d_out_channels,
kernel_size=3,
stride=1)
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
self.nonlinearity = get_activation(non_linearity)
@@ -384,12 +339,10 @@ class ResnetBlockCausal3D(nn.Module):
if self.up:
self.upsample = UpsampleCausal3D(in_channels, use_conv=False)
elif self.down:
self.downsample = DownsampleCausal3D(in_channels,
use_conv=False,
name="op")
self.downsample = DownsampleCausal3D(in_channels, use_conv=False, name="op")
self.use_in_shortcut = (self.in_channels != conv_3d_out_channels if
use_in_shortcut is None else use_in_shortcut)
self.use_in_shortcut = (self.in_channels != conv_3d_out_channels
if use_in_shortcut is None else use_in_shortcut)
self.conv_shortcut = None
if self.use_in_shortcut:
@@ -409,8 +362,7 @@ class ResnetBlockCausal3D(nn.Module):
) -> torch.FloatTensor:
hidden_states = input_tensor
if (self.time_embedding_norm == "ada_group"
or self.time_embedding_norm == "spatial"):
if (self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial"):
hidden_states = self.norm1(hidden_states, temb)
else:
hidden_states = self.norm1(hidden_states)
@@ -438,8 +390,7 @@ class ResnetBlockCausal3D(nn.Module):
if temb is not None and self.time_embedding_norm == "default":
hidden_states = hidden_states + temb
if (self.time_embedding_norm == "ada_group"
or self.time_embedding_norm == "spatial"):
if (self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial"):
hidden_states = self.norm2(hidden_states, temb)
else:
hidden_states = self.norm2(hidden_states)
@@ -456,8 +407,7 @@ class ResnetBlockCausal3D(nn.Module):
if self.conv_shortcut is not None:
input_tensor = self.conv_shortcut(input_tensor)
output_tensor = (input_tensor +
hidden_states) / self.output_scale_factor
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
@@ -497,9 +447,7 @@ def get_down_block3d(
)
attention_head_dim = num_attention_heads
down_block_type = (down_block_type[7:]
if down_block_type.startswith("UNetRes") else
down_block_type)
down_block_type = (down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type)
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
@@ -553,8 +501,7 @@ def get_up_block3d(
)
attention_head_dim = num_attention_heads
up_block_type = (up_block_type[7:]
if up_block_type.startswith("UNetRes") else up_block_type)
up_block_type = (up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type)
if up_block_type == "UpDecoderBlockCausal3D":
return UpDecoderBlockCausal3D(
num_layers=num_layers,
@@ -595,13 +542,11 @@ class UNetMidBlockCausal3D(nn.Module):
output_scale_factor: float = 1.0,
):
super().__init__()
resnet_groups = (resnet_groups if resnet_groups is not None else min(
in_channels // 4, 32))
resnet_groups = (resnet_groups if resnet_groups is not None else min(in_channels // 4, 32))
self.add_attention = add_attention
if attn_groups is None:
attn_groups = (resnet_groups
if resnet_time_scale_shift == "default" else None)
attn_groups = (resnet_groups if resnet_time_scale_shift == "default" else None)
# there is always at least one resnet
resnets = [
@@ -636,9 +581,7 @@ class UNetMidBlockCausal3D(nn.Module):
rescale_output_factor=output_scale_factor,
eps=resnet_eps,
norm_num_groups=attn_groups,
spatial_norm_dim=(temb_channels
if resnet_time_scale_shift
== "spatial" else None),
spatial_norm_dim=(temb_channels if resnet_time_scale_shift == "spatial" else None),
residual_connection=True,
bias=True,
upcast_softmax=True,
@@ -664,29 +607,19 @@ class UNetMidBlockCausal3D(nn.Module):
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
def forward(self,
hidden_states: torch.FloatTensor,
temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
def forward(self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
hidden_states = self.resnets[0](hidden_states, temb)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
B, C, T, H, W = hidden_states.shape
hidden_states = rearrange(hidden_states,
"b c f h w -> b (f h w) c")
attention_mask = prepare_causal_attention_mask(
T,
H * W,
hidden_states.dtype,
hidden_states.device,
batch_size=B)
hidden_states = attn(hidden_states,
temb=temb,
attention_mask=attention_mask)
hidden_states = rearrange(hidden_states,
"b (f h w) c -> b c f h w",
f=T,
h=H,
w=W)
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
attention_mask = prepare_causal_attention_mask(T,
H * W,
hidden_states.dtype,
hidden_states.device,
batch_size=B)
hidden_states = attn(hidden_states, temb=temb, attention_mask=attention_mask)
hidden_states = rearrange(hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W)
hidden_states = resnet(hidden_states, temb)
return hidden_states
@@ -745,9 +678,7 @@ class DownEncoderBlockCausal3D(nn.Module):
else:
self.downsamplers = None
def forward(self,
hidden_states: torch.FloatTensor,
scale: float = 1.0) -> torch.FloatTensor:
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
@@ -761,21 +692,21 @@ class DownEncoderBlockCausal3D(nn.Module):
class UpDecoderBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
resolution_idx: Optional[int] = None,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default", # default, spatial
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
resnet_pre_norm: bool = True,
output_scale_factor: float = 1.0,
add_upsample: bool = True,
upsample_scale_factor=(2, 2, 2),
temb_channels: Optional[int] = None,
self,
in_channels: int,
out_channels: int,
resolution_idx: Optional[int] = None,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default", # default, spatial
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
resnet_pre_norm: bool = True,
output_scale_factor: float = 1.0,
add_upsample: bool = True,
upsample_scale_factor=(2, 2, 2),
temb_channels: Optional[int] = None,
):
super().__init__()
resnets = []
+34 -77
View File
@@ -8,8 +8,7 @@ from diffusers.models.attention_processor import SpatialNorm
from diffusers.utils import BaseOutput, is_torch_version
from diffusers.utils.torch_utils import randn_tensor
from .unet_causal_3d_blocks import (CausalConv3d, UNetMidBlockCausal3D,
get_down_block3d, get_up_block3d)
from .unet_causal_3d_blocks import CausalConv3d, UNetMidBlockCausal3D, get_down_block3d, get_up_block3d
@dataclass
@@ -47,10 +46,7 @@ class EncoderCausal3D(nn.Module):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels,
block_out_channels[0],
kernel_size=3,
stride=1)
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
@@ -60,33 +56,25 @@ class EncoderCausal3D(nn.Module):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(
i >=
(len(block_out_channels) - 1 - num_time_downsample_layers)
and not is_final_block)
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (len(block_out_channels) - 1 - num_time_downsample_layers)
and not is_final_block)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}."
)
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(downsample_stride_T +
downsample_stride_HW)
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(add_spatial_downsample
or add_time_downsample),
add_downsample=bool(add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
@@ -111,20 +99,15 @@ class EncoderCausal3D(nn.Module):
)
# out
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1],
num_groups=norm_num_groups,
eps=1e-6)
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
self.conv_act = nn.SiLU()
conv_out_channels = 2 * out_channels if double_z else out_channels
self.conv_out = CausalConv3d(block_out_channels[-1],
conv_out_channels,
kernel_size=3)
self.conv_out = CausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
def forward(self, sample: torch.FloatTensor) -> torch.FloatTensor:
r"""The forward method of the `EncoderCausal3D` class."""
assert len(
sample.shape) == 5, "The input tensor should have 5 dimensions"
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
@@ -165,10 +148,7 @@ class DecoderCausal3D(nn.Module):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels,
block_out_channels[-1],
kernel_size=3,
stride=1)
self.conv_in = CausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
self.mid_block = None
self.up_blocks = nn.ModuleList([])
@@ -180,8 +160,7 @@ class DecoderCausal3D(nn.Module):
resnet_eps=1e-6,
resnet_act_fn=act_fn,
output_scale_factor=1,
resnet_time_scale_shift="default"
if norm_type == "group" else norm_type,
resnet_time_scale_shift="default" if norm_type == "group" else norm_type,
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
temb_channels=temb_channels,
@@ -195,25 +174,19 @@ class DecoderCausal3D(nn.Module):
prev_output_channel = output_channel
output_channel = reversed_block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_upsample_layers = int(
np.log2(spatial_compression_ratio))
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
num_time_upsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
add_time_upsample = bool(
i >= len(block_out_channels) - 1 - num_time_upsample_layers
and not is_final_block)
add_time_upsample = bool(i >= len(block_out_channels) - 1 - num_time_upsample_layers
and not is_final_block)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}."
)
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1,
1)
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
upsample_scale_factor_T = (2, ) if add_time_upsample else (1, )
upsample_scale_factor = tuple(upsample_scale_factor_T +
upsample_scale_factor_HW)
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
up_block = get_up_block3d(
up_block_type,
num_layers=self.layers_per_block + 1,
@@ -234,17 +207,11 @@ class DecoderCausal3D(nn.Module):
# out
if norm_type == "spatial":
self.conv_norm_out = SpatialNorm(block_out_channels[0],
temb_channels)
self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels)
else:
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[0],
num_groups=norm_num_groups,
eps=1e-6)
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
self.conv_act = nn.SiLU()
self.conv_out = CausalConv3d(block_out_channels[0],
out_channels,
kernel_size=3)
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
self.gradient_checkpointing = False
@@ -254,8 +221,7 @@ class DecoderCausal3D(nn.Module):
latent_embeds: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
r"""The forward method of the `DecoderCausal3D` class."""
assert len(
sample.shape) == 5, "The input tensor should have 5 dimensions."
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions."
sample = self.conv_in(sample)
@@ -289,15 +255,12 @@ class DecoderCausal3D(nn.Module):
)
else:
# middle
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block), sample,
latent_embeds)
sample = torch.utils.checkpoint.checkpoint(create_custom_forward(self.mid_block), sample, latent_embeds)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(up_block), sample, latent_embeds)
sample = torch.utils.checkpoint.checkpoint(create_custom_forward(up_block), sample, latent_embeds)
else:
# middle
sample = self.mid_block(sample, latent_embeds)
@@ -334,14 +297,11 @@ class DiagonalGaussianDistribution(object):
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(
self.mean,
device=self.parameters.device,
dtype=self.parameters.dtype)
self.var = self.std = torch.zeros_like(self.mean,
device=self.parameters.device,
dtype=self.parameters.dtype)
def sample(
self,
generator: Optional[torch.Generator] = None) -> torch.FloatTensor:
def sample(self, generator: Optional[torch.Generator] = None) -> torch.FloatTensor:
# make sure sample is on the same device as the parameters and has same dtype
sample = randn_tensor(
self.mean.shape,
@@ -364,20 +324,17 @@ class DiagonalGaussianDistribution(object):
)
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var +
self.var / other.var - 1.0 - self.logvar + other.logvar,
torch.pow(self.mean - other.mean, 2) / other.var + self.var / other.var - 1.0 - self.logvar +
other.logvar,
dim=reduce_dim,
)
def nll(self,
sample: torch.Tensor,
dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(
logtwopi + self.logvar +
torch.pow(sample - self.mean, 2) / self.var,
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims,
)
+85 -201
View File
@@ -21,29 +21,23 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.models.attention import FeedForward
from diffusers.models.attention_processor import Attention, AttentionProcessor
from diffusers.models.embeddings import (
CombinedTimestepGuidanceTextProjEmbeddings,
CombinedTimestepTextProjEmbeddings, get_1d_rotary_pos_embed)
from diffusers.models.embeddings import (CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings,
get_1d_rotary_pos_embed)
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import (AdaLayerNormContinuous,
AdaLayerNormZero,
AdaLayerNormZeroSingle)
from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, logging,
scale_lora_layers, unscale_lora_layers)
from diffusers.models.normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle
from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads,
local_heads)
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
class HunyuanVideoAttnProcessor2_0:
@@ -51,8 +45,7 @@ class HunyuanVideoAttnProcessor2_0:
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"HunyuanVideoAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0."
)
"HunyuanVideoAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
@@ -66,8 +59,7 @@ class HunyuanVideoAttnProcessor2_0:
sequence_length = hidden_states.size(1)
encoder_sequence_length = encoder_hidden_states.size(1)
if attn.add_q_proj is None and encoder_hidden_states is not None:
hidden_states = torch.cat([hidden_states, encoder_hidden_states],
dim=1)
hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)
# 1. QKV projections
query = attn.to_q(hidden_states)
@@ -96,18 +88,14 @@ class HunyuanVideoAttnProcessor2_0:
if attn.add_q_proj is None and encoder_hidden_states is not None:
query = torch.cat(
[
apply_rotary_emb(
query[:, :, :-encoder_hidden_states.shape[1]],
image_rotary_emb),
apply_rotary_emb(query[:, :, :-encoder_hidden_states.shape[1]], image_rotary_emb),
query[:, :, -encoder_hidden_states.shape[1]:],
],
dim=2,
)
key = torch.cat(
[
apply_rotary_emb(
key[:, :, :-encoder_hidden_states.shape[1]],
image_rotary_emb),
apply_rotary_emb(key[:, :, :-encoder_hidden_states.shape[1]], image_rotary_emb),
key[:, :, -encoder_hidden_states.shape[1]:],
],
dim=2,
@@ -122,16 +110,12 @@ class HunyuanVideoAttnProcessor2_0:
encoder_key = attn.add_k_proj(encoder_hidden_states)
encoder_value = attn.add_v_proj(encoder_hidden_states)
encoder_query = encoder_query.unflatten(
2, (attn.heads, -1)).transpose(1, 2)
encoder_key = encoder_key.unflatten(2, (attn.heads, -1)).transpose(
1, 2)
encoder_value = encoder_value.unflatten(
2, (attn.heads, -1)).transpose(1, 2)
encoder_query = encoder_query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
encoder_key = encoder_key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
encoder_value = encoder_value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
if attn.norm_added_q is not None:
encoder_query = attn.norm_added_q(encoder_query).to(
encoder_value)
encoder_query = attn.norm_added_q(encoder_query).to(encoder_value)
if attn.norm_added_k is not None:
encoder_key = attn.norm_added_k(encoder_key).to(encoder_value)
@@ -140,17 +124,10 @@ class HunyuanVideoAttnProcessor2_0:
value = torch.cat([value, encoder_value], dim=2)
if get_sequence_parallel_state():
query_img, query_txt = query[:, :, :
sequence_length, :], query[:, :,
sequence_length:, :]
key_img, key_txt = key[:, :, :
sequence_length, :], key[:, :,
sequence_length:, :]
value_img, value_txt = value[:, :, :
sequence_length, :], value[:, :,
sequence_length:, :]
query_img = all_to_all_4D(query_img, scatter_dim=1,
gather_dim=2) #
query_img, query_txt = query[:, :, :sequence_length, :], query[:, :, sequence_length:, :]
key_img, key_txt = key[:, :, :sequence_length, :], key[:, :, sequence_length:, :]
value_img, value_txt = value[:, :, :sequence_length, :], value[:, :, sequence_length:, :]
query_img = all_to_all_4D(query_img, scatter_dim=1, gather_dim=2) #
key_img = all_to_all_4D(key_img, scatter_dim=1, gather_dim=2)
value_img = all_to_all_4D(value_img, scatter_dim=1, gather_dim=2)
@@ -171,24 +148,15 @@ class HunyuanVideoAttnProcessor2_0:
attention_mask = attention_mask[:, 0, :]
seq_len = qkv.shape[1]
attn_len = attention_mask.shape[1]
attention_mask = F.pad(attention_mask, (seq_len - attn_len, 0),
value=True)
attention_mask = F.pad(attention_mask, (seq_len - attn_len, 0), value=True)
hidden_states = flash_attn_no_pad(qkv,
attention_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None)
hidden_states = flash_attn_no_pad(qkv, attention_mask, causal=False, dropout_p=0.0, softmax_scale=None)
if get_sequence_parallel_state():
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length * nccl_info.sp_size, encoder_sequence_length),
dim=1)
hidden_states = all_to_all_4D(hidden_states,
scatter_dim=1,
gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states,
dim=2).contiguous()
(sequence_length * nccl_info.sp_size, encoder_sequence_length), dim=1)
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.flatten(2, 3)
@@ -225,35 +193,26 @@ class HunyuanVideoPatchEmbed(nn.Module):
) -> None:
super().__init__()
patch_size = (patch_size, patch_size, patch_size) if isinstance(
patch_size, int) else patch_size
self.proj = nn.Conv3d(in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size)
patch_size = (patch_size, patch_size, patch_size) if isinstance(patch_size, int) else patch_size
self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.proj(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1,
2) # BCFHW -> BNC
hidden_states = hidden_states.flatten(2).transpose(1, 2) # BCFHW -> BNC
return hidden_states
class HunyuanVideoAdaNorm(nn.Module):
def __init__(self,
in_features: int,
out_features: Optional[int] = None) -> None:
def __init__(self, in_features: int, out_features: Optional[int] = None) -> None:
super().__init__()
out_features = out_features or 2 * in_features
self.linear = nn.Linear(in_features, out_features)
self.nonlinearity = nn.SiLU()
def forward(
self, temb: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor,
torch.Tensor]:
def forward(self,
temb: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
temb = self.linear(self.nonlinearity(temb))
gate_msa, gate_mlp = temb.chunk(2, dim=1)
gate_msa, gate_mlp = gate_msa.unsqueeze(1), gate_mlp.unsqueeze(1)
@@ -274,9 +233,7 @@ class HunyuanVideoIndividualTokenRefinerBlock(nn.Module):
hidden_size = num_attention_heads * attention_head_dim
self.norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=True,
eps=1e-6)
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
self.attn = Attention(
query_dim=hidden_size,
cross_attention_dim=None,
@@ -285,13 +242,8 @@ class HunyuanVideoIndividualTokenRefinerBlock(nn.Module):
bias=attention_bias,
)
self.norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=True,
eps=1e-6)
self.ff = FeedForward(hidden_size,
mult=mlp_width_ratio,
activation_fn="linear-silu",
dropout=mlp_drop_rate)
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
self.ff = FeedForward(hidden_size, mult=mlp_width_ratio, activation_fn="linear-silu", dropout=mlp_drop_rate)
self.norm_out = HunyuanVideoAdaNorm(hidden_size, 2 * hidden_size)
@@ -352,9 +304,7 @@ class HunyuanVideoIndividualTokenRefiner(nn.Module):
batch_size = attention_mask.shape[0]
seq_len = attention_mask.shape[1]
attention_mask = attention_mask.to(hidden_states.device).bool()
self_attn_mask_1 = attention_mask.view(batch_size, 1, 1,
seq_len).repeat(
1, 1, seq_len, 1)
self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1)
self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)
self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool()
self_attn_mask[:, :, :, 0] = True
@@ -381,8 +331,8 @@ class HunyuanVideoTokenRefiner(nn.Module):
hidden_size = num_attention_heads * attention_head_dim
self.time_text_embed = CombinedTimestepTextProjEmbeddings(
embedding_dim=hidden_size, pooled_projection_dim=in_channels)
self.time_text_embed = CombinedTimestepTextProjEmbeddings(embedding_dim=hidden_size,
pooled_projection_dim=in_channels)
self.proj_in = nn.Linear(in_channels, hidden_size, bias=True)
self.token_refiner = HunyuanVideoIndividualTokenRefiner(
num_attention_heads=num_attention_heads,
@@ -404,8 +354,7 @@ class HunyuanVideoTokenRefiner(nn.Module):
else:
original_dtype = hidden_states.dtype
mask_float = attention_mask.float().unsqueeze(-1)
pooled_projections = (hidden_states * mask_float).sum(
dim=1) / mask_float.sum(dim=1)
pooled_projections = (hidden_states * mask_float).sum(dim=1) / mask_float.sum(dim=1)
pooled_projections = pooled_projections.to(original_dtype)
temb = self.time_text_embed(timestep, pooled_projections)
@@ -417,11 +366,7 @@ class HunyuanVideoTokenRefiner(nn.Module):
class HunyuanVideoRotaryPosEmbed(nn.Module):
def __init__(self,
patch_size: int,
patch_size_t: int,
rope_dim: List[int],
theta: float = 256.0) -> None:
def __init__(self, patch_size: int, patch_size_t: int, rope_dim: List[int], theta: float = 256.0) -> None:
super().__init__()
self.patch_size = patch_size
@@ -432,8 +377,7 @@ class HunyuanVideoRotaryPosEmbed(nn.Module):
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
rope_sizes = [
num_frames * nccl_info.sp_size // self.patch_size_t,
height // self.patch_size, width // self.patch_size
num_frames * nccl_info.sp_size // self.patch_size_t, height // self.patch_size, width // self.patch_size
]
axes_grids = []
@@ -441,26 +385,18 @@ class HunyuanVideoRotaryPosEmbed(nn.Module):
# Note: The following line diverges from original behaviour. We create the grid on the device, whereas
# original implementation creates it on CPU and then moves it to device. This results in numerical
# differences in layerwise debugging outputs, but visually it is the same.
grid = torch.arange(0,
rope_sizes[i],
device=hidden_states.device,
dtype=torch.float32)
grid = torch.arange(0, rope_sizes[i], device=hidden_states.device, dtype=torch.float32)
axes_grids.append(grid)
grid = torch.meshgrid(*axes_grids, indexing="ij") # [W, H, T]
grid = torch.stack(grid, dim=0) # [3, W, H, T]
freqs = []
for i in range(3):
freq = get_1d_rotary_pos_embed(self.rope_dim[i],
grid[i].reshape(-1),
self.theta,
use_real=True)
freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True)
freqs.append(freq)
freqs_cos = torch.cat([f[0] for f in freqs],
dim=1) # (W * H * T, D / 2)
freqs_sin = torch.cat([f[1] for f in freqs],
dim=1) # (W * H * T, D / 2)
freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2)
freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2)
return freqs_cos, freqs_sin
@@ -505,8 +441,7 @@ class HunyuanVideoSingleTransformerBlock(nn.Module):
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.shape[1]
hidden_states = torch.cat([hidden_states, encoder_hidden_states],
dim=1)
hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)
residual = hidden_states
@@ -554,8 +489,7 @@ class HunyuanVideoTransformerBlock(nn.Module):
hidden_size = num_attention_heads * attention_head_dim
self.norm1 = AdaLayerNormZero(hidden_size, norm_type="layer_norm")
self.norm1_context = AdaLayerNormZero(hidden_size,
norm_type="layer_norm")
self.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm")
self.attn = Attention(
query_dim=hidden_size,
@@ -571,19 +505,11 @@ class HunyuanVideoTransformerBlock(nn.Module):
eps=1e-6,
)
self.norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.ff = FeedForward(hidden_size,
mult=mlp_ratio,
activation_fn="gelu-approximate")
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate")
self.norm2_context = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.ff_context = FeedForward(hidden_size,
mult=mlp_ratio,
activation_fn="gelu-approximate")
self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate")
def forward(
self,
@@ -594,8 +520,7 @@ class HunyuanVideoTransformerBlock(nn.Module):
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
# 1. Input normalization
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(
hidden_states, emb=temb)
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb)
norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context(
encoder_hidden_states, emb=temb)
@@ -609,30 +534,25 @@ class HunyuanVideoTransformerBlock(nn.Module):
# 3. Modulation and residual connection
hidden_states = hidden_states + attn_output * gate_msa.unsqueeze(1)
encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(
1)
encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1)
norm_hidden_states = self.norm2(hidden_states)
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
norm_hidden_states = norm_hidden_states * (
1 + scale_mlp[:, None]) + shift_mlp[:, None]
norm_encoder_hidden_states = norm_encoder_hidden_states * (
1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None]
norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None]
# 4. Feed-forward
ff_output = self.ff(norm_hidden_states)
context_ff_output = self.ff_context(norm_encoder_hidden_states)
hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output
encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(
1) * context_ff_output
encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output
return hidden_states, encoder_hidden_states
class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
FromOriginalModelMixin):
class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r"""
A Transformer model for video-like data used in [HunyuanVideo](https://huggingface.co/tencent/HunyuanVideo).
@@ -699,26 +619,19 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
out_channels = out_channels or in_channels
# 1. Latent and condition embedders
self.x_embedder = HunyuanVideoPatchEmbed(
(patch_size_t, patch_size, patch_size), in_channels, inner_dim)
self.context_embedder = HunyuanVideoTokenRefiner(
text_embed_dim,
num_attention_heads,
attention_head_dim,
num_layers=num_refiner_layers)
self.time_text_embed = CombinedTimestepGuidanceTextProjEmbeddings(
inner_dim, pooled_projection_dim)
self.x_embedder = HunyuanVideoPatchEmbed((patch_size_t, patch_size, patch_size), in_channels, inner_dim)
self.context_embedder = HunyuanVideoTokenRefiner(text_embed_dim,
num_attention_heads,
attention_head_dim,
num_layers=num_refiner_layers)
self.time_text_embed = CombinedTimestepGuidanceTextProjEmbeddings(inner_dim, pooled_projection_dim)
# 2. RoPE
self.rope = HunyuanVideoRotaryPosEmbed(patch_size, patch_size_t,
rope_axes_dim, rope_theta)
self.rope = HunyuanVideoRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta)
# 3. Dual stream transformer blocks
self.transformer_blocks = nn.ModuleList([
HunyuanVideoTransformerBlock(num_attention_heads,
attention_head_dim,
mlp_ratio=mlp_ratio,
qk_norm=qk_norm)
HunyuanVideoTransformerBlock(num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm)
for _ in range(num_layers)
])
@@ -727,17 +640,12 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
HunyuanVideoSingleTransformerBlock(num_attention_heads,
attention_head_dim,
mlp_ratio=mlp_ratio,
qk_norm=qk_norm)
for _ in range(num_single_layers)
qk_norm=qk_norm) for _ in range(num_single_layers)
])
# 5. Output projection
self.norm_out = AdaLayerNormContinuous(inner_dim,
inner_dim,
elementwise_affine=False,
eps=1e-6)
self.proj_out = nn.Linear(
inner_dim, patch_size_t * patch_size * patch_size * out_channels)
self.norm_out = AdaLayerNormContinuous(inner_dim, inner_dim, elementwise_affine=False, eps=1e-6)
self.proj_out = nn.Linear(inner_dim, patch_size_t * patch_size * patch_size * out_channels)
self.gradient_checkpointing = False
@@ -752,15 +660,12 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
# set recursively
processors = {}
def fn_recursive_add_processors(name: str, module: torch.nn.Module,
processors: Dict[str,
AttentionProcessor]):
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor()
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child,
processors)
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
return processors
@@ -770,9 +675,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
return processors
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
def set_attn_processor(self, processor: Union[AttentionProcessor,
Dict[str,
AttentionProcessor]]):
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
r"""
Sets the attention processor to use to compute attention.
@@ -790,11 +693,9 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
f" number of attention layers: {count}. Please make sure to pass {count} processor classes.")
def fn_recursive_attn_processor(name: str, module: torch.nn.Module,
processor):
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor)
@@ -802,8 +703,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
module.set_processor(processor.pop(f"{name}.processor"))
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child,
processor)
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
@@ -823,9 +723,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
return_dict: bool = True,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance is None:
guidance = torch.tensor([6016.0],
device=hidden_states.device,
dtype=torch.bfloat16)
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
@@ -837,11 +735,8 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if attention_kwargs is not None and attention_kwargs.get(
"scale", None) is not None:
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
logger.warning("Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective.")
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p, p_t = self.config.patch_size, self.config.patch_size_t
@@ -849,8 +744,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
post_patch_height = height // p
post_patch_width = width // p
pooled_projections = encoder_hidden_states[:, 0, :self.config.
pooled_projection_dim]
pooled_projections = encoder_hidden_states[:, 0, :self.config.pooled_projection_dim]
encoder_hidden_states = encoder_hidden_states[:, 1:]
# 1. RoPE
@@ -859,9 +753,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
# 2. Conditional embeddings
temb = self.time_text_embed(timestep, guidance, pooled_projections)
hidden_states = self.x_embedder(hidden_states)
encoder_hidden_states = self.context_embedder(encoder_hidden_states,
timestep,
encoder_attention_mask)
encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask)
# 3. Attention mask preparation
latent_sequence_length = hidden_states.shape[1]
@@ -873,13 +765,11 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
device=hidden_states.device,
dtype=torch.bool) # [B, N, N]
effective_condition_sequence_length = encoder_attention_mask.sum(
dim=1, dtype=torch.int)
effective_condition_sequence_length = encoder_attention_mask.sum(dim=1, dtype=torch.int)
effective_sequence_length = latent_sequence_length + effective_condition_sequence_length
for i in range(batch_size):
attention_mask[i, :effective_sequence_length[i], :
effective_sequence_length[i]] = True
attention_mask[i, :effective_sequence_length[i], :effective_sequence_length[i]] = True
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
@@ -894,9 +784,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
return custom_forward
ckpt_kwargs: Dict[str, Any] = {
"use_reentrant": False
} if is_torch_version(">=", "1.11.0") else {}
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
for block in self.transformer_blocks:
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
@@ -922,23 +810,19 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
else:
for block in self.transformer_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states, encoder_hidden_states, temb, attention_mask,
image_rotary_emb)
hidden_states, encoder_hidden_states = block(hidden_states, encoder_hidden_states, temb, attention_mask,
image_rotary_emb)
for block in self.single_transformer_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states, encoder_hidden_states, temb, attention_mask,
image_rotary_emb)
hidden_states, encoder_hidden_states = block(hidden_states, encoder_hidden_states, temb, attention_mask,
image_rotary_emb)
# 5. Output projection
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size,
post_patch_num_frames,
post_patch_height,
post_patch_width, -1, p_t, p, p)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames, post_patch_height, post_patch_width,
-1, p_t, p, p)
hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7)
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
+59 -124
View File
@@ -20,22 +20,18 @@ import torch
import torch.nn.functional as F
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.loaders import HunyuanVideoLoraLoaderMixin
from diffusers.models import (AutoencoderKLHunyuanVideo,
HunyuanVideoTransformer3DModel)
from diffusers.pipelines.hunyuan_video.pipeline_output import \
HunyuanVideoPipelineOutput
from diffusers.models import AutoencoderKLHunyuanVideo, HunyuanVideoTransformer3DModel
from diffusers.pipelines.hunyuan_video.pipeline_output import HunyuanVideoPipelineOutput
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from einops import rearrange
from transformers import (CLIPTextModel, CLIPTokenizer, LlamaModel,
LlamaTokenizerFast)
from transformers import CLIPTextModel, CLIPTokenizer, LlamaModel, LlamaTokenizerFast
from fastvideo.utils.communications import all_gather
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -66,14 +62,13 @@ EXAMPLE_DOC_STRING = """
"""
DEFAULT_PROMPT_TEMPLATE = {
"template":
("<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"),
"template": ("<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"),
"crop_start":
95,
}
@@ -112,28 +107,22 @@ def retrieve_timesteps(
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError(
"Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
)
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
if timesteps is not None:
accepts_timesteps = "timesteps" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys())
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
f" timestep schedules. Please check whether you are using the correct scheduler.")
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys())
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
f" sigmas schedules. Please check whether you are using the correct scheduler.")
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
@@ -195,13 +184,10 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
)
self.vae_scale_factor_temporal = (self.vae.temporal_compression_ratio
if hasattr(self, "vae")
and self.vae is not None else 4)
if hasattr(self, "vae") and self.vae is not None else 4)
self.vae_scale_factor_spatial = (self.vae.spatial_compression_ratio
if hasattr(self, "vae")
and self.vae is not None else 8)
self.video_processor = VideoProcessor(
vae_scale_factor=self.vae_scale_factor_spatial)
if hasattr(self, "vae") and self.vae is not None else 8)
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
def _get_llama_prompt_embeds(
self,
@@ -263,12 +249,9 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt,
seq_len, -1)
prompt_attention_mask = prompt_attention_mask.repeat(
1, num_videos_per_prompt)
prompt_attention_mask = prompt_attention_mask.view(
batch_size * num_videos_per_prompt, seq_len)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
prompt_attention_mask = prompt_attention_mask.repeat(1, num_videos_per_prompt)
prompt_attention_mask = prompt_attention_mask.view(batch_size * num_videos_per_prompt, seq_len)
return prompt_embeds, prompt_attention_mask
@@ -295,25 +278,17 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
)
text_input_ids = text_inputs.input_ids
untruncated_ids = self.tokenizer_2(prompt,
padding="longest",
return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[
-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer_2.batch_decode(
untruncated_ids[:, max_sequence_length - 1:-1])
logger.warning(
"The following part of your input was truncated because CLIP can only handle sequences up to"
f" {max_sequence_length} tokens: {removed_text}")
untruncated_ids = self.tokenizer_2(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer_2.batch_decode(untruncated_ids[:, max_sequence_length - 1:-1])
logger.warning("The following part of your input was truncated because CLIP can only handle sequences up to"
f" {max_sequence_length} tokens: {removed_text}")
prompt_embeds = self.text_encoder_2(
text_input_ids.to(device),
output_hidden_states=False).pooler_output
prompt_embeds = self.text_encoder_2(text_input_ids.to(device), output_hidden_states=False).pooler_output
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt,
-1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, -1)
return prompt_embeds
@@ -365,13 +340,10 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
prompt_template=None,
):
if height % 16 != 0 or width % 16 != 0:
raise ValueError(
f"`height` and `width` have to be divisible by 16 but are {height} and {width}."
)
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs):
if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
@@ -386,28 +358,18 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
" only forward one of the two.")
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str)
and not isinstance(prompt, list)):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
)
elif prompt_2 is not None and (not isinstance(prompt_2, str)
and not isinstance(prompt_2, list)):
raise ValueError(
f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}"
)
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined.")
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
elif prompt_2 is not None and (not isinstance(prompt_2, str) and not isinstance(prompt_2, list)):
raise ValueError(f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}")
if prompt_template is not None:
if not isinstance(prompt_template, dict):
raise ValueError(
f"`prompt_template` has to be of type `dict` but is {type(prompt_template)}"
)
raise ValueError(f"`prompt_template` has to be of type `dict` but is {type(prompt_template)}")
if "template" not in prompt_template:
raise ValueError(
f"`prompt_template` has to contain a key `template` but only found {prompt_template.keys()}"
)
f"`prompt_template` has to contain a key `template` but only found {prompt_template.keys()}")
def prepare_latents(
self,
@@ -418,8 +380,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
num_frames: int = 129,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
generator: Optional[Union[torch.Generator,
List[torch.Generator]]] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if latents is not None:
@@ -435,13 +396,9 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
f" size of {batch_size}. Make sure the batch size matches the length of the generators.")
latents = randn_tensor(shape,
generator=generator,
device=device,
dtype=dtype)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
def enable_vae_slicing(self):
@@ -502,8 +459,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
sigmas: List[float] = None,
guidance_scale: float = 6.0,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator,
List[torch.Generator]]] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
pooled_prompt_embeds: Optional[torch.Tensor] = None,
@@ -511,8 +467,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
output_type: Optional[str] = "pil",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[Union[Callable[[int, int, Dict],
None], PipelineCallback,
callback_on_step_end: Optional[Union[Callable[[int, int, Dict], None], PipelineCallback,
MultiPipelineCallbacks]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
prompt_template: Dict[str, Any] = DEFAULT_PROMPT_TEMPLATE,
@@ -591,8 +546,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
"""
if isinstance(callback_on_step_end,
(PipelineCallback, MultiPipelineCallbacks)):
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
# 1. Check inputs. Raise error if not correct
@@ -640,8 +594,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
pooled_prompt_embeds = pooled_prompt_embeds.to(transformer_dtype)
# 4. Prepare timesteps
sigmas = np.linspace(1.0, 0.0, num_inference_steps +
1)[:-1] if sigmas is None else sigmas
sigmas = np.linspace(1.0, 0.0, num_inference_steps + 1)[:-1] if sigmas is None else sigmas
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -651,8 +604,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
# 5. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
num_latent_frames = (num_frames -
1) // self.vae_scale_factor_temporal + 1
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
@@ -668,19 +620,14 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
# check sequence_parallel
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()
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
# 6. Prepare guidance condition
guidance = torch.tensor([guidance_scale] * latents.shape[0],
dtype=transformer_dtype,
device=device) * 1000.0
guidance = torch.tensor([guidance_scale] * latents.shape[0], dtype=transformer_dtype, device=device) * 1000.0
# 7. Denoising loop
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
with self.progress_bar(total=num_inference_steps) as progress_bar:
@@ -694,17 +641,14 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
if pooled_prompt_embeds.shape[-1] != prompt_embeds.shape[-1]:
pooled_prompt_embeds_padding = F.pad(
pooled_prompt_embeds,
(0, prompt_embeds.shape[2] -
pooled_prompt_embeds.shape[1]),
(0, prompt_embeds.shape[2] - pooled_prompt_embeds.shape[1]),
value=0,
).unsqueeze(1)
encoder_hidden_states = torch.cat(
[pooled_prompt_embeds_padding, prompt_embeds], dim=1)
encoder_hidden_states = torch.cat([pooled_prompt_embeds_padding, prompt_embeds], dim=1)
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=
encoder_hidden_states, # [1, 257, 4096]
encoder_hidden_states=encoder_hidden_states, # [1, 257, 4096]
timestep=timestep,
encoder_attention_mask=prompt_attention_mask,
guidance=guidance,
@@ -713,37 +657,28 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
)[0]
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(
self, i, t, callback_kwargs)
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop(
"prompt_embeds", prompt_embeds)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if get_sequence_parallel_state():
latents = all_gather(latents, dim=2)
if not output_type == "latent":
latents = latents.to(
self.vae.dtype) / self.vae.config.scaling_factor
latents = latents.to(self.vae.dtype) / self.vae.config.scaling_factor
video = self.vae.decode(latents, return_dict=False)[0]
video = self.video_processor.postprocess_video(
video, output_type=output_type)
video = self.video_processor.postprocess_video(video, output_type=output_type)
else:
video = latents
@@ -6,18 +6,9 @@ from safetensors.torch import save_file
parser = argparse.ArgumentParser()
parser.add_argument("--diffusers_path", required=True, type=str)
parser.add_argument("--transformer_path",
type=str,
default=None,
help="Path to save transformer model")
parser.add_argument("--vae_encoder_path",
type=str,
default=None,
help="Path to save VAE encoder model")
parser.add_argument("--vae_decoder_path",
type=str,
default=None,
help="Path to save VAE decoder model")
parser.add_argument("--transformer_path", type=str, default=None, help="Path to save transformer model")
parser.add_argument("--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model")
parser.add_argument("--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model")
args = parser.parse_args()
@@ -39,36 +30,22 @@ def convert_diffusers_transformer_to_mochi(state_dict):
new_state_dict = {}
# Convert patch_embed
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop(
"patch_embed.proj.weight")
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop(
"patch_embed.proj.bias")
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop("patch_embed.proj.weight")
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop("patch_embed.proj.bias")
# Convert time_embed
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_1.weight")
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_1.bias")
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_2.weight")
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_2.bias")
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop(
"time_embed.pooler.to_kv.weight")
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop(
"time_embed.pooler.to_kv.bias")
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop(
"time_embed.pooler.to_q.weight")
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop(
"time_embed.pooler.to_q.bias")
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop(
"time_embed.pooler.to_out.weight")
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop(
"time_embed.pooler.to_out.bias")
new_state_dict["t5_yproj.weight"] = original_state_dict.pop(
"time_embed.caption_proj.weight")
new_state_dict["t5_yproj.bias"] = original_state_dict.pop(
"time_embed.caption_proj.bias")
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.weight")
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.bias")
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.weight")
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.bias")
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop("time_embed.pooler.to_kv.weight")
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop("time_embed.pooler.to_kv.bias")
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop("time_embed.pooler.to_q.weight")
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop("time_embed.pooler.to_q.bias")
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop("time_embed.pooler.to_out.weight")
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop("time_embed.pooler.to_out.bias")
new_state_dict["t5_yproj.weight"] = original_state_dict.pop("time_embed.caption_proj.weight")
new_state_dict["t5_yproj.bias"] = original_state_dict.pop("time_embed.caption_proj.bias")
# Convert transformer blocks
num_layers = 48
@@ -77,25 +54,19 @@ def convert_diffusers_transformer_to_mochi(state_dict):
new_prefix = f"blocks.{i}."
# norm1
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(
block_prefix + "norm1.linear.weight")
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(
block_prefix + "norm1.linear.bias")
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(block_prefix + "norm1.linear.weight")
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(block_prefix + "norm1.linear.bias")
if i < num_layers - 1:
new_state_dict[new_prefix +
"mod_y.weight"] = original_state_dict.pop(
block_prefix + "norm1_context.linear.weight")
new_state_dict[new_prefix +
"mod_y.bias"] = original_state_dict.pop(
block_prefix + "norm1_context.linear.bias")
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(block_prefix +
"norm1_context.linear.weight")
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(block_prefix +
"norm1_context.linear.bias")
else:
new_state_dict[new_prefix +
"mod_y.weight"] = original_state_dict.pop(
block_prefix + "norm1_context.linear_1.weight")
new_state_dict[new_prefix +
"mod_y.bias"] = original_state_dict.pop(
block_prefix + "norm1_context.linear_1.bias")
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(block_prefix +
"norm1_context.linear_1.weight")
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(block_prefix +
"norm1_context.linear_1.bias")
# Visual attention
q = original_state_dict.pop(block_prefix + "attn1.to_q.weight")
@@ -104,18 +75,13 @@ def convert_diffusers_transformer_to_mochi(state_dict):
qkv_weight = torch.cat([q, k, v], dim=0)
new_state_dict[new_prefix + "attn.qkv_x.weight"] = qkv_weight
new_state_dict[new_prefix +
"attn.q_norm_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_q.weight")
new_state_dict[new_prefix +
"attn.k_norm_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_k.weight")
new_state_dict[new_prefix +
"attn.proj_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.to_out.0.weight")
new_state_dict[new_prefix +
"attn.proj_x.bias"] = original_state_dict.pop(
block_prefix + "attn1.to_out.0.bias")
new_state_dict[new_prefix + "attn.q_norm_x.weight"] = original_state_dict.pop(block_prefix +
"attn1.norm_q.weight")
new_state_dict[new_prefix + "attn.k_norm_x.weight"] = original_state_dict.pop(block_prefix +
"attn1.norm_k.weight")
new_state_dict[new_prefix + "attn.proj_x.weight"] = original_state_dict.pop(block_prefix +
"attn1.to_out.0.weight")
new_state_dict[new_prefix + "attn.proj_x.bias"] = original_state_dict.pop(block_prefix + "attn1.to_out.0.bias")
# Context attention
q = original_state_dict.pop(block_prefix + "attn1.add_q_proj.weight")
@@ -124,46 +90,34 @@ def convert_diffusers_transformer_to_mochi(state_dict):
qkv_weight = torch.cat([q, k, v], dim=0)
new_state_dict[new_prefix + "attn.qkv_y.weight"] = qkv_weight
new_state_dict[new_prefix +
"attn.q_norm_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_added_q.weight")
new_state_dict[new_prefix +
"attn.k_norm_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_added_k.weight")
new_state_dict[new_prefix + "attn.q_norm_y.weight"] = original_state_dict.pop(block_prefix +
"attn1.norm_added_q.weight")
new_state_dict[new_prefix + "attn.k_norm_y.weight"] = original_state_dict.pop(block_prefix +
"attn1.norm_added_k.weight")
if i < num_layers - 1:
new_state_dict[new_prefix +
"attn.proj_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.to_add_out.weight")
new_state_dict[new_prefix +
"attn.proj_y.bias"] = original_state_dict.pop(
block_prefix + "attn1.to_add_out.bias")
new_state_dict[new_prefix + "attn.proj_y.weight"] = original_state_dict.pop(block_prefix +
"attn1.to_add_out.weight")
new_state_dict[new_prefix + "attn.proj_y.bias"] = original_state_dict.pop(block_prefix +
"attn1.to_add_out.bias")
# MLP
new_state_dict[new_prefix + "mlp_x.w1.weight"] = reverse_proj_gate(
original_state_dict.pop(block_prefix + "ff.net.0.proj.weight"))
new_state_dict[new_prefix +
"mlp_x.w2.weight"] = original_state_dict.pop(
block_prefix + "ff.net.2.weight")
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(block_prefix + "ff.net.2.weight")
if i < num_layers - 1:
new_state_dict[new_prefix + "mlp_y.w1.weight"] = reverse_proj_gate(
original_state_dict.pop(block_prefix +
"ff_context.net.0.proj.weight"))
new_state_dict[new_prefix +
"mlp_y.w2.weight"] = original_state_dict.pop(
block_prefix + "ff_context.net.2.weight")
original_state_dict.pop(block_prefix + "ff_context.net.0.proj.weight"))
new_state_dict[new_prefix + "mlp_y.w2.weight"] = original_state_dict.pop(block_prefix +
"ff_context.net.2.weight")
# Output layers
new_state_dict["final_layer.mod.weight"] = reverse_scale_shift(
original_state_dict.pop("norm_out.linear.weight"), dim=0)
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(
original_state_dict.pop("norm_out.linear.bias"), dim=0)
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop(
"proj_out.weight")
new_state_dict["final_layer.linear.bias"] = original_state_dict.pop(
"proj_out.bias")
new_state_dict["final_layer.mod.weight"] = reverse_scale_shift(original_state_dict.pop("norm_out.linear.weight"),
dim=0)
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(original_state_dict.pop("norm_out.linear.bias"), dim=0)
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop("proj_out.weight")
new_state_dict["final_layer.linear.bias"] = original_state_dict.pop("proj_out.bias")
new_state_dict["pos_frequencies"] = original_state_dict.pop(
"pos_frequencies")
new_state_dict["pos_frequencies"] = original_state_dict.pop("pos_frequencies")
print("Remaining Keys:", original_state_dict.keys())
@@ -178,271 +132,182 @@ def convert_diffusers_vae_to_mochi(state_dict):
# Convert encoder
prefix = "encoder."
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(
f"{prefix}proj_in.weight")
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(
f"{prefix}proj_in.bias")
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(f"{prefix}proj_in.weight")
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(f"{prefix}proj_in.bias")
# Convert block_in
for i in range(3):
encoder_state_dict[
f"layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight")
encoder_state_dict[
f"layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias")
encoder_state_dict[
f"layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight")
encoder_state_dict[
f"layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias")
encoder_state_dict[
f"layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight")
encoder_state_dict[
f"layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias")
encoder_state_dict[
f"layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight")
encoder_state_dict[
f"layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
encoder_state_dict[f"layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight")
encoder_state_dict[f"layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias")
encoder_state_dict[f"layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight")
encoder_state_dict[f"layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias")
encoder_state_dict[f"layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight")
encoder_state_dict[f"layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias")
encoder_state_dict[f"layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight")
encoder_state_dict[f"layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
# Convert down_blocks
down_block_layers = [3, 4, 6]
for block in range(3):
encoder_state_dict[
f"layers.{block+4}.layers.0.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.weight")
encoder_state_dict[
f"layers.{block+4}.layers.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.bias")
encoder_state_dict[f"layers.{block+4}.layers.0.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.weight")
encoder_state_dict[f"layers.{block+4}.layers.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.bias")
for i in range(down_block_layers[block]):
# Convert resnets
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias")
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias")
# Convert attentions
q = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight")
k = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight")
v = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight")
q = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight")
k = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight")
v = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight")
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight")
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias")
# Convert block_out
for i in range(3):
encoder_state_dict[
f"layers.{i+7}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight")
encoder_state_dict[
f"layers.{i+7}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias")
encoder_state_dict[
f"layers.{i+7}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight")
encoder_state_dict[
f"layers.{i+7}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias")
encoder_state_dict[
f"layers.{i+7}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight")
encoder_state_dict[
f"layers.{i+7}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias")
encoder_state_dict[
f"layers.{i+7}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight")
encoder_state_dict[
f"layers.{i+7}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
encoder_state_dict[f"layers.{i+7}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight")
encoder_state_dict[f"layers.{i+7}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias")
encoder_state_dict[f"layers.{i+7}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight")
encoder_state_dict[f"layers.{i+7}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias")
encoder_state_dict[f"layers.{i+7}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight")
encoder_state_dict[f"layers.{i+7}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias")
encoder_state_dict[f"layers.{i+7}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight")
encoder_state_dict[f"layers.{i+7}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
q = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_q.weight")
k = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_k.weight")
v = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_v.weight")
q = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_q.weight")
k = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_k.weight")
v = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
encoder_state_dict[
f"layers.{i+7}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[f"layers.{i+7}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[
f"layers.{i+7}.attn_block.attn.out.weight"] = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_out.0.weight")
encoder_state_dict[
f"layers.{i+7}.attn_block.attn.out.bias"] = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_out.0.bias")
encoder_state_dict[
f"layers.{i+7}.attn_block.norm.weight"] = original_state_dict.pop(
f"{prefix}block_out.norms.{i}.norm_layer.weight")
encoder_state_dict[
f"layers.{i+7}.attn_block.norm.bias"] = original_state_dict.pop(
f"{prefix}block_out.norms.{i}.norm_layer.bias")
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.weight"] = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_out.0.weight")
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.bias"] = original_state_dict.pop(
f"{prefix}block_out.attentions.{i}.to_out.0.bias")
encoder_state_dict[f"layers.{i+7}.attn_block.norm.weight"] = original_state_dict.pop(
f"{prefix}block_out.norms.{i}.norm_layer.weight")
encoder_state_dict[f"layers.{i+7}.attn_block.norm.bias"] = original_state_dict.pop(
f"{prefix}block_out.norms.{i}.norm_layer.bias")
# Convert output layers
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(
f"{prefix}norm_out.norm_layer.weight")
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(
f"{prefix}norm_out.norm_layer.bias")
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(
f"{prefix}proj_out.weight")
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.weight")
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.bias")
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
# Convert decoder
prefix = "decoder."
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(
f"{prefix}conv_in.weight")
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(
f"{prefix}conv_in.bias")
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(f"{prefix}conv_in.weight")
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(f"{prefix}conv_in.bias")
# Convert block_in
for i in range(3):
decoder_state_dict[
f"blocks.0.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight")
decoder_state_dict[
f"blocks.0.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias")
decoder_state_dict[
f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight")
decoder_state_dict[
f"blocks.0.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias")
decoder_state_dict[
f"blocks.0.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight")
decoder_state_dict[
f"blocks.0.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias")
decoder_state_dict[
f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight")
decoder_state_dict[
f"blocks.0.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
decoder_state_dict[f"blocks.0.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight")
decoder_state_dict[f"blocks.0.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias")
decoder_state_dict[f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight")
decoder_state_dict[f"blocks.0.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias")
decoder_state_dict[f"blocks.0.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight")
decoder_state_dict[f"blocks.0.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias")
decoder_state_dict[f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight")
decoder_state_dict[f"blocks.0.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
# Convert up_blocks
up_block_layers = [6, 4, 3]
for block in range(3):
for i in range(up_block_layers[block]):
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight")
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias")
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight")
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias")
decoder_state_dict[
f"blocks.{block+1}.proj.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.weight")
decoder_state_dict[
f"blocks.{block+1}.proj.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.bias")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight")
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias")
decoder_state_dict[f"blocks.{block+1}.proj.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.weight")
decoder_state_dict[f"blocks.{block+1}.proj.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.bias")
# Convert block_out
for i in range(3):
decoder_state_dict[
f"blocks.4.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight")
decoder_state_dict[
f"blocks.4.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias")
decoder_state_dict[
f"blocks.4.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight")
decoder_state_dict[
f"blocks.4.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias")
decoder_state_dict[
f"blocks.4.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight")
decoder_state_dict[
f"blocks.4.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias")
decoder_state_dict[
f"blocks.4.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight")
decoder_state_dict[
f"blocks.4.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
decoder_state_dict[f"blocks.4.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight")
decoder_state_dict[f"blocks.4.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias")
decoder_state_dict[f"blocks.4.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight")
decoder_state_dict[f"blocks.4.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias")
decoder_state_dict[f"blocks.4.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight")
decoder_state_dict[f"blocks.4.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias")
decoder_state_dict[f"blocks.4.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight")
decoder_state_dict[f"blocks.4.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
# Convert output layers
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(
f"{prefix}proj_out.weight")
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(
f"{prefix}proj_out.bias")
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(f"{prefix}proj_out.bias")
return encoder_state_dict, decoder_state_dict
@@ -469,8 +334,7 @@ def main(args):
ensure_directory_exists(transformer_path)
print("Converting transformer model...")
transformer_state_dict = convert_diffusers_transformer_to_mochi(
pipe.transformer.state_dict())
transformer_state_dict = convert_diffusers_transformer_to_mochi(pipe.transformer.state_dict())
save_file(transformer_state_dict, transformer_path)
print(f"Saved transformer to {transformer_path}")
@@ -482,8 +346,7 @@ def main(args):
ensure_directory_exists(decoder_path)
print("Converting VAE models...")
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(
pipe.vae.state_dict())
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(pipe.vae.state_dict())
save_file(encoder_state_dict, encoder_path)
print(f"Saved VAE encoder to {encoder_path}")
@@ -491,9 +354,7 @@ def main(args):
save_file(decoder_state_dict, decoder_path)
print(f"Saved VAE decoder to {decoder_path}")
elif args.vae_encoder_path or args.vae_decoder_path:
print(
"Warning: Both VAE encoder and decoder paths must be specified to convert VAE models."
)
print("Warning: Both VAE encoder and decoder paths must be specified to convert VAE models.")
if __name__ == "__main__":
+44 -110
View File
@@ -21,22 +21,18 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import PeftAdapterMixin
from diffusers.models.attention import FeedForward as HF_FeedForward
from diffusers.models.attention_processor import Attention
from diffusers.models.embeddings import (MochiCombinedTimestepCaptionEmbedding,
PatchEmbed)
from diffusers.models.embeddings import MochiCombinedTimestepCaptionEmbedding, PatchEmbed
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import AdaLayerNormContinuous
from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, logging,
scale_lora_layers, unscale_lora_layers)
from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers
from diffusers.utils.torch_utils import maybe_allow_in_graph
from liger_kernel.ops.swiglu import LigerSiLUMulFunction
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.models.mochi_hf.norm import (MochiLayerNormContinuous,
MochiModulatedRMSNorm,
MochiRMSNorm, MochiRMSNormZero)
from fastvideo.models.mochi_hf.norm import (MochiLayerNormContinuous, MochiModulatedRMSNorm, MochiRMSNorm,
MochiRMSNormZero)
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -54,8 +50,7 @@ class FeedForward(HF_FeedForward):
inner_dim=None,
bias: bool = True,
):
super().__init__(dim, dim_out, mult, dropout, activation_fn,
final_dropout, inner_dim, bias)
super().__init__(dim, dim_out, mult, dropout, activation_fn, final_dropout, inner_dim, bias)
assert activation_fn == "swiglu"
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
@@ -100,26 +95,17 @@ class MochiAttention(nn.Module):
self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.add_k_proj = nn.Linear(added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias)
self.add_v_proj = nn.Linear(added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias)
self.add_k_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
self.add_v_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
if self.context_pre_only is not None:
self.add_q_proj = nn.Linear(added_kv_proj_dim,
self.inner_dim,
bias=added_proj_bias)
self.add_q_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
self.to_out = nn.ModuleList([])
self.to_out.append(
nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Dropout(dropout))
if not self.context_pre_only:
self.to_add_out = nn.Linear(self.inner_dim,
self.out_context_dim,
bias=out_bias)
self.to_add_out = nn.Linear(self.inner_dim, self.out_context_dim, bias=out_bias)
self.processor = processor
@@ -144,9 +130,7 @@ class MochiAttnProcessor2_0:
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0."
)
raise ImportError("MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
@@ -198,9 +182,7 @@ class MochiAttnProcessor2_0:
def shrink_head(encoder_state, dim):
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
return encoder_state.narrow(
dim, nccl_info.rank_within_group * local_heads,
local_heads)
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
encoder_query = shrink_head(encoder_query, dim=2)
encoder_key = shrink_head(encoder_key, dim=2)
@@ -241,11 +223,7 @@ class MochiAttnProcessor2_0:
attn_mask = encoder_attention_mask[:, :].bool()
attn_mask = F.pad(attn_mask, (sequence_length, 0), value=True)
hidden_states = flash_attn_no_pad(qkv,
attn_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None)
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
# hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask = None, dropout_p=0.0, is_causal=False)
@@ -258,11 +236,8 @@ class MochiAttnProcessor2_0:
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1)
# B, S, H, D
hidden_states = all_to_all_4D(hidden_states,
scatter_dim=1,
gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states,
dim=2).contiguous()
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.flatten(2, 3)
@@ -324,16 +299,10 @@ class MochiTransformerBlock(nn.Module):
self.ff_inner_dim = (4 * dim * 2) // 3
self.ff_context_inner_dim = (4 * pooled_projection_dim * 2) // 3
self.norm1 = MochiRMSNormZero(dim,
4 * dim,
eps=eps,
elementwise_affine=False)
self.norm1 = MochiRMSNormZero(dim, 4 * dim, eps=eps, elementwise_affine=False)
if not context_pre_only:
self.norm1_context = MochiRMSNormZero(dim,
4 * pooled_projection_dim,
eps=eps,
elementwise_affine=False)
self.norm1_context = MochiRMSNormZero(dim, 4 * pooled_projection_dim, eps=eps, elementwise_affine=False)
else:
self.norm1_context = MochiLayerNormContinuous(
embedding_dim=pooled_projection_dim,
@@ -357,17 +326,12 @@ class MochiTransformerBlock(nn.Module):
# TODO(aryan): norm_context layers are not needed when `context_pre_only` is True
self.norm2 = MochiModulatedRMSNorm(eps=eps)
self.norm2_context = (MochiModulatedRMSNorm(
eps=eps) if not self.context_pre_only else None)
self.norm2_context = (MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None)
self.norm3 = MochiModulatedRMSNorm(eps)
self.norm3_context = (MochiModulatedRMSNorm(
eps=eps) if not self.context_pre_only else None)
self.norm3_context = (MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None)
self.ff = FeedForward(dim,
inner_dim=self.ff_inner_dim,
activation_fn=activation_fn,
bias=False)
self.ff = FeedForward(dim, inner_dim=self.ff_inner_dim, activation_fn=activation_fn, bias=False)
self.ff_context = None
if not context_pre_only:
self.ff_context = FeedForward(
@@ -389,8 +353,7 @@ class MochiTransformerBlock(nn.Module):
image_rotary_emb: Optional[torch.Tensor] = None,
output_attn=False,
) -> Tuple[torch.Tensor, torch.Tensor]:
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(
hidden_states, temb)
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb)
if not self.context_pre_only:
(
@@ -400,8 +363,7 @@ class MochiTransformerBlock(nn.Module):
enc_gate_mlp,
) = self.norm1_context(encoder_hidden_states, temb)
else:
norm_encoder_hidden_states = self.norm1_context(
encoder_hidden_states, temb)
norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states, temb)
attn_hidden_states, context_attn_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
@@ -410,28 +372,21 @@ class MochiTransformerBlock(nn.Module):
encoder_attention_mask=encoder_attention_mask,
)
hidden_states = hidden_states + self.norm2(
attn_hidden_states,
torch.tanh(gate_msa).unsqueeze(1))
norm_hidden_states = self.norm3(
hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32)))
hidden_states = hidden_states + self.norm2(attn_hidden_states, torch.tanh(gate_msa).unsqueeze(1))
norm_hidden_states = self.norm3(hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32)))
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + self.norm4(
ff_output,
torch.tanh(gate_mlp).unsqueeze(1))
hidden_states = hidden_states + self.norm4(ff_output, torch.tanh(gate_mlp).unsqueeze(1))
if not self.context_pre_only:
encoder_hidden_states = encoder_hidden_states + self.norm2_context(
context_attn_hidden_states,
torch.tanh(enc_gate_msa).unsqueeze(1))
encoder_hidden_states = encoder_hidden_states + self.norm2_context(context_attn_hidden_states,
torch.tanh(enc_gate_msa).unsqueeze(1))
norm_encoder_hidden_states = self.norm3_context(
encoder_hidden_states,
(1 + enc_scale_mlp.unsqueeze(1).to(torch.float32)),
)
context_ff_output = self.ff_context(norm_encoder_hidden_states)
encoder_hidden_states = encoder_hidden_states + self.norm4_context(
context_ff_output,
torch.tanh(enc_gate_mlp).unsqueeze(1))
encoder_hidden_states = encoder_hidden_states + self.norm4_context(context_ff_output,
torch.tanh(enc_gate_mlp).unsqueeze(1))
if not output_attn:
attn_hidden_states = None
@@ -455,11 +410,7 @@ class MochiRoPE(nn.Module):
self.target_area = base_height * base_width
def _centers(self, start, stop, num, device, dtype) -> torch.Tensor:
edges = torch.linspace(start,
stop,
num + 1,
device=device,
dtype=dtype)
edges = torch.linspace(start, stop, num + 1, device=device, dtype=dtype)
return (edges[:-1] + edges[1:]) / 2
def _get_positions(
@@ -471,21 +422,16 @@ class MochiRoPE(nn.Module):
dtype: Optional[torch.dtype] = None,
) -> torch.Tensor:
scale = (self.target_area / (height * width))**0.5
t = torch.arange(num_frames * nccl_info.sp_size,
device=device,
dtype=dtype)
h = self._centers(-height * scale / 2, height * scale / 2, height,
device, dtype)
w = self._centers(-width * scale / 2, width * scale / 2, width, device,
dtype)
t = torch.arange(num_frames * nccl_info.sp_size, device=device, dtype=dtype)
h = self._centers(-height * scale / 2, height * scale / 2, height, device, dtype)
w = self._centers(-width * scale / 2, width * scale / 2, width, device, dtype)
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
positions = torch.stack([grid_t, grid_h, grid_w], dim=-1).view(-1, 3)
return positions
def _create_rope(self, freqs: torch.Tensor,
pos: torch.Tensor) -> torch.Tensor:
def _create_rope(self, freqs: torch.Tensor, pos: torch.Tensor) -> torch.Tensor:
with torch.autocast(freqs.device.type, enabled=False):
# Always run ROPE freqs computation in FP32
freqs = torch.einsum(
@@ -578,8 +524,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
num_attention_heads=8,
)
self.pos_frequencies = nn.Parameter(
torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0))
self.pos_frequencies = nn.Parameter(torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0))
self.rope = MochiRoPE()
self.transformer_blocks = nn.ModuleList([
@@ -601,8 +546,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
eps=1e-6,
norm_type="layer_norm",
)
self.proj_out = nn.Linear(inner_dim,
patch_size * patch_size * out_channels)
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
self.gradient_checkpointing = False
@@ -621,8 +565,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
) -> torch.Tensor:
assert (return_dict is False
), "return_dict is not supported in MochiTransformer3DModel"
assert (return_dict is False), "return_dict is not supported in MochiTransformer3DModel"
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
@@ -634,11 +577,8 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if (attention_kwargs is not None
and attention_kwargs.get("scale", None) is not None):
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
if (attention_kwargs is not None and attention_kwargs.get("scale", None) is not None):
logger.warning("Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective.")
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p = self.config.patch_size
@@ -656,8 +596,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1)
hidden_states = self.patch_embed(hidden_states)
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(
1, 2)
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2)
image_rotary_emb = self.rope(
self.pos_frequencies,
@@ -678,9 +617,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
return custom_forward
ckpt_kwargs: Dict[str, Any] = ({
"use_reentrant": False
} if is_torch_version(">=", "1.11.0") else {})
ckpt_kwargs: Dict[str, Any] = ({"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {})
(
hidden_states,
encoder_hidden_states,
@@ -710,12 +647,9 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, num_frames,
post_patch_height,
post_patch_width, p, p, -1)
hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1)
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5)
output = hidden_states.reshape(batch_size, -1, num_frames, height,
width)
output = hidden_states.reshape(batch_size, -1, num_frames, height, width)
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
+4 -8
View File
@@ -78,9 +78,7 @@ class MochiLayerNormContinuous(nn.Module):
# AdaLN
self.silu = nn.SiLU()
self.linear_1 = nn.Linear(conditioning_embedding_dim,
embedding_dim,
bias=bias)
self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias)
self.norm = MochiModulatedRMSNorm(eps=eps)
def forward(
@@ -117,16 +115,14 @@ class MochiRMSNormZero(nn.Module):
self.linear = nn.Linear(embedding_dim, hidden_dim)
self.norm = MochiModulatedRMSNorm(eps=eps)
def forward(
self, hidden_states: torch.Tensor, emb: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
def forward(self, hidden_states: torch.Tensor,
emb: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
hidden_states_dtype = hidden_states.dtype
emb = self.linear(self.silu(emb))
scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1)
hidden_states = self.norm(hidden_states,
(1 + scale_msa[:, None].to(torch.float32)))
hidden_states = self.norm(hidden_states, (1 + scale_msa[:, None].to(torch.float32)))
hidden_states = hidden_states.to(hidden_states_dtype)
return hidden_states, gate_msa, scale_mlp, gate_mlp
+60 -132
View File
@@ -24,8 +24,7 @@ from diffusers.models.autoencoders import AutoencoderKL
from diffusers.pipelines.mochi.pipeline_output import MochiPipelineOutput
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import (is_torch_xla_available, logging,
replace_example_docstring)
from diffusers.utils import is_torch_xla_available, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from einops import rearrange
@@ -33,8 +32,7 @@ from transformers import T5EncoderModel, T5TokenizerFast
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.utils.communications import all_gather
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
@@ -78,19 +76,14 @@ def calculate_shift(
def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
if linear_steps is None:
linear_steps = num_steps // 2
linear_sigma_schedule = [
i * threshold_noise / linear_steps for i in range(linear_steps)
]
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
quadratic_steps = num_steps - linear_steps
quadratic_coef = threshold_noise_step_diff / (linear_steps *
quadratic_steps**2)
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (
quadratic_steps**2)
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
const = quadratic_coef * (linear_steps**2)
quadratic_sigma_schedule = [
quadratic_coef * (i**2) + linear_coef * i + const
for i in range(linear_steps, num_steps)
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, num_steps)
]
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule
sigma_schedule = [1.0 - x for x in sigma_schedule]
@@ -130,28 +123,22 @@ def retrieve_timesteps(
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError(
"Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
)
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
if timesteps is not None:
accepts_timesteps = "timesteps" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys())
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
f" timestep schedules. Please check whether you are using the correct scheduler.")
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys())
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
f" sigmas schedules. Please check whether you are using the correct scheduler.")
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
@@ -187,9 +174,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
model_cpu_offload_seq = "text_encoder->transformer->vae"
_optional_components = []
_callback_tensor_inputs = [
"latents", "prompt_embeds", "negative_prompt_embeds"
]
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(
self,
@@ -212,11 +197,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
self.vae_temporal_scale_factor = 6
self.patch_size = 2
self.video_processor = VideoProcessor(
vae_scale_factor=self.vae_spatial_scale_factor)
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_spatial_scale_factor)
self.tokenizer_max_length = (self.tokenizer.model_max_length
if hasattr(self, "tokenizer")
and self.tokenizer is not None else 77)
if hasattr(self, "tokenizer") and self.tokenizer is not None else 77)
self.default_height = 480
self.default_width = 848
@@ -247,31 +230,23 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
untruncated_ids = self.tokenizer(prompt,
padding="longest",
return_tensors="pt").input_ids
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[
-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(
untruncated_ids[:, max_sequence_length - 1:-1])
logger.warning(
"The following part of your input was truncated because `max_sequence_length` is set to "
f" {max_sequence_length} tokens: {removed_text}")
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1:-1])
logger.warning("The following part of your input was truncated because `max_sequence_length` is set to "
f" {max_sequence_length} tokens: {removed_text}")
prompt_embeds = self.text_encoder(
text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt,
seq_len, -1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
prompt_attention_mask = prompt_attention_mask.repeat(
num_videos_per_prompt, 1)
prompt_attention_mask = prompt_attention_mask.repeat(num_videos_per_prompt, 1)
return prompt_embeds, prompt_attention_mask
@@ -335,11 +310,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
if do_classifier_free_guidance and negative_prompt_embeds is None:
negative_prompt = negative_prompt or ""
negative_prompt = (batch_size * [negative_prompt] if isinstance(
negative_prompt, str) else negative_prompt)
negative_prompt = (batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt)
if prompt is not None and type(prompt) is not type(
negative_prompt):
if prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}.")
@@ -379,13 +352,10 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
negative_prompt_attention_mask=None,
):
if height % 8 != 0 or width % 8 != 0:
raise ValueError(
f"`height` and `width` have to be divisible by 8 but are {height} and {width}."
)
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs):
if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs
for k in callback_on_step_end_tensor_inputs):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
@@ -396,24 +366,15 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
" only forward one of the two.")
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str)
and not isinstance(prompt, list)):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
)
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined.")
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if prompt_embeds is not None and prompt_attention_mask is None:
raise ValueError(
"Must provide `prompt_attention_mask` when specifying `prompt_embeds`."
)
raise ValueError("Must provide `prompt_attention_mask` when specifying `prompt_embeds`.")
if (negative_prompt_embeds is not None
and negative_prompt_attention_mask is None):
raise ValueError(
"Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`."
)
if (negative_prompt_embeds is not None and negative_prompt_attention_mask is None):
raise ValueError("Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`.")
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
@@ -479,13 +440,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
f" size of {batch_size}. Make sure the batch size matches the length of the generators.")
latents = randn_tensor(shape,
generator=generator,
device=device,
dtype=torch.float32)
latents = randn_tensor(shape, generator=generator, device=device, dtype=torch.float32)
latents = latents.to(dtype)
return latents
@@ -522,8 +479,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
timesteps: List[int] = None,
guidance_scale: float = 4.5,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator,
List[torch.Generator]]] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None,
@@ -532,8 +488,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
output_type: Optional[str] = "pil",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[Callable[[int, int, Dict],
None]] = None,
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 256,
return_all_states=False,
@@ -612,8 +567,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
is returned where the first element is a list with the generated images.
"""
if isinstance(callback_on_step_end,
(PipelineCallback, MultiPipelineCallbacks)):
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
height = height or self.default_height
@@ -624,8 +578,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
prompt=prompt,
height=height,
width=width,
callback_on_step_end_tensor_inputs=
callback_on_step_end_tensor_inputs,
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
@@ -665,10 +618,8 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
device=device,
)
if self.do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds],
dim=0)
prompt_attention_mask = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask], dim=0)
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
# 4. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
@@ -685,17 +636,14 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
)
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()
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
original_noise = copy.deepcopy(latents)
# 5. Prepare timestep
# from https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
threshold_noise = 0.025
sigmas = linear_quadratic_schedule(num_inference_steps,
threshold_noise)
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
sigmas = np.array(sigmas)
# check if of type FlowMatchEulerDiscreteScheduler
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
@@ -712,25 +660,19 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
num_inference_steps,
device,
)
num_warmup_steps = max(
len(timesteps) - num_inference_steps * self.scheduler.order, 0)
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
self._num_timesteps = len(timesteps)
# 6. Denoising loop
self._progress_bar_config = {
"disable": nccl_info.rank_within_group != 0
}
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
latent_model_input = (torch.cat(
[latents] *
2) if self.do_classifier_free_guidance else latents)
latent_model_input = (torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0]).to(
latents.dtype)
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
noise_pred = self.transformer(
hidden_states=latent_model_input,
@@ -745,15 +687,11 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
noise_pred = noise_pred.to(torch.float32)
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (
noise_pred_text - noise_pred_uncond)
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
# compute the previous noisy sample x_t -> x_t-1
latents_dtype = latents.dtype
latents = self.scheduler.step(noise_pred,
t,
latents.to(torch.float32),
return_dict=False)[0]
latents = self.scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
latents = latents.to(latents_dtype)
if latents.dtype != latents_dtype:
@@ -765,17 +703,13 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(
self, i, t, callback_kwargs)
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop(
"prompt_embeds", prompt_embeds)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
@@ -795,25 +729,19 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
else:
# unscale/denormalize the latents
# denormalize with the mean and std if available and not None
has_latents_mean = (hasattr(self.vae.config, "latents_mean")
and self.vae.config.latents_mean is not None)
has_latents_std = (hasattr(self.vae.config, "latents_std")
and self.vae.config.latents_std is not None)
has_latents_mean = (hasattr(self.vae.config, "latents_mean") and self.vae.config.latents_mean is not None)
has_latents_std = (hasattr(self.vae.config, "latents_std") and self.vae.config.latents_std is not None)
if has_latents_mean and has_latents_std:
latents_mean = (torch.tensor(
self.vae.config.latents_mean).view(1, 12, 1, 1, 1).to(
latents.device, latents.dtype))
latents_std = (torch.tensor(self.vae.config.latents_std).view(
1, 12, 1, 1, 1).to(latents.device, latents.dtype))
latents = (
latents * latents_std / self.vae.config.scaling_factor +
latents_mean)
latents_mean = (torch.tensor(self.vae.config.latents_mean).view(1, 12, 1, 1,
1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(self.vae.config.latents_std).view(1, 12, 1, 1,
1).to(latents.device, latents.dtype))
latents = (latents * latents_std / self.vae.config.scaling_factor + latents_mean)
else:
latents = latents / self.vae.config.scaling_factor
video = self.vae.decode(latents, return_dict=False)[0]
video = self.video_processor.postprocess_video(
video, output_type=output_type)
video = self.video_processor.postprocess_video(video, output_type=output_type)
# Offload all models
self.maybe_free_model_hooks()
+7
View File
@@ -0,0 +1,7 @@
import os
os.environ["NCCL_DEBUG"] = "ERROR"
from .diffusion.scheduler import *
from .diffusion.video_pipeline import *
from .modules.model import *
@@ -0,0 +1 @@
__version__ = "0.1.0"
+174
View File
@@ -0,0 +1,174 @@
import argparse
def parse_args(namespace=None):
parser = argparse.ArgumentParser(description="StepVideo inference script")
parser = add_extra_models_args(parser)
parser = add_denoise_schedule_args(parser)
parser = add_inference_args(parser)
parser = add_parallel_args(parser)
args = parser.parse_args(namespace=namespace)
return args
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
group.add_argument(
"--vae_url",
type=str,
default='127.0.0.1',
help="vae url.",
)
group.add_argument(
"--caption_url",
type=str,
default='127.0.0.1',
help="caption url.",
)
return parser
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Denoise schedule args")
# Flow Matching
group.add_argument(
"--time_shift",
type=float,
default=7.0,
help="Shift factor for flow matching schedulers.",
)
group.add_argument(
"--flow_reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
group.add_argument(
"--flow_solver",
type=str,
default="euler",
help="Solver for flow matching.",
)
return parser
def add_inference_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Inference args")
# ======================== Model loads ========================
group.add_argument(
"--model_dir",
type=str,
default="./ckpts",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--model_resolution",
type=str,
default="540p",
choices=["540p"],
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
# ======================== Inference general setting ========================
group.add_argument(
"--batch_size",
type=int,
default=1,
help="Batch size for inference and evaluation.",
)
group.add_argument(
"--infer_steps",
type=int,
default=50,
help="Number of denoising steps for inference.",
)
group.add_argument(
"--save_path",
type=str,
default="./results",
help="Path to save the generated samples.",
)
group.add_argument(
"--name_suffix",
type=str,
default="",
help="Suffix for the names of saved samples.",
)
group.add_argument(
"--num_videos",
type=int,
default=1,
help="Number of videos to generate for each prompt.",
)
# ---sample size---
group.add_argument(
"--num_frames",
type=int,
default=204,
help="How many frames to sample from a video. ",
)
group.add_argument(
"--height",
type=int,
default=544,
help="The height of video sample",
)
group.add_argument(
"--width",
type=int,
default=992,
help="The width of video sample",
)
# --- prompt ---
group.add_argument(
"--prompt",
type=str,
default=None,
help="Prompt for sampling during evaluation.",
)
group.add_argument("--seed", type=int, default=1234, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument("--pos_magic",
type=str,
default="超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。",
help="Positive magic prompt for sampling.")
group.add_argument("--neg_magic",
type=str,
default="画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。",
help="Negative magic prompt for sampling.")
group.add_argument("--cfg_scale", type=float, default=9.0, help="Classifier free guidance scale.")
return parser
def add_parallel_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Parallel args")
# ======================== Model loads ========================
group.add_argument(
"--ulysses_degree",
type=int,
default=8,
help="Ulysses degree.",
)
group.add_argument(
"--ring_degree",
type=int,
default=1,
help="Ulysses degree.",
)
return parser
@@ -0,0 +1,220 @@
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
"""
Euler scheduler.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
"""
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
reverse: bool = False,
solver: str = "euler",
device: Union[str, torch.device] = None,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
self._step_index = None
self._begin_index = None
self.device = device
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def set_timesteps(
self,
num_inference_steps: int,
time_shift: float = 13.0,
device: Union[str, torch.device] = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
"""
device = device or self.device
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1, device=device)
sigmas = self.sd3_time_shift(sigmas, time_shift)
if not self.config.reverse:
sigmas = 1 - sigmas
self.sigmas = sigmas
self.timesteps = sigmas[:-1]
# Reset step index
self._step_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
return indices[pos].item()
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor, time_shift: float = 13.0):
return (time_shift * t) / (1 + (time_shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = False,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
Args:
model_output (`torch.FloatTensor`):
The direct output from learned diffusion model.
timestep (`float`):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)):
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
self._init_step_index(timestep)
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
# upon completion increase step index by one
self._step_index += 1
if not return_dict:
return prev_sample
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
def __len__(self):
return self.config.num_train_timesteps
+325
View File
@@ -0,0 +1,325 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
import asyncio
import pickle
from dataclasses import dataclass
from typing import Dict, List, Optional, Union
import numpy as np
import torch
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput
from fastvideo.models.stepvideo.diffusion.scheduler import FlowMatchDiscreteScheduler
from fastvideo.models.stepvideo.modules.model import StepVideoModel
from fastvideo.models.stepvideo.utils import VideoProcessor
def call_api_gen(url, api, port=8080):
url = f"http://{url}:{port}/{api}-api"
import aiohttp
async def _fn(samples, *args, **kwargs):
if api == 'vae':
data = {
"samples": samples,
}
elif api == 'caption':
data = {
"prompts": samples,
}
else:
raise Exception(f"Not supported api: {api}...")
async with aiohttp.ClientSession() as sess:
data_bytes = pickle.dumps(data)
async with sess.get(url, data=data_bytes, timeout=12000) as response:
result = bytearray()
while not response.content.at_eof():
chunk = await response.content.read(1024)
result += chunk
response_data = pickle.loads(result)
return response_data
return _fn
@dataclass
class StepVideoPipelineOutput(BaseOutput):
video: Union[torch.Tensor, np.ndarray]
class StepVideoPipeline(DiffusionPipeline):
r"""
Pipeline for text-to-video generation using StepVideo.
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
Args:
transformer ([`StepVideoModel`]):
Conditional Transformer to denoise the encoded image latents.
scheduler ([`FlowMatchDiscreteScheduler`]):
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
vae_url:
remote vae server's url.
caption_url:
remote caption (stepllm and clip) server's url.
"""
def __init__(
self,
transformer: StepVideoModel,
scheduler: FlowMatchDiscreteScheduler,
vae_url: str = '127.0.0.1',
caption_url: str = '127.0.0.1',
save_path: str = './results',
name_suffix: str = '',
):
super().__init__()
self.register_modules(
transformer=transformer,
scheduler=scheduler,
)
self.vae_scale_factor_temporal = self.vae.temporal_compression_ratio if getattr(self, "vae", None) else 8
self.vae_scale_factor_spatial = self.vae.spatial_compression_ratio if getattr(self, "vae", None) else 16
self.video_processor = VideoProcessor(save_path, name_suffix)
self.vae_url = vae_url
self.caption_url = caption_url
self.setup_api(self.vae_url, self.caption_url)
def setup_api(self, vae_url, caption_url):
self.vae_url = vae_url
self.caption_url = caption_url
self.caption = call_api_gen(caption_url, 'caption')
self.vae = call_api_gen(vae_url, 'vae')
return self
def encode_prompt(
self,
prompt: str,
neg_magic: str = '',
pos_magic: str = '',
):
device = self._execution_device
prompts = [prompt + pos_magic]
bs = len(prompts)
prompts += [neg_magic] * bs
data = asyncio.run(self.caption(prompts))
prompt_embeds, prompt_attention_mask, clip_embedding = data['y'].to(device), data['y_mask'].to(
device), data['clip_embedding'].to(device)
return prompt_embeds, clip_embedding, prompt_attention_mask
def decode_vae(self, samples):
samples = asyncio.run(self.vae(samples.cpu()))
return samples
def check_inputs(self, num_frames, width, height):
num_frames = max(num_frames // 17 * 17, 1)
width = max(width // 16 * 16, 16)
height = max(height // 16 * 16, 16)
return num_frames, width, height
def prepare_latents(
self,
batch_size: int,
num_channels_latents: 64,
height: int = 544,
width: int = 992,
num_frames: int = 204,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if latents is not None:
return latents.to(device=device, dtype=dtype)
num_frames, width, height = self.check_inputs(num_frames, width, height)
shape = (
batch_size,
max(num_frames // 17 * 3, 1),
num_channels_latents,
int(height) // self.vae_scale_factor_spatial,
int(width) // self.vae_scale_factor_spatial,
) # b,f,c,h,w
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators.")
if generator is None:
generator = torch.Generator(device=self._execution_device)
latents = torch.randn(shape, generator=generator, device=device, dtype=dtype)
return latents
@torch.inference_mode()
def __call__(
self,
prompt: Union[str, List[str]] = None,
height: int = 544,
width: int = 992,
num_frames: int = 204,
num_inference_steps: int = 50,
guidance_scale: float = 9.0,
time_shift: float = 13.0,
neg_magic: str = "",
pos_magic: str = "",
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
output_type: Optional[str] = "mp4",
output_file_name: Optional[str] = "",
return_dict: bool = True,
mask_strategy: Optional[Dict[str, list]] = None,
):
r"""
The call function to the pipeline for generation.
Args:
prompt (`str` or `List[str]`, *optional*):
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
instead.
height (`int`, defaults to `544`):
The height in pixels of the generated image.
width (`int`, defaults to `992`):
The width in pixels of the generated image.
num_frames (`int`, defaults to `204`):
The number of frames in the generated video.
num_inference_steps (`int`, defaults to `50`):
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
expense of slower inference.
guidance_scale (`float`, defaults to `9.0`):
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
`guidance_scale` is defined as `w` of equation 2. of [Imagen
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
usually at the expense of lower image quality.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
The number of images to generate per prompt.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
generation deterministic.
latents (`torch.Tensor`, *optional*):
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
tensor is generated by sampling using the supplied random `generator`.
output_type (`str`, *optional*, defaults to `"pil"`):
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
output_file_name(`str`, *optional*`):
The output mp4 file name.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`StepVideoPipelineOutput`] instead of a plain tuple.
Examples:
Returns:
[`~StepVideoPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`StepVideoPipelineOutput`] is returned, otherwise a `tuple` is returned
where the first element is a list with the generated images and the second element is a list of `bool`s
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
"""
# 1. Check inputs. Raise error if not correct
device = self._execution_device
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
do_classifier_free_guidance = guidance_scale > 1.0
# 3. Encode input prompt
prompt_embeds, prompt_embeds_2, prompt_attention_mask = self.encode_prompt(
prompt=prompt,
neg_magic=neg_magic,
pos_magic=pos_magic,
)
transformer_dtype = self.transformer.dtype
prompt_embeds = prompt_embeds.to(transformer_dtype)
prompt_attention_mask = prompt_attention_mask.to(transformer_dtype)
prompt_embeds_2 = prompt_embeds_2.to(transformer_dtype)
# 4. Prepare timesteps
self.scheduler.set_timesteps(num_inference_steps=num_inference_steps, time_shift=time_shift, device=device)
# 5. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
num_channels_latents,
height,
width,
num_frames,
torch.bfloat16,
device,
generator,
latents,
)
def dict_to_3d_list(best_masks, t_max=50, l_max=48, h_max=48):
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
if best_masks is None:
return result
for key, value in best_masks.items():
timestep, layer, head = map(int, key.split('_'))
result[timestep][layer][head] = value
return result
mask_strategy = dict_to_3d_list(mask_strategy)
#best_mask_selections = None
# 7. Denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(self.scheduler.timesteps):
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
latent_model_input = latent_model_input.to(transformer_dtype)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype)
noise_pred = self.transformer(
hidden_states=latent_model_input,
timestep=timestep,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
encoder_hidden_states_2=prompt_embeds_2,
return_dict=False,
mask_strategy=mask_strategy[i],
)
# perform guidance
if do_classifier_free_guidance:
noise_pred_text, noise_pred_uncond = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(model_output=noise_pred, timestep=t, sample=latents)
progress_bar.update()
if not torch.distributed.is_initialized() or int(torch.distributed.get_rank()) == 0:
if not output_type == "latent":
video = self.decode_vae(latents)
video = self.video_processor.postprocess_video(video,
output_file_name=output_file_name,
output_type=output_type)
else:
video = latents
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return (video, )
return StepVideoPipelineOutput(video=video)
+96
View File
@@ -0,0 +1,96 @@
import torch
import torch.nn as nn
from einops import rearrange
from flash_attn import flash_attn_func
try:
from st_attn import sliding_tile_attention
except ImportError:
print("Could not load Sliding Tile Attention.")
sliding_tile_attention = None
from fastvideo.utils.communications import all_to_all_4D
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
class Attention(nn.Module):
def __init__(self):
super().__init__()
def attn_processor(self, attn_type):
if attn_type == 'torch':
return self.torch_attn_func
elif attn_type == 'parallel':
return self.parallel_attn_func
else:
raise Exception('Not supported attention type...')
def tile(self, x, sp_size):
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=36 // sp_size, h=48, w=48)
return rearrange(x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
n_t=6,
n_h=6,
n_w=6,
ts_t=6,
ts_h=8,
ts_w=8)
def untile(self, x, sp_size):
x = rearrange(x,
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
n_t=6,
n_h=6,
n_w=6,
ts_t=6,
ts_h=8,
ts_w=8)
return rearrange(x, "b (t sp h w) head d -> b (sp t h w) head d", sp=sp_size, t=36 // sp_size, h=48, w=48)
def torch_attn_func(self, q, k, v, attn_mask=None, causal=False, drop_rate=0.0, **kwargs):
if attn_mask is not None and attn_mask.dtype != torch.bool:
attn_mask = attn_mask.to(q.dtype)
if attn_mask is not None and attn_mask.ndim == 3: ## no head
n_heads = q.shape[2]
attn_mask = attn_mask.unsqueeze(1).repeat(1, n_heads, 1, 1)
q, k, v = map(lambda x: rearrange(x, 'b s h d -> b h s d'), (q, k, v))
x = torch.nn.functional.scaled_dot_product_attention(q,
k,
v,
attn_mask=attn_mask,
dropout_p=drop_rate,
is_causal=causal)
x = rearrange(x, 'b h s d -> b s h d')
return x
def parallel_attn_func(self, q, k, v, causal=False, mask_strategy=None, **kwargs):
if get_sequence_parallel_state():
q = all_to_all_4D(q, scatter_dim=2, gather_dim=1)
k = all_to_all_4D(k, scatter_dim=2, gather_dim=1)
v = all_to_all_4D(v, scatter_dim=2, gather_dim=1)
if mask_strategy[0] is not None:
q = self.tile(q, nccl_info.sp_size).transpose(1, 2).contiguous()
k = self.tile(k, nccl_info.sp_size).transpose(1, 2).contiguous()
v = self.tile(v, nccl_info.sp_size).transpose(1, 2).contiguous()
head_num = q.size(1) # 48 // sp_size
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)]
x = sliding_tile_attention(q, k, v, windows, 0, False).transpose(1, 2).contiguous()
x = self.untile(x, nccl_info.sp_size)
else:
x = flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=False)
if get_sequence_parallel_state():
x = all_to_all_4D(x, scatter_dim=1, gather_dim=2)
x = x.to(q.dtype)
return x
+296
View File
@@ -0,0 +1,296 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Optional
import torch
import torch.nn as nn
from einops import rearrange
from fastvideo.models.stepvideo.modules.attentions import Attention
from fastvideo.models.stepvideo.modules.normalization import RMSNorm
from fastvideo.models.stepvideo.modules.rope import RoPE3D
class SelfAttention(Attention):
def __init__(self, hidden_dim, head_dim, bias=False, with_rope=True, with_qk_norm=True, attn_type='torch'):
super().__init__()
self.head_dim = head_dim
self.n_heads = hidden_dim // head_dim
self.wqkv = nn.Linear(hidden_dim, hidden_dim * 3, bias=bias)
self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)
self.with_rope = with_rope
self.with_qk_norm = with_qk_norm
if self.with_qk_norm:
self.q_norm = RMSNorm(head_dim, elementwise_affine=True)
self.k_norm = RMSNorm(head_dim, elementwise_affine=True)
if self.with_rope:
self.rope_3d = RoPE3D(freq=1e4, F0=1.0, scaling_factor=1.0)
self.rope_ch_split = [64, 32, 32]
self.core_attention = self.attn_processor(attn_type=attn_type)
self.parallel = attn_type == 'parallel'
def apply_rope3d(self, x, fhw_positions, rope_ch_split, parallel=True):
x = self.rope_3d(x, fhw_positions, rope_ch_split, parallel)
return x
def forward(self, x, cu_seqlens=None, max_seqlen=None, rope_positions=None, attn_mask=None, mask_strategy=None):
xqkv = self.wqkv(x)
xqkv = xqkv.view(*x.shape[:-1], self.n_heads, 3 * self.head_dim)
xq, xk, xv = torch.split(xqkv, [self.head_dim] * 3, dim=-1) ## seq_len, n, dim
if self.with_qk_norm:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
if self.with_rope:
xq = self.apply_rope3d(xq, rope_positions, self.rope_ch_split, parallel=self.parallel)
xk = self.apply_rope3d(xk, rope_positions, self.rope_ch_split, parallel=self.parallel)
output = self.core_attention(xq,
xk,
xv,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
attn_mask=attn_mask,
mask_strategy=mask_strategy)
output = rearrange(output, 'b s h d -> b s (h d)')
output = self.wo(output)
return output
class CrossAttention(Attention):
def __init__(self, hidden_dim, head_dim, bias=False, with_qk_norm=True, attn_type='torch'):
super().__init__()
self.head_dim = head_dim
self.n_heads = hidden_dim // head_dim
self.wq = nn.Linear(hidden_dim, hidden_dim, bias=bias)
self.wkv = nn.Linear(hidden_dim, hidden_dim * 2, bias=bias)
self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)
self.with_qk_norm = with_qk_norm
if self.with_qk_norm:
self.q_norm = RMSNorm(head_dim, elementwise_affine=True)
self.k_norm = RMSNorm(head_dim, elementwise_affine=True)
self.core_attention = self.attn_processor(attn_type=attn_type)
def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, attn_mask=None):
xq = self.wq(x)
xq = xq.view(*xq.shape[:-1], self.n_heads, self.head_dim)
xkv = self.wkv(encoder_hidden_states)
xkv = xkv.view(*xkv.shape[:-1], self.n_heads, 2 * self.head_dim)
xk, xv = torch.split(xkv, [self.head_dim] * 2, dim=-1) ## seq_len, n, dim
if self.with_qk_norm:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
output = self.core_attention(xq, xk, xv, attn_mask=attn_mask)
output = rearrange(output, 'b s h d -> b s (h d)')
output = self.wo(output)
return output
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
self.approximate = approximate
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
return torch.nn.functional.gelu(gate, approximate=self.approximate)
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states = self.gelu(hidden_states)
return hidden_states
class FeedForward(nn.Module):
def __init__(
self,
dim: int,
inner_dim: Optional[int] = None,
dim_out: Optional[int] = None,
mult: int = 4,
bias: bool = False,
):
super().__init__()
inner_dim = dim * mult if inner_dim is None else inner_dim
dim_out = dim if dim_out is None else dim_out
self.net = nn.ModuleList([
GELU(dim, inner_dim, approximate="tanh", bias=bias),
nn.Identity(),
nn.Linear(inner_dim, dim_out, bias=bias)
])
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
def modulate(x, scale, shift):
x = x * (1 + scale) + shift
return x
def gate(x, gate):
x = gate * x
return x
class StepVideoTransformerBlock(nn.Module):
r"""
A basic Transformer block.
Parameters:
dim (`int`): The number of channels in the input and output.
num_attention_heads (`int`): The number of heads to use for multi-head attention.
attention_head_dim (`int`): The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
num_embeds_ada_norm (:
obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
attention_bias (:
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
only_cross_attention (`bool`, *optional*):
Whether to use only cross-attention layers. In this case two cross attention layers are used.
double_self_attention (`bool`, *optional*):
Whether to use two self-attention layers. In this case no cross attention layers are used.
upcast_attention (`bool`, *optional*):
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
final_dropout (`bool` *optional*, defaults to False):
Whether to apply a final dropout after the last feed-forward layer.
attention_type (`str`, *optional*, defaults to `"default"`):
The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
positional_embeddings (`str`, *optional*, defaults to `None`):
The type of positional embeddings to apply to.
num_positional_embeddings (`int`, *optional*, defaults to `None`):
The maximum number of positional embeddings to apply.
"""
def __init__(self,
dim: int,
attention_head_dim: int,
norm_eps: float = 1e-5,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = False,
attention_type: str = 'parallel'):
super().__init__()
self.dim = dim
self.norm1 = nn.LayerNorm(dim, eps=norm_eps)
self.attn1 = SelfAttention(dim,
attention_head_dim,
bias=False,
with_rope=True,
with_qk_norm=True,
attn_type=attention_type)
self.norm2 = nn.LayerNorm(dim, eps=norm_eps)
self.attn2 = CrossAttention(dim, attention_head_dim, bias=False, with_qk_norm=True, attn_type='torch')
self.ff = FeedForward(dim=dim, inner_dim=ff_inner_dim, dim_out=dim, bias=ff_bias)
self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5)
@torch.no_grad()
def forward(self,
q: torch.Tensor,
kv: Optional[torch.Tensor] = None,
timestep: Optional[torch.LongTensor] = None,
attn_mask=None,
rope_positions: list = None,
mask_strategy=None) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (torch.clone(chunk) for chunk in (
self.scale_shift_table[None] + timestep.reshape(-1, 6, self.dim)).chunk(6, dim=1))
scale_shift_q = modulate(self.norm1(q), scale_msa, shift_msa)
attn_q = self.attn1(scale_shift_q, rope_positions=rope_positions, mask_strategy=mask_strategy)
q = gate(attn_q, gate_msa) + q
attn_q = self.attn2(q, kv, attn_mask)
q = attn_q + q
scale_shift_q = modulate(self.norm2(q), scale_mlp, shift_mlp)
ff_output = self.ff(scale_shift_q)
q = gate(ff_output, gate_mlp) + q
return q
class PatchEmbed(nn.Module):
"""2D Image to Patch Embedding"""
def __init__(
self,
patch_size=64,
in_channels=3,
embed_dim=768,
layer_norm=False,
flatten=True,
bias=True,
):
super().__init__()
self.flatten = flatten
self.layer_norm = layer_norm
self.proj = nn.Conv2d(in_channels,
embed_dim,
kernel_size=(patch_size, patch_size),
stride=patch_size,
bias=bias)
def forward(self, latent):
latent = self.proj(latent).to(latent.dtype)
if self.flatten:
latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC
if self.layer_norm:
latent = self.norm(latent)
return latent
+198
View File
@@ -0,0 +1,198 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Dict, Optional
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
from einops import rearrange, repeat
from torch import nn
from fastvideo.models.stepvideo.modules.blocks import PatchEmbed, StepVideoTransformerBlock
from fastvideo.models.stepvideo.modules.normalization import AdaLayerNormSingle, PixArtAlphaTextProjection
from fastvideo.models.stepvideo.parallel import parallel_forward
from fastvideo.models.stepvideo.utils import with_empty_init
class StepVideoModel(ModelMixin, ConfigMixin):
_no_split_modules = ["StepVideoTransformerBlock", "PatchEmbed"]
@with_empty_init
@register_to_config
def __init__(
self,
num_attention_heads: int = 48,
attention_head_dim: int = 128,
in_channels: int = 64,
out_channels: Optional[int] = 64,
num_layers: int = 48,
dropout: float = 0.0,
patch_size: int = 1,
norm_type: str = "ada_norm_single",
norm_elementwise_affine: bool = False,
norm_eps: float = 1e-6,
use_additional_conditions: Optional[bool] = False,
caption_channels: Optional[int] | list | tuple = [6144, 1024],
attention_type: Optional[str] = "parallel",
):
super().__init__()
# Set some common variables used across the board.
self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim
self.out_channels = in_channels if out_channels is None else out_channels
self.use_additional_conditions = use_additional_conditions
self.pos_embed = PatchEmbed(
patch_size=patch_size,
in_channels=self.config.in_channels,
embed_dim=self.inner_dim,
)
self.transformer_blocks = nn.ModuleList([
StepVideoTransformerBlock(dim=self.inner_dim,
attention_head_dim=self.config.attention_head_dim,
attention_type=attention_type) for _ in range(self.config.num_layers)
])
# 3. Output blocks.
self.norm_out = nn.LayerNorm(self.inner_dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine)
self.scale_shift_table = nn.Parameter(torch.randn(2, self.inner_dim) / self.inner_dim**0.5)
self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels)
self.patch_size = patch_size
self.adaln_single = AdaLayerNormSingle(self.inner_dim, use_additional_conditions=self.use_additional_conditions)
if isinstance(self.config.caption_channels, int):
caption_channel = self.config.caption_channels
else:
caption_channel, clip_channel = self.config.caption_channels
self.clip_projection = nn.Linear(clip_channel, self.inner_dim)
self.caption_norm = nn.LayerNorm(caption_channel, eps=norm_eps, elementwise_affine=norm_elementwise_affine)
self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channel, hidden_size=self.inner_dim)
self.parallel = attention_type == 'parallel'
def patchfy(self, hidden_states):
hidden_states = rearrange(hidden_states, 'b f c h w -> (b f) c h w')
hidden_states = self.pos_embed(hidden_states)
return hidden_states
def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states, q_seqlen):
kv_seqlens = encoder_attention_mask.sum(dim=1).int()
mask = torch.zeros([len(kv_seqlens), q_seqlen, max(kv_seqlens)],
dtype=torch.bool,
device=encoder_attention_mask.device)
encoder_hidden_states = encoder_hidden_states[:, :max(kv_seqlens)]
for i, kv_len in enumerate(kv_seqlens):
mask[i, :, :kv_len] = 1
return encoder_hidden_states, mask
@parallel_forward
def block_forward(self,
hidden_states,
encoder_hidden_states=None,
timestep=None,
rope_positions=None,
attn_mask=None,
parallel=True,
mask_strategy=None):
for i, block in enumerate(self.transformer_blocks):
hidden_states = block(hidden_states,
encoder_hidden_states,
timestep=timestep,
attn_mask=attn_mask,
rope_positions=rope_positions,
mask_strategy=mask_strategy[i])
return hidden_states
@torch.inference_mode()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
encoder_hidden_states_2: Optional[torch.Tensor] = None,
timestep: Optional[torch.LongTensor] = None,
added_cond_kwargs: Dict[str, torch.Tensor] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
fps: torch.Tensor = None,
return_dict: bool = True,
mask_strategy=None,
):
assert hidden_states.ndim == 5
"hidden_states's shape should be (bsz, f, ch, h ,w)"
bsz, frame, _, height, width = hidden_states.shape
height, width = height // self.patch_size, width // self.patch_size
hidden_states = self.patchfy(hidden_states)
len_frame = hidden_states.shape[1]
if self.use_additional_conditions:
added_cond_kwargs = {
"resolution": torch.tensor([(height, width)] * bsz,
device=hidden_states.device,
dtype=hidden_states.dtype),
"nframe": torch.tensor([frame] * bsz, device=hidden_states.device, dtype=hidden_states.dtype),
"fps": fps
}
else:
added_cond_kwargs = {}
timestep, embedded_timestep = self.adaln_single(timestep, added_cond_kwargs=added_cond_kwargs)
encoder_hidden_states = self.caption_projection(self.caption_norm(encoder_hidden_states))
if encoder_hidden_states_2 is not None and hasattr(self, 'clip_projection'):
clip_embedding = self.clip_projection(encoder_hidden_states_2)
encoder_hidden_states = torch.cat([clip_embedding, encoder_hidden_states], dim=1)
hidden_states = rearrange(hidden_states, '(b f) l d-> b (f l) d', b=bsz, f=frame, l=len_frame).contiguous()
encoder_hidden_states, attn_mask = self.prepare_attn_mask(encoder_attention_mask,
encoder_hidden_states,
q_seqlen=frame * len_frame)
hidden_states = self.block_forward(hidden_states,
encoder_hidden_states,
timestep=timestep,
rope_positions=[frame, height, width],
attn_mask=attn_mask,
parallel=self.parallel,
mask_strategy=mask_strategy)
hidden_states = rearrange(hidden_states, 'b (f l) d -> (b f) l d', b=bsz, f=frame, l=len_frame)
embedded_timestep = repeat(embedded_timestep, 'b d -> (b f) d', f=frame).contiguous()
shift, scale = (self.scale_shift_table[None] + embedded_timestep[:, None]).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states)
# Modulation
hidden_states = hidden_states * (1 + scale) + shift
hidden_states = self.proj_out(hidden_states)
# unpatchify
hidden_states = hidden_states.reshape(shape=(-1, height, width, self.patch_size, self.patch_size,
self.out_channels))
hidden_states = rearrange(hidden_states, 'n h w p q c -> n c h p w q')
output = hidden_states.reshape(shape=(-1, self.out_channels, height * self.patch_size, width * self.patch_size))
output = rearrange(output, '(b f) c h w -> b f c h w', f=frame)
if return_dict:
return {'x': output}
return output
+312
View File
@@ -0,0 +1,312 @@
import math
from typing import Dict, Optional, Tuple
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
ACTIVATION_FUNCTIONS = {
"swish": nn.SiLU(),
"silu": nn.SiLU(),
"mish": nn.Mish(),
"gelu": nn.GELU(),
"relu": nn.ReLU(),
}
def get_activation(act_fn: str) -> nn.Module:
"""Helper function to get activation function from string.
Args:
act_fn (str): Name of activation function.
Returns:
nn.Module: Activation function.
"""
act_fn = act_fn.lower()
if act_fn in ACTIVATION_FUNCTIONS:
return ACTIVATION_FUNCTIONS[act_fn]
else:
raise ValueError(f"Unsupported activation function: {act_fn}")
def get_timestep_embedding(
timesteps: torch.Tensor,
embedding_dim: int,
flip_sin_to_cos: bool = False,
downscale_freq_shift: float = 1,
scale: float = 1,
max_period: int = 10000,
):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param embedding_dim: the dimension of the output. :param max_period: controls the minimum frequency of the
embeddings. :return: an [N x dim] Tensor of positional embeddings.
"""
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
half_dim = embedding_dim // 2
exponent = -math.log(max_period) * torch.arange(start=0, end=half_dim, dtype=torch.float32, device=timesteps.device)
exponent = exponent / (half_dim - downscale_freq_shift)
emb = torch.exp(exponent)
emb = timesteps[:, None].float() * emb[None, :]
# scale embeddings
emb = scale * emb
# concat sine and cosine embeddings
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
# flip sine and cosine embeddings
if flip_sin_to_cos:
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
# zero pad
if embedding_dim % 2 == 1:
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
return emb
class Timesteps(nn.Module):
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float):
super().__init__()
self.num_channels = num_channels
self.flip_sin_to_cos = flip_sin_to_cos
self.downscale_freq_shift = downscale_freq_shift
def forward(self, timesteps):
t_emb = get_timestep_embedding(
timesteps,
self.num_channels,
flip_sin_to_cos=self.flip_sin_to_cos,
downscale_freq_shift=self.downscale_freq_shift,
)
return t_emb
class TimestepEmbedding(nn.Module):
def __init__(self,
in_channels: int,
time_embed_dim: int,
act_fn: str = "silu",
out_dim: int = None,
post_act_fn: Optional[str] = None,
cond_proj_dim=None,
sample_proj_bias=True):
super().__init__()
linear_cls = nn.Linear
self.linear_1 = linear_cls(
in_channels,
time_embed_dim,
bias=sample_proj_bias,
)
if cond_proj_dim is not None:
self.cond_proj = linear_cls(
cond_proj_dim,
in_channels,
bias=False,
)
else:
self.cond_proj = None
self.act = get_activation(act_fn)
if out_dim is not None:
time_embed_dim_out = out_dim
else:
time_embed_dim_out = time_embed_dim
self.linear_2 = linear_cls(
time_embed_dim,
time_embed_dim_out,
bias=sample_proj_bias,
)
if post_act_fn is None:
self.post_act = None
else:
self.post_act = get_activation(post_act_fn)
def forward(self, sample, condition=None):
if condition is not None:
sample = sample + self.cond_proj(condition)
sample = self.linear_1(sample)
if self.act is not None:
sample = self.act(sample)
sample = self.linear_2(sample)
if self.post_act is not None:
sample = self.post_act(sample)
return sample
class PixArtAlphaCombinedTimestepSizeEmbeddings(nn.Module):
def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool = False):
super().__init__()
self.outdim = size_emb_dim
self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
self.use_additional_conditions = use_additional_conditions
if self.use_additional_conditions:
self.additional_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.resolution_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim)
self.nframe_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
self.fps_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
def forward(self, timestep, resolution=None, nframe=None, fps=None):
hidden_dtype = next(self.timestep_embedder.parameters()).dtype
timesteps_proj = self.time_proj(timestep)
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D)
if self.use_additional_conditions:
batch_size = timestep.shape[0]
resolution_emb = self.additional_condition_proj(resolution.flatten()).to(hidden_dtype)
resolution_emb = self.resolution_embedder(resolution_emb).reshape(batch_size, -1)
nframe_emb = self.additional_condition_proj(nframe.flatten()).to(hidden_dtype)
nframe_emb = self.nframe_embedder(nframe_emb).reshape(batch_size, -1)
conditioning = timesteps_emb + resolution_emb + nframe_emb
if fps is not None:
fps_emb = self.additional_condition_proj(fps.flatten()).to(hidden_dtype)
fps_emb = self.fps_embedder(fps_emb).reshape(batch_size, -1)
conditioning = conditioning + fps_emb
else:
conditioning = timesteps_emb
return conditioning
class AdaLayerNormSingle(nn.Module):
r"""
Norm layer adaptive layer norm single (adaLN-single).
As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).
Parameters:
embedding_dim (`int`): The size of each embedding vector.
use_additional_conditions (`bool`): To use additional conditions for normalization or not.
"""
def __init__(self, embedding_dim: int, use_additional_conditions: bool = False, time_step_rescale=1000):
super().__init__()
self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings(embedding_dim,
size_emb_dim=embedding_dim // 2,
use_additional_conditions=use_additional_conditions)
self.silu = nn.SiLU()
self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True)
self.time_step_rescale = time_step_rescale ## timestep usually in [0, 1], we rescale it to [0,1000] for stability
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: Dict[str, torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale, **added_cond_kwargs)
out = self.linear(self.silu(embedded_timestep))
return out, embedded_timestep
class PixArtAlphaTextProjection(nn.Module):
"""
Projects caption embeddings. Also handles dropout for classifier-free guidance.
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
"""
def __init__(self, in_features, hidden_size):
super().__init__()
self.linear_1 = nn.Linear(
in_features,
hidden_size,
bias=True,
)
self.act_1 = nn.GELU(approximate="tanh")
self.linear_2 = nn.Linear(
hidden_size,
hidden_size,
bias=True,
)
def forward(self, caption):
hidden_states = self.linear_1(caption)
hidden_states = self.act_1(hidden_states)
hidden_states = self.linear_2(hidden_states)
return hidden_states
+90
View File
@@ -0,0 +1,90 @@
import torch
from fastvideo.utils.parallel_states import nccl_info
class RoPE1D:
def __init__(self, freq=1e4, F0=1.0, scaling_factor=1.0):
self.base = freq
self.F0 = F0
self.scaling_factor = scaling_factor
self.cache = {}
def get_cos_sin(self, D, seq_len, device, dtype):
if (D, seq_len, device, dtype) not in self.cache:
inv_freq = 1.0 / (self.base**(torch.arange(0, D, 2).float().to(device) / D))
t = torch.arange(seq_len, device=device, dtype=inv_freq.dtype)
freqs = torch.einsum("i,j->ij", t, inv_freq).to(dtype)
freqs = torch.cat((freqs, freqs), dim=-1)
cos = freqs.cos() # (Seq, Dim)
sin = freqs.sin()
self.cache[D, seq_len, device, dtype] = (cos, sin)
return self.cache[D, seq_len, device, dtype]
@staticmethod
def rotate_half(x):
x1, x2 = x[..., :x.shape[-1] // 2], x[..., x.shape[-1] // 2:]
return torch.cat((-x2, x1), dim=-1)
def apply_rope1d(self, tokens, pos1d, cos, sin):
assert pos1d.ndim == 2
cos = torch.nn.functional.embedding(pos1d, cos)[:, :, None, :]
sin = torch.nn.functional.embedding(pos1d, sin)[:, :, None, :]
return (tokens * cos) + (self.rotate_half(tokens) * sin)
def __call__(self, tokens, positions):
"""
input:
* tokens: batch_size x ntokens x nheads x dim
* positions: batch_size x ntokens (t position of each token)
output:
* tokens after applying RoPE2D (batch_size x ntokens x nheads x dim)
"""
D = tokens.size(3)
assert positions.ndim == 2 # Batch, Seq
cos, sin = self.get_cos_sin(D, int(positions.max()) + 1, tokens.device, tokens.dtype)
tokens = self.apply_rope1d(tokens, positions, cos, sin)
return tokens
class RoPE3D(RoPE1D):
def __init__(self, freq=1e4, F0=1.0, scaling_factor=1.0):
super(RoPE3D, self).__init__(freq, F0, scaling_factor)
self.position_cache = {}
def get_mesh_3d(self, rope_positions, bsz):
f, h, w = rope_positions
if f"{f}-{h}-{w}" not in self.position_cache:
x = torch.arange(f, device='cpu')
y = torch.arange(h, device='cpu')
z = torch.arange(w, device='cpu')
self.position_cache[f"{f}-{h}-{w}"] = torch.cartesian_prod(x, y, z).view(1, f * h * w, 3).expand(bsz, -1, 3)
return self.position_cache[f"{f}-{h}-{w}"]
def __call__(self, tokens, rope_positions, ch_split, parallel=False):
"""
input:
* tokens: batch_size x ntokens x nheads x dim
* rope_positions: list of (f, h, w)
output:
* tokens after applying RoPE2D (batch_size x ntokens x nheads x dim)
"""
assert sum(ch_split) == tokens.size(-1)
mesh_grid = self.get_mesh_3d(rope_positions, bsz=tokens.shape[0])
out = []
for i, (D, x) in enumerate(zip(ch_split, torch.split(tokens, ch_split, dim=-1))):
cos, sin = self.get_cos_sin(D, int(mesh_grid.max()) + 1, tokens.device, tokens.dtype)
if parallel:
mesh = torch.chunk(mesh_grid[:, :, i], nccl_info.sp_size, dim=1)[nccl_info.rank_within_group].clone()
else:
mesh = mesh_grid[:, :, i].clone()
x = self.apply_rope1d(x, mesh.to(tokens.device), cos, sin)
out.append(x)
tokens = torch.cat(out, dim=-1)
return tokens
+21
View File
@@ -0,0 +1,21 @@
import torch
from fastvideo.utils.communications import all_gather
from fastvideo.utils.parallel_states import nccl_info
def parallel_forward(fn_):
def wrapTheFunction(_, hidden_states, *args, **kwargs):
if kwargs['parallel']:
hidden_states = torch.chunk(hidden_states, nccl_info.sp_size, dim=-2)[nccl_info.rank_within_group]
kwargs['attn_mask'] = torch.chunk(kwargs['attn_mask'], nccl_info.sp_size,
dim=-2)[nccl_info.rank_within_group]
output = fn_(_, hidden_states, *args, **kwargs)
if kwargs['parallel']:
output = all_gather(output.contiguous(), dim=-2)
return output
return wrapTheFunction
@@ -0,0 +1,12 @@
import os
import torch
from fastvideo.models.stepvideo.config import parse_args
try:
args = parse_args()
torch.ops.load_library(
os.path.join(args.model_dir, 'lib/liboptimus_ths-torch2.5-cu124.cpython-310-x86_64-linux-gnu.so'))
except Exception as err:
print(err)
+36
View File
@@ -0,0 +1,36 @@
import os
import torch
import torch.nn as nn
from transformers import BertModel, BertTokenizer
class HunyuanClip(nn.Module):
"""
Hunyuan clip code copied from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py
hunyuan's clip used BertModel and BertTokenizer, so we copy it.
"""
def __init__(self, model_dir, max_length=77):
super(HunyuanClip, self).__init__()
self.max_length = max_length
self.tokenizer = BertTokenizer.from_pretrained(os.path.join(model_dir, 'tokenizer'))
self.text_encoder = BertModel.from_pretrained(os.path.join(model_dir, 'clip_text_encoder'))
@torch.no_grad
def forward(self, prompts, with_mask=True):
self.device = next(self.text_encoder.parameters()).device
text_inputs = self.tokenizer(
prompts,
padding="max_length",
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
prompt_embeds = self.text_encoder(
text_inputs.input_ids.to(self.device),
attention_mask=text_inputs.attention_mask.to(self.device) if with_mask else None,
)
return prompt_embeds.last_hidden_state, prompt_embeds.pooler_output
+45
View File
@@ -0,0 +1,45 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
import torch
def flash_attn_func(q,
k,
v,
dropout_p=0.0,
softmax_scale=None,
causal=True,
return_attn_probs=False,
tp_group_rank=0,
tp_group_size=1):
softmax_scale = q.size(-1)**(-0.5) if softmax_scale is None else softmax_scale
return torch.ops.Optimus.fwd(q, k, v, None, dropout_p, softmax_scale, causal, return_attn_probs, None,
tp_group_rank, tp_group_size)[0]
class FlashSelfAttention(torch.nn.Module):
def __init__(
self,
attention_dropout=0.0,
):
super().__init__()
self.dropout_p = attention_dropout
def forward(self, q, k, v, cu_seqlens=None, max_seq_len=None):
if cu_seqlens is None:
output = flash_attn_func(q, k, v, dropout_p=self.dropout_p)
else:
raise ValueError('cu_seqlens is not supported!')
return output
+291
View File
@@ -0,0 +1,291 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
import os
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from transformers.modeling_utils import PretrainedConfig, PreTrainedModel
from fastvideo.models.stepvideo.modules.normalization import RMSNorm
from fastvideo.models.stepvideo.text_encoder.flashattention import FlashSelfAttention
from fastvideo.models.stepvideo.text_encoder.tokenizer import LLaMaEmbedding, Wrapped_StepChatTokenizer
from fastvideo.models.stepvideo.utils import with_empty_init
def safediv(n, d):
q, r = divmod(n, d)
assert r == 0
return q
class MultiQueryAttention(nn.Module):
def __init__(self, cfg, layer_id=None):
super().__init__()
self.head_dim = cfg.hidden_size // cfg.num_attention_heads
self.max_seq_len = cfg.seq_length
self.use_flash_attention = cfg.use_flash_attn
assert self.use_flash_attention, 'FlashAttention is required!'
self.n_groups = cfg.num_attention_groups
self.tp_size = 1
self.n_local_heads = cfg.num_attention_heads
self.n_local_groups = self.n_groups
self.wqkv = nn.Linear(
cfg.hidden_size,
cfg.hidden_size + self.head_dim * 2 * self.n_groups,
bias=False,
)
self.wo = nn.Linear(
cfg.hidden_size,
cfg.hidden_size,
bias=False,
)
assert self.use_flash_attention, 'non-Flash attention not supported yet.'
self.core_attention = FlashSelfAttention(attention_dropout=cfg.attention_dropout)
self.layer_id = layer_id
def forward(
self,
x: torch.Tensor,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
):
seqlen, bsz, dim = x.shape
xqkv = self.wqkv(x)
xq, xkv = torch.split(
xqkv,
(dim // self.tp_size, self.head_dim * 2 * self.n_groups // self.tp_size),
dim=-1,
)
# gather on 1st dimension
xq = xq.view(seqlen, bsz, self.n_local_heads, self.head_dim)
xkv = xkv.view(seqlen, bsz, self.n_local_groups, 2 * self.head_dim)
xk, xv = xkv.chunk(2, -1)
# rotary embedding + flash attn
xq = rearrange(xq, "s b h d -> b s h d")
xk = rearrange(xk, "s b h d -> b s h d")
xv = rearrange(xv, "s b h d -> b s h d")
q_per_kv = self.n_local_heads // self.n_local_groups
if q_per_kv > 1:
b, s, h, d = xk.size()
if h == 1:
xk = xk.expand(b, s, q_per_kv, d)
xv = xv.expand(b, s, q_per_kv, d)
else:
''' To cover the cases where h > 1, we have
the following implementation, which is equivalent to:
xk = xk.repeat_interleave(q_per_kv, dim=-2)
xv = xv.repeat_interleave(q_per_kv, dim=-2)
but can avoid calling aten::item() that involves cpu.
'''
idx = torch.arange(q_per_kv * h, device=xk.device).reshape(q_per_kv, -1).permute(1, 0).flatten()
xk = torch.index_select(xk.repeat(1, 1, q_per_kv, 1), 2, idx).contiguous()
xv = torch.index_select(xv.repeat(1, 1, q_per_kv, 1), 2, idx).contiguous()
if self.use_flash_attention:
output = self.core_attention(xq, xk, xv, cu_seqlens=cu_seqlens, max_seq_len=max_seq_len)
# reduce-scatter only support first dimension now
output = rearrange(output, "b s h d -> s b (h d)").contiguous()
else:
xq, xk, xv = [rearrange(x, "b s ... -> s b ...").contiguous() for x in (xq, xk, xv)]
output = self.core_attention(xq, xk, xv, mask)
output = self.wo(output)
return output
class FeedForward(nn.Module):
def __init__(
self,
cfg,
dim: int,
hidden_dim: int,
layer_id: int,
multiple_of: int = 256,
):
super().__init__()
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
def swiglu(x):
x = torch.chunk(x, 2, dim=-1)
return F.silu(x[0]) * x[1]
self.swiglu = swiglu
self.w1 = nn.Linear(
dim,
2 * hidden_dim,
bias=False,
)
self.w2 = nn.Linear(
hidden_dim,
dim,
bias=False,
)
def forward(self, x):
x = self.swiglu(self.w1(x))
output = self.w2(x)
return output
class TransformerBlock(nn.Module):
def __init__(self, cfg, layer_id: int):
super().__init__()
self.n_heads = cfg.num_attention_heads
self.dim = cfg.hidden_size
self.head_dim = cfg.hidden_size // cfg.num_attention_heads
self.attention = MultiQueryAttention(
cfg,
layer_id=layer_id,
)
self.feed_forward = FeedForward(
cfg,
dim=cfg.hidden_size,
hidden_dim=cfg.ffn_hidden_size,
layer_id=layer_id,
)
self.layer_id = layer_id
self.attention_norm = RMSNorm(
cfg.hidden_size,
eps=cfg.layernorm_epsilon,
)
self.ffn_norm = RMSNorm(
cfg.hidden_size,
eps=cfg.layernorm_epsilon,
)
def forward(
self,
x: torch.Tensor,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
):
residual = self.attention.forward(self.attention_norm(x), mask, cu_seqlens, max_seq_len)
h = x + residual
ffn_res = self.feed_forward.forward(self.ffn_norm(h))
out = h + ffn_res
return out
class Transformer(nn.Module):
def __init__(
self,
config,
max_seq_size=8192,
):
super().__init__()
self.num_layers = config.num_layers
self.layers = self._build_layers(config)
def _build_layers(self, config):
layers = torch.nn.ModuleList()
for layer_id in range(self.num_layers):
layers.append(TransformerBlock(
config,
layer_id=layer_id + 1,
))
return layers
def forward(
self,
hidden_states,
attention_mask,
cu_seqlens=None,
max_seq_len=None,
):
if max_seq_len is not None and not isinstance(max_seq_len, torch.Tensor):
max_seq_len = torch.tensor(max_seq_len, dtype=torch.int32, device="cpu")
for lid, layer in enumerate(self.layers):
hidden_states = layer(
hidden_states,
attention_mask,
cu_seqlens,
max_seq_len,
)
return hidden_states
class Step1Model(PreTrainedModel):
config_class = PretrainedConfig
@with_empty_init
def __init__(
self,
config,
):
super().__init__(config)
self.tok_embeddings = LLaMaEmbedding(config)
self.transformer = Transformer(config)
def forward(
self,
input_ids=None,
attention_mask=None,
):
hidden_states = self.tok_embeddings(input_ids)
hidden_states = self.transformer(
hidden_states,
attention_mask,
)
return hidden_states
class STEP1TextEncoder(torch.nn.Module):
def __init__(self, model_dir, max_length=320):
super(STEP1TextEncoder, self).__init__()
self.max_length = max_length
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
text_encoder = Step1Model.from_pretrained(model_dir)
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
@torch.no_grad
def forward(self, prompts, with_mask=True, max_length=None):
self.device = next(self.text_encoder.parameters()).device
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
if type(prompts) is str:
prompts = [prompts]
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
+209
View File
@@ -0,0 +1,209 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import List
import torch
import torch.nn as nn
class LLaMaEmbedding(nn.Module):
"""Language model embeddings.
Arguments:
hidden_size: hidden size
vocab_size: vocabulary size
max_sequence_length: maximum size of sequence. This
is used for positional embedding
embedding_dropout_prob: dropout probability for embeddings
init_method: weight initialization method
num_tokentypes: size of the token-type embeddings. 0 value
will ignore this embedding
"""
def __init__(
self,
cfg,
):
super().__init__()
self.hidden_size = cfg.hidden_size
self.params_dtype = cfg.params_dtype
self.fp32_residual_connection = cfg.fp32_residual_connection
self.embedding_weights_in_fp32 = cfg.embedding_weights_in_fp32
self.word_embeddings = torch.nn.Embedding(
cfg.padded_vocab_size,
self.hidden_size,
)
self.embedding_dropout = torch.nn.Dropout(cfg.hidden_dropout)
def forward(self, input_ids):
# Embeddings.
if self.embedding_weights_in_fp32:
self.word_embeddings = self.word_embeddings.to(torch.float32)
embeddings = self.word_embeddings(input_ids)
if self.embedding_weights_in_fp32:
embeddings = embeddings.to(self.params_dtype)
self.word_embeddings = self.word_embeddings.to(self.params_dtype)
# Data format change to avoid explicit transposes : [b s h] --> [s b h].
embeddings = embeddings.transpose(0, 1).contiguous()
# If the input flag for fp32 residual connection is set, convert for float.
if self.fp32_residual_connection:
embeddings = embeddings.float()
# Dropout.
embeddings = self.embedding_dropout(embeddings)
return embeddings
class StepChatTokenizer:
"""Step Chat Tokenizer"""
def __init__(
self,
model_file,
name="StepChatTokenizer",
bot_token="<|BOT|>", # Begin of Turn
eot_token="<|EOT|>", # End of Turn
call_start_token="<|CALL_START|>", # Call Start
call_end_token="<|CALL_END|>", # Call End
think_start_token="<|THINK_START|>", # Think Start
think_end_token="<|THINK_END|>", # Think End
mask_start_token="<|MASK_1e69f|>", # Mask start
mask_end_token="<|UNMASK_1e69f|>", # Mask end
):
import sentencepiece
self._tokenizer = sentencepiece.SentencePieceProcessor(model_file=model_file)
self._vocab = {}
self._inv_vocab = {}
self._special_tokens = {}
self._inv_special_tokens = {}
self._t5_tokens = []
for idx in range(self._tokenizer.get_piece_size()):
text = self._tokenizer.id_to_piece(idx)
self._inv_vocab[idx] = text
self._vocab[text] = idx
if self._tokenizer.is_control(idx) or self._tokenizer.is_unknown(idx):
self._special_tokens[text] = idx
self._inv_special_tokens[idx] = text
self._unk_id = self._tokenizer.unk_id()
self._bos_id = self._tokenizer.bos_id()
self._eos_id = self._tokenizer.eos_id()
for token in [bot_token, eot_token, call_start_token, call_end_token, think_start_token, think_end_token]:
assert token in self._vocab, f"Token '{token}' not found in tokenizer"
assert token in self._special_tokens, f"Token '{token}' is not a special token"
for token in [mask_start_token, mask_end_token]:
assert token in self._vocab, f"Token '{token}' not found in tokenizer"
self._bot_id = self._tokenizer.piece_to_id(bot_token)
self._eot_id = self._tokenizer.piece_to_id(eot_token)
self._call_start_id = self._tokenizer.piece_to_id(call_start_token)
self._call_end_id = self._tokenizer.piece_to_id(call_end_token)
self._think_start_id = self._tokenizer.piece_to_id(think_start_token)
self._think_end_id = self._tokenizer.piece_to_id(think_end_token)
self._mask_start_id = self._tokenizer.piece_to_id(mask_start_token)
self._mask_end_id = self._tokenizer.piece_to_id(mask_end_token)
self._underline_id = self._tokenizer.piece_to_id("\u2581")
@property
def vocab(self):
return self._vocab
@property
def inv_vocab(self):
return self._inv_vocab
@property
def vocab_size(self):
return self._tokenizer.vocab_size()
def tokenize(self, text: str) -> List[int]:
return self._tokenizer.encode_as_ids(text)
def detokenize(self, token_ids: List[int]) -> str:
return self._tokenizer.decode_ids(token_ids)
class Tokens:
def __init__(self, input_ids, cu_input_ids, attention_mask, cu_seqlens, max_seq_len) -> None:
self.input_ids = input_ids
self.attention_mask = attention_mask
self.cu_input_ids = cu_input_ids
self.cu_seqlens = cu_seqlens
self.max_seq_len = max_seq_len
def to(self, device):
self.input_ids = self.input_ids.to(device)
self.attention_mask = self.attention_mask.to(device)
self.cu_input_ids = self.cu_input_ids.to(device)
self.cu_seqlens = self.cu_seqlens.to(device)
return self
class Wrapped_StepChatTokenizer(StepChatTokenizer):
def __call__(self, text, max_length=320, padding="max_length", truncation=True, return_tensors="pt"):
# [bos, ..., eos, pad, pad, ..., pad]
self.BOS = 1
self.EOS = 2
self.PAD = 2
out_tokens = []
attn_mask = []
if len(text) == 0:
part_tokens = [self.BOS] + [self.EOS]
valid_size = len(part_tokens)
if len(part_tokens) < max_length:
part_tokens += [self.PAD] * (max_length - valid_size)
out_tokens.append(part_tokens)
attn_mask.append([1] * valid_size + [0] * (max_length - valid_size))
else:
for part in text:
part_tokens = self.tokenize(part)
part_tokens = part_tokens[:(max_length - 2)] # leave 2 space for bos and eos
part_tokens = [self.BOS] + part_tokens + [self.EOS]
valid_size = len(part_tokens)
if len(part_tokens) < max_length:
part_tokens += [self.PAD] * (max_length - valid_size)
out_tokens.append(part_tokens)
attn_mask.append([1] * valid_size + [0] * (max_length - valid_size))
out_tokens = torch.tensor(out_tokens, dtype=torch.long)
attn_mask = torch.tensor(attn_mask, dtype=torch.long)
# padding y based on tp size
padded_len = 0
padded_flag = True if padded_len > 0 else False
if padded_flag:
pad_tokens = torch.tensor([[self.PAD] * max_length], device=out_tokens.device)
pad_attn_mask = torch.tensor([[1] * padded_len + [0] * (max_length - padded_len)], device=attn_mask.device)
out_tokens = torch.cat([out_tokens, pad_tokens], dim=0)
attn_mask = torch.cat([attn_mask, pad_attn_mask], dim=0)
# cu_seqlens
cu_out_tokens = out_tokens.masked_select(attn_mask != 0).unsqueeze(0)
seqlen = attn_mask.sum(dim=1).tolist()
cu_seqlens = torch.cumsum(torch.tensor([0] + seqlen), 0).to(device=out_tokens.device, dtype=torch.int32)
max_seq_len = max(seqlen)
return Tokens(out_tokens, cu_out_tokens, attn_mask, cu_seqlens, max_seq_len)
+2
View File
@@ -0,0 +1,2 @@
from .utils import *
from .video_process import *
@@ -0,0 +1,117 @@
# from stepvideo.diffusion.video_pipeline import StepVideoPipeline
import torch
import torch.nn as nn
from torch.nn import functional as F
def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
_bits = torch.tensor(bits)
_mantissa_bit = torch.tensor(mantissa_bit)
_sign_bits = torch.tensor(sign_bits)
M = torch.clamp(torch.round(_mantissa_bit), 1, _bits - _sign_bits)
E = _bits - _sign_bits - M
bias = 2**(E - 1) - 1
mantissa = 1
for i in range(mantissa_bit - 1):
mantissa += 1 / (2**(i + 1))
maxval = mantissa * 2**(2**E - 1 - bias)
return maxval
def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
"""
Default is E4M3.
"""
bits = torch.tensor(bits)
mantissa_bit = torch.tensor(mantissa_bit)
sign_bits = torch.tensor(sign_bits)
M = torch.clamp(torch.round(mantissa_bit), 1, bits - sign_bits)
E = bits - sign_bits - M
bias = 2**(E - 1) - 1
mantissa = 1
for i in range(mantissa_bit - 1):
mantissa += 1 / (2**(i + 1))
maxval = mantissa * 2**(2**E - 1 - bias)
minval = -maxval
minval = -maxval if sign_bits == 1 else torch.zeros_like(maxval)
input_clamp = torch.min(torch.max(x, minval), maxval)
log_scales = torch.clamp((torch.floor(torch.log2(torch.abs(input_clamp)) + bias)).detach(), 1.0)
log_scales = 2.0**(log_scales - M - bias.type(x.dtype))
# dequant
qdq_out = torch.round(input_clamp / log_scales) * log_scales
return qdq_out, log_scales
def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
for i in range(len(x.shape) - 1):
scale = scale.unsqueeze(-1)
new_x = x / scale
quant_dequant_x, log_scales = quantize_to_fp8(new_x, bits=bits, mantissa_bit=mantissa_bit, sign_bits=sign_bits)
return quant_dequant_x, scale, log_scales
def fp8_activation_dequant(qdq_out, scale, dtype):
qdq_out = qdq_out.type(dtype)
quant_dequant_x = qdq_out * scale.to(dtype)
return quant_dequant_x
def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype
#####
if cls.weight.dtype != torch.float8_e4m3fn:
assert False
maxval = get_fp_maxval()
scale = torch.max(torch.abs(cls.weight.flatten())) / maxval
linear_weight, scale, log_scales = fp8_tensor_quant(cls.weight, scale)
linear_weight = linear_weight.to(torch.float8_e4m3fn)
weight_dtype = linear_weight.dtype
else:
scale = cls.fp8_scale.to(cls.weight.device)
linear_weight = cls.weight
#####
if weight_dtype == torch.float8_e4m3fn:
if True or len(input.shape) == 3:
cls_dequant = fp8_activation_dequant(linear_weight, scale, original_dtype)
if cls.bias is not None:
print(f"input dtype: {input.dtype}")
print(f"cls_dequant dtype: {cls_dequant.dtype}")
print(f"cls.bias dtype: {cls.bias.dtype}")
output = F.linear(input, cls_dequant, cls.bias)
else:
output = F.linear(input, cls_dequant)
return output
else:
return cls.original_forward(input.to(original_dtype))
else:
return cls.original_forward(input)
def convert_fp8_linear(module, original_dtype, params_to_keep={}):
setattr(module, "fp8_matmul_enabled", True)
fp8_layers = []
scale_dict = {}
counter = 0
for key, layer in module.named_modules():
if isinstance(layer, nn.Linear) and 'transformer_blocks' in key:
print(f"Converting {key} to FP8")
fp8_layers.append(key)
original_forward = layer.forward
maxval = get_fp_maxval()
scale = torch.max(torch.abs(layer.weight.flatten())) / maxval
original_weight = layer.weight.data # Store a reference to the original weights
quantized_weight, scale, _ = fp8_tensor_quant(original_weight, scale)
scale_dict[key] = scale
layer.weight = torch.nn.Parameter(quantized_weight.to(torch.float8_e4m3fn))
del original_weight # Delete the reference to the original weights
torch.cuda.empty_cache()
# print(f"layer weight dtype: {layer.weight.dtype} for layer {key}")
setattr(layer, "fp8_scale", scale.to(dtype=original_dtype))
setattr(layer, "original_forward", original_forward)
setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input))
counter += 1
return scale_dict
+60
View File
@@ -0,0 +1,60 @@
import random
from functools import wraps
import numpy as np
import torch
import torch.utils._device
def setup_seed(seed):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
class EmptyInitOnDevice(torch.overrides.TorchFunctionMode):
def __init__(self, device=None):
self.device = device
def __torch_function__(self, func, types, args=(), kwargs=None):
kwargs = kwargs or {}
if getattr(func, '__module__', None) == 'torch.nn.init':
if 'tensor' in kwargs:
return kwargs['tensor']
else:
return args[0]
if self.device is not None and func in torch.utils._device._device_constructors(
) and kwargs.get('device') is None:
kwargs['device'] = self.device
return func(*args, **kwargs)
def with_empty_init(func):
@wraps(func)
def wrapper(*args, **kwargs):
with EmptyInitOnDevice('cpu'):
return func(*args, **kwargs)
return wrapper
def culens2mask(cu_seqlens=None, cu_seqlens_kv=None, max_seqlen=None, max_seqlen_kv=None, is_causal=False):
assert len(cu_seqlens) == len(cu_seqlens_kv)
"q k v should have same bsz..."
bsz = len(cu_seqlens) - 1
seqlens = cu_seqlens[1:] - cu_seqlens[:-1]
seqlens_kv = cu_seqlens_kv[1:] - cu_seqlens_kv[:-1]
attn_mask = torch.zeros(bsz, max_seqlen, max_seqlen_kv, dtype=torch.bool)
for i, (seq_len, seq_len_kv) in enumerate(zip(seqlens, seqlens_kv)):
if is_causal:
attn_mask[i, :seq_len, :seq_len_kv] = torch.triu(torch.ones(seq_len, seq_len_kv), diagonal=1).bool()
else:
attn_mask[i, :seq_len, :seq_len_kv] = torch.ones([seq_len, seq_len_kv], dtype=torch.bool)
return attn_mask
@@ -0,0 +1,51 @@
import os
import imageio
import numpy as np
import torch
class VideoProcessor:
def __init__(self, save_path: str = './results', name_suffix: str = ''):
self.save_path = save_path
os.makedirs(self.save_path, exist_ok=True)
self.name_suffix = name_suffix
def crop2standard540p(self, vid_array):
_, height, width, _ = vid_array.shape
height_center = height // 2
width_center = width // 2
if width_center > height_center: ## horizon mode
return vid_array[:, height_center - 270:height_center + 270, width_center - 480:width_center + 480]
elif width_center < height_center: ## portrait mode
return vid_array[:, height_center - 480:height_center + 480, width_center - 270:width_center + 270]
else:
return vid_array
def save_imageio_video(self, video_array: np.array, output_filename: str, fps=25, codec='libx264'):
ffmpeg_params = [
"-vf",
"atadenoise=0a=0.1:0b=0.1:1a=0.1:1b=0.1", # denoise
]
with imageio.get_writer(output_filename, fps=fps, codec=codec, ffmpeg_params=ffmpeg_params) as vid_writer:
for img_array in video_array:
vid_writer.append_data(img_array)
def postprocess_video(self, video_tensor, output_file_name='', output_type="mp4", crop2standard540p=True):
if len(self.name_suffix) == 0:
video_path = os.path.join(self.save_path, f"{output_file_name}.{output_type}")
else:
video_path = os.path.join(self.save_path, f"{output_file_name}-{self.name_suffix}.{output_type}")
video_tensor = torch.cat([t for t in video_tensor], dim=-2)
video_tensor = (video_tensor.cpu().clamp(-1, 1) + 1) * 127.5
video_array = video_tensor.clamp(0, 255).to(torch.uint8).numpy().transpose(0, 2, 3, 1)
if crop2standard540p:
video_array = self.crop2standard540p(video_array)
self.save_imageio_video(video_array, video_path)
print(f"Saved the generated video in {video_path}")
+975
View File
@@ -0,0 +1,975 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
import torch
from einops import rearrange
from torch import nn
from torch.nn import functional as F
from fastvideo.models.stepvideo.utils import with_empty_init
def base_group_norm(x, norm_layer, act_silu=False, channel_last=False):
if hasattr(base_group_norm, 'spatial') and base_group_norm.spatial:
assert channel_last
x_shape = x.shape
x = x.flatten(0, 1)
if channel_last:
# Permute to NCHW format
x = x.permute(0, 3, 1, 2)
out = F.group_norm(x.contiguous(), norm_layer.num_groups, norm_layer.weight, norm_layer.bias, norm_layer.eps)
if act_silu:
out = F.silu(out)
if channel_last:
# Permute back to NHWC format
out = out.permute(0, 2, 3, 1)
out = out.view(x_shape)
else:
if channel_last:
# Permute to NCHW format
x = x.permute(0, 3, 1, 2)
out = F.group_norm(x.contiguous(), norm_layer.num_groups, norm_layer.weight, norm_layer.bias, norm_layer.eps)
if act_silu:
out = F.silu(out)
if channel_last:
# Permute back to NHWC format
out = out.permute(0, 2, 3, 1)
return out
def base_conv2d(x, conv_layer, channel_last=False, residual=None):
if channel_last:
x = x.permute(0, 3, 1, 2) # NHWC to NCHW
out = F.conv2d(x, conv_layer.weight, conv_layer.bias, stride=conv_layer.stride, padding=conv_layer.padding)
if residual is not None:
if channel_last:
residual = residual.permute(0, 3, 1, 2) # NHWC to NCHW
out += residual
if channel_last:
out = out.permute(0, 2, 3, 1) # NCHW to NHWC
return out
def base_conv3d(x, conv_layer, channel_last=False, residual=None, only_return_output=False):
if only_return_output:
size = cal_outsize(x.shape, conv_layer.weight.shape, conv_layer.stride, conv_layer.padding)
return torch.empty(size, device=x.device, dtype=x.dtype)
if channel_last:
x = x.permute(0, 4, 1, 2, 3) # NDHWC to NCDHW
out = F.conv3d(x, conv_layer.weight, conv_layer.bias, stride=conv_layer.stride, padding=conv_layer.padding)
if residual is not None:
if channel_last:
residual = residual.permute(0, 4, 1, 2, 3) # NDHWC to NCDHW
out += residual
if channel_last:
out = out.permute(0, 2, 3, 4, 1) # NCDHW to NDHWC
return out
def cal_outsize(input_sizes, kernel_sizes, stride, padding):
stride_d, stride_h, stride_w = stride
padding_d, padding_h, padding_w = padding
dilation_d, dilation_h, dilation_w = 1, 1, 1
in_d = input_sizes[1]
in_h = input_sizes[2]
in_w = input_sizes[3]
kernel_d = kernel_sizes[2]
kernel_h = kernel_sizes[3]
kernel_w = kernel_sizes[4]
out_channels = kernel_sizes[0]
out_d = calc_out_(in_d, padding_d, dilation_d, kernel_d, stride_d)
out_h = calc_out_(in_h, padding_h, dilation_h, kernel_h, stride_h)
out_w = calc_out_(in_w, padding_w, dilation_w, kernel_w, stride_w)
size = [input_sizes[0], out_d, out_h, out_w, out_channels]
return size
def calc_out_(in_size, padding, dilation, kernel, stride):
return (in_size + 2 * padding - dilation * (kernel - 1) - 1) // stride + 1
def base_conv3d_channel_last(x, conv_layer, residual=None):
in_numel = x.numel()
out_numel = int(x.numel() * conv_layer.out_channels / conv_layer.in_channels)
if (in_numel >= 2**30) or (out_numel >= 2**30):
assert conv_layer.stride[0] == 1, "time split asks time stride = 1"
B, T, H, W, C = x.shape
K = conv_layer.kernel_size[0]
chunks = 4
chunk_size = T // chunks
if residual is None:
out_nhwc = base_conv3d(x, conv_layer, channel_last=True, residual=residual, only_return_output=True)
else:
out_nhwc = residual
assert B == 1
for i in range(chunks):
if i == chunks - 1:
xi = x[:1, chunk_size * i:]
out_nhwci = out_nhwc[:1, chunk_size * i:]
else:
xi = x[:1, chunk_size * i:chunk_size * (i + 1) + K - 1]
out_nhwci = out_nhwc[:1, chunk_size * i:chunk_size * (i + 1)]
if residual is not None:
if i == chunks - 1:
ri = residual[:1, chunk_size * i:]
else:
ri = residual[:1, chunk_size * i:chunk_size * (i + 1)]
else:
ri = None
out_nhwci.copy_(base_conv3d(xi, conv_layer, channel_last=True, residual=ri))
else:
out_nhwc = base_conv3d(x, conv_layer, channel_last=True, residual=residual)
return out_nhwc
class Upsample2D(nn.Module):
def __init__(self, channels, use_conv=False, use_conv_transpose=False, out_channels=None):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.use_conv_transpose = use_conv_transpose
if use_conv:
self.conv = nn.Conv2d(self.channels, self.out_channels, 3, padding=1)
else:
assert "Not Supported"
self.conv = nn.ConvTranspose2d(channels, self.out_channels, 4, 2, 1)
def forward(self, x, output_size=None):
assert x.shape[-1] == self.channels
if self.use_conv_transpose:
return self.conv(x)
if output_size is None:
x = F.interpolate(x.permute(0, 3, 1, 2).to(memory_format=torch.channels_last),
scale_factor=2.0,
mode='nearest').permute(0, 2, 3, 1).contiguous()
else:
x = F.interpolate(x.permute(0, 3, 1, 2).to(memory_format=torch.channels_last),
size=output_size,
mode='nearest').permute(0, 2, 3, 1).contiguous()
# x = self.conv(x)
x = base_conv2d(x, self.conv, channel_last=True)
return x
class Downsample2D(nn.Module):
def __init__(self, channels, use_conv=False, out_channels=None, padding=1):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.padding = padding
stride = 2
if use_conv:
self.conv = nn.Conv2d(self.channels, self.out_channels, 3, stride=stride, padding=padding)
else:
assert self.channels == self.out_channels
self.conv = nn.AvgPool2d(kernel_size=stride, stride=stride)
def forward(self, x):
assert x.shape[-1] == self.channels
if self.use_conv and self.padding == 0:
pad = (0, 0, 0, 1, 0, 1)
x = F.pad(x, pad, mode="constant", value=0)
assert x.shape[-1] == self.channels
# x = self.conv(x)
x = base_conv2d(x, self.conv, channel_last=True)
return x
class CausalConv(nn.Module):
def __init__(self, chan_in, chan_out, kernel_size, **kwargs):
super().__init__()
if isinstance(kernel_size, int):
kernel_size = kernel_size if isinstance(kernel_size, tuple) else ((kernel_size, ) * 3)
time_kernel_size, height_kernel_size, width_kernel_size = kernel_size
self.dilation = kwargs.pop('dilation', 1)
self.stride = kwargs.pop('stride', 1)
if isinstance(self.stride, int):
self.stride = (self.stride, 1, 1)
time_pad = self.dilation * (time_kernel_size - 1) + max((1 - self.stride[0]), 0)
height_pad = height_kernel_size // 2
width_pad = width_kernel_size // 2
self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0)
self.time_uncausal_padding = (width_pad, width_pad, height_pad, height_pad, 0, 0)
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=self.stride, dilation=self.dilation, **kwargs)
self.is_first_run = True
def forward(self, x, is_init=True, residual=None):
x = nn.functional.pad(x, self.time_causal_padding if is_init else self.time_uncausal_padding)
x = self.conv(x)
if residual is not None:
x.add_(residual)
return x
class ChannelDuplicatingPixelUnshuffleUpSampleLayer3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
factor: int,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.factor = factor
assert out_channels * factor**3 % in_channels == 0
self.repeats = out_channels * factor**3 // in_channels
def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:
x = x.repeat_interleave(self.repeats, dim=1)
x = x.view(x.size(0), self.out_channels, self.factor, self.factor, self.factor, x.size(2), x.size(3), x.size(4))
x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
x = x.view(x.size(0), self.out_channels,
x.size(2) * self.factor,
x.size(4) * self.factor,
x.size(6) * self.factor)
x = x[:, :, self.factor - 1:, :, :]
return x
class ConvPixelShuffleUpSampleLayer3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int,
factor: int,
):
super().__init__()
self.factor = factor
out_ratio = factor**3
self.conv = CausalConv(in_channels, out_channels * out_ratio, kernel_size=kernel_size)
def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:
x = self.conv(x, is_init)
x = self.pixel_shuffle_3d(x, self.factor)
return x
@staticmethod
def pixel_shuffle_3d(x: torch.Tensor, factor: int) -> torch.Tensor:
batch_size, channels, depth, height, width = x.size()
new_channels = channels // (factor**3)
new_depth = depth * factor
new_height = height * factor
new_width = width * factor
x = x.view(batch_size, new_channels, factor, factor, factor, depth, height, width)
x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
x = x.view(batch_size, new_channels, new_depth, new_height, new_width)
x = x[:, :, factor - 1:, :, :]
return x
class ConvPixelUnshuffleDownSampleLayer3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int,
factor: int,
):
super().__init__()
self.factor = factor
out_ratio = factor**3
assert out_channels % out_ratio == 0
self.conv = CausalConv(in_channels, out_channels // out_ratio, kernel_size=kernel_size)
def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:
x = self.conv(x, is_init)
x = self.pixel_unshuffle_3d(x, self.factor)
return x
@staticmethod
def pixel_unshuffle_3d(x: torch.Tensor, factor: int) -> torch.Tensor:
pad = (0, 0, 0, 0, factor - 1, 0) # (left, right, top, bottom, front, back)
x = F.pad(x, pad)
B, C, D, H, W = x.shape
x = x.view(B, C, D // factor, factor, H // factor, factor, W // factor, factor)
x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
x = x.view(B, C * factor**3, D // factor, H // factor, W // factor)
return x
class PixelUnshuffleChannelAveragingDownSampleLayer3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
factor: int,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.factor = factor
assert in_channels * factor**3 % out_channels == 0
self.group_size = in_channels * factor**3 // out_channels
def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:
pad = (0, 0, 0, 0, self.factor - 1, 0) # (left, right, top, bottom, front, back)
x = F.pad(x, pad)
B, C, D, H, W = x.shape
x = x.view(B, C, D // self.factor, self.factor, H // self.factor, self.factor, W // self.factor, self.factor)
x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
x = x.view(B, C * self.factor**3, D // self.factor, H // self.factor, W // self.factor)
x = x.view(B, self.out_channels, self.group_size, D // self.factor, H // self.factor, W // self.factor)
x = x.mean(dim=2)
return x
def base_group_norm_with_zero_pad(x, norm_layer, act_silu=True, pad_size=2):
out_shape = list(x.shape)
out_shape[1] += pad_size
out = torch.empty(out_shape, dtype=x.dtype, device=x.device)
out[:, pad_size:] = base_group_norm(x, norm_layer, act_silu=act_silu, channel_last=True)
out[:, :pad_size] = 0
return out
class CausalConvChannelLast(CausalConv):
def __init__(self, chan_in, chan_out, kernel_size, **kwargs):
super().__init__(chan_in, chan_out, kernel_size, **kwargs)
self.time_causal_padding = (0, 0) + self.time_causal_padding
self.time_uncausal_padding = (0, 0) + self.time_uncausal_padding
def forward(self, x, is_init=True, residual=None):
if self.is_first_run:
self.is_first_run = False
# self.conv.weight = nn.Parameter(self.conv.weight.permute(0,2,3,4,1).contiguous())
x = nn.functional.pad(x, self.time_causal_padding if is_init else self.time_uncausal_padding)
x = base_conv3d_channel_last(x, self.conv, residual=residual)
return x
class CausalConvAfterNorm(CausalConv):
def __init__(self, chan_in, chan_out, kernel_size, **kwargs):
super().__init__(chan_in, chan_out, kernel_size, **kwargs)
if self.time_causal_padding == (1, 1, 1, 1, 2, 0):
self.conv = nn.Conv3d(chan_in,
chan_out,
kernel_size,
stride=self.stride,
dilation=self.dilation,
padding=(0, 1, 1),
**kwargs)
else:
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=self.stride, dilation=self.dilation, **kwargs)
self.is_first_run = True
def forward(self, x, is_init=True, residual=None):
if self.is_first_run:
self.is_first_run = False
if self.time_causal_padding == (1, 1, 1, 1, 2, 0):
pass
else:
x = nn.functional.pad(x, self.time_causal_padding).contiguous()
x = base_conv3d_channel_last(x, self.conv, residual=residual)
return x
class AttnBlock(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.norm = nn.GroupNorm(num_groups=32, num_channels=in_channels)
self.q = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)
self.k = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)
self.v = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)
self.proj_out = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)
def attention(self, x, is_init=True):
x = base_group_norm(x, self.norm, act_silu=False, channel_last=True)
q = self.q(x, is_init)
k = self.k(x, is_init)
v = self.v(x, is_init)
b, t, h, w, c = q.shape
q, k, v = map(lambda x: rearrange(x, "b t h w c -> b 1 (t h w) c"), (q, k, v))
x = nn.functional.scaled_dot_product_attention(q, k, v, is_causal=True)
x = rearrange(x, "b 1 (t h w) c -> b t h w c", t=t, h=h, w=w)
return x
def forward(self, x):
x = x.permute(0, 2, 3, 4, 1).contiguous()
h = self.attention(x)
x = self.proj_out(h, residual=x)
x = x.permute(0, 4, 1, 2, 3)
return x
class Resnet3DBlock(nn.Module):
def __init__(
self,
in_channels,
out_channels=None,
temb_channels=512,
conv_shortcut=False,
):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.norm1 = nn.GroupNorm(num_groups=32, num_channels=in_channels)
self.conv1 = CausalConvAfterNorm(in_channels, out_channels, kernel_size=3)
if temb_channels > 0:
self.temb_proj = nn.Linear(temb_channels, out_channels)
self.norm2 = nn.GroupNorm(num_groups=32, num_channels=out_channels)
self.conv2 = CausalConvAfterNorm(out_channels, out_channels, kernel_size=3)
assert conv_shortcut is False
self.use_conv_shortcut = conv_shortcut
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = CausalConvAfterNorm(in_channels, out_channels, kernel_size=3)
else:
self.nin_shortcut = CausalConvAfterNorm(in_channels, out_channels, kernel_size=1)
def forward(self, x, temb=None, is_init=True):
x = x.permute(0, 2, 3, 4, 1).contiguous()
h = base_group_norm_with_zero_pad(x, self.norm1, act_silu=True, pad_size=2)
h = self.conv1(h)
if temb is not None:
h = h + self.temb_proj(nn.functional.silu(temb))[:, :, None, None]
x = self.nin_shortcut(x) if self.in_channels != self.out_channels else x
h = base_group_norm_with_zero_pad(h, self.norm2, act_silu=True, pad_size=2)
x = self.conv2(h, residual=x)
x = x.permute(0, 4, 1, 2, 3)
return x
class Downsample3D(nn.Module):
def __init__(self, in_channels, with_conv, stride):
super().__init__()
self.with_conv = with_conv
if with_conv:
self.conv = CausalConv(in_channels, in_channels, kernel_size=3, stride=stride)
def forward(self, x, is_init=True):
if self.with_conv:
x = self.conv(x, is_init)
else:
x = nn.functional.avg_pool3d(x, kernel_size=2, stride=2)
return x
class VideoEncoder(nn.Module):
def __init__(
self,
ch=32,
ch_mult=(4, 8, 16, 16),
num_res_blocks=2,
in_channels=3,
z_channels=16,
double_z=True,
down_sampling_layer=[1, 2],
resamp_with_conv=True,
version=1,
):
super().__init__()
temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
# downsampling
self.conv_in = CausalConv(in_channels, ch, kernel_size=3)
self.down_sampling_layer = down_sampling_layer
in_ch_mult = (1, ) + tuple(ch_mult)
self.down = nn.ModuleList()
for i_level in range(self.num_resolutions):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = ch * in_ch_mult[i_level]
block_out = ch * ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(Resnet3DBlock(in_channels=block_in, out_channels=block_out, temb_channels=temb_ch))
block_in = block_out
down = nn.Module()
down.block = block
down.attn = attn
if i_level != self.num_resolutions - 1:
if i_level in self.down_sampling_layer:
down.downsample = Downsample3D(block_in, resamp_with_conv, stride=(2, 2, 2))
else:
down.downsample = Downsample2D(block_in, resamp_with_conv, padding=0) #DIFF
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)
# end
self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_in)
self.version = version
if version == 2:
channels = 4 * z_channels * 2**3
self.conv_patchify = ConvPixelUnshuffleDownSampleLayer3D(block_in, channels, kernel_size=3, factor=2)
self.shortcut_pathify = PixelUnshuffleChannelAveragingDownSampleLayer3D(block_in, channels, 2)
self.shortcut_out = PixelUnshuffleChannelAveragingDownSampleLayer3D(
channels, 2 * z_channels if double_z else z_channels, 1)
self.conv_out = CausalConvChannelLast(channels, 2 * z_channels if double_z else z_channels, kernel_size=3)
else:
self.conv_out = CausalConvAfterNorm(block_in, 2 * z_channels if double_z else z_channels, kernel_size=3)
@torch.inference_mode()
def forward(self, x, video_frame_num, is_init=True):
# timestep embedding
temb = None
t = video_frame_num
# downsampling
h = self.conv_in(x, is_init)
# make it real channel last, but behave like normal layout
h = h.permute(0, 2, 3, 4, 1).contiguous().permute(0, 4, 1, 2, 3)
for i_level in range(self.num_resolutions):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](h, temb, is_init)
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
if i_level != self.num_resolutions - 1:
if isinstance(self.down[i_level].downsample, Downsample2D):
_, _, t, _, _ = h.shape
h = rearrange(h, "b c t h w -> (b t) h w c", t=t)
h = self.down[i_level].downsample(h)
h = rearrange(h, "(b t) h w c -> b c t h w", t=t)
else:
h = self.down[i_level].downsample(h, is_init)
h = self.mid.block_1(h, temb, is_init)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb, is_init)
h = h.permute(0, 2, 3, 4, 1).contiguous() # b c l h w -> b l h w c
if self.version == 2:
h = base_group_norm(h, self.norm_out, act_silu=True, channel_last=True)
h = h.permute(0, 4, 1, 2, 3).contiguous()
shortcut = self.shortcut_pathify(h, is_init)
h = self.conv_patchify(h, is_init)
h = h.add_(shortcut)
shortcut = self.shortcut_out(h, is_init).permute(0, 2, 3, 4, 1)
h = self.conv_out(h.permute(0, 2, 3, 4, 1).contiguous(), is_init)
h = h.add_(shortcut)
else:
h = base_group_norm_with_zero_pad(h, self.norm_out, act_silu=True, pad_size=2)
h = self.conv_out(h, is_init)
h = h.permute(0, 4, 1, 2, 3) # b l h w c -> b c l h w
h = rearrange(h, "b c t h w -> b t c h w")
return h
class Res3DBlockUpsample(nn.Module):
def __init__(self, input_filters, num_filters, down_sampling_stride, down_sampling=False):
super().__init__()
self.input_filters = input_filters
self.num_filters = num_filters
self.act_ = nn.SiLU(inplace=True)
self.conv1 = CausalConvChannelLast(num_filters, num_filters, kernel_size=[3, 3, 3])
self.norm1 = nn.GroupNorm(32, num_filters)
self.conv2 = CausalConvChannelLast(num_filters, num_filters, kernel_size=[3, 3, 3])
self.norm2 = nn.GroupNorm(32, num_filters)
self.down_sampling = down_sampling
if down_sampling:
self.down_sampling_stride = down_sampling_stride
else:
self.down_sampling_stride = [1, 1, 1]
if num_filters != input_filters or down_sampling:
self.conv3 = CausalConvChannelLast(input_filters,
num_filters,
kernel_size=[1, 1, 1],
stride=self.down_sampling_stride)
self.norm3 = nn.GroupNorm(32, num_filters)
def forward(self, x, is_init=False):
x = x.permute(0, 2, 3, 4, 1).contiguous()
residual = x
h = self.conv1(x, is_init)
h = base_group_norm(h, self.norm1, act_silu=True, channel_last=True)
h = self.conv2(h, is_init)
h = base_group_norm(h, self.norm2, act_silu=False, channel_last=True)
if self.down_sampling or self.num_filters != self.input_filters:
x = self.conv3(x, is_init)
x = base_group_norm(x, self.norm3, act_silu=False, channel_last=True)
h.add_(x)
h = self.act_(h)
if residual is not None:
h.add_(residual)
h = h.permute(0, 4, 1, 2, 3)
return h
class Upsample3D(nn.Module):
def __init__(self, in_channels, scale_factor=2):
super().__init__()
self.scale_factor = scale_factor
self.conv3d = Res3DBlockUpsample(input_filters=in_channels,
num_filters=in_channels,
down_sampling_stride=(1, 1, 1),
down_sampling=False)
def forward(self, x, is_init=True, is_split=True):
b, c, t, h, w = x.shape
# x = x.permute(0,2,3,4,1).contiguous().permute(0,4,1,2,3).to(memory_format=torch.channels_last_3d)
if is_split:
split_size = c // 8
x_slices = torch.split(x, split_size, dim=1)
x = [nn.functional.interpolate(x, scale_factor=self.scale_factor) for x in x_slices]
x = torch.cat(x, dim=1)
else:
x = nn.functional.interpolate(x, scale_factor=self.scale_factor)
x = self.conv3d(x, is_init)
return x
class VideoDecoder(nn.Module):
def __init__(
self,
ch=128,
z_channels=16,
out_channels=3,
ch_mult=(1, 2, 4, 4),
num_res_blocks=2,
temporal_up_layers=[2, 3],
temporal_downsample=4,
resamp_with_conv=True,
version=1,
):
super().__init__()
temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.temporal_downsample = temporal_downsample
block_in = ch * ch_mult[self.num_resolutions - 1]
self.version = version
if version == 2:
channels = 4 * z_channels * 2**3
self.conv_in = CausalConv(z_channels, channels, kernel_size=3)
self.shortcut_in = ChannelDuplicatingPixelUnshuffleUpSampleLayer3D(z_channels, channels, 1)
self.conv_unpatchify = ConvPixelShuffleUpSampleLayer3D(channels, block_in, kernel_size=3, factor=2)
self.shortcut_unpathify = ChannelDuplicatingPixelUnshuffleUpSampleLayer3D(channels, block_in, 2)
else:
self.conv_in = CausalConv(z_channels, block_in, kernel_size=3)
# middle
self.mid = nn.Module()
self.mid.block_1 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)
# upsampling
self.up_id = len(temporal_up_layers)
self.video_frame_num = 1
self.cur_video_frame_num = self.video_frame_num // 2**self.up_id + 1
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch * ch_mult[i_level]
for i_block in range(self.num_res_blocks + 1):
block.append(Resnet3DBlock(in_channels=block_in, out_channels=block_out, temb_channels=temb_ch))
block_in = block_out
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
if i_level in temporal_up_layers:
up.upsample = Upsample3D(block_in)
self.cur_video_frame_num = self.cur_video_frame_num * 2
else:
up.upsample = Upsample2D(block_in, resamp_with_conv)
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_in)
self.conv_out = CausalConvAfterNorm(block_in, out_channels, kernel_size=3)
@torch.inference_mode()
def forward(self, z, is_init=True):
z = rearrange(z, "b t c h w -> b c t h w")
h = self.conv_in(z, is_init=is_init)
if self.version == 2:
shortcut = self.shortcut_in(z, is_init=is_init)
h = h.add_(shortcut)
shortcut = self.shortcut_unpathify(h, is_init=is_init)
h = self.conv_unpatchify(h, is_init=is_init)
h = h.add_(shortcut)
temb = None
h = h.permute(0, 2, 3, 4, 1).contiguous().permute(0, 4, 1, 2, 3)
h = self.mid.block_1(h, temb, is_init=is_init)
h = self.mid.attn_1(h)
h = h.permute(0, 2, 3, 4, 1).contiguous().permute(0, 4, 1, 2, 3)
h = self.mid.block_2(h, temb, is_init=is_init)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks + 1):
h = h.permute(0, 2, 3, 4, 1).contiguous().permute(0, 4, 1, 2, 3)
h = self.up[i_level].block[i_block](h, temb, is_init=is_init)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
if isinstance(self.up[i_level].upsample, Upsample2D):
B = h.size(0)
h = h.permute(0, 2, 3, 4, 1).flatten(0, 1)
h = self.up[i_level].upsample(h)
h = h.unflatten(0, (B, -1)).permute(0, 4, 1, 2, 3)
else:
h = self.up[i_level].upsample(h, is_init=is_init)
# end
h = h.permute(0, 2, 3, 4, 1) # b c l h w -> b l h w c
h = base_group_norm_with_zero_pad(h, self.norm_out, act_silu=True, pad_size=2)
h = self.conv_out(h)
h = h.permute(0, 4, 1, 2, 3)
if is_init:
h = h[:, :, (self.temporal_downsample - 1):]
return h
def rms_norm(input, normalized_shape, eps=1e-6):
dtype = input.dtype
input = input.to(torch.float32)
variance = input.pow(2).flatten(-len(normalized_shape)).mean(-1)[(..., ) + (None, ) * len(normalized_shape)]
input = input * torch.rsqrt(variance + eps)
return input.to(dtype)
class DiagonalGaussianDistribution(object):
def __init__(self, parameters, deterministic=False, rms_norm_mean=False, only_return_mean=False):
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=-3) #N,[X],C,H,W
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
self.deterministic = deterministic
if self.deterministic:
self.var = self.std = torch.zeros_like(self.mean,
device=self.parameters.device,
dtype=self.parameters.dtype)
if rms_norm_mean:
self.mean = rms_norm(self.mean, self.mean.size()[1:])
self.only_return_mean = only_return_mean
def sample(self, generator=None):
# make sure sample is on the same device
# as the parameters and has same dtype
sample = torch.randn(self.mean.shape, generator=generator, device=self.parameters.device)
sample = sample.to(dtype=self.parameters.dtype)
x = self.mean + self.std * sample
if self.only_return_mean:
return self.mean
else:
return x
class AutoencoderKL(nn.Module):
@with_empty_init
def __init__(
self,
in_channels=3,
out_channels=3,
z_channels=16,
num_res_blocks=2,
model_path=None,
weight_dict={},
world_size=1,
version=1,
):
super().__init__()
self.frame_len = 17
self.latent_len = 3 if version == 2 else 5
base_group_norm.spatial = True if version == 2 else False
self.encoder = VideoEncoder(
in_channels=in_channels,
z_channels=z_channels,
num_res_blocks=num_res_blocks,
version=version,
)
self.decoder = VideoDecoder(
z_channels=z_channels,
out_channels=out_channels,
num_res_blocks=num_res_blocks,
version=version,
)
if model_path is not None:
weight_dict = self.init_from_ckpt(model_path)
if len(weight_dict) != 0:
self.load_from_dict(weight_dict)
self.convert_channel_last()
self.world_size = world_size
def init_from_ckpt(self, model_path):
from safetensors import safe_open
p = {}
with safe_open(model_path, framework="pt", device="cpu") as f:
for k in f.keys():
tensor = f.get_tensor(k)
if k.startswith("decoder.conv_out."):
k = k.replace("decoder.conv_out.", "decoder.conv_out.conv.")
p[k] = tensor
return p
def load_from_dict(self, p):
self.load_state_dict(p)
def convert_channel_last(self):
#Conv2d NCHW->NHWC
pass
def naive_encode(self, x, is_init_image=True):
b, len, c, h, w = x.size()
x = rearrange(x, 'b l c h w -> b c l h w').contiguous()
z = self.encoder(x, len, True) # 下采样[1, 4, 8, 16, 16]
return z
@torch.inference_mode()
def encode(self, x):
# b (nc cf) c h w -> (b nc) cf c h w -> encode -> (b nc) cf c h w -> b (nc cf) c h w
chunks = list(x.split(self.frame_len, dim=1))
for i in range(len(chunks)):
chunks[i] = self.naive_encode(chunks[i], True)
z = torch.cat(chunks, dim=1)
posterior = DiagonalGaussianDistribution(z)
return posterior.sample()
def decode_naive(self, z, is_init=True):
z = z.to(next(self.decoder.parameters()).dtype)
dec = self.decoder(z, is_init)
return dec
@torch.inference_mode()
def decode(self, z):
# b (nc cf) c h w -> (b nc) cf c h w -> decode -> (b nc) c cf h w -> b (nc cf) c h w
chunks = list(z.split(self.latent_len, dim=1))
if self.world_size > 1:
chunks_total_num = len(chunks)
max_num_per_rank = (chunks_total_num + self.world_size - 1) // self.world_size
rank = torch.distributed.get_rank()
chunks_ = chunks[max_num_per_rank * rank:max_num_per_rank * (rank + 1)]
if len(chunks_) < max_num_per_rank:
chunks_.extend(chunks[:max_num_per_rank - len(chunks_)])
chunks = chunks_
for i in range(len(chunks)):
chunks[i] = self.decode_naive(chunks[i], True).permute(0, 2, 1, 3, 4)
x = torch.cat(chunks, dim=1)
if self.world_size > 1:
x_ = torch.empty([x.size(0), (self.world_size * max_num_per_rank) * self.frame_len, *x.shape[2:]],
dtype=x.dtype,
device=x.device)
torch.distributed.all_gather_into_tensor(x_, x)
x = x_[:, :chunks_total_num * self.frame_len]
x = self.mix(x)
return x
def mix(self, x):
remain_scale = 0.6
mix_scale = 1. - remain_scale
front = slice(self.frame_len - 1, x.size(1) - 1, self.frame_len)
back = slice(self.frame_len, x.size(1), self.frame_len)
x[:, back] = x[:, back] * remain_scale + x[:, front] * mix_scale
x[:, front] = x[:, front] * remain_scale + x[:, back] * mix_scale
return x
@@ -0,0 +1,179 @@
import argparse
import os
import pickle
import threading
import torch
from flask import Blueprint, Flask, Response, request
from flask_restful import Api, Resource
device = f'cuda:{torch.cuda.device_count()-1}'
dtype = torch.bfloat16
def parsed_args():
parser = argparse.ArgumentParser(description="StepVideo API Functions")
parser.add_argument('--model_dir', type=str)
parser.add_argument('--clip_dir', type=str, default='hunyuan_clip')
parser.add_argument('--llm_dir', type=str, default='step_llm')
parser.add_argument('--vae_dir', type=str, default='vae')
parser.add_argument('--port', type=str, default='8080')
args = parser.parse_args()
return args
class StepVaePipeline(Resource):
def __init__(self, vae_dir, version=2):
self.vae = self.build_vae(vae_dir, version)
self.scale_factor = 1.0
def build_vae(self, vae_dir, version=2):
from fastvideo.models.stepvideo.vae.vae import AutoencoderKL
(model_name, z_channels) = ("vae_v2.safetensors", 64) if version == 2 else ("vae.safetensors", 16)
model_path = os.path.join(vae_dir, model_name)
model = AutoencoderKL(
z_channels=z_channels,
model_path=model_path,
version=version,
).to(dtype).to(device).eval()
print("Initialized vae...")
return model
def decode(self, samples, *args, **kwargs):
with torch.no_grad():
try:
dtype = next(self.vae.parameters()).dtype
device = next(self.vae.parameters()).device
samples = self.vae.decode(samples.to(dtype).to(device) / self.scale_factor)
if hasattr(samples, 'sample'):
samples = samples.sample
return samples
except:
torch.cuda.empty_cache()
return None
lock = threading.Lock()
class VAEapi(Resource):
def __init__(self, vae_pipeline):
self.vae_pipeline = vae_pipeline
def get(self):
with lock:
try:
feature = pickle.loads(request.get_data())
feature['api'] = 'vae'
feature = {k: v for k, v in feature.items() if v is not None}
video_latents = self.vae_pipeline.decode(**feature)
response = pickle.dumps(video_latents)
except Exception as e:
print("Caught Exception: ", e)
return Response(e)
return Response(response)
class CaptionPipeline(Resource):
def __init__(self, llm_dir, clip_dir):
self.text_encoder = self.build_llm(llm_dir)
self.clip = self.build_clip(clip_dir)
def build_llm(self, model_dir):
from fastvideo.models.stepvideo.text_encoder.stepllm import STEP1TextEncoder
text_encoder = STEP1TextEncoder(model_dir, max_length=320).to(dtype).to(device).eval()
print("Initialized text encoder...")
return text_encoder
def build_clip(self, model_dir):
from fastvideo.models.stepvideo.text_encoder.clip import HunyuanClip
clip = HunyuanClip(model_dir, max_length=77).to(device).eval()
print("Initialized clip encoder...")
return clip
def embedding(self, prompts, *args, **kwargs):
with torch.no_grad():
try:
y, y_mask = self.text_encoder(prompts)
clip_embedding, _ = self.clip(prompts)
len_clip = clip_embedding.shape[1]
y_mask = torch.nn.functional.pad(y_mask, (len_clip, 0),
value=1) ## pad attention_mask with clip's length
data = {
'y': y.detach().cpu(),
'y_mask': y_mask.detach().cpu(),
'clip_embedding': clip_embedding.to(torch.bfloat16).detach().cpu()
}
return data
except Exception as err:
print(f"{err}")
return None
lock = threading.Lock()
class Captionapi(Resource):
def __init__(self, caption_pipeline):
self.caption_pipeline = caption_pipeline
def get(self):
with lock:
try:
feature = pickle.loads(request.get_data())
feature['api'] = 'caption'
feature = {k: v for k, v in feature.items() if v is not None}
embeddings = self.caption_pipeline.embedding(**feature)
response = pickle.dumps(embeddings)
except Exception as e:
print("Caught Exception: ", e)
return Response(e)
return Response(response)
class RemoteServer(object):
def __init__(self, args) -> None:
self.app = Flask(__name__)
root = Blueprint("root", __name__)
self.app.register_blueprint(root)
api = Api(self.app)
self.vae_pipeline = StepVaePipeline(vae_dir=os.path.join(args.model_dir, args.vae_dir))
api.add_resource(
VAEapi,
"/vae-api",
resource_class_args=[self.vae_pipeline],
)
self.caption_pipeline = CaptionPipeline(llm_dir=os.path.join(args.model_dir, args.llm_dir),
clip_dir=os.path.join(args.model_dir, args.clip_dir))
api.add_resource(
Captionapi,
"/caption-api",
resource_class_args=[self.caption_pipeline],
)
def run(self, host="0.0.0.0", port=8080):
self.app.run(host, port=port, threaded=True, debug=False)
if __name__ == "__main__":
args = parsed_args()
flask_server = RemoteServer(args)
flask_server.run(host="0.0.0.0", port=args.port)
+14 -33
View File
@@ -9,8 +9,7 @@ from diffusers.utils import export_to_video
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
def generate_video_and_latent(pipe, prompt, height, width, num_frames,
num_inference_steps, guidance_scale):
def generate_video_and_latent(pipe, prompt, height, width, num_frames, num_inference_steps, guidance_scale):
# Set the random seed for reproducibility
generator = torch.Generator("cuda").manual_seed(12345)
# Generate videos from the input prompt
@@ -25,8 +24,7 @@ def generate_video_and_latent(pipe, prompt, height, width, num_frames,
output_type="latent_and_video",
)
# prompt_embed has negative prompt at index 0
return noise[0], video[0], latent[0], prompt_embed[
1], prompt_attention_mask[1]
return noise[0], video[0], latent[0], prompt_embed[1], prompt_attention_mask[1]
# return dummy tensor to debug first
# return torch.zeros(1, 3, 480, 848), torch.zeros(1, 256, 16, 16)
@@ -40,22 +38,15 @@ if __name__ == "__main__":
parser.add_argument("--num_inference_steps", type=int, default=64)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--prompt_path",
type=str,
default="data/dummyVid/videos2caption.json")
parser.add_argument("--dataset_output_dir",
type=str,
default="data/dummySynthetic")
parser.add_argument("--prompt_path", type=str, default="data/dummyVid/videos2caption.json")
parser.add_argument("--dataset_output_dir", type=str, default="data/dummySynthetic")
args = parser.parse_args()
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size, "local rank", local_rank)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
if not isinstance(args.prompt_path, list):
args.prompt_path = [args.prompt_path]
@@ -63,8 +54,7 @@ if __name__ == "__main__":
text_prompt = open(args.prompt_path[0], "r").readlines()
text_prompt = [i.strip() for i in text_prompt]
pipe = MochiPipeline.from_pretrained(args.model_path,
torch_dtype=torch.bfloat16)
pipe = MochiPipeline.from_pretrained(args.model_path, torch_dtype=torch.bfloat16)
pipe.enable_vae_tiling()
pipe.enable_model_cpu_offload(gpu_id=local_rank)
# make dir if not exist
@@ -73,10 +63,8 @@ if __name__ == "__main__":
os.makedirs(os.path.join(args.dataset_output_dir, "noise"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_embed"),
exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_attention_mask"),
exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_attention_mask"), exist_ok=True)
data = []
for i, prompt in enumerate(text_prompt):
if i % world_size != local_rank:
@@ -98,17 +86,11 @@ if __name__ == "__main__":
)
# save latent
video_name = str(i)
noise_path = os.path.join(args.dataset_output_dir, "noise",
video_name + ".pt")
latent_path = os.path.join(args.dataset_output_dir, "latent",
video_name + ".pt")
prompt_embed_path = os.path.join(args.dataset_output_dir,
"prompt_embed", video_name + ".pt")
video_path = os.path.join(args.dataset_output_dir, "video",
video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.dataset_output_dir,
"prompt_attention_mask",
video_name + ".pt")
noise_path = os.path.join(args.dataset_output_dir, "noise", video_name + ".pt")
latent_path = os.path.join(args.dataset_output_dir, "latent", video_name + ".pt")
prompt_embed_path = os.path.join(args.dataset_output_dir, "prompt_embed", video_name + ".pt")
video_path = os.path.join(args.dataset_output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.dataset_output_dir, "prompt_attention_mask", video_name + ".pt")
# save latent
torch.save(noise, noise_path)
torch.save(latent, latent_path)
@@ -132,6 +114,5 @@ if __name__ == "__main__":
# save json
if local_rank == 0:
all_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.dataset_output_dir, "videos2caption.json"),
"w") as f:
with open(os.path.join(args.dataset_output_dir, "videos2caption.json"), "w") as f:
json.dump(all_data, f, indent=4)
+29 -62
View File
@@ -10,8 +10,7 @@ import torchvision
from einops import rearrange
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state, nccl_info)
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
def initialize_distributed():
@@ -19,10 +18,7 @@ def initialize_distributed():
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
initialize_sequence_parallel_state(world_size)
@@ -40,14 +36,16 @@ def main(args):
os.makedirs(os.path.dirname(save_path), exist_ok=True)
# Load models
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(
models_root_path, args=args)
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
# Get the updated args
args = hunyuan_video_sampler.args
with open(args.prompt) as f:
prompts = f.readlines()
if args.prompt.endswith('.txt'):
with open(args.prompt) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [args.prompt]
for prompt in prompts:
outputs = hunyuan_video_sampler.predict(
@@ -71,9 +69,7 @@ def main(args):
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
outputs.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
imageio.mimsave(os.path.join(args.output_path, f"{prompt[:100]}.mp4"),
outputs,
fps=args.fps)
imageio.mimsave(os.path.join(args.output_path, f"{prompt[:100]}.mp4"), outputs, fps=args.fps)
if __name__ == "__main__":
@@ -96,14 +92,8 @@ if __name__ == "__main__":
default="flow",
help="Denoise type for noised inputs.",
)
parser.add_argument("--seed",
type=int,
default=None,
help="Seed for evaluation.")
parser.add_argument("--neg_prompt",
type=str,
default=None,
help="Negative prompt for sampling.")
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
parser.add_argument(
"--guidance_scale",
type=float,
@@ -116,14 +106,8 @@ if __name__ == "__main__":
default=6.0,
help="Embedded classifier free guidance scale.",
)
parser.add_argument("--flow_shift",
type=int,
default=7,
help="Flow shift parameter.")
parser.add_argument("--batch_size",
type=int,
default=1,
help="Batch size for inference.")
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
parser.add_argument(
"--num_videos",
type=int,
@@ -134,8 +118,7 @@ if __name__ == "__main__":
"--load-key",
type=str,
default="module",
help=
"Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
parser.add_argument(
"--use-cpu-offload",
@@ -145,20 +128,17 @@ if __name__ == "__main__":
parser.add_argument(
"--dit-weight",
type=str,
default=
"data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
)
parser.add_argument(
"--reproduce",
action="store_true",
help=
"Enable reproducibility by setting random seeds and deterministic algorithms.",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
help=
"Disable autocast for denoising loop and vae decoding in pipeline sampling.",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
# Flow Matching
@@ -167,10 +147,7 @@ if __name__ == "__main__":
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
parser.add_argument("--flow-solver",
type=str,
default="euler",
help="Solver for flow matching.")
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
parser.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
@@ -187,20 +164,11 @@ if __name__ == "__main__":
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument("--precision",
type=str,
default="bf16",
choices=["fp32", "fp16", "bf16"])
parser.add_argument("--rope-theta",
type=int,
default=256,
help="Theta used in RoPE.")
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument("--vae-precision",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--vae-sp", action="store_true", default=False)
@@ -214,12 +182,8 @@ if __name__ == "__main__":
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template",
type=str,
default="dit-llm-encode")
parser.add_argument("--prompt-template-video",
type=str,
default="dit-llm-encode-video")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
@@ -230,6 +194,11 @@ if __name__ == "__main__":
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
parser.add_argument(
"--enable_torch_compile",
action="store_true",
help="Use torch.compile for speeding up STA inference without teacache",
)
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)
@@ -237,7 +206,5 @@ if __name__ == "__main__":
args = parser.parse_args()
# process for vae sequence parallel
if args.vae_sp and not args.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
main(args)
+407
View File
@@ -0,0 +1,407 @@
import argparse
import json
import os
from pathlib import Path
from typing import Any, Dict, Optional, Union
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torchvision
from einops import rearrange
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
from fastvideo.models.hunyuan.modules.modulate_layers import modulate
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
def teacache_forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
mask_strategy=None,
output_features=False,
output_features_stride=8,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
guidance=None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance is None:
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
img = x = hidden_states
text_mask = encoder_attention_mask
t = timestep
txt = encoder_hidden_states[:, 1:]
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
_, _, ot, oh, ow = x.shape # codespell:ignore
tt, th, tw = (
ot // self.patch_size[0], # codespell:ignore
oh // self.patch_size[1], # codespell:ignore
ow // self.patch_size[2], # codespell:ignore
)
original_tt = nccl_info.sp_size * tt
freqs_cos, freqs_sin = self.get_rotary_pos_embed((original_tt, th, tw))
# Prepare modulation vectors.
vec = self.time_in(t)
# text modulation
vec = vec + self.vector_in(text_states_2)
# guidance modulation
if self.guidance_embed:
if guidance is None:
raise ValueError("Didn't get guidance strength for guidance distilled model.")
# our timestep_embedding is merged into guidance_in(TimestepEmbedder)
vec = vec + self.guidance_in(guidance)
# Embed image and text.
img = self.img_in(img)
if self.text_projection == "linear":
txt = self.txt_in(txt)
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
else:
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
if self.enable_teacache:
inp = img.clone()
vec_ = vec.clone()
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.double_blocks[0].img_mod(vec_).chunk(6, dim=-1)
normed_inp = self.double_blocks[0].img_norm1(inp)
modulated_inp = modulate(normed_inp, shift=img_mod1_shift, scale=img_mod1_scale)
if self.cnt == 0 or self.cnt == self.num_steps - 1:
should_calc = True
self.accumulated_rel_l1_distance = 0
else:
coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02]
rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance += rescale_func(
((modulated_inp - self.previous_modulated_input).abs().mean() /
self.previous_modulated_input.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
should_calc = False
else:
should_calc = True
self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = modulated_inp
self.cnt += 1
if self.cnt == self.num_steps:
self.cnt = 0
if self.enable_teacache:
if not should_calc:
img += self.previous_residual
else:
ori_img = img.clone()
# --------------------- Pass through DiT blocks ------------------------
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if output_features:
features_list = []
if len(self.single_blocks) > 0:
for index, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
(freqs_cos, freqs_sin),
text_mask,
mask_strategy[index + len(self.double_blocks)],
]
x = block(*single_block_args)
if output_features and _ % output_features_stride == 0:
features_list.append(x[:, :img_seq_len, ...])
img = x[:, :img_seq_len, ...]
self.previous_residual = img - ori_img
else:
# --------------------- Pass through DiT blocks ------------------------
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if output_features:
features_list = []
if len(self.single_blocks) > 0:
for index, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
(freqs_cos, freqs_sin),
text_mask,
mask_strategy[index + len(self.double_blocks)],
]
x = block(*single_block_args)
if output_features and _ % output_features_stride == 0:
features_list.append(x[:, :img_seq_len, ...])
img = x[:, :img_seq_len, ...]
# ---------------------------- Final layer ------------------------------
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
img = self.unpatchify(img, tt, th, tw)
assert not return_dict, "return_dict is not supported."
if output_features:
features_list = torch.stack(features_list, dim=0)
else:
features_list = None
return (img, features_list)
def initialize_distributed():
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
initialize_sequence_parallel_state(world_size)
def main(args):
initialize_distributed()
print(nccl_info.sp_size)
print(args)
models_root_path = Path(args.model_path)
if not models_root_path.exists():
raise ValueError(f"`models_root` not exists: {models_root_path}")
# Create save folder to save the samples
save_path = args.output_path
os.makedirs(os.path.dirname(save_path), exist_ok=True)
# Load models
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
# Get the updated args
args = hunyuan_video_sampler.args
# teacache
hunyuan_video_sampler.pipeline.transformer.__class__.enable_teacache = args.enable_teacache
hunyuan_video_sampler.pipeline.transformer.__class__.cnt = 0
hunyuan_video_sampler.pipeline.transformer.__class__.num_steps = args.num_inference_steps
hunyuan_video_sampler.pipeline.transformer.__class__.rel_l1_thresh = args.rel_l1_thresh # 0.1 for 1.6x speedup, 0.15 for 2.1x speedup
hunyuan_video_sampler.pipeline.transformer.__class__.accumulated_rel_l1_distance = 0
hunyuan_video_sampler.pipeline.transformer.__class__.previous_modulated_input = None
hunyuan_video_sampler.pipeline.transformer.__class__.previous_residual = None
hunyuan_video_sampler.pipeline.transformer.__class__.forward = teacache_forward
with open(args.mask_strategy_file_path, 'r') as f:
mask_strategy = json.load(f)
if args.prompt.endswith('.txt'):
with open(args.prompt) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [args.prompt]
for prompt in prompts:
outputs = hunyuan_video_sampler.predict(
prompt=prompt,
height=args.height,
width=args.width,
video_length=args.num_frames,
seed=args.seed,
negative_prompt=args.neg_prompt,
infer_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
num_videos_per_prompt=args.num_videos,
flow_shift=args.flow_shift,
batch_size=args.batch_size,
embedded_guidance_scale=args.embedded_cfg_scale,
mask_strategy=mask_strategy,
)
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
outputs = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
outputs.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
imageio.mimsave(os.path.join(args.output_path, f"{prompt[:100]}.mp4"), outputs, fps=args.fps)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# Basic parameters
parser.add_argument("--prompt", type=str, help="prompt file for inference")
parser.add_argument("--num_frames", type=int, default=16)
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--model_path", type=str, default="data/hunyuan")
parser.add_argument("--output_path", type=str, default="./outputs/video")
parser.add_argument("--fps", type=int, default=24)
# Additional parameters
parser.add_argument(
"--sliding_block_size",
type=str,
default="8,6,10",
help="Sliding block size for sliding block attention.",
)
parser.add_argument(
"--denoise-type",
type=str,
default="flow",
help="Denoise type for noised inputs.",
)
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
parser.add_argument(
"--guidance_scale",
type=float,
default=1.0,
help="Classifier free guidance scale.",
)
parser.add_argument(
"--embedded_cfg_scale",
type=float,
default=6.0,
help="Embedded classifier free guidance scale.",
)
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
parser.add_argument(
"--num_videos",
type=int,
default=1,
help="Number of videos to generate per prompt.",
)
parser.add_argument(
"--load-key",
type=str,
default="module",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
parser.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
parser.add_argument(
"--dit-weight",
type=str,
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
)
parser.add_argument(
"--reproduce",
action="store_true",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
# Flow Matching
parser.add_argument(
"--flow-reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
parser.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
help=
"Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
)
parser.add_argument(
"--linear-schedule-end",
type=int,
default=25,
help="End step for linear quadratic schedule for flow matching.",
)
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--vae-sp", action="store_true", default=False)
parser.add_argument("--text-encoder", type=str, default="llm")
parser.add_argument(
"--text-encoder-precision",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
parser.add_argument("--text-encoder-2", type=str, default="clipL")
parser.add_argument(
"--text-encoder-precision-2",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
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("--skip_time_steps", type=int, default=10)
parser.add_argument(
"--mask_strategy_selected",
type=lambda x: [int(i) for i in x.strip('[]').split(',')], # Convert string to list of integers
default=[1, 2, 6], # Now can be directly set as a list
help="order of candidates")
parser.add_argument(
"--rel_l1_thresh",
type=float,
default=0.15,
help="0.1 for 1.6x speedup, 0.15 for 2.1x speedup",
)
parser.add_argument(
"--enable_teacache",
action="store_true",
help="Use teacache for speeding up inference",
)
parser.add_argument(
"--enable_torch_compile",
action="store_true",
help="Use torch.compile for speeding up STA inference without teacache",
)
parser.add_argument("--mask_strategy_file_path", type=str, default="assets/mask_strategy.json")
args = parser.parse_args()
# process for vae sequence parallel
if args.vae_sp and not args.vae_tiling:
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
if args.enable_teacache and args.enable_torch_compile:
raise ValueError(
"--enable_teacache and --enable_torch_compile cannot be used simultaneously. Please enable only one of these options."
)
main(args)
+55 -116
View File
@@ -8,11 +8,9 @@ import torch.distributed as dist
from diffusers import BitsAndBytesConfig
from diffusers.utils import export_to_video
from fastvideo.models.hunyuan_hf.modeling_hunyuan import \
HunyuanVideoTransformer3DModel
from fastvideo.models.hunyuan_hf.modeling_hunyuan import HunyuanVideoTransformer3DModel
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state, nccl_info)
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
def initialize_distributed():
@@ -21,10 +19,7 @@ def initialize_distributed():
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
initialize_sequence_parallel_state(world_size)
@@ -36,35 +31,27 @@ def inference(args):
weight_dtype = torch.bfloat16
if args.transformer_path is not None:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
args.transformer_path)
transformer = HunyuanVideoTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
args.model_path,
subfolder="transformer/",
torch_dtype=weight_dtype)
transformer = HunyuanVideoTransformer3DModel.from_pretrained(args.model_path,
subfolder="transformer/",
torch_dtype=weight_dtype)
pipe = HunyuanVideoPipeline.from_pretrained(args.model_path,
transformer=transformer,
torch_dtype=weight_dtype)
pipe = HunyuanVideoPipeline.from_pretrained(args.model_path, transformer=transformer, torch_dtype=weight_dtype)
pipe.enable_vae_tiling()
if args.lora_checkpoint_dir is not None:
print(f"Loading LoRA weights from {args.lora_checkpoint_dir}")
config_path = os.path.join(args.lora_checkpoint_dir,
"lora_config.json")
config_path = os.path.join(args.lora_checkpoint_dir, "lora_config.json")
with open(config_path, "r") as f:
lora_config_dict = json.load(f)
rank = lora_config_dict["lora_params"]["lora_rank"]
lora_alpha = lora_config_dict["lora_params"]["lora_alpha"]
lora_scaling = lora_alpha / rank
pipe.load_lora_weights(args.lora_checkpoint_dir,
adapter_name="default")
pipe.load_lora_weights(args.lora_checkpoint_dir, adapter_name="default")
pipe.set_adapters(["default"], [lora_scaling])
print(
f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}"
)
print(f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}")
if args.cpu_offload:
pipe.enable_model_cpu_offload(device)
else:
@@ -73,13 +60,10 @@ def inference(args):
# Generate videos from the input prompt
if args.prompt_embed_path is not None:
prompt_embeds = (torch.load(args.prompt_embed_path,
map_location="cpu",
prompt_embeds = (torch.load(args.prompt_embed_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
encoder_attention_mask = (torch.load(
args.encoder_attention_mask_path,
map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
encoder_attention_mask = (torch.load(args.encoder_attention_mask_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
prompts = None
elif args.prompt_path is not None:
prompts = [line.strip() for line in open(args.prompt_path, "r")]
@@ -133,52 +117,44 @@ def inference_quantization(args):
model_id = args.model_path
if args.quantization == "nf4":
quantization_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_quant_type="nf4",
llm_int8_skip_modules=["proj_out", "norm_out"])
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id,
subfolder="transformer/",
torch_dtype=torch.bfloat16,
quantization_config=quantization_config)
quantization_config = BitsAndBytesConfig(load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_quant_type="nf4",
llm_int8_skip_modules=["proj_out", "norm_out"])
transformer = HunyuanVideoTransformer3DModel.from_pretrained(model_id,
subfolder="transformer/",
torch_dtype=torch.bfloat16,
quantization_config=quantization_config)
if args.quantization == "int8":
quantization_config = BitsAndBytesConfig(
load_in_8bit=True, llm_int8_skip_modules=["proj_out", "norm_out"])
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id,
subfolder="transformer/",
torch_dtype=torch.bfloat16,
quantization_config=quantization_config)
quantization_config = BitsAndBytesConfig(load_in_8bit=True, llm_int8_skip_modules=["proj_out", "norm_out"])
transformer = HunyuanVideoTransformer3DModel.from_pretrained(model_id,
subfolder="transformer/",
torch_dtype=torch.bfloat16,
quantization_config=quantization_config)
elif not args.quantization:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer/",
torch_dtype=torch.bfloat16).to(device)
transformer = HunyuanVideoTransformer3DModel.from_pretrained(model_id,
subfolder="transformer/",
torch_dtype=torch.bfloat16).to(device)
print("Max vram for read transformer:",
round(torch.cuda.max_memory_allocated(device="cuda") / 1024**3, 3),
"GiB")
print("Max vram for read transformer:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024**3, 3), "GiB")
torch.cuda.reset_max_memory_allocated(device)
if not args.cpu_offload:
pipe = HunyuanVideoPipeline.from_pretrained(
model_id, torch_dtype=torch.bfloat16).to(device)
pipe = HunyuanVideoPipeline.from_pretrained(model_id, torch_dtype=torch.bfloat16).to(device)
pipe.transformer = transformer
else:
pipe = HunyuanVideoPipeline.from_pretrained(model_id,
transformer=transformer,
torch_dtype=torch.bfloat16)
pipe = HunyuanVideoPipeline.from_pretrained(model_id, transformer=transformer, torch_dtype=torch.bfloat16)
torch.cuda.reset_max_memory_allocated(device)
pipe.scheduler._shift = args.flow_shift
pipe.vae.enable_tiling()
if args.cpu_offload:
pipe.enable_model_cpu_offload()
print("Max vram for init pipeline:",
round(torch.cuda.max_memory_allocated(device="cuda") / 1024**3, 3),
"GiB")
with open(args.prompt) as f:
prompts = f.readlines()
print("Max vram for init pipeline:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024**3, 3), "GiB")
if args.prompt.endswith('.txt'):
with open(args.prompt) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [args.prompt]
generator = torch.Generator("cpu").manual_seed(args.seed)
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
@@ -193,14 +169,9 @@ def inference_quantization(args):
num_inference_steps=args.num_inference_steps,
generator=generator,
).frames[0]
export_to_video(output,
os.path.join(args.output_path, f"{prompt[:100]}.mp4"),
fps=args.fps)
export_to_video(output, os.path.join(args.output_path, f"{prompt[:100]}.mp4"), fps=args.fps)
print("Time:", round(time.perf_counter() - start_time, 2), "seconds")
print(
"Max vram for denoise:",
round(torch.cuda.max_memory_allocated(device="cuda") / 1024**3, 3),
"GiB")
print("Max vram for denoise:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024**3, 3), "GiB")
if __name__ == "__main__":
@@ -233,14 +204,8 @@ if __name__ == "__main__":
default="flow",
help="Denoise type for noised inputs.",
)
parser.add_argument("--seed",
type=int,
default=None,
help="Seed for evaluation.")
parser.add_argument("--neg_prompt",
type=str,
default=None,
help="Negative prompt for sampling.")
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
parser.add_argument(
"--guidance_scale",
type=float,
@@ -253,14 +218,8 @@ if __name__ == "__main__":
default=6.0,
help="Embedded classifier free guidance scale.",
)
parser.add_argument("--flow_shift",
type=int,
default=7,
help="Flow shift parameter.")
parser.add_argument("--batch_size",
type=int,
default=1,
help="Batch size for inference.")
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
parser.add_argument(
"--num_videos",
type=int,
@@ -271,26 +230,22 @@ if __name__ == "__main__":
"--load-key",
type=str,
default="module",
help=
"Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
parser.add_argument(
"--dit-weight",
type=str,
default=
"data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
)
parser.add_argument(
"--reproduce",
action="store_true",
help=
"Enable reproducibility by setting random seeds and deterministic algorithms.",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
help=
"Disable autocast for denoising loop and vae decoding in pipeline sampling.",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
# Flow Matching
@@ -299,10 +254,7 @@ if __name__ == "__main__":
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
parser.add_argument("--flow-solver",
type=str,
default="euler",
help="Solver for flow matching.")
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
parser.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
@@ -319,20 +271,11 @@ if __name__ == "__main__":
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument("--precision",
type=str,
default="bf16",
choices=["fp32", "fp16", "bf16", "fp8"])
parser.add_argument("--rope-theta",
type=int,
default=256,
help="Theta used in RoPE.")
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16", "fp8"])
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument("--vae-precision",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--text-encoder", type=str, default="llm")
@@ -345,12 +288,8 @@ if __name__ == "__main__":
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template",
type=str,
default="dit-llm-encode")
parser.add_argument("--prompt-template-video",
type=str,
default="dit-llm-encode-video")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
+12 -29
View File
@@ -10,8 +10,7 @@ from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state, nccl_info)
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
def initialize_distributed():
@@ -19,10 +18,7 @@ def initialize_distributed():
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
initialize_sequence_parallel_state(world_size)
@@ -45,33 +41,25 @@ def main(args):
args.linear_range,
)
if args.transformer_path is not None:
transformer = MochiTransformer3DModel.from_pretrained(
args.transformer_path)
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/")
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(args.model_path,
transformer=transformer,
scheduler=scheduler)
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
if args.lora_checkpoint_dir is not None:
print(f"Loading LoRA weights from {args.lora_checkpoint_dir}")
config_path = os.path.join(args.lora_checkpoint_dir,
"lora_config.json")
config_path = os.path.join(args.lora_checkpoint_dir, "lora_config.json")
with open(config_path, "r") as f:
lora_config_dict = json.load(f)
rank = lora_config_dict["lora_params"]["lora_rank"]
lora_alpha = lora_config_dict["lora_params"]["lora_alpha"]
lora_scaling = lora_alpha / rank
pipe.load_lora_weights(args.lora_checkpoint_dir,
adapter_name="default")
pipe.load_lora_weights(args.lora_checkpoint_dir, adapter_name="default")
pipe.set_adapters(["default"], [lora_scaling])
print(
f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}"
)
print(f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}")
# pipe.to(device)
pipe.enable_model_cpu_offload(device)
@@ -79,13 +67,10 @@ def main(args):
# Generate videos from the input prompt
if args.prompt_embed_path is not None:
prompt_embeds = (torch.load(args.prompt_embed_path,
map_location="cpu",
prompt_embeds = (torch.load(args.prompt_embed_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
encoder_attention_mask = (torch.load(
args.encoder_attention_mask_path,
map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
encoder_attention_mask = (torch.load(args.encoder_attention_mask_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
prompts = None
elif args.prompt_path is not None:
prompts = [line.strip() for line in open(args.prompt_path, "r")]
@@ -151,9 +136,7 @@ if __name__ == "__main__":
parser.add_argument("--prompt_embed_path", type=str, default=None)
parser.add_argument("--prompt_path", type=str, default=None)
parser.add_argument("--scheduler_type", type=str, default="euler")
parser.add_argument("--encoder_attention_mask_path",
type=str,
default=None)
parser.add_argument("--encoder_attention_mask_path", type=str, default=None)
parser.add_argument(
"--lora_checkpoint_dir",
type=str,
+3 -7
View File
@@ -14,14 +14,10 @@ def main(args):
# do not invert
scheduler = FlowMatchEulerDiscreteScheduler()
if args.transformer_path is not None:
transformer = MochiTransformer3DModel.from_pretrained(
args.transformer_path)
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(args.model_path,
transformer=transformer,
scheduler=scheduler)
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
# pipe.to("cuda:1")
pipe.enable_model_cpu_offload()
+246
View File
@@ -0,0 +1,246 @@
import argparse
import os
import torch
import torch.distributed as dist
import torch.nn as nn
from fastvideo.models.stepvideo.diffusion.scheduler import FlowMatchDiscreteScheduler
from fastvideo.models.stepvideo.diffusion.video_pipeline import StepVideoPipeline
from fastvideo.models.stepvideo.modules.model import StepVideoModel
from fastvideo.models.stepvideo.utils import setup_seed
from fastvideo.models.stepvideo.utils.quantization import convert_fp8_linear, fp8_linear_forward
from fastvideo.utils.logging_ import main_print
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
def initialize_distributed():
os.environ["TOKENIZERS_PARALLELISM"] = "false"
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
print("world_size", world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
initialize_sequence_parallel_state(world_size)
def parse_args(namespace=None):
parser = argparse.ArgumentParser(description="StepVideo inference script")
parser = add_extra_models_args(parser)
parser = add_denoise_schedule_args(parser)
parser = add_inference_args(parser)
args = parser.parse_args(namespace=namespace)
return args
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
group.add_argument(
"--vae_url",
type=str,
default='127.0.0.1',
help="vae url.",
)
group.add_argument(
"--caption_url",
type=str,
default='127.0.0.1',
help="caption url.",
)
return parser
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Denoise schedule args")
# Flow Matching
group.add_argument(
"--time_shift",
type=float,
default=13,
help="Shift factor for flow matching schedulers.",
)
group.add_argument(
"--flow_reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
group.add_argument(
"--flow_solver",
type=str,
default="euler",
help="Solver for flow matching.",
)
return parser
def add_inference_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Inference args")
# ======================== Model loads ========================
group.add_argument(
"--model_dir",
type=str,
default="./ckpts",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--model_resolution",
type=str,
default="540p",
choices=["540p"],
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
group.add_argument(
"--use-fp8",
action="store_true",
help="FP8 Quantization for single GPU support.",
)
# ======================== Inference general setting ========================
group.add_argument(
"--batch_size",
type=int,
default=1,
help="Batch size for inference and evaluation.",
)
group.add_argument(
"--infer_steps",
type=int,
default=50,
help="Number of denoising steps for inference.",
)
group.add_argument(
"--save_path",
type=str,
default="./results",
help="Path to save the generated samples.",
)
group.add_argument(
"--name_suffix",
type=str,
default="",
help="Suffix for the names of saved samples.",
)
group.add_argument(
"--num_videos",
type=int,
default=1,
help="Number of videos to generate for each prompt.",
)
# ---sample size---
group.add_argument(
"--num_frames",
type=int,
default=204,
help="How many frames to sample from a video. ",
)
group.add_argument(
"--height",
type=int,
default=768,
help="The height of video sample",
)
group.add_argument(
"--width",
type=int,
default=768,
help="The width of video sample",
)
# --- prompt ---
group.add_argument(
"--prompt",
type=str,
default=None,
help="Prompt for sampling during evaluation.",
)
group.add_argument("--seed", type=int, default=1234, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument("--pos_magic",
type=str,
default="超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。",
help="Positive magic prompt for sampling.")
group.add_argument("--neg_magic",
type=str,
default="画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。",
help="Negative magic prompt for sampling.")
group.add_argument("--cfg_scale", type=float, default=9.0, help="Classifier free guidance scale.")
return parser
if __name__ == "__main__":
args = parse_args()
initialize_distributed()
main_print(f"sequence parallel size: {nccl_info.sp_size}")
device = torch.cuda.current_device()
setup_seed(args.seed)
main_print("Loading model, this might take a while...")
scheduler = FlowMatchDiscreteScheduler()
if args.use_fp8:
assert int(os.getenv("WORLD_SIZE", 1)) == 1
transformer = StepVideoModel.from_pretrained(os.path.join(args.model_dir, "transformer"),
torch_dtype=torch.bfloat16,
device="cpu")
if not os.path.exists(args.model_dir + "/fp8_transformer.pth"):
print("no_fp8 weight, creating...")
scale_dict = convert_fp8_linear(transformer, torch.bfloat16)
torch.save(transformer.state_dict(), args.model_dir + "/fp8_transformer.pth")
torch.save(scale_dict, args.model_dir + "/fp8_scale_dict.pth")
else:
transformer.load_state_dict(torch.load(args.model_dir + "/fp8_transformer.pth"))
scale_dict = torch.load(args.model_dir + "/fp8_scale_dict.pth")
original_dtype = torch.bfloat16
for key, layer in transformer.named_modules():
if isinstance(layer, nn.Linear) and 'transformer_blocks' in key and key in scale_dict:
layer.weight.data = layer.weight.data.to(torch.float8_e4m3fn)
print(f"{key}, layer.weight.dtype: {layer.weight.dtype}")
original_forward = layer.forward
scale = scale_dict[key]
setattr(layer, "fp8_scale", scale.to(dtype=original_dtype))
setattr(layer, "original_forward", original_forward)
setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input))
else:
transformer = StepVideoModel.from_pretrained(os.path.join(args.model_dir, "transformer"),
torch_dtype=torch.bfloat16,
device=device)
transformer = transformer.to(device)
pipeline = StepVideoPipeline(transformer, scheduler, save_path=args.save_path)
pipeline.setup_api(
vae_url=args.vae_url,
caption_url=args.caption_url,
)
if args.prompt.endswith('.txt'):
with open(args.prompt) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [args.prompt]
for prompt in prompts:
videos = pipeline(prompt=prompt,
num_frames=args.num_frames,
height=args.height,
width=args.width,
num_inference_steps=args.infer_steps,
guidance_scale=args.cfg_scale,
time_shift=args.time_shift,
pos_magic=args.pos_magic,
neg_magic=args.neg_magic,
output_file_name=prompt[:50])
dist.destroy_process_group()
@@ -0,0 +1,374 @@
import argparse
import json
import os
import types
from typing import Dict, Optional
import numpy as np
import torch
import torch.distributed as dist
from einops import rearrange, repeat
from fastvideo.models.stepvideo.diffusion.scheduler import FlowMatchDiscreteScheduler
from fastvideo.models.stepvideo.diffusion.video_pipeline import StepVideoPipeline
from fastvideo.models.stepvideo.modules.model import StepVideoModel
from fastvideo.models.stepvideo.utils import setup_seed
from fastvideo.utils.logging_ import main_print
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
def initialize_distributed():
os.environ["TOKENIZERS_PARALLELISM"] = "false"
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
main_print(f"world_size: {world_size}")
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
initialize_sequence_parallel_state(world_size)
def parse_args(namespace=None):
parser = argparse.ArgumentParser(description="StepVideo inference script")
parser = add_extra_models_args(parser)
parser = add_denoise_schedule_args(parser)
parser = add_inference_args(parser)
args = parser.parse_args(namespace=namespace)
return args
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
group.add_argument(
"--vae_url",
type=str,
default='127.0.0.1',
help="vae url.",
)
group.add_argument(
"--caption_url",
type=str,
default='127.0.0.1',
help="caption url.",
)
return parser
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Denoise schedule args")
# Flow Matching
group.add_argument(
"--time_shift",
type=float,
default=13,
help="Shift factor for flow matching schedulers.",
)
group.add_argument(
"--flow_reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
group.add_argument(
"--flow_solver",
type=str,
default="euler",
help="Solver for flow matching.",
)
return parser
def add_inference_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Inference args")
# ======================== Model loads ========================
group.add_argument(
"--model_dir",
type=str,
default="./ckpts",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--model_resolution",
type=str,
default="540p",
choices=["540p"],
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
# ======================== Inference general setting ========================
group.add_argument(
"--batch_size",
type=int,
default=1,
help="Batch size for inference and evaluation.",
)
group.add_argument(
"--infer_steps",
type=int,
default=50,
help="Number of denoising steps for inference.",
)
group.add_argument(
"--save_path",
type=str,
default="./results",
help="Path to save the generated samples.",
)
group.add_argument(
"--name_suffix",
type=str,
default="",
help="Suffix for the names of saved samples.",
)
group.add_argument(
"--num_videos",
type=int,
default=1,
help="Number of videos to generate for each prompt.",
)
# ---sample size---
group.add_argument(
"--num_frames",
type=int,
default=204,
help="How many frames to sample from a video. ",
)
group.add_argument(
"--height",
type=int,
default=768,
help="The height of video sample",
)
group.add_argument(
"--width",
type=int,
default=768,
help="The width of video sample",
)
# --- prompt ---
group.add_argument(
"--prompt",
type=str,
default=None,
help="Prompt for sampling during evaluation.",
)
group.add_argument("--seed", type=int, default=1234, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument("--pos_magic",
type=str,
default="超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。",
help="Positive magic prompt for sampling.")
group.add_argument("--neg_magic",
type=str,
default="画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。",
help="Negative magic prompt for sampling.")
group.add_argument("--cfg_scale", type=float, default=9.0, help="Classifier free guidance scale.")
group.add_argument("--mask_search_files_path", type=str, default="assets/mask_strategy.json")
group.add_argument("--mask_strategy_file_path", type=str, default="assets/mask_strategy_stepvideo.json")
group.add_argument("--skip_time_steps", type=int, default=10)
group.add_argument(
"--mask_strategy_selected",
type=lambda x: [int(i) for i in x.strip('[]').split(',')], # Convert string to list of integers
default=[1, 2, 6], # Now can be directly set as a list
help="order of candidates")
parser.add_argument(
"--rel_l1_thresh",
type=float,
default=0,
help="0.22 for 1.67x speedup, 0.23 for 2.1x speedup",
)
parser.add_argument(
"--enable_teacache",
action="store_true",
help="Use teacache for speeding up inference",
)
return parser
def teacache_forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
encoder_hidden_states_2: Optional[torch.Tensor] = None,
timestep: Optional[torch.LongTensor] = None,
added_cond_kwargs: Dict[str, torch.Tensor] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
fps: torch.Tensor = None,
return_dict: bool = True,
mask_strategy=None,
):
assert hidden_states.ndim == 5
"hidden_states's shape should be (bsz, f, ch, h ,w)"
bsz, frame, _, height, width = hidden_states.shape
height, width = height // self.patch_size, width // self.patch_size
hidden_states = self.patchfy(hidden_states)
len_frame = hidden_states.shape[1]
if self.use_additional_conditions:
added_cond_kwargs = {
"resolution": torch.tensor([(height, width)] * bsz, device=hidden_states.device, dtype=hidden_states.dtype),
"nframe": torch.tensor([frame] * bsz, device=hidden_states.device, dtype=hidden_states.dtype),
"fps": fps
}
else:
added_cond_kwargs = {}
timestep, embedded_timestep = self.adaln_single(timestep, added_cond_kwargs=added_cond_kwargs)
encoder_hidden_states = self.caption_projection(self.caption_norm(encoder_hidden_states))
if encoder_hidden_states_2 is not None and hasattr(self, 'clip_projection'):
clip_embedding = self.clip_projection(encoder_hidden_states_2)
encoder_hidden_states = torch.cat([clip_embedding, encoder_hidden_states], dim=1)
hidden_states = rearrange(hidden_states, '(b f) l d-> b (f l) d', b=bsz, f=frame, l=len_frame).contiguous()
embedded_timestep = repeat(embedded_timestep, 'b d -> (b f) d', f=frame).contiguous()
shift, scale = (self.scale_shift_table[None] + embedded_timestep[:, None]).chunk(2, dim=1)
encoder_hidden_states, attn_mask = self.prepare_attn_mask(encoder_attention_mask,
encoder_hidden_states,
q_seqlen=frame * len_frame)
if self.enable_teacache:
hidden_states_ = hidden_states.clone()
normed_hidden_states = self.transformer_blocks[0].norm1(hidden_states_)
normed_hidden_states = rearrange(normed_hidden_states, 'b (f l) d -> (b f) l d', b=bsz, f=frame, l=len_frame)
modulated_inp = normed_hidden_states * (1 + scale) + shift
if self.cnt == 0 or self.cnt == self.num_steps - 1:
should_calc = True
self.accumulated_rel_l1_distance = 0
else:
coefficients = [6.74352814e+03, -2.22814115e+03, 2.55029094e+02, -1.12338285e+01, 2.84921593e-01]
rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance += rescale_func(
((modulated_inp - self.previous_modulated_input).abs().mean() /
self.previous_modulated_input.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
# print(f"accumulated_rel_l1_distance: {self.accumulated_rel_l1_distance}")
should_calc = False
else:
# print(f"accumulated_rel_l1_distance: {self.accumulated_rel_l1_distance}")
should_calc = True
self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = modulated_inp
self.cnt += 1
if self.cnt == self.num_steps:
self.cnt = 0
if self.enable_teacache:
if not should_calc:
# print(f"skip step {self.cnt}")
hidden_states += self.previous_residual
else:
# print(f"calc step {self.cnt}")
ori_hidden_states = hidden_states.clone()
hidden_states = self.block_forward(hidden_states,
encoder_hidden_states,
timestep=timestep,
rope_positions=[frame, height, width],
attn_mask=attn_mask,
parallel=self.parallel,
mask_strategy=mask_strategy)
self.previous_residual = hidden_states - ori_hidden_states
else:
# --------------------- Pass through DiT blocks ------------------------
hidden_states = self.block_forward(hidden_states,
encoder_hidden_states,
timestep=timestep,
rope_positions=[frame, height, width],
attn_mask=attn_mask,
parallel=self.parallel,
mask_strategy=mask_strategy)
# ---------------------------- Final layer ------------------------------
hidden_states = rearrange(hidden_states, 'b (f l) d -> (b f) l d', b=bsz, f=frame, l=len_frame)
hidden_states = self.norm_out(hidden_states)
# Modulation
hidden_states = hidden_states * (1 + scale) + shift
hidden_states = self.proj_out(hidden_states)
# unpatchify
hidden_states = hidden_states.reshape(shape=(-1, height, width, self.patch_size, self.patch_size,
self.out_channels))
hidden_states = rearrange(hidden_states, 'n h w p q c -> n c h p w q')
output = hidden_states.reshape(shape=(-1, self.out_channels, height * self.patch_size, width * self.patch_size))
output = rearrange(output, '(b f) c h w -> b f c h w', f=frame)
if return_dict:
return {'x': output}
return output
if __name__ == "__main__":
args = parse_args()
initialize_distributed()
main_print(f"sequence parallel size: {nccl_info.sp_size}")
device = torch.cuda.current_device()
setup_seed(args.seed)
main_print("Loading model, this might take a while...")
transformer = StepVideoModel.from_pretrained(os.path.join(args.model_dir, "transformer"),
torch_dtype=torch.bfloat16,
device_map=device)
if args.enable_teacache:
transformer.forward = types.MethodType(teacache_forward, transformer)
scheduler = FlowMatchDiscreteScheduler()
pipeline = StepVideoPipeline(transformer, scheduler, save_path=args.save_path)
pipeline.setup_api(
vae_url=args.vae_url,
caption_url=args.caption_url,
)
# TeaCache
pipeline.transformer.__class__.enable_teacache = True
pipeline.transformer.__class__.cnt = 0
pipeline.transformer.__class__.num_steps = args.infer_steps
pipeline.transformer.__class__.rel_l1_thresh = args.rel_l1_thresh # 0.1 for 1.6x speedup, 0.15 for 2.1x speedup
pipeline.transformer.__class__.accumulated_rel_l1_distance = 0
pipeline.transformer.__class__.previous_modulated_input = None
pipeline.transformer.__class__.previous_residual = None
with open(args.mask_strategy_file_path, 'r') as f:
mask_strategy = json.load(f)
if args.prompt.endswith('.txt'):
with open(args.prompt) as f:
prompts = [line.strip() for line in f.readlines()]
else:
prompts = [args.prompt]
for prompt in prompts:
main_print(f"Generating video for prompt: {prompt}")
videos = pipeline(prompt=prompt,
num_frames=args.num_frames,
height=args.height,
width=args.width,
num_inference_steps=args.infer_steps,
guidance_scale=args.cfg_scale,
time_shift=args.time_shift,
pos_magic=args.pos_magic,
neg_magic=args.neg_magic,
output_file_name=prompt[:150],
mask_strategy=mask_strategy)
dist.destroy_process_group()
+77 -173
View File
@@ -19,25 +19,19 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint,
save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing,
get_dit_fsdp_kwargs)
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_dit_fsdp_kwargs)
from fastvideo.utils.load import load_transformer
from fastvideo.utils.logging_ import main_print
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
get_sequence_parallel_state,
initialize_sequence_parallel_state
)
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
@@ -77,16 +71,11 @@ def compute_density_for_timestep_sampling(
return u
def get_sigmas(noise_scheduler,
device,
timesteps,
n_dim=4,
dtype=torch.float32):
def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32):
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
schedule_timesteps = noise_scheduler.timesteps.to(device)
timesteps = timesteps.to(device)
step_indices = [(schedule_timesteps == t).nonzero().item()
for t in timesteps]
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < n_dim:
@@ -132,8 +121,7 @@ def train_one_step(
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
timesteps = noise_scheduler.timesteps[indices].to(device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
broadcast(timesteps)
@@ -154,10 +142,7 @@ def train_one_step(
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
input_kwargs["guidance"] = torch.tensor([1000.0], device=noisy_model_input.device, dtype=torch.bfloat16)
model_pred = transformer(**input_kwargs)[0]
if precondition_outputs:
@@ -167,8 +152,7 @@ def train_one_step(
else:
target = noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
loss = (torch.mean((model_pred.float() - target.float())**2) / gradient_accumulation_steps)
loss.backward()
@@ -234,32 +218,23 @@ def main(args):
transformer.add_adapter(transformer_lora_config)
if args.resume_from_lora_checkpoint:
lora_state_dict = pipe.lora_state_dict(
args.resume_from_lora_checkpoint)
lora_state_dict = pipe.lora_state_dict(args.resume_from_lora_checkpoint)
transformer_state_dict = {
f'{k.replace("transformer.", "")}': v
for k, v in lora_state_dict.items() if k.startswith("transformer.")
}
transformer_state_dict = convert_unet_state_dict_to_peft(
transformer_state_dict)
incompatible_keys = set_peft_model_state_dict(transformer,
transformer_state_dict,
adapter_name="default")
transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict)
incompatible_keys = set_peft_model_state_dict(transformer, transformer_state_dict, adapter_name="default")
if incompatible_keys is not None:
# check only for unexpected keys
unexpected_keys = getattr(incompatible_keys, "unexpected_keys",
None)
unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None)
if unexpected_keys:
main_print(
f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
f" {unexpected_keys}. ")
main_print(f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
f" {unexpected_keys}. ")
main_print(
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
transformer,
args.fsdp_sharding_startegy,
@@ -271,14 +246,9 @@ def main(args):
if args.use_lora:
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = [
"to_k", "to_q", "to_v", "to_out.0"
]
transformer._no_split_modules = [
no_split_module.__name__ for no_split_module in no_split_modules
]
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
transformer)
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = [no_split_module.__name__ for no_split_module in no_split_modules]
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
@@ -287,8 +257,7 @@ def main(args):
main_print("--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
@@ -296,8 +265,7 @@ def main(args):
noise_scheduler = FlowMatchEulerDiscreteScheduler()
params_to_optimize = transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
@@ -309,8 +277,8 @@ def main(args):
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer)
transformer, optimizer, init_steps = resume_lora_optimizer(transformer, args.resume_from_lora_checkpoint,
optimizer)
main_print(f"optimizer: {optimizer}")
lr_scheduler = get_scheduler(
@@ -323,8 +291,7 @@ def main(args):
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
@@ -346,42 +313,33 @@ def main(args):
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
total_batch_size = (world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(
f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
assert NotImplementedError("resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
@@ -449,26 +407,18 @@ def main(args):
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step, pipe)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step, pipe)
else:
# Your existing checkpoint saving code
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(args,
transformer,
device,
torch.bfloat16,
step,
shift=args.shift)
log_validation(args, transformer, device, torch.bfloat16, step, shift=args.shift)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps, pipe)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps, pipe)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
@@ -476,13 +426,10 @@ def main(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_type",
type=str,
default="mochi",
help=
"The type of model to train. Currentlt support [mochi, hunyuan_hf, hunyuan]"
)
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train. Currentlt support [mochi, hunyuan_hf, hunyuan]")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
@@ -492,8 +439,7 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
@@ -501,10 +447,7 @@ if __name__ == "__main__":
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
@@ -541,16 +484,12 @@ if __name__ == "__main__":
parser.add_argument("--validation_steps", type=int, default=50)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
@@ -562,40 +501,31 @@ if __name__ == "__main__":
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
help=("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift",
type=float,
default=1.0,
help=("Set shift to 7 for hunyuan model."))
parser.add_argument("--shift", type=float, default=1.0, help=("Set shift to 7 for hunyuan model."))
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
help=("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
@@ -604,29 +534,25 @@ if __name__ == "__main__":
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
help="Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help=
"Initial learning rate (after the potential warmup period) to use.",
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
@@ -634,47 +560,36 @@ if __name__ == "__main__":
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
help="Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
@@ -688,14 +603,8 @@ if __name__ == "__main__":
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
parser.add_argument(
@@ -720,17 +629,15 @@ if __name__ == "__main__":
"--mode_scale",
type=float,
default=1.29,
help=
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
help=('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument(
"--lr_num_cycles",
@@ -744,10 +651,7 @@ if __name__ == "__main__":
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.01,
help="Weight decay to apply.")
parser.add_argument("--weight_decay", type=float, default=0.01, help="Weight decay to apply.")
parser.add_argument(
"--master_weight_type",
type=str,
+30 -57
View File
@@ -6,24 +6,16 @@ import torch
import torch.distributed.checkpoint as dist_cp
from peft import get_peft_model_state_dict
from safetensors.torch import load_file, save_file
from torch.distributed.checkpoint.default_planner import (DefaultLoadPlanner,
DefaultSavePlanner)
from torch.distributed.checkpoint.optimizer import \
load_sharded_optimizer_state_dict
from torch.distributed.fsdp import (FullOptimStateDictConfig,
FullStateDictConfig)
from torch.distributed.checkpoint.default_planner import DefaultLoadPlanner, DefaultSavePlanner
from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_dict
from torch.distributed.fsdp import FullOptimStateDictConfig, FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
from fastvideo.utils.logging_ import main_print
def save_checkpoint_optimizer(model,
optimizer,
rank,
output_dir,
step,
discriminator=False):
def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discriminator=False):
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
@@ -41,8 +33,7 @@ def save_checkpoint_optimizer(model,
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
if rank <= 0 and not discriminator:
weight_path = os.path.join(save_dir,
"diffusion_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(model.config)
config_dict.pop('dtype')
@@ -53,8 +44,7 @@ def save_checkpoint_optimizer(model,
optimizer_path = os.path.join(save_dir, "optimizer.pt")
torch.save(optim_state, optimizer_path)
else:
weight_path = os.path.join(save_dir,
"discriminator_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
torch.save(optim_state, optimizer_path)
@@ -74,8 +64,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
weight_path = os.path.join(save_dir,
"diffusion_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(transformer.config)
if "dtype" in config_dict:
@@ -115,8 +104,7 @@ def save_checkpoint_generator_discriminator(
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
weight_path = os.path.join(hf_weight_dir,
"diffusion_pytorch_model.safetensors")
weight_path = os.path.join(hf_weight_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
main_print(f"--> saved HF weight checkpoint at path {hf_weight_dir}")
@@ -140,8 +128,7 @@ def save_checkpoint_generator_discriminator(
planner=DefaultSavePlanner(),
)
discriminator_fsdp_state_dir = os.path.join(save_dir,
"discriminator_fsdp_state")
discriminator_fsdp_state_dir = os.path.join(save_dir, "discriminator_fsdp_state")
os.makedirs(discriminator_fsdp_state_dir, exist_ok=True)
with FSDP.state_dict_type(
discriminator,
@@ -149,13 +136,11 @@ def save_checkpoint_generator_discriminator(
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
optim_state = FSDP.optim_state_dict(discriminator,
discriminator_optimizer)
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
model_state = discriminator.state_dict()
state_dict = {"optimizer": optim_state, "model": model_state}
if rank <= 0:
discriminator_fsdp_state_fil = os.path.join(
discriminator_fsdp_state_dir, "discriminator_state.pt")
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
torch.save(state_dict, discriminator_fsdp_state_fil)
main_print("--> saved FSDP state checkpoint")
@@ -171,8 +156,7 @@ def load_sharded_model(model, optimizer, model_dir, optimizer_dir):
storage_reader=dist_cp.FileSystemReader(optimizer_dir),
)
optim_state = optim_state["optimizer"]
flattened_osd = FSDP.optim_state_dict_to_load(
model=model, optim=optimizer, optim_state_dict=optim_state)
flattened_osd = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=optim_state)
optimizer.load_state_dict(flattened_osd)
dist_cp.load_state_dict(
state_dict=weight_state_dict,
@@ -199,37 +183,30 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
else:
optim_state = None
model.load_state_dict(model_state)
discriminator_optim_state = FSDP.optim_state_dict_to_load(
model=model, optim=optimizer, optim_state_dict=optim_state)
discriminator_optim_state = FSDP.optim_state_dict_to_load(model=model,
optim=optimizer,
optim_state_dict=optim_state)
optimizer.load_state_dict(discriminator_optim_state)
main_print(
f"--> loaded discriminator and discriminator optimizer from path {checkpoint_file}"
)
main_print(f"--> loaded discriminator and discriminator optimizer from path {checkpoint_file}")
return model, optimizer
def resume_training_generator_discriminator(model, optimizer, discriminator,
discriminator_optimizer,
checkpoint_dir, rank):
def resume_training_generator_discriminator(model, optimizer, discriminator, discriminator_optimizer, checkpoint_dir,
rank):
step = int(checkpoint_dir.split("-")[-1])
model_weight_dir = os.path.join(checkpoint_dir, "model_weights_state")
model_optimizer_dir = os.path.join(checkpoint_dir, "model_optimizer_state")
model, optimizer = load_sharded_model(model, optimizer, model_weight_dir,
model_optimizer_dir)
discriminator_ckpt_file = os.path.join(checkpoint_dir,
"discriminator_fsdp_state",
"discriminator_state.pt")
discriminator, discriminator_optimizer = load_full_state_model(
discriminator, discriminator_optimizer, discriminator_ckpt_file, rank)
model, optimizer = load_sharded_model(model, optimizer, model_weight_dir, model_optimizer_dir)
discriminator_ckpt_file = os.path.join(checkpoint_dir, "discriminator_fsdp_state", "discriminator_state.pt")
discriminator, discriminator_optimizer = load_full_state_model(discriminator, discriminator_optimizer,
discriminator_ckpt_file, rank)
return model, optimizer, discriminator, discriminator_optimizer, step
def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
weight_path = os.path.join(checkpoint_dir,
"diffusion_pytorch_model.safetensors")
weight_path = os.path.join(checkpoint_dir, "diffusion_pytorch_model.safetensors")
if discriminator:
weight_path = os.path.join(checkpoint_dir,
"discriminator_pytorch_model.safetensors")
weight_path = os.path.join(checkpoint_dir, "discriminator_pytorch_model.safetensors")
model_weights = load_file(weight_path)
with FSDP.state_dict_type(
@@ -246,15 +223,13 @@ def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
else:
optim_path = os.path.join(checkpoint_dir, "optimizer.pt")
optimizer_state_dict = torch.load(optim_path, weights_only=False)
optim_state = FSDP.optim_state_dict_to_load(
model=model, optim=optimizer, optim_state_dict=optimizer_state_dict)
optim_state = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=optimizer_state_dict)
optimizer.load_state_dict(optim_state)
step = int(checkpoint_dir.split("-")[-1])
return model, optimizer, step
def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step,
pipeline):
def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipeline):
with FSDP.state_dict_type(
transformer,
StateDictType.FULL_STATE_DICT,
@@ -275,8 +250,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step,
torch.save(lora_optim_state, optim_path)
# save lora weight
main_print(f"--> saving LoRA checkpoint at step {step}")
transformer_lora_layers = get_peft_model_state_dict(
model=transformer, state_dict=full_state_dict)
transformer_lora_layers = get_peft_model_state_dict(model=transformer, state_dict=full_state_dict)
pipeline.save_lora_weights(
save_directory=save_dir,
transformer_lora_layers=transformer_lora_layers,
@@ -303,10 +277,9 @@ def resume_lora_optimizer(transformer, checkpoint_dir, optimizer):
config_dict = json.load(f)
optim_path = os.path.join(checkpoint_dir, "lora_optimizer.pt")
optimizer_state_dict = torch.load(optim_path, weights_only=False)
optim_state = FSDP.optim_state_dict_to_load(
model=transformer,
optim=optimizer,
optim_state_dict=optimizer_state_dict)
optim_state = FSDP.optim_state_dict_to_load(model=transformer,
optim=optimizer,
optim_state_dict=optimizer_state_dict)
optimizer.load_state_dict(optim_state)
step = config_dict["step"]
main_print(f"--> Successfully resuming LoRA optimizer from step {step}")
+25 -51
View File
@@ -17,10 +17,7 @@ def broadcast(input_: torch.Tensor):
dist.broadcast(input_, src=src, group=nccl_info.group)
def _all_to_all_4D(input: torch.tensor,
scatter_idx: int = 2,
gather_idx: int = 1,
group=None) -> torch.tensor:
def _all_to_all_4D(input: torch.tensor, scatter_idx: int = 2, gather_idx: int = 1, group=None) -> torch.tensor:
"""
all-to-all for QKV
@@ -33,9 +30,7 @@ def _all_to_all_4D(input: torch.tensor,
Returns:
torch.tensor: resharded tensor (bs, seqlen/P, hc, hs)
"""
assert (
input.dim() == 4
), f"input must be 4D tensor, got {input.dim()} and shape {input.shape}"
assert (input.dim() == 4), f"input must be 4D tensor, got {input.dim()} and shape {input.shape}"
seq_world_size = dist.get_world_size(group)
@@ -47,8 +42,7 @@ def _all_to_all_4D(input: torch.tensor,
# transpose groups of heads with the seq-len parallel dimension, so that we can scatter them!
# (bs, seqlen/P, hc, hs) -reshape-> (bs, seq_len/P, P, hc/P, hs) -transpose(0,2)-> (P, seq_len/P, bs, hc/P, hs)
input_t = (input.reshape(bs, shard_seqlen, seq_world_size, shard_hc,
hs).transpose(0, 2).contiguous())
input_t = (input.reshape(bs, shard_seqlen, seq_world_size, shard_hc, hs).transpose(0, 2).contiguous())
output = torch.empty_like(input_t)
# https://pytorch.org/docs/stable/distributed.html#torch.distributed.all_to_all_single
@@ -62,8 +56,7 @@ def _all_to_all_4D(input: torch.tensor,
output = output.reshape(seqlen, bs, shard_hc, hs)
# (seq_len, bs, hc/P, hs) -reshape-> (bs, seq_len, hc/P, hs)
output = output.transpose(0, 1).contiguous().reshape(
bs, seqlen, shard_hc, hs)
output = output.transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
return output
@@ -76,10 +69,11 @@ def _all_to_all_4D(input: torch.tensor,
# transpose groups of heads with the seq-len parallel dimension, so that we can scatter them!
# (bs, seqlen, hc/P, hs) -reshape-> (bs, P, seq_len/P, hc/P, hs) -transpose(0, 3)-> (hc/P, P, seqlen/P, bs, hs) -transpose(0, 1) -> (P, hc/P, seqlen/P, bs, hs)
input_t = (input.reshape(
bs, seq_world_size, shard_seqlen, shard_hc,
hs).transpose(0, 3).transpose(0, 1).contiguous().reshape(
seq_world_size, shard_hc, shard_seqlen, bs, hs))
input_t = (input.reshape(bs, seq_world_size, shard_seqlen, shard_hc,
hs).transpose(0,
3).transpose(0,
1).contiguous().reshape(seq_world_size, shard_hc,
shard_seqlen, bs, hs))
output = torch.empty_like(input_t)
# https://pytorch.org/docs/stable/distributed.html#torch.distributed.all_to_all_single
@@ -94,13 +88,11 @@ def _all_to_all_4D(input: torch.tensor,
output = output.reshape(hc, shard_seqlen, bs, hs)
# (hc, seqlen/N, bs, hs) -tranpose(0,2)-> (bs, seqlen/N, hc, hs)
output = output.transpose(0, 2).contiguous().reshape(
bs, shard_seqlen, hc, hs)
output = output.transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError(
"scatter_idx must be 1 or 2 and gather_idx must be 1 or 2")
raise RuntimeError("scatter_idx must be 1 or 2 and gather_idx must be 1 or 2")
class SeqAllToAll4D(torch.autograd.Function):
@@ -120,12 +112,10 @@ class SeqAllToAll4D(torch.autograd.Function):
return _all_to_all_4D(input, scatter_idx, gather_idx, group=group)
@staticmethod
def backward(ctx: Any,
*grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
def backward(ctx: Any, *grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
return (
None,
SeqAllToAll4D.apply(ctx.group, *grad_output, ctx.gather_idx,
ctx.scatter_idx),
SeqAllToAll4D.apply(ctx.group, *grad_output, ctx.gather_idx, ctx.scatter_idx),
None,
None,
)
@@ -136,8 +126,7 @@ def all_to_all_4D(
scatter_dim: int = 2,
gather_dim: int = 1,
):
return SeqAllToAll4D.apply(nccl_info.group, input_, scatter_dim,
gather_dim)
return SeqAllToAll4D.apply(nccl_info.group, input_, scatter_dim, gather_dim)
def _all_to_all(
@@ -147,10 +136,7 @@ def _all_to_all(
scatter_dim: int,
gather_dim: int,
):
input_list = [
t.contiguous()
for t in torch.tensor_split(input_, world_size, scatter_dim)
]
input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
dist.all_to_all(output_list, input_list, group=group)
return torch.cat(output_list, dim=gather_dim).contiguous()
@@ -172,8 +158,7 @@ class _AllToAll(torch.autograd.Function):
ctx.scatter_dim = scatter_dim
ctx.gather_dim = gather_dim
ctx.world_size = dist.get_world_size(process_group)
output = _all_to_all(input_, ctx.world_size, process_group,
scatter_dim, gather_dim)
output = _all_to_all(input_, ctx.world_size, process_group, scatter_dim, gather_dim)
return output
@staticmethod
@@ -253,8 +238,7 @@ def all_gather(input_: torch.Tensor, dim: int = 1):
return _AllGather.apply(input_, dim)
def prepare_sequence_parallel_data(hidden_states, encoder_hidden_states,
attention_mask, encoder_attention_mask):
def prepare_sequence_parallel_data(hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask):
if nccl_info.sp_size == 1:
return (
hidden_states,
@@ -263,18 +247,11 @@ def prepare_sequence_parallel_data(hidden_states, encoder_hidden_states,
encoder_attention_mask,
)
def prepare(hidden_states, encoder_hidden_states, attention_mask,
encoder_attention_mask):
def prepare(hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask):
hidden_states = all_to_all(hidden_states, scatter_dim=2, gather_dim=0)
encoder_hidden_states = all_to_all(encoder_hidden_states,
scatter_dim=1,
gather_dim=0)
attention_mask = all_to_all(attention_mask,
scatter_dim=1,
gather_dim=0)
encoder_attention_mask = all_to_all(encoder_attention_mask,
scatter_dim=1,
gather_dim=0)
encoder_hidden_states = all_to_all(encoder_hidden_states, scatter_dim=1, gather_dim=0)
attention_mask = all_to_all(attention_mask, scatter_dim=1, gather_dim=0)
encoder_attention_mask = all_to_all(encoder_attention_mask, scatter_dim=1, gather_dim=0)
return (
hidden_states,
encoder_hidden_states,
@@ -301,8 +278,7 @@ def prepare_sequence_parallel_data(hidden_states, encoder_hidden_states,
return hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
def sp_parallel_dataloader_wrapper(dataloader, device, train_batch_size,
sp_size, train_sp_batch_size):
def sp_parallel_dataloader_wrapper(dataloader, device, train_batch_size, sp_size, train_sp_batch_size):
while True:
for data_item in dataloader:
latents, cond, attn_mask, cond_mask = data_item
@@ -316,11 +292,9 @@ def sp_parallel_dataloader_wrapper(dataloader, device, train_batch_size,
else:
latents, cond, attn_mask, cond_mask = prepare_sequence_parallel_data(
latents, cond, attn_mask, cond_mask)
assert (
train_batch_size * sp_size >= train_sp_batch_size
), "train_batch_size * sp_size should be greater than train_sp_batch_size"
for iter in range(train_batch_size * sp_size //
train_sp_batch_size):
assert (train_batch_size * sp_size >=
train_sp_batch_size), "train_batch_size * sp_size should be greater than train_sp_batch_size"
for iter in range(train_batch_size * sp_size // train_sp_batch_size):
st_idx = iter * train_sp_batch_size
ed_idx = (iter + 1) * train_sp_batch_size
encoder_hidden_states = cond[st_idx:ed_idx]
+18 -51
View File
@@ -30,9 +30,7 @@ class DecordInit(object):
results (dict): The resulting dict to be modified and passed
to the next transform in pipeline.
"""
reader = decord.VideoReader(filename,
ctx=self.ctx,
num_threads=self.num_threads)
reader = decord.VideoReader(filename, ctx=self.ctx, num_threads=self.num_threads)
return reader
def __repr__(self):
@@ -94,8 +92,7 @@ class Collate:
self.max_thw,
self.ae_stride_thw,
)
assert not torch.any(
torch.isnan(pad_batch_tubes)), "after pad_batch_tubes"
assert not torch.any(torch.isnan(pad_batch_tubes)), "after pad_batch_tubes"
return pad_batch_tubes, attention_mask, input_ids, cond_mask
def process(
@@ -109,25 +106,18 @@ class Collate:
ae_stride_thw,
):
# pad to max multiple of ds_stride
batch_input_size = [i.shape
for i in batch_tubes] # [(c t h w), (c t h w)]
batch_input_size = [i.shape for i in batch_tubes] # [(c t h w), (c t h w)]
assert len(batch_input_size) == self.batch_size
if self.group_frame or self.group_resolution or self.batch_size == 1: #
len_each_batch = batch_input_size
idx_length_dict = dict(
[*zip(list(range(self.batch_size)), len_each_batch)])
idx_length_dict = dict([*zip(list(range(self.batch_size)), len_each_batch)])
count_dict = Counter(len_each_batch)
if len(count_dict) != 1:
sorted_by_value = sorted(count_dict.items(),
key=lambda item: item[1])
sorted_by_value = sorted(count_dict.items(), key=lambda item: item[1])
pick_length = sorted_by_value[-1][0] # the highest frequency
candidate_batch = [
idx for idx, length in idx_length_dict.items()
if length == pick_length
]
candidate_batch = [idx for idx, length in idx_length_dict.items() if length == pick_length]
random_select_batch = [
random.choice(candidate_batch)
for _ in range(len(len_each_batch) - len(candidate_batch))
random.choice(candidate_batch) for _ in range(len(len_each_batch) - len(candidate_batch))
]
print(
batch_input_size,
@@ -141,8 +131,7 @@ class Collate:
pick_idx = candidate_batch + random_select_batch
batch_tubes = [batch_tubes[i] for i in pick_idx]
batch_input_size = [i.shape for i in batch_tubes
] # [(c t h w), (c t h w)]
batch_input_size = [i.shape for i in batch_tubes] # [(c t h w), (c t h w)]
input_ids = [input_ids[i] for i in pick_idx] # b [1, l]
cond_mask = [cond_mask[i] for i in pick_idx] # b [1, l]
@@ -159,10 +148,7 @@ class Collate:
pad_to_multiple(max_w, ds_stride),
)
pad_max_t = pad_max_t + 1 - self.ae_stride_t
each_pad_t_h_w = [[
pad_max_t - i.shape[1], pad_max_h - i.shape[2],
pad_max_w - i.shape[3]
] for i in batch_tubes]
each_pad_t_h_w = [[pad_max_t - i.shape[1], pad_max_h - i.shape[2], pad_max_w - i.shape[3]] for i in batch_tubes]
pad_batch_tubes = [
F.pad(im, (0, pad_w, 0, pad_h, 0, pad_t), value=0)
for (pad_t, pad_h, pad_w), im in zip(each_pad_t_h_w, batch_tubes)
@@ -229,10 +215,7 @@ def split_to_even_chunks(indices, lengths, num_chunks, batch_size):
if batch_size != len(chunk):
assert batch_size > len(chunk)
if len(chunk) != 0:
chunk = chunk + [
random.choice(chunk)
for _ in range(batch_size - len(chunk))
]
chunk = chunk + [random.choice(chunk) for _ in range(batch_size - len(chunk))]
else:
chunk = random.choice(pad_chunks)
print(chunks[idx], "->", chunk)
@@ -256,16 +239,11 @@ def megabatch_frame_alignment(megabatches, lengths):
# mixed frame length, align megabatch inside
if len(count_dict) != 1:
sorted_by_value = sorted(count_dict.items(),
key=lambda item: item[1])
sorted_by_value = sorted(count_dict.items(), key=lambda item: item[1])
pick_length = sorted_by_value[-1][0] # the highest frequency
candidate_batch = [
idx for idx, length in idx_length_dict.items()
if length == pick_length
]
candidate_batch = [idx for idx, length in idx_length_dict.items() if length == pick_length]
random_select_batch = [
random.choice(candidate_batch)
for i in range(len(idx_length_dict) - len(candidate_batch))
random.choice(candidate_batch) for i in range(len(idx_length_dict) - len(candidate_batch))
]
aligned_magabatch = candidate_batch + random_select_batch
aligned_magabatches.append(aligned_magabatch)
@@ -287,8 +265,7 @@ def get_length_grouped_indices(
):
# We need to use torch for the random part as a distributed sampler will set the random seed for torch.
if generator is None:
generator = torch.Generator().manual_seed(
seed) # every rank will generate a fixed order but random index
generator = torch.Generator().manual_seed(seed) # every rank will generate a fixed order but random index
indices = torch.randperm(len(lengths), generator=generator).tolist()
@@ -297,29 +274,20 @@ def get_length_grouped_indices(
# chunk dataset to megabatches
megabatch_size = world_size * batch_size
megabatches = [
indices[i:i + megabatch_size]
for i in range(0, len(lengths), megabatch_size)
]
megabatches = [indices[i:i + megabatch_size] for i in range(0, len(lengths), megabatch_size)]
# make sure the length in each magabatch is align with each other
megabatches = megabatch_frame_alignment(megabatches, lengths)
# aplit aligned megabatch into batches
megabatches = [
split_to_even_chunks(megabatch, lengths, world_size, batch_size)
for megabatch in megabatches
]
megabatches = [split_to_even_chunks(megabatch, lengths, world_size, batch_size) for megabatch in megabatches]
# random megabatches to do video-image mix training
indices = torch.randperm(len(megabatches), generator=generator).tolist()
shuffled_megabatches = [megabatches[i] for i in indices]
# expand indices and return
return [
i for megabatch in shuffled_megabatches for batch in megabatch
for i in batch
]
return [i for megabatch in shuffled_megabatches for batch in megabatch for i in batch]
class LengthGroupedSampler(Sampler):
@@ -370,6 +338,5 @@ class LengthGroupedSampler(Sampler):
index += batch_size * world_size
return result
indices = distributed_sampler(indices, self.rank, self.batch_size,
self.world_size)
indices = distributed_sampler(indices, self.rank, self.batch_size, self.world_size)
return iter(indices)
+1 -3
View File
@@ -35,6 +35,4 @@ if __name__ == "__main__":
except Exception:
pass
print("\n" +
"\n".join([f"- {key}: {value}"
for key, value in info.items()]) + "\n")
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
+3 -4
View File
@@ -4,8 +4,8 @@ from functools import partial
import torch
from peft.utils.other import fsdp_auto_wrap_policy
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
CheckpointImpl, apply_activation_checkpointing, checkpoint_wrapper)
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (CheckpointImpl, apply_activation_checkpointing,
checkpoint_wrapper)
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
@@ -93,8 +93,7 @@ def get_dit_fsdp_kwargs(
sharding_strategy = ShardingStrategy._HYBRID_SHARD_ZERO2
device_id = torch.cuda.current_device()
cpu_offload = (torch.distributed.fsdp.CPUOffload(
offload_params=True) if cpu_offload else None)
cpu_offload = (torch.distributed.fsdp.CPUOffload(offload_params=True) if cpu_offload else None)
fsdp_kwargs = {
"auto_wrap_policy": auto_wrap_policy,
"mixed_precision": mixed_precision,
+41 -79
View File
@@ -7,16 +7,13 @@ from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
from torch import nn
from transformers import AutoTokenizer, T5EncoderModel
from fastvideo.models.hunyuan.modules.models import (
HYVideoDiffusionTransformer, MMDoubleStreamBlock, MMSingleStreamBlock)
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
MMSingleStreamBlock)
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import \
AutoencoderKLCausal3D
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (
HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformer3DModel,
HunyuanVideoTransformerBlock)
from fastvideo.models.mochi_hf.modeling_mochi import (MochiTransformer3DModel,
MochiTransformerBlock)
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (HunyuanVideoSingleTransformerBlock,
HunyuanVideoTransformer3DModel, HunyuanVideoTransformerBlock)
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
from fastvideo.utils.logging_ import main_print
hunyuan_config = {
@@ -62,8 +59,7 @@ class HunyuanTextEncoderWrapper(nn.Module):
super().__init__()
text_len = 256
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get(
"crop_start", 0)
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get("crop_start", 0)
max_length = text_len + crop_start
@@ -72,8 +68,7 @@ class HunyuanTextEncoderWrapper(nn.Module):
# prompt_template_video
prompt_template_video = PROMPT_TEMPLATE["dit-llm-encode-video"]
text_encoder_path = os.path.join(pretrained_model_name_or_path,
"text_encoder")
text_encoder_path = os.path.join(pretrained_model_name_or_path, "text_encoder")
self.text_encoder = TextEncoder(
text_encoder_type="llm",
text_encoder_path=text_encoder_path,
@@ -88,8 +83,7 @@ class HunyuanTextEncoderWrapper(nn.Module):
logger=None,
device=device,
)
text_encoder_path_2 = os.path.join(pretrained_model_name_or_path,
"text_encoder_2")
text_encoder_path_2 = os.path.join(pretrained_model_name_or_path, "text_encoder_2")
self.text_encoder_2 = TextEncoder(
text_encoder_type="clipL",
text_encoder_path=text_encoder_path_2,
@@ -110,9 +104,7 @@ class HunyuanTextEncoderWrapper(nn.Module):
text_inputs = text_encoder.text2tokens(prompt, data_type=data_type)
if clip_skip is None:
prompt_outputs = text_encoder.encode(text_inputs,
data_type="video",
device=device)
prompt_outputs = text_encoder.encode(text_inputs, data_type="video", device=device)
prompt_embeds = prompt_outputs.hidden_state
else:
prompt_outputs = text_encoder.encode(
@@ -123,16 +115,14 @@ class HunyuanTextEncoderWrapper(nn.Module):
)
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip + 1)]
prompt_embeds = text_encoder.model.text_model.final_layer_norm(
prompt_embeds)
prompt_embeds = text_encoder.model.text_model.final_layer_norm(prompt_embeds)
attention_mask = prompt_outputs.attention_mask
if attention_mask is not None:
attention_mask = attention_mask.to(device)
bs_embed, seq_len = attention_mask.shape
attention_mask = attention_mask.repeat(1, num_videos_per_prompt)
attention_mask = attention_mask.view(
bs_embed * num_videos_per_prompt, seq_len)
attention_mask = attention_mask.view(bs_embed * num_videos_per_prompt, seq_len)
if text_encoder is not None:
prompt_embeds_dtype = text_encoder.dtype
@@ -141,27 +131,23 @@ class HunyuanTextEncoderWrapper(nn.Module):
else:
prompt_embeds_dtype = prompt_embeds.dtype
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype,
device=device)
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
if prompt_embeds.ndim == 2:
bs_embed, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, -1)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
else:
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, seq_len, -1)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, seq_len, -1)
return (prompt_embeds, attention_mask)
def encode_prompt(self, prompt):
prompt_embeds, attention_mask = self.encode_(prompt, self.text_encoder)
prompt_embeds_2, attention_mask_2 = self.encode_(
prompt, self.text_encoder_2)
prompt_embeds_2, attention_mask_2 = self.encode_(prompt, self.text_encoder_2)
prompt_embeds_2 = F.pad(
prompt_embeds_2,
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
@@ -175,11 +161,9 @@ class MochiTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
self.text_encoder = T5EncoderModel.from_pretrained(
os.path.join(pretrained_model_name_or_path,
"text_encoder")).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(
os.path.join(pretrained_model_name_or_path, "tokenizer"))
self.text_encoder = T5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path,
"text_encoder")).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "tokenizer"))
self.max_sequence_length = 256
def encode_prompt(self, prompt):
@@ -201,19 +185,12 @@ class MochiTextEncoderWrapper(nn.Module):
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
untruncated_ids = self.tokenizer(prompt,
padding="longest",
return_tensors="pt").input_ids
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[
-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(
untruncated_ids[:, self.max_sequence_length - 1:-1])
main_print(
f"Truncated text input: {prompt} to: {removed_text} for model input."
)
prompt_embeds = self.text_encoder(
text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.max_sequence_length - 1:-1])
main_print(f"Truncated text input: {prompt} to: {removed_text} for model input.")
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
@@ -229,20 +206,16 @@ def load_hunyuan_state_dict(model, dit_model_name_or_path):
model_path = dit_model_name_or_path
bare_model = "unknown"
state_dict = torch.load(model_path,
map_location=lambda storage, loc: storage,
weights_only=True)
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage, weights_only=True)
if bare_model == "unknown" and ("ema" in state_dict
or "module" in state_dict):
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
bare_model = False
if bare_model is False:
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}.")
raise KeyError(f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}.")
model.load_state_dict(state_dict, strict=True)
return model
@@ -288,8 +261,7 @@ def load_transformer(
**hunyuan_config,
dtype=master_weight_type,
)
transformer = load_hunyuan_state_dict(transformer,
dit_model_name_or_path)
transformer = load_hunyuan_state_dict(transformer, dit_model_name_or_path)
if master_weight_type == torch.bfloat16:
transformer = transformer.bfloat16()
else:
@@ -300,23 +272,20 @@ def load_transformer(
def load_vae(model_type, pretrained_model_name_or_path):
weight_dtype = torch.float32
if model_type == "mochi":
vae = AutoencoderKLMochi.from_pretrained(
pretrained_model_name_or_path,
subfolder="vae",
torch_dtype=weight_dtype).to("cuda")
vae = AutoencoderKLMochi.from_pretrained(pretrained_model_name_or_path,
subfolder="vae",
torch_dtype=weight_dtype).to("cuda")
autocast_type = torch.bfloat16
fps = 30
elif model_type == "hunyuan_hf":
vae = AutoencoderKLHunyuanVideo.from_pretrained(
pretrained_model_name_or_path,
subfolder="vae",
torch_dtype=weight_dtype).to("cuda")
vae = AutoencoderKLHunyuanVideo.from_pretrained(pretrained_model_name_or_path,
subfolder="vae",
torch_dtype=weight_dtype).to("cuda")
autocast_type = torch.bfloat16
fps = 24
elif model_type == "hunyuan":
vae_precision = torch.float32
vae_path = os.path.join(pretrained_model_name_or_path,
"hunyuan-video-t2v-720p/vae")
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
config = AutoencoderKLCausal3D.load_config(vae_path)
vae = AutoencoderKLCausal3D.from_config(config)
@@ -328,10 +297,7 @@ def load_vae(model_type, pretrained_model_name_or_path):
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if any(k.startswith("vae.") for k in ckpt.keys()):
ckpt = {
k.replace("vae.", ""): v
for k, v in ckpt.items() if k.startswith("vae.")
}
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
vae.load_state_dict(ckpt)
vae = vae.to(dtype=vae_precision)
vae.requires_grad_(False)
@@ -344,11 +310,9 @@ def load_vae(model_type, pretrained_model_name_or_path):
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
if model_type == "mochi":
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path,
device)
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "hunyuan" or "hunyuan_hf":
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path,
device)
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
else:
raise ValueError(f"Unsupported model type: {model_type}")
return text_encoder
@@ -359,8 +323,7 @@ def get_no_split_modules(transformer):
if isinstance(transformer, MochiTransformer3DModel):
return (MochiTransformerBlock, )
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
return (HunyuanVideoSingleTransformerBlock,
HunyuanVideoTransformerBlock)
return (HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformerBlock)
elif isinstance(transformer, HYVideoDiffusionTransformer):
return (MMDoubleStreamBlock, MMSingleStreamBlock)
else:
@@ -371,7 +334,6 @@ if __name__ == "__main__":
# test encode prompt
device = torch.cuda.current_device()
pretrained_model_name_or_path = "data/hunyuan"
text_encoder = load_text_encoder("hunyuan", pretrained_model_name_or_path,
device)
text_encoder = load_text_encoder("hunyuan", pretrained_model_name_or_path, device)
prompt = "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."
prompt_embeds, attention_mask = text_encoder.encode_prompt(prompt)
+7 -15
View File
@@ -13,23 +13,18 @@ def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
)
args.optimizer = "adamw"
if args.use_8bit_adam and not (args.optimizer.lower()
not in ["adam", "adamw"]):
logger.warning(
f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was "
f"set to {args.optimizer.lower()}")
if args.use_8bit_adam and not (args.optimizer.lower() not in ["adam", "adamw"]):
logger.warning(f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was "
f"set to {args.optimizer.lower()}")
if args.use_8bit_adam:
try:
import bitsandbytes as bnb
except ImportError:
raise ImportError(
"To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`."
)
raise ImportError("To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`.")
if args.optimizer.lower() == "adamw":
optimizer_class = (bnb.optim.AdamW8bit
if args.use_8bit_adam else torch.optim.AdamW)
optimizer_class = (bnb.optim.AdamW8bit if args.use_8bit_adam else torch.optim.AdamW)
optimizer = optimizer_class(
params_to_optimize,
@@ -50,16 +45,13 @@ def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
try:
import prodigyopt
except ImportError:
raise ImportError(
"To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`"
)
raise ImportError("To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`")
optimizer_class = prodigyopt.Prodigy
if args.learning_rate <= 0.1:
logger.warning(
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0"
)
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0")
optimizer = optimizer_class(
params_to_optimize,
+1 -2
View File
@@ -50,8 +50,7 @@ def initialize_sequence_parallel_group(sequence_parallel_size):
nccl_info.global_rank = rank
num_sequence_parallel_groups: int = world_size // sequence_parallel_size
for i in range(num_sequence_parallel_groups):
ranks = range(i * sequence_parallel_size,
(i + 1) * sequence_parallel_size)
ranks = range(i * sequence_parallel_size, (i + 1) * sequence_parallel_size)
group = dist.new_group(ranks)
if rank in ranks:
nccl_info.group = group
+34 -71
View File
@@ -1,3 +1,4 @@
# isort: skip_file
import gc
import os
from typing import List, Optional, Union
@@ -13,12 +14,10 @@ from tqdm import tqdm
import wandb
from fastvideo.distill.solver import PCMFMScheduler
from fastvideo.models.mochi_hf.pipeline_mochi import (
linear_quadratic_schedule, retrieve_timesteps)
from fastvideo.models.mochi_hf.pipeline_mochi import (linear_quadratic_schedule, retrieve_timesteps)
from fastvideo.utils.communications import all_gather
from fastvideo.utils.load import load_vae
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
nccl_info)
from fastvideo.utils.parallel_states import (get_sequence_parallel_state, nccl_info)
def prepare_latents(
@@ -39,10 +38,7 @@ def prepare_latents(
shape = (batch_size, num_channels_latents, num_frames, height, width)
latents = randn_tensor(shape,
generator=generator,
device=device,
dtype=dtype)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
@@ -75,10 +71,8 @@ def sample_validation_video(
do_classifier_free_guidance = guidance_scale > 1.0
if do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds],
dim=0)
prompt_attention_mask = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask], dim=0)
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
# 4. Prepare latent variables
# TODO: Remove hardcore
@@ -96,9 +90,7 @@ def sample_validation_video(
)
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()
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
# 5. Prepare timestep
@@ -120,8 +112,7 @@ def sample_validation_video(
num_inference_steps,
device,
)
num_warmup_steps = max(
len(timesteps) - num_inference_steps * scheduler.order, 0)
num_warmup_steps = max(len(timesteps) - num_inference_steps * scheduler.order, 0)
# 6. Denoising loop
# with self.progress_bar(total=num_inference_steps) as progress_bar:
@@ -134,8 +125,7 @@ def sample_validation_video(
desc="Validation sampling...",
) as progress_bar:
for i, t in enumerate(timesteps):
latent_model_input = (torch.cat([latents] * 2)
if do_classifier_free_guidance else latents)
latent_model_input = (torch.cat([latents] * 2) if do_classifier_free_guidance else latents)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
with torch.autocast("cuda", dtype=torch.bfloat16):
@@ -151,15 +141,11 @@ def sample_validation_video(
noise_pred = noise_pred.to(torch.float32)
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (
noise_pred_text - noise_pred_uncond)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
# compute the previous noisy sample x_t -> x_t-1
latents_dtype = latents.dtype
latents = scheduler.step(noise_pred,
t,
latents.to(torch.float32),
return_dict=False)[0]
latents = scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
latents = latents.to(latents_dtype)
if latents.dtype != latents_dtype:
@@ -167,8 +153,7 @@ def sample_validation_video(
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
latents = latents.to(latents_dtype)
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and
(i + 1) % scheduler.order == 0):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % scheduler.order == 0):
progress_bar.update()
if get_sequence_parallel_state():
@@ -179,24 +164,19 @@ def sample_validation_video(
else:
# unscale/denormalize the latents
# denormalize with the mean and std if available and not None
has_latents_mean = (hasattr(vae.config, "latents_mean")
and vae.config.latents_mean is not None)
has_latents_std = (hasattr(vae.config, "latents_std")
and vae.config.latents_std is not None)
has_latents_mean = (hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None)
has_latents_std = (hasattr(vae.config, "latents_std") and vae.config.latents_std is not None)
if has_latents_mean and has_latents_std:
latents_mean = (torch.tensor(vae.config.latents_mean).view(
1, 12, 1, 1, 1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(vae.config.latents_std).view(
1, 12, 1, 1, 1).to(latents.device, latents.dtype))
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1,
1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype))
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
else:
latents = latents / vae.config.scaling_factor
with torch.autocast("cuda", dtype=vae.dtype):
video = vae.decode(latents, return_dict=False)[0]
video_processor = VideoProcessor(
vae_scale_factor=vae_spatial_scale_factor)
video = video_processor.postprocess_video(video,
output_type=output_type)
video_processor = VideoProcessor(vae_scale_factor=vae_spatial_scale_factor)
video = video_processor.postprocess_video(video, output_type=output_type)
return (video, )
@@ -228,8 +208,7 @@ def log_validation(
num_channels_latents = 16
else:
raise ValueError(f"Model type {args.model_type} not supported")
vae, autocast_type, fps = load_vae(args.model_type,
args.pretrained_model_name_or_path)
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
vae.enable_tiling()
if scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
@@ -246,9 +225,7 @@ def log_validation(
# args.validation_prompt_dir
validation_guidance_scale_ls = args.validation_guidance_scale.split(",")
validation_guidance_scale_ls = [
float(scale) for scale in validation_guidance_scale_ls
]
validation_guidance_scale_ls = [float(scale) for scale in validation_guidance_scale_ls]
for validation_sampling_step in args.validation_sampling_steps.split(","):
validation_sampling_step = int(validation_sampling_step)
for validation_guidance_scale in validation_guidance_scale_ls:
@@ -256,37 +233,29 @@ def log_validation(
# prompt_embed are named embed0 to embedN
# check how many embeds are there
embe_dir = os.path.join(args.validation_prompt_dir, "prompt_embed")
mask_dir = os.path.join(args.validation_prompt_dir,
"prompt_attention_mask")
mask_dir = os.path.join(args.validation_prompt_dir, "prompt_attention_mask")
embeds = sorted([f for f in os.listdir(embe_dir)])
masks = sorted([f for f in os.listdir(mask_dir)])
num_embeds = len(embeds)
validation_prompt_ids = list(range(num_embeds))
num_sp_groups = int(os.getenv("WORLD_SIZE",
"1")) // nccl_info.sp_size
num_sp_groups = int(os.getenv("WORLD_SIZE", "1")) // nccl_info.sp_size
# pad to multiple of groups
if num_embeds % num_sp_groups != 0:
validation_prompt_ids += [0] * (num_sp_groups -
num_embeds % num_sp_groups)
validation_prompt_ids += [0] * (num_sp_groups - num_embeds % num_sp_groups)
num_embeds_per_group = len(validation_prompt_ids) // num_sp_groups
local_prompt_ids = validation_prompt_ids[nccl_info.group_id *
num_embeds_per_group:
(nccl_info.group_id + 1) *
num_embeds_per_group:(nccl_info.group_id + 1) *
num_embeds_per_group]
for i in local_prompt_ids:
prompt_embed_path = os.path.join(embe_dir, f"{embeds[i]}")
prompt_mask_path = os.path.join(mask_dir, f"{masks[i]}")
prompt_embeds = (torch.load(
prompt_embed_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
prompt_attention_mask = (torch.load(
prompt_mask_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
negative_prompt_embeds = torch.zeros(
256, 4096).to(device).unsqueeze(0)
negative_prompt_attention_mask = (
torch.zeros(256).bool().to(device).unsqueeze(0))
prompt_embeds = (torch.load(prompt_embed_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
prompt_attention_mask = (torch.load(prompt_mask_path, map_location="cpu",
weights_only=True).to(device).unsqueeze(0))
negative_prompt_embeds = torch.zeros(256, 4096).to(device).unsqueeze(0)
negative_prompt_attention_mask = (torch.zeros(256).bool().to(device).unsqueeze(0))
generator = torch.Generator(device="cpu").manual_seed(12345)
video = sample_validation_video(
args.model_type,
@@ -303,8 +272,7 @@ def log_validation(
prompt_embeds=prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_embeds=negative_prompt_embeds,
negative_prompt_attention_mask=
negative_prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
vae_spatial_scale_factor=vae_spatial_scale_factor,
vae_temporal_scale_factor=vae_temporal_scale_factor,
num_channels_latents=num_channels_latents,
@@ -317,9 +285,7 @@ def log_validation(
torch.cuda.empty_cache()
# log if main process
torch.distributed.barrier()
all_videos = [
None for i in range(int(os.getenv("WORLD_SIZE", "1")))
] # remove padded videos
all_videos = [None for i in range(int(os.getenv("WORLD_SIZE", "1")))] # remove padded videos
torch.distributed.all_gather_object(all_videos, videos)
if nccl_info.global_rank == 0:
# remove padding
@@ -337,9 +303,6 @@ def log_validation(
logs = {
f"{'ema_' if ema else ''}validation_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}":
[
wandb.Video(filename)
for i, filename in enumerate(video_filenames)
]
[wandb.Video(filename) for i, filename in enumerate(video_filenames)]
}
wandb.log(logs, step=global_step)

Some files were not shown because too many files have changed in this diff Show More