Compare commits
76
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0d12c41fc8 | ||
|
|
9aebc4ada1 | ||
|
|
b53cf7425c | ||
|
|
d9ce056901 | ||
|
|
218449c54d | ||
|
|
221958bcde | ||
|
|
4a1f1e35bb | ||
|
|
e0e05f97f2 | ||
|
|
dd75ee8509 | ||
|
|
0aed1868df | ||
|
|
d467c7cd35 | ||
|
|
88b2583c2c | ||
|
|
a730e43d5f | ||
|
|
edf116fa46 | ||
|
|
de3cefb5e5 | ||
|
|
e087e85e09 | ||
|
|
e1b998b6ef | ||
|
|
fb49c93dbc | ||
|
|
172f4802b4 | ||
|
|
24e57fafc9 | ||
|
|
6debd46482 | ||
|
|
f7dc36f7ea | ||
|
|
a0fb954f56 | ||
|
|
053106922c | ||
|
|
a57122c519 | ||
|
|
b393570e45 | ||
|
|
285635e8c0 | ||
|
|
58cfd71b5e | ||
|
|
3bf892b6ab | ||
|
|
85639d1101 | ||
|
|
6ab2263f3a | ||
|
|
b421c2e183 | ||
|
|
de1e8d868e | ||
|
|
98b92be25e | ||
|
|
8d41d505fe | ||
|
|
8cfdf58a17 | ||
|
|
cf15594055 | ||
|
|
ce95c2df29 | ||
|
|
d417e4c7c4 | ||
|
|
2a70d05b4f | ||
|
|
03187fd83a | ||
|
|
8a128ad815 | ||
|
|
13f665e455 | ||
|
|
94ba0ab6ea | ||
|
|
5a5d0ef1a0 | ||
|
|
45e4adca4d | ||
|
|
7106eadffc | ||
|
|
2f3a8661bf | ||
|
|
6d0082c1a9 | ||
|
|
3d9189571a | ||
|
|
b042e321a1 | ||
|
|
44bda9f8a3 | ||
|
|
035ba5f5cf | ||
|
|
52ba538e8e | ||
|
|
f0bc297260 | ||
|
|
7413b1dd5f | ||
|
|
8ec82cbc37 | ||
|
|
54e74aec3f | ||
|
|
c1c276b616 | ||
|
|
8e5fa4d383 | ||
|
|
06ff1c912f | ||
|
|
857f5df51b | ||
|
|
18c5ca131d | ||
|
|
c16625242e | ||
|
|
fcc45701c9 | ||
|
|
6c87e003aa | ||
|
|
23181aac96 | ||
|
|
26b65a8baf | ||
|
|
25498a9d85 | ||
|
|
afb3ac9fc1 | ||
|
|
3d385600e9 | ||
|
|
5680039dbe | ||
|
|
2c1eb3b0e2 | ||
|
|
5e2e3ab06a | ||
|
|
f5ca624aff | ||
|
|
d1ea86d351 |
@@ -0,0 +1,29 @@
|
||||
name: 🐞 Bug report
|
||||
description: Create a report to help us reproduce and fix the bug
|
||||
title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
description: A clear and concise description of what the bug is.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Reproduction
|
||||
description: |
|
||||
What command or script did you run? Which **model** are you using?
|
||||
placeholder: |
|
||||
A placeholder for the command.
|
||||
validations:
|
||||
required: true
|
||||
@@ -0,0 +1,17 @@
|
||||
name: 🚀 Feature request
|
||||
description: Suggest an idea for this project
|
||||
title: "[Feature] "
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Motivation
|
||||
description: |
|
||||
A clear and concise description of the motivation of the feature.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Related resources
|
||||
description: |
|
||||
If there is an official code release or third-party implementations, please also provide the information here, which would be very helpful.
|
||||
@@ -0,0 +1 @@
|
||||
blank_issues_enabled: false
|
||||
@@ -0,0 +1,45 @@
|
||||
name: codespell
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- "**/*.md"
|
||||
- "**/*.rst"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/codespell.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- "**/*.md"
|
||||
- "**/*.rst"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/codespell.yml
|
||||
|
||||
jobs:
|
||||
codespell:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-lint.txt
|
||||
- name: Spelling check with codespell
|
||||
run: |
|
||||
# Refer to the above environment variable here
|
||||
codespell --toml pyproject.toml $CODESPELL_EXCLUDES
|
||||
@@ -0,0 +1,50 @@
|
||||
name: ruff
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/matchers/ruff.json
|
||||
- .github/workflows/ruff.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
# This workflow is only relevant when one of the following files changes.
|
||||
# However, we have github configured to expect and require this workflow
|
||||
# to run and pass before github with auto-merge a pull request. Until github
|
||||
# allows more flexible auto-merge policy, we can just run this on every PR.
|
||||
# It doesn't take that long to run, anyway.
|
||||
#paths:
|
||||
# - "**/*.py"
|
||||
# - pyproject.toml
|
||||
# - requirements-lint.txt
|
||||
# - .github/workflows/matchers/ruff.json
|
||||
# - .github/workflows/ruff.yml
|
||||
|
||||
jobs:
|
||||
ruff:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-lint.txt
|
||||
- name: Analysing the code with ruff
|
||||
run: |
|
||||
ruff check .
|
||||
- name: Run isort
|
||||
run: |
|
||||
isort . --check-only
|
||||
@@ -0,0 +1,30 @@
|
||||
name: Run Tests
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [ main ]
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install torch
|
||||
pip install packaging ninja
|
||||
pip install -e .
|
||||
pip install pytest
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest
|
||||
@@ -0,0 +1,38 @@
|
||||
name: yapf
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- .github/workflows/yapf.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- .github/workflows/yapf.yml
|
||||
|
||||
jobs:
|
||||
yapf:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install yapf==0.32.0
|
||||
pip install toml==0.10.2
|
||||
- name: Running yapf
|
||||
run: |
|
||||
yapf --diff --recursive .
|
||||
@@ -20,7 +20,6 @@ wandb/
|
||||
*.pt
|
||||
cache_dir/
|
||||
wandb/
|
||||
test*
|
||||
sample_video*
|
||||
sample_image*
|
||||
512*
|
||||
|
||||
@@ -1,21 +1,201 @@
|
||||
MIT License
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
Copyright (c) 2024 PKU-YUAN's Group (袁粒课题组-北大信工) and Rabbitpre AI
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
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:
|
||||
1. Definitions.
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
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.
|
||||
@@ -1,108 +1,187 @@
|
||||
# Fast Video
|
||||
This is currently based on Open-Sora-1.2.0: https://github.com/PKU-YuanGroup/Open-Sora-Plan/tree/294993ca78bf65dec1c3b6fb25541432c545eda9
|
||||
<div align="center">
|
||||
<img src=assets/logo.jpg width="30%"/>
|
||||
</div>
|
||||
|
||||
## Envrironment
|
||||
Change the index-url cuda version according to your system.
|
||||
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>
|
||||
</p>
|
||||
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- 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.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
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/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.
|
||||
```
|
||||
conda create -n fastvideo python=3.10.12
|
||||
conda activate fastvideo
|
||||
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
pip install git+https://github.com/huggingface/diffusers.git@76b7d86a9a5c0c2186efa09c4a67b5f5666ac9e3
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
./env_setup.sh fastvideo
|
||||
```
|
||||
|
||||
## 🚀 Inference
|
||||
|
||||
### 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.
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
pip install -e . && pip install -e ".[train]"
|
||||
sudo apt-get update && apt install screen && pip install watch gpustat
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
### FastHunyuan
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
### FastMochi
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
|
||||
## Prepare Data & Models
|
||||
We've prepared some debug data to facilitate development. To make sure the training pipeline is correct, train on the debug data and make sure the model overfit on it (feed it the same text prompt and see if the output video is the same as the training data)
|
||||
|
||||
## 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we preprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
mkdir data && mkdir data/outputs/
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/mochi_diffuser --local_dir=data/mochi --repo_type=model
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/Merge-30k-Data --local_dir=data/Merge-30k-Data --repo_type=dataset
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/validation_embeddings --local_dir=data/validation_embeddings --repo_type=dataset
|
||||
cd data/Merge-30k-Data
|
||||
cat Merged30K.tar.gz.part.* > Merged30K.tar.gz
|
||||
rm Merged30K.tar.gz.part.*
|
||||
tar --use-compress-program="pigz --processes 64" -xvf Merged30K.tar.gz
|
||||
mv ephemeral/hao.zhang/codefolder/FastVideo-OSP/data/Merged-30K-Data/* .
|
||||
rm -r ephemeral
|
||||
rm Merged30K.tar.gz
|
||||
cd ../..
|
||||
Next, download the original model weights with:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
```
|
||||
To launch the distillation process, use the following commands:
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
## Finetune
|
||||
### ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
Download the original model weights as specified in [Distill Section](#-distill):
|
||||
|
||||
## Things Learned
|
||||
1. shift8 clear but got structural artifacts
|
||||
2. lq, 0.025 vague
|
||||
3. adv not really helpful
|
||||
4. shift8 euler steps 50 v.s. 100 very similar
|
||||
5. 为啥image不会越distill越炸
|
||||
6. EMA, 大batchsize, 1.5,2.5,3.5,4.5
|
||||
7. Must have schedule
|
||||
8. phase 1, 2 learning rate 5e-6不行
|
||||
Then you can run the finetune with:
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
### ⚡ Lora Finetune
|
||||
|
||||
## Experiments
|
||||
Scripts are located at scripts/experiment_N.sh
|
||||
|
||||
1. pcm_linear_quadratic, euler_steps 50, 0.025
|
||||
2. pcm_linear_quadratic, euler_steps 50, 0.05
|
||||
3. shift 8, euler_steps 100
|
||||
4. shift 8, euler_steps 50
|
||||
5. shift 8, euler_steps 100, adv
|
||||
6. pcm_linear_quadratic, euler_steps 50, 0.025, adv
|
||||
7. pcm_linear_quadratic, euler_steps 50, 0.05, multiphase 125
|
||||
8. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
9. pcm_linear_quadratic, euler_steps 50, 0.05, range 0.75
|
||||
10. pcm_linear_quadratic, euler_steps 50, 0.05, batchsize 32
|
||||
11. pcm_linear_quadratic, euler_steps 50, learning rate,1e-7
|
||||
12. shift1, euler_steps 50
|
||||
|
||||
13. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1
|
||||
14. 4.5 cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
15. pcm_linear_quadratic, euler_steps 50, 0.15, linear_range 0.75
|
||||
16. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75 ema 0.95, decay 0.0
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
#### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
|
||||
17. no cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
18. shift16, euler_steps 50
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
19. 4step_infer_shift16_euler_50
|
||||
20. 4step_infer_shift12_euler_50
|
||||
21. 4step_infer_lq_euler_50_thresh0.1_lrg_0.75
|
||||
22. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, lr 1e-7
|
||||
23. lq_euler_50_thres0.1_lrg_0.75_bs_64
|
||||
24. lq_euler_50_thres0.1_lrg_0.75_lr5e-7
|
||||
#### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
Also, we provide script to resize your videos:
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
#### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
#### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
#### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
25. shift1_euler_50_0.75_phase1
|
||||
26. kill
|
||||
27. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, ema 0.95, cfg 4.5
|
||||
- More distillation methods
|
||||
- [ ] Add Distribution Matching Distillation
|
||||
- More models support
|
||||
- [ ] Add CogvideoX model
|
||||
- Code update
|
||||
- [ ] fp8 support
|
||||
- [ ] faster load model and save model support
|
||||
|
||||
28. lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.95
|
||||
29. lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95_cfg7
|
||||
30. lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.98_cfg4.5
|
||||
31. lq_euler_50_thresh0.1_lrg_0.75_phase1_lr_3e-7
|
||||
32. lq_euler_50_thresh0.15_lrg_0.75_phase1_ema0.95_cfg4.5
|
||||
33. lq_euler_50_thres0.1_linear_range_0.75_repro
|
||||
34. lq_euler_50_thres0.1_lrg_0.75_reproduc
|
||||
## 🤝 Contributing
|
||||
|
||||
35. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, learning rate 5e-6
|
||||
36. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 2, learning rate 1e-6
|
||||
37. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 2, learning rate 5e-6
|
||||
38. lq_euler_50_thres0.1_linear_range_0.75, learning rate 5e-6
|
||||
39. lq_euler_50_thres0.1_linear_range_0.75, learning rate 1e-5
|
||||
40. lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6_repro
|
||||
We welcome all contributions. Please run `bash format.sh` 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.
|
||||
|
||||
41. lq_euler_50_thres0.1_lrg_0.75_reproduce
|
||||
42. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 4, learning rate 1e-6
|
||||
43. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, learning rate 1e-6, cfg 6.0
|
||||
44. lq_euler_50_thres0.1_lrg_0.75_phase1_lr_5e-6_test_norm
|
||||
45. lq_euler_50_thres0.1_lrg_0.75_phase1_lr_5e-6_pred_decay_0.1_latent14
|
||||
46-48. lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6, l2 or l1, decay weight 0.1 to 0.001
|
||||
## Acknowledgement
|
||||
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).
|
||||
|
||||
49.
|
||||
We thank MBZUAI and Anyscale for their support throughout this project.
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
After Width: | Height: | Size: 22 MiB |
Binary file not shown.
Binary file not shown.
|
After Width: | Height: | Size: 149 KiB |
@@ -0,0 +1,8 @@
|
||||
Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.
|
||||
A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature.
|
||||
A hand with delicate fingers picks up a bright yellow lemon from a wooden bowl filled with lemons and sprigs of mint against a peach-colored background. The hand gently tosses the lemon up and catches it, showcasing its smooth texture. A beige string bag sits beside the bowl, adding a rustic touch to the scene. Additional lemons, one halved, are scattered around the base of the bowl. The even lighting enhances the vibrant colors and creates a fresh, inviting atmosphere.
|
||||
A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.
|
||||
A superintelligent humanoid robot waking up. The robot has a sleek metallic body with futuristic design features. Its glowing red eyes are the focal point, emanating a sharp, intense light as it powers on. The scene is set in a dimly lit, high-tech laboratory filled with glowing control panels, robotic arms, and holographic screens. The setting emphasizes advanced technology and an atmosphere of mystery. The ambiance is eerie and dramatic, highlighting the moment of awakening and the robots immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
|
||||
fox in the forest close-up quickly turned its head to the left
|
||||
Man walking his dog in the woods on a hot sunny day
|
||||
A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.
|
||||
@@ -0,0 +1,24 @@
|
||||
# Configuration for Cog ⚙️
|
||||
# Reference: https://cog.run/yaml
|
||||
|
||||
build:
|
||||
gpu: true
|
||||
cuda: "12.1"
|
||||
python_version: "3.10"
|
||||
python_packages:
|
||||
- "torch==2.4.0"
|
||||
- "torchvision"
|
||||
- "ninja==1.11.1.3"
|
||||
- "transformers==4.46.1"
|
||||
- "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
|
||||
- "accelerate==1.0.1"
|
||||
- "safetensors==0.4.5"
|
||||
- "peft==0.13.2"
|
||||
- "packaging==24.2"
|
||||
- "git+https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
run:
|
||||
- FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation
|
||||
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/latest/download/pget_$(uname -s)_$(uname -m)" && chmod +x /usr/local/bin/pget
|
||||
|
||||
predict: "predict.py:Predictor"
|
||||
@@ -0,0 +1,210 @@
|
||||
import argparse
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
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
|
||||
|
||||
|
||||
def init_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=25)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=8)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
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("--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)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.1)
|
||||
parser.add_argument("--linear_range", type=float, default=0.75)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def load_model(args):
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
|
||||
scheduler = PCMFMScheduler(
|
||||
1000,
|
||||
args.shift,
|
||||
args.num_euler_timesteps,
|
||||
linear_quadratic,
|
||||
args.linear_threshold,
|
||||
args.linear_range,
|
||||
)
|
||||
|
||||
if 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)
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to(device)
|
||||
# if args.cpu_offload:
|
||||
pipe.enable_sequential_cpu_offload()
|
||||
return pipe
|
||||
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
):
|
||||
if randomize_seed:
|
||||
seed = torch.randint(0, 1000000, (1, )).item()
|
||||
|
||||
generator = torch.Generator(device="cuda").manual_seed(seed)
|
||||
|
||||
if not use_negative_prompt:
|
||||
negative_prompt = None
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
output = pipe(
|
||||
prompt=[prompt],
|
||||
negative_prompt=negative_prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
generator=generator,
|
||||
).frames[0]
|
||||
|
||||
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
|
||||
export_to_video(output, output_path, fps=30)
|
||||
return output_path, seed
|
||||
|
||||
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
args = init_args()
|
||||
pipe = load_model(args)
|
||||
print("load model successfully")
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
|
||||
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=args.height,
|
||||
)
|
||||
width = gr.Slider(label="Width",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=args.width)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=21,
|
||||
maximum=163,
|
||||
value=args.num_frames,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=args.guidance_scale,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=4,
|
||||
maximum=100,
|
||||
value=args.num_inference_steps,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(label="Use negative prompt",
|
||||
value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
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")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=negative_prompt,
|
||||
)
|
||||
|
||||
run_button.click(
|
||||
fn=generate_video,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output],
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
@@ -0,0 +1,68 @@
|
||||
|
||||
|
||||
|
||||
## 🧱 Data Preprocess
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
```
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
```
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
### Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
For image media,
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
```
|
||||
For video media,
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
```
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
Executable
+12
@@ -0,0 +1,12 @@
|
||||
#!/bin/bash
|
||||
|
||||
# install torch
|
||||
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
|
||||
# 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
|
||||
|
||||
# install fastvideo
|
||||
pip install -e .
|
||||
@@ -0,0 +1,173 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.utils.load import load_text_encoder, load_vae
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class T5dataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
vae_debug,
|
||||
):
|
||||
self.json_path = json_path
|
||||
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"])
|
||||
|
||||
def __getitem__(self, idx):
|
||||
caption = self.train_dataset[idx]["caption"]
|
||||
filename = self.train_dataset[idx]["latent_path"].split(".")[0]
|
||||
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"]),
|
||||
map_location="cpu",
|
||||
)
|
||||
else:
|
||||
latents = []
|
||||
|
||||
return dict(caption=caption,
|
||||
latents=latents,
|
||||
filename=filename,
|
||||
length=length)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.train_dataset)
|
||||
|
||||
|
||||
def main(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)
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
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)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
json_data = []
|
||||
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"], )
|
||||
if args.vae_debug:
|
||||
latents = data["latents"]
|
||||
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")
|
||||
# save latent
|
||||
torch.save(prompt_embeds[idx], prompt_embed_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)
|
||||
item = {}
|
||||
item["length"] = int(data["length"][idx])
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["prompt_embed_path"] = video_name + ".pt"
|
||||
item["prompt_attention_mask"] = video_name + ".pt"
|
||||
item["caption"] = data["caption"][idx]
|
||||
json_data.append(item)
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
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:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument(
|
||||
"--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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
type=int,
|
||||
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("--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.",
|
||||
)
|
||||
parser.add_argument("--vae_debug", action="store_true")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,135 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import getdataset
|
||||
from fastvideo.utils.load import load_vae
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def main(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)
|
||||
train_dataset = getdataset(args)
|
||||
sampler = DistributedSampler(train_dataset,
|
||||
rank=local_rank,
|
||||
num_replicas=world_size,
|
||||
shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
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)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
json_data = []
|
||||
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()
|
||||
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")
|
||||
torch.save(latents[idx].to(torch.bfloat16), latent_path)
|
||||
item = {}
|
||||
item["length"] = latents[idx].shape[1]
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["caption"] = data["text"][idx]
|
||||
json_data.append(item)
|
||||
print(f"{video_name} processed")
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
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:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
type=int,
|
||||
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("--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("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
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("--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.",
|
||||
)
|
||||
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***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,78 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
|
||||
from fastvideo.utils.load import load_text_encoder
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def main(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)
|
||||
|
||||
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)
|
||||
|
||||
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
|
||||
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
|
||||
os.makedirs(
|
||||
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)
|
||||
|
||||
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
prompts = [line.strip() for line in lines]
|
||||
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)
|
||||
file_name = prompt.split(".")[0]
|
||||
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",
|
||||
"prompt_attention_mask",
|
||||
f"{file_name}.pt",
|
||||
)
|
||||
torch.save(prompt_embeds[0], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[0],
|
||||
prompt_attention_mask_path)
|
||||
print(f"sample {file_name} saved")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,68 +1,82 @@
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset
|
||||
from fastvideo.dataset.transform import Normalize255, TemporalRandomCrop,CenterCropResizeVideo
|
||||
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. * x - 1.)
|
||||
resize_topcrop = [CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True), ]
|
||||
resize = [CenterCropResizeVideo((args.max_height, args.max_width)), ]
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width),
|
||||
top_crop=True),
|
||||
]
|
||||
resize = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width)),
|
||||
]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
# norm_fun
|
||||
*resize,
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
norm_fun
|
||||
*resize_topcrop,
|
||||
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)
|
||||
if args.dataset == 't2v':
|
||||
return T2V_dataset(args, transform=transform, temporal_sample=temporal_sample, tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name,
|
||||
cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(
|
||||
args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from accelerate import Accelerator
|
||||
from fastvideo.dataset.t2v_datasets import dataset_prog
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
args = type('args', (),
|
||||
{
|
||||
'ae': 'CausalVAEModel_4x8x8',
|
||||
'dataset': 't2v',
|
||||
'attention_mode': 'xformers',
|
||||
'use_rope': True,
|
||||
'text_max_length': 300,
|
||||
'max_height': 320,
|
||||
'max_width': 240,
|
||||
'num_frames': 1,
|
||||
'use_image_num': 0,
|
||||
'interpolation_scale_t': 1,
|
||||
'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',
|
||||
'video_data': '1',
|
||||
'train_fps': 24,
|
||||
'drop_short_ratio': 1.0,
|
||||
'use_img_from_vid': False,
|
||||
'speed_factor': 1.0,
|
||||
'cfg': 0.1,
|
||||
'text_encoder_name': 'google/mt5-xxl',
|
||||
'dataloader_num_workers': 10,
|
||||
|
||||
}
|
||||
from accelerate import Accelerator
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import dataset_prog
|
||||
|
||||
args = type(
|
||||
"args",
|
||||
(),
|
||||
{
|
||||
"ae": "CausalVAEModel_4x8x8",
|
||||
"dataset": "t2v",
|
||||
"attention_mode": "xformers",
|
||||
"use_rope": True,
|
||||
"text_max_length": 300,
|
||||
"max_height": 320,
|
||||
"max_width": 240,
|
||||
"num_frames": 1,
|
||||
"use_image_num": 0,
|
||||
"interpolation_scale_t": 1,
|
||||
"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",
|
||||
"video_data": "1",
|
||||
"train_fps": 24,
|
||||
"drop_short_ratio": 1.0,
|
||||
"use_img_from_vid": False,
|
||||
"speed_factor": 1.0,
|
||||
"cfg": 0.1,
|
||||
"text_encoder_name": "google/mt5-xxl",
|
||||
"dataloader_num_workers": 10,
|
||||
},
|
||||
)
|
||||
accelerator = Accelerator()
|
||||
dataset = getdataset(args)
|
||||
@@ -70,7 +84,10 @@ 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:
|
||||
@@ -81,5 +98,7 @@ if __name__ == "__main__":
|
||||
continue
|
||||
assert caps[0] is not None and len(caps[0]) > 0
|
||||
print(num, zero)
|
||||
import ipdb;ipdb.set_trace()
|
||||
print('end')
|
||||
import ipdb
|
||||
|
||||
ipdb.set_trace()
|
||||
print("end")
|
||||
|
||||
@@ -1,25 +1,30 @@
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
):
|
||||
):
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
self.cfg_rate = cfg_rate
|
||||
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")
|
||||
with open(self.json_path, 'r') as f:
|
||||
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
|
||||
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
|
||||
@@ -28,27 +33,44 @@ 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"]
|
||||
# load
|
||||
latent = torch.load(os.path.join(self.latent_dir, latent_file), map_location="cpu", weights_only=True)
|
||||
# TODO: Hack
|
||||
|
||||
prompt_attention_mask_file = self.data_anno[idx][
|
||||
"prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(
|
||||
os.path.join(self.latent_dir, latent_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
if random.random() < self.cfg_rate:
|
||||
prompt_embed = self.uncond_prompt_embed
|
||||
prompt_attention_mask = self.uncond_prompt_mask
|
||||
else:
|
||||
prompt_embed = torch.load(os.path.join(self.prompt_embed_dir, prompt_embed_file), map_location="cpu", weights_only=True)
|
||||
prompt_attention_mask = torch.load(os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file), map_location="cpu", weights_only=True)
|
||||
prompt_embed = torch.load(
|
||||
os.path.join(self.prompt_embed_dir, prompt_embed_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
prompt_attention_mask = torch.load(
|
||||
os.path.join(self.prompt_attention_mask_dir,
|
||||
prompt_attention_mask_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
return latent, prompt_embed, prompt_attention_mask
|
||||
|
||||
|
||||
def __len__(self):
|
||||
return len(self.data_anno)
|
||||
|
||||
|
||||
|
||||
def latent_collate_function(batch):
|
||||
# return latent, prompt, latent_attn_mask, text_attn_mask
|
||||
# latent_attn_mask: # b t h w
|
||||
@@ -59,9 +81,21 @@ def latent_collate_function(batch):
|
||||
max_t = max([latent.shape[1] for latent in latents])
|
||||
max_h = max([latent.shape[2] for latent in latents])
|
||||
max_w = max([latent.shape[3] for latent in latents])
|
||||
|
||||
|
||||
# padding
|
||||
latents = [torch.nn.functional.pad(latent, (0, max_t - latent.shape[1], 0, max_h - latent.shape[2], 0, max_w - latent.shape[3])) for latent in latents]
|
||||
latents = [
|
||||
torch.nn.functional.pad(
|
||||
latent,
|
||||
(
|
||||
0,
|
||||
max_t - latent.shape[1],
|
||||
0,
|
||||
max_h - latent.shape[2],
|
||||
0,
|
||||
max_w - latent.shape[3],
|
||||
),
|
||||
) for latent in latents
|
||||
]
|
||||
# attn mask
|
||||
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
|
||||
# set to 0 if padding
|
||||
@@ -69,15 +103,28 @@ def latent_collate_function(batch):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
|
||||
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
latents = torch.stack(latents, dim=0)
|
||||
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
|
||||
|
||||
|
||||
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, prompt_embed.shape, latent_attn_mask.shape, prompt_attention_mask.shape)
|
||||
import pdb; pdb.set_trace()
|
||||
print(
|
||||
latent.shape,
|
||||
prompt_embed.shape,
|
||||
latent_attn_mask.shape,
|
||||
prompt_attention_mask.shape,
|
||||
)
|
||||
import pdb
|
||||
|
||||
pdb.set_trace()
|
||||
|
||||
+141
-104
@@ -1,27 +1,22 @@
|
||||
import json
|
||||
import os, io, csv, math, random
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
from decord import VideoReader
|
||||
from os.path import join as opj
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data.dataset import Dataset
|
||||
from torch.utils.data import DataLoader, Dataset, get_worker_info
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from accelerate.logging import get_logger
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.utils import text_preprocessing
|
||||
import torchvision
|
||||
logger = get_logger(__name__)
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
"""
|
||||
这是一个元类,用于创建单例类。
|
||||
"""
|
||||
_instances = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
@@ -32,6 +27,7 @@ class SingletonMeta(type):
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self):
|
||||
self.cap_list = []
|
||||
self.elements = []
|
||||
@@ -50,10 +46,11 @@ 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]
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info):
|
||||
if work_info is None:
|
||||
@@ -61,22 +58,29 @@ 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
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
|
||||
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
|
||||
@@ -95,17 +99,18 @@ class T2V_dataset(Dataset):
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if not ('mt5' in args.text_encoder_name):
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
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)
|
||||
|
||||
@@ -117,39 +122,37 @@ class T2V_dataset(Dataset):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
try:
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.info(f'Error with {e}')
|
||||
if idx in dataset_prog.cap_list:
|
||||
logger.info(f"Caught an exception! {dataset_prog.cap_list[idx]}")
|
||||
return self.__getitem__(random.randint(0, self.__len__() - 1))
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx):
|
||||
path = dataset_prog.cap_list[idx]['path']
|
||||
if path.endswith('.mp4'):
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
|
||||
def get_video(self, idx):
|
||||
video_path = dataset_prog.cap_list[idx]['path']
|
||||
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")
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
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')
|
||||
video = video.to(torch.uint8)
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}'
|
||||
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]['cap']
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
@@ -158,51 +161,66 @@ class T2V_dataset(Dataset):
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding='max_length',
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors='pt'
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(
|
||||
pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
)
|
||||
input_ids = text_tokens_and_mask['input_ids']
|
||||
cond_mask = text_tokens_and_mask['attention_mask']
|
||||
return dict(pixel_values=video, text=text, input_ids=input_ids, cond_mask=cond_mask, path=video_path)
|
||||
|
||||
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 = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, 'h w c -> c h w').unsqueeze(0) # [1 c h w]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# 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) # [1 C H W] -> num_img [1 C H W]
|
||||
|
||||
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 = text_preprocessing(caps, support_Chinese=self.support_Chinese)
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
text = text if random.random() > self.cfg else ""
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding='max_length',
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors='pt'
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
input_ids = text_tokens_and_mask['input_ids'] # 1, l
|
||||
cond_mask = text_tokens_and_mask['attention_mask'] # 1, l
|
||||
return dict(pixel_values=image, text=text, input_ids=input_ids, cond_mask=cond_mask, path=image_data['path'])
|
||||
|
||||
def define_frame_index(self, cap_list):
|
||||
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
@@ -213,100 +231,119 @@ class T2V_dataset(Dataset):
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i['path']
|
||||
cap = i.get('cap', None)
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith('.mp4'):
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get('duration', None)
|
||||
fps = i.get('fps', None)
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get('resolution', None)
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
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(height, width, max_h_div_w_ratio=hw_aspect_thr*aspect,
|
||||
min_h_div_w_ratio=1/hw_aspect_thr*aspect)
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
start_frame_idx = 0
|
||||
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))
|
||||
frame_indices = frame_indices[begin_index: end_index]
|
||||
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()
|
||||
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)
|
||||
sample_num_frames.append(i['sample_num_frames'])
|
||||
elif path.endswith('.jpg'): # image
|
||||
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
|
||||
new_cap_list.append(i)
|
||||
i['sample_num_frames'] = 1
|
||||
sample_num_frames.append(i['sample_num_frames'])
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(f"Unknown file extention {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
logger.info(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)}')
|
||||
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)}"
|
||||
)
|
||||
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]
|
||||
with open(data, "r") as f:
|
||||
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:
|
||||
with open(anno, "r") as f:
|
||||
sub_list = json.load(f)
|
||||
logger.info(f'Building {anno}...')
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]['path'] = opj(folder, sub_list[i]['path'])
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
|
||||
def get_cap_list(self):
|
||||
cap_lists = self.read_jsons(self.data)
|
||||
return cap_lists
|
||||
|
||||
+136
-82
@@ -1,7 +1,8 @@
|
||||
import torch
|
||||
import random
|
||||
import numbers
|
||||
from torchvision.transforms import RandomCrop, RandomResizedCrop
|
||||
import random
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip):
|
||||
@@ -20,19 +21,19 @@ 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):
|
||||
@@ -42,21 +43,37 @@ def crop(clip, i, j, h, w):
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i: i + h, j: j + w]
|
||||
return clip[..., i:i + h, j:j + 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}")
|
||||
return torch.nn.functional.interpolate(clip, size=target_size, mode=interpolation_mode, align_corners=True, antialias=True)
|
||||
raise ValueError(
|
||||
f"target size should be tuple (height, width), instead got {target_size}"
|
||||
)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
size=target_size,
|
||||
mode=interpolation_mode,
|
||||
align_corners=True,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
|
||||
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(clip, scale_factor=scale_, mode=interpolation_mode, align_corners=True, antialias=True)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
scale_factor=scale_,
|
||||
mode=interpolation_mode,
|
||||
align_corners=True,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
|
||||
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
|
||||
@@ -107,11 +124,10 @@ def center_crop_using_short_edge(clip):
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
|
||||
def center_crop_th_tw(clip, th, tw, top_crop):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
tr = th / tw
|
||||
@@ -121,30 +137,29 @@ def center_crop_th_tw(clip, th, tw, top_crop):
|
||||
else:
|
||||
new_h = h
|
||||
new_w = int(h / tr)
|
||||
|
||||
|
||||
i = 0 if top_crop else int(round((h - new_h) / 2.0))
|
||||
j = int(round((w - new_w) / 2.0))
|
||||
return crop(clip, i, j, new_h, new_w)
|
||||
|
||||
|
||||
def random_shift_crop(clip):
|
||||
'''
|
||||
"""
|
||||
Slide along the long edge, with the short edge as crop size
|
||||
'''
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
|
||||
if h <= w:
|
||||
long_edge = w
|
||||
short_edge = h
|
||||
else:
|
||||
long_edge = h
|
||||
short_edge = w
|
||||
|
||||
th, tw = short_edge, short_edge
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
i = torch.randint(0, h - th + 1, size=(1, )).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1, )).item()
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
@@ -159,7 +174,8 @@ 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
|
||||
|
||||
@@ -197,6 +213,7 @@ def hflip(clip):
|
||||
|
||||
|
||||
class RandomCropVideo:
|
||||
|
||||
def __init__(self, size):
|
||||
if isinstance(size, numbers.Number):
|
||||
self.size = (int(size), int(size))
|
||||
@@ -219,13 +236,15 @@ 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
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
i = torch.randint(0, h - th + 1, size=(1, )).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1, )).item()
|
||||
|
||||
return i, j, th, tw
|
||||
|
||||
@@ -234,8 +253,9 @@ class RandomCropVideo:
|
||||
|
||||
|
||||
class SpatialStrideCropVideo:
|
||||
|
||||
def __init__(self, stride):
|
||||
self.stride = stride
|
||||
self.stride = stride
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
@@ -258,17 +278,18 @@ class SpatialStrideCropVideo:
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
|
||||
class LongSideResizeVideo:
|
||||
'''
|
||||
"""
|
||||
First use the long side,
|
||||
then resize to the specified size
|
||||
'''
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
skip_low_resolution=False,
|
||||
interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
skip_low_resolution=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
self.size = size
|
||||
self.skip_low_resolution = skip_low_resolution
|
||||
@@ -291,27 +312,30 @@ 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:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class CenterCropResizeVideo:
|
||||
'''
|
||||
"""
|
||||
First use the short side for cropping length,
|
||||
center crop video, then resize to the specified size
|
||||
'''
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
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
|
||||
@@ -325,10 +349,16 @@ 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, target_size=self.size,
|
||||
interpolation_mode=self.interpolation_mode)
|
||||
clip_center_crop_resize = resize(
|
||||
clip_center_crop,
|
||||
target_size=self.size,
|
||||
interpolation_mode=self.interpolation_mode,
|
||||
)
|
||||
return clip_center_crop_resize
|
||||
|
||||
def __repr__(self) -> str:
|
||||
@@ -336,19 +366,21 @@ class CenterCropResizeVideo:
|
||||
|
||||
|
||||
class UCFCenterCropVideo:
|
||||
'''
|
||||
"""
|
||||
First scale to the specified size in equal proportion to the short edge,
|
||||
then center cropping
|
||||
'''
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
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)
|
||||
@@ -363,7 +395,9 @@ 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
|
||||
|
||||
@@ -372,18 +406,20 @@ class UCFCenterCropVideo:
|
||||
|
||||
|
||||
class KineticsRandomCropResizeVideo:
|
||||
'''
|
||||
"""
|
||||
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
|
||||
'''
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
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)
|
||||
@@ -392,19 +428,23 @@ 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
|
||||
|
||||
|
||||
class CenterCropVideo:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
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)
|
||||
@@ -516,6 +556,7 @@ class TemporalRandomCrop(object):
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
|
||||
|
||||
class DynamicSampleDuration(object):
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
@@ -530,31 +571,34 @@ 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
|
||||
return 0, truncate_t
|
||||
|
||||
if __name__ == '__main__':
|
||||
from torchvision import transforms
|
||||
import torchvision.io as io
|
||||
import numpy as np
|
||||
from torchvision.utils import save_image
|
||||
|
||||
if __name__ == "__main__":
|
||||
import os
|
||||
|
||||
vframes, aframes, info = io.read_video(
|
||||
filename='./v_Archery_g01_c03.avi',
|
||||
pts_unit='sec',
|
||||
output_format='TCHW'
|
||||
)
|
||||
import numpy as np
|
||||
import torchvision.io as io
|
||||
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")
|
||||
|
||||
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
|
||||
@@ -569,7 +613,10 @@ 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]
|
||||
@@ -580,12 +627,19 @@ 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(select_vframes_trans[i], os.path.join('./test000', '%04d.png' % i), normalize=True,
|
||||
value_range=(-1, 1))
|
||||
save_image(
|
||||
select_vframes_trans[i],
|
||||
os.path.join("./test000", "%04d.png" % i),
|
||||
normalize=True,
|
||||
value_range=(-1, 1),
|
||||
)
|
||||
|
||||
+602
-425
File diff suppressed because it is too large
Load Diff
@@ -1,30 +1,11 @@
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.models.attention import JointTransformerBlock
|
||||
from diffusers.models.attention_processor import Attention, AttentionProcessor
|
||||
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.models.embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed
|
||||
from diffusers.models.transformers.transformer_2d import Transformer2DModelOutput
|
||||
from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel
|
||||
from diffusers.utils import logging
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
|
||||
class DiscriminatorHead(nn.Module):
|
||||
|
||||
def __init__(self, input_channel, output_channel=1):
|
||||
super().__init__()
|
||||
inner_channel = 1024
|
||||
@@ -48,9 +29,9 @@ class DiscriminatorHead(nn.Module):
|
||||
def forward(self, x):
|
||||
b, twh, c = x.shape
|
||||
t = twh // (30 * 53)
|
||||
x = x.view(-1, 30 *53, c)
|
||||
x = x.view(-1, 30 * 53, c)
|
||||
x = x.permute(0, 2, 1)
|
||||
x = x.view(b*t, c, 30, 53)
|
||||
x = x.view(b * t, c, 30, 53)
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x) + x
|
||||
x = self.conv_out(x)
|
||||
@@ -61,45 +42,41 @@ class Discriminator(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stride = 8,
|
||||
stride=8,
|
||||
num_h_per_head=1,
|
||||
adapter_channel_dims=[3072],
|
||||
total_layers=48,
|
||||
):
|
||||
super().__init__()
|
||||
adapter_channel_dims = adapter_channel_dims * (48 // stride)
|
||||
adapter_channel_dims = adapter_channel_dims * (total_layers // stride)
|
||||
self.stride = stride
|
||||
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
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
self.heads = nn.ModuleList([
|
||||
nn.ModuleList([
|
||||
DiscriminatorHead(adapter_channel)
|
||||
for _ in range(self.num_h_per_head)
|
||||
]) for adapter_channel in adapter_channel_dims
|
||||
])
|
||||
|
||||
def forward(self, features):
|
||||
outputs = []
|
||||
|
||||
def create_custom_forward(module):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
assert len(features) // self.stride == len(self.heads)
|
||||
for i in range(0, len(features), self.stride):
|
||||
for h in self.heads[i//self.stride]:
|
||||
|
||||
assert len(features) == len(self.heads)
|
||||
for i in range(0, len(features)):
|
||||
for h in self.heads[i]:
|
||||
# out = torch.utils.checkpoint.checkpoint(
|
||||
# create_custom_forward(h),
|
||||
# features[i],
|
||||
# use_reentrant=False
|
||||
# )
|
||||
out=h(features[i])
|
||||
out = h(features[i])
|
||||
outputs.append(out)
|
||||
return outputs
|
||||
|
||||
|
||||
|
||||
+62
-62
@@ -3,12 +3,11 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -17,13 +16,14 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
class PCMFMSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
return out.reshape(b, *((1, ) * (len(x_shape) - 1)))
|
||||
|
||||
|
||||
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@@ -34,30 +34,33 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
shift: float = 1.0,
|
||||
pcm_timesteps: int = 50,
|
||||
linear_quadratic=False,
|
||||
linear_quadratic_threshold=0.025,
|
||||
linear_quadratic_threshold=0.025,
|
||||
linear_range=0.5,
|
||||
):
|
||||
|
||||
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
|
||||
self.euler_timesteps = (np.arange(1, pcm_timesteps + 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()
|
||||
|
||||
@@ -116,9 +119,9 @@ 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).
|
||||
|
||||
@@ -129,9 +132,10 @@ 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()
|
||||
|
||||
@@ -139,8 +143,8 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
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.zeros(1, device=self.sigmas_.device)])
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
@@ -202,18 +206,12 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
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 (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)
|
||||
@@ -231,27 +229,30 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample,)
|
||||
return (prev_sample, )
|
||||
|
||||
return PCMFMSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
|
||||
|
||||
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.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)
|
||||
|
||||
@@ -264,10 +265,10 @@ 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
|
||||
|
||||
@@ -279,30 +280,29 @@ 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)
|
||||
)
|
||||
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)
|
||||
)
|
||||
-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
|
||||
|
||||
|
||||
+620
-359
File diff suppressed because it is too large
Load Diff
@@ -1,123 +0,0 @@
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
import numpy as np
|
||||
from torch.nn.utils.parametrizations import spectral_norm
|
||||
import os
|
||||
class DummyDiscriminator(nn.Module):
|
||||
def __init__(self, dim_in, num_layers):
|
||||
super().__init__()
|
||||
self.layers = nn.ModuleList()
|
||||
for _ in range(num_layers):
|
||||
self.layers.append(nn.Linear(dim_in, 1))
|
||||
|
||||
def forward(self, features):
|
||||
logits = []
|
||||
for layer, feature in zip(self.layers, features):
|
||||
mean = feature.mean(dim=1)
|
||||
logits.append(layer(mean))
|
||||
return torch.cat(logits, dim=1)
|
||||
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
def __init__(self, fn):
|
||||
super().__init__()
|
||||
self.fn = fn
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return (self.fn(x) + x) / np.sqrt(2)
|
||||
|
||||
|
||||
class SpectralConv1d(nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
self.conv = spectral_norm(nn.Conv1d(*args, **kwargs))
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
class BatchNormLocal(nn.Module):
|
||||
def __init__(self, num_features: int, affine: bool = True, virtual_bs: int = 8, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.virtual_bs = virtual_bs
|
||||
self.eps = eps
|
||||
self.affine = affine
|
||||
|
||||
if self.affine:
|
||||
self.weight = nn.Parameter(torch.ones(num_features))
|
||||
self.bias = nn.Parameter(torch.zeros(num_features))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
shape = x.size()
|
||||
|
||||
# Calculate stats.
|
||||
mean = x.mean([0, 2], keepdim=True)
|
||||
var = x.var([0, 2], keepdim=True, unbiased=False)
|
||||
x = (x - mean) / (torch.sqrt(var + self.eps))
|
||||
|
||||
if self.affine:
|
||||
x = x * self.weight[None, :, None] + self.bias[None, :, None]
|
||||
|
||||
return x.view(shape)
|
||||
|
||||
def make_block(channels: int, kernel_size: int) -> nn.Module:
|
||||
return nn.Sequential(
|
||||
SpectralConv1d(
|
||||
channels,
|
||||
channels,
|
||||
kernel_size = kernel_size,
|
||||
padding = kernel_size//2,
|
||||
padding_mode = 'circular',
|
||||
),
|
||||
BatchNormLocal(channels),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
)
|
||||
|
||||
class DiscHead(nn.Module):
|
||||
def __init__(self, feature_dim: int, text_c_dim: int, cmap_dim: int = 64, cnn_dim=512):
|
||||
super().__init__()
|
||||
self.channels = feature_dim
|
||||
self.text_c_dim = text_c_dim
|
||||
self.cmap_dim = cmap_dim
|
||||
self.down_proj = SpectralConv1d(feature_dim, cnn_dim, kernel_size=1, padding=0)
|
||||
self.main = nn.Sequential(
|
||||
make_block(cnn_dim, kernel_size=1),
|
||||
ResidualBlock(make_block(cnn_dim, kernel_size=9))
|
||||
)
|
||||
|
||||
self.cmapper = nn.Linear(self.text_c_dim, cmap_dim)
|
||||
self.cls = SpectralConv1d(cnn_dim, cmap_dim, kernel_size=1, padding=0)
|
||||
|
||||
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
|
||||
h = self.down_proj(x)
|
||||
h = self.main(h)
|
||||
out = self.cls(h)
|
||||
|
||||
cmap = self.cmapper(c).unsqueeze(-1)
|
||||
out = (out * cmap).sum(1, keepdim=True) * (1 / np.sqrt(self.cmap_dim))
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class LADDDiscriminator(nn.Module):
|
||||
def __init__(self, feature_dim, text_cond_dim, num_layers, layers_stride):
|
||||
super().__init__()
|
||||
heads = []
|
||||
for i in range(0, num_layers, layers_stride):
|
||||
heads.append(DiscHead(feature_dim, text_cond_dim))
|
||||
self.heads = nn.ModuleList(heads)
|
||||
self.layers_stride = layers_stride
|
||||
self.num_layers = num_layers
|
||||
|
||||
def forward(self, features, text_conditions) -> torch.Tensor:
|
||||
text_conditions = text_conditions.mean(1)
|
||||
# layer, B, L, C -> layer, B, C, L
|
||||
features = features.transpose(2, 3)
|
||||
logits = []
|
||||
for i in range(0, self.num_layers, self.layers_stride):
|
||||
head = self.heads[i//self.layers_stride]
|
||||
feat = features[i]
|
||||
logits.append(head(feat, text_conditions).view(feat.size(0), -1))
|
||||
logits = torch.cat(logits, dim=1)
|
||||
|
||||
|
||||
return logits
|
||||
@@ -1,155 +0,0 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
import torch.distributed as dist
|
||||
|
||||
from diffusers.utils import export_to_video
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
|
||||
import argparse
|
||||
import os
|
||||
from diffusers.models.transformers.transformer_mochi import MochiTransformerBlock
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
|
||||
|
||||
|
||||
import sys
|
||||
import pdb
|
||||
|
||||
class ForkedPdb(pdb.Pdb):
|
||||
"""
|
||||
PDB Subclass for debugging multi-processed code
|
||||
Suggested in: https://stackoverflow.com/questions/4716533/how-to-attach-debugger-to-a-python-subproccess
|
||||
"""
|
||||
def interaction(self, *args, **kwargs):
|
||||
_stdin = sys.stdin
|
||||
try:
|
||||
sys.stdin = open('/dev/stdin')
|
||||
pdb.Pdb.interaction(self, *args, **kwargs)
|
||||
finally:
|
||||
sys.stdin = _stdin
|
||||
def assert_all_close_list(input_list):
|
||||
for i in range(len(input_list) - 1):
|
||||
assert torch.allclose(input_list[i], input_list[i + 1]), f"input_list[{i}]: {input_list[i]}, input_list[{i+1}]: {input_list[i+1]}"
|
||||
|
||||
weight_dtype = torch.float32
|
||||
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)
|
||||
return world_size
|
||||
|
||||
def main_print(content):
|
||||
if int(os.getenv('RANK', 0)) <= 0:
|
||||
print(content)
|
||||
|
||||
@torch.inference_mode
|
||||
def test_single_block(batch_size, device, seed):
|
||||
# set manual seed
|
||||
torch.manual_seed(seed)
|
||||
device = torch.cuda.current_device()
|
||||
block = MochiTransformerBlock(
|
||||
dim=768,
|
||||
num_attention_heads=12,
|
||||
attention_head_dim=64,
|
||||
pooled_projection_dim=256,
|
||||
qk_norm="rms_norm",
|
||||
activation_fn="swiglu",
|
||||
context_pre_only=False,
|
||||
).to(device)
|
||||
hidden_states = torch.randn(1, 16, 768).to(device).repeat(batch_size, 1, 1)
|
||||
encoder_hidden_states = torch.randn(1, 4, 256).to(device).repeat(batch_size, 1, 1)
|
||||
temb = torch.randn(1, 768).to(device).repeat(batch_size, 1)
|
||||
# shard hiddent_states according to world_size
|
||||
local_seq_length = hidden_states.shape[1] // nccl_info.sp_size
|
||||
hidden_states = hidden_states.narrow(1, nccl_info.global_rank * local_seq_length, local_seq_length)
|
||||
main_print(hidden_states.shape)
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=temb,
|
||||
)
|
||||
mean = hidden_states[0].mean()
|
||||
torch.distributed.all_reduce(mean, op=torch.distributed.ReduceOp.SUM)
|
||||
mean = mean / nccl_info.sp_size
|
||||
return mean
|
||||
|
||||
@torch.inference_mode
|
||||
def test_DiT(batch_size, transformer, seed):
|
||||
generator = torch.Generator(torch.cuda.current_device()).manual_seed(seed)
|
||||
device = torch.cuda.current_device()
|
||||
latent = torch.randn((1, 12, 8, 12, 8), device=device, dtype=weight_dtype, generator=generator).repeat(batch_size, 1, 1, 1, 1)
|
||||
prompt_embeds = torch.randn((1, 20, 4096), device=device, dtype=weight_dtype, generator=generator).repeat(batch_size, 1, 1)
|
||||
prompt_attention_mask = torch.ones((1, 20), device=device, dtype=weight_dtype).repeat(batch_size, 1)
|
||||
timestep = 0
|
||||
timestep = torch.tensor(timestep, device=device, dtype=weight_dtype).unsqueeze(0).repeat(batch_size)
|
||||
local_seq_length = latent.shape[2] // nccl_info.sp_size
|
||||
latent = latent.narrow(2, nccl_info.global_rank * local_seq_length, local_seq_length)
|
||||
# main_print(latent.shape)
|
||||
hidden_states = transformer(
|
||||
hidden_states=latent,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
def calculate_mean(states):
|
||||
mean = states.mean()
|
||||
torch.distributed.all_reduce(mean, op=torch.distributed.ReduceOp.SUM)
|
||||
mean = mean / int(os.getenv('WORLD_SIZE', 1))
|
||||
return mean
|
||||
mean1 = calculate_mean(hidden_states[0])
|
||||
main_print(hidden_states.shape)
|
||||
if hidden_states.shape[0] > 1:
|
||||
mean2 = calculate_mean(hidden_states[1])
|
||||
return mean1, mean2
|
||||
return mean1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
world_size = initialize_distributed()
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
parser.add_argument("--test_single_block", action="store_true")
|
||||
args = parser.parse_args()
|
||||
seed = args.seed
|
||||
|
||||
if args.test_single_block:
|
||||
pass
|
||||
single_no_patch_bs_1 = test_single_block(1)
|
||||
single_no_patch_bs_2 = test_single_block(2)
|
||||
# check all close
|
||||
assert torch.allclose(single_no_patch_bs_1, single_no_patch_bs_2)
|
||||
single_patch_bs_1 = test_single_block(1)
|
||||
single_patch_bs_2 = test_single_block(2)
|
||||
assert torch.allclose(single_patch_bs_1, single_patch_bs_2)
|
||||
assert torch.allclose(single_no_patch_bs_1, single_patch_bs_2)
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
|
||||
sp_patch_bs_1 = test_single_block(1)
|
||||
sp_patch_bs_2 = test_single_block(2)
|
||||
|
||||
assert torch.allclose(sp_patch_bs_1, sp_patch_bs_2)
|
||||
assert torch.allclose(single_no_patch_bs_1, sp_patch_bs_2)
|
||||
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained("data/mochi/transformer", torch_dtype=weight_dtype).to(device)
|
||||
single_no_patch_bs_1 = test_DiT(1, transformer, seed)
|
||||
single_no_patch_bs_2_a, single_no_patch_bs_2_b = test_DiT(2, transformer, seed)
|
||||
|
||||
|
||||
single_patch_bs_1 = test_DiT(1, transformer, seed)
|
||||
single_patch_bs_2_a, single_patch_bs_2_b = test_DiT(2, transformer, seed)
|
||||
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
|
||||
sp_patch_bs_1 = test_DiT(1, transformer, seed)
|
||||
sp_patch_bs_2_a, sp_patch_bs_2_b = test_DiT(2, transformer, seed)
|
||||
|
||||
assert_all_close_list([single_no_patch_bs_1, single_no_patch_bs_2_a, single_no_patch_bs_2_b, single_patch_bs_1, single_patch_bs_2_a, single_patch_bs_2_b, sp_patch_bs_1, sp_patch_bs_2_a, sp_patch_bs_2_b])
|
||||
|
||||
main_print(sp_patch_bs_1)
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
from einops import rearrange
|
||||
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):
|
||||
# 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 = 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,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
)
|
||||
output = rearrange(
|
||||
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,
|
||||
)
|
||||
return output
|
||||
@@ -0,0 +1,89 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
__all__ = [
|
||||
"C_SCALE",
|
||||
"PROMPT_TEMPLATE",
|
||||
"MODEL_BASE",
|
||||
"PRECISIONS",
|
||||
"NORMALIZATION_TYPE",
|
||||
"ACTIVATION_TYPE",
|
||||
"VAE_PATH",
|
||||
"TEXT_ENCODER_PATH",
|
||||
"TOKENIZER_PATH",
|
||||
"TEXT_PROJECTION",
|
||||
"DATA_TYPE",
|
||||
"NEGATIVE_PROMPT",
|
||||
]
|
||||
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
}
|
||||
|
||||
# =================== Constant Values =====================
|
||||
# Computation scale factor, 1P = 1_000_000_000_000_000. Tensorboard will display the value in PetaFLOPS to avoid
|
||||
# overflow error when tensorboard logging values.
|
||||
C_SCALE = 1_000_000_000_000_000
|
||||
|
||||
# When using decoder-only models, we must provide a prompt template to instruct the text encoder
|
||||
# on how to generate the text.
|
||||
# --------------------------------------------------------------------
|
||||
PROMPT_TEMPLATE_ENCODE = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|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|>")
|
||||
|
||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"dit-llm-encode": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE,
|
||||
"crop_start": 36,
|
||||
},
|
||||
"dit-llm-encode-video": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
},
|
||||
}
|
||||
|
||||
# ======================= Model ======================
|
||||
PRECISIONS = {"fp32", "fp16", "bf16"}
|
||||
NORMALIZATION_TYPE = {"layer", "rms"}
|
||||
ACTIVATION_TYPE = {"relu", "silu", "gelu", "gelu_tanh"}
|
||||
|
||||
# =================== Model Path =====================
|
||||
MODEL_BASE = os.getenv("MODEL_BASE", "./data/hunyuan")
|
||||
|
||||
# =================== Data =======================
|
||||
DATA_TYPE = {"image", "video", "image_video"}
|
||||
|
||||
# 3D VAE
|
||||
VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
|
||||
|
||||
# Text Encoder
|
||||
TEXT_ENCODER_PATH = {
|
||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||
"llm": f"{MODEL_BASE}/text_encoder",
|
||||
}
|
||||
|
||||
# Tokenizer
|
||||
TOKENIZER_PATH = {
|
||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||
"llm": f"{MODEL_BASE}/text_encoder",
|
||||
}
|
||||
|
||||
TEXT_PROJECTION = {
|
||||
"linear", # Default, an nn.Linear() layer
|
||||
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
# ruff: noqa: F401
|
||||
from .pipelines import HunyuanVideoPipeline
|
||||
from .schedulers import FlowMatchDiscreteScheduler
|
||||
@@ -0,0 +1,2 @@
|
||||
# ruff: noqa: F401
|
||||
from .pipeline_hunyuan_video import HunyuanVideoPipeline
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,2 @@
|
||||
# ruff: noqa: F401
|
||||
from .scheduling_flow_match_discrete import FlowMatchDiscreteScheduler
|
||||
@@ -0,0 +1,248 @@
|
||||
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
#
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
|
||||
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.
|
||||
shift (`float`, defaults to 1.0):
|
||||
The shift value for the timestep schedule.
|
||||
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,
|
||||
shift: float = 1.0,
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: Optional[int] = 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.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,
|
||||
device: Union[str, torch.device] = None,
|
||||
n_tokens: int = 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.
|
||||
"""
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
|
||||
sigmas = self.sd3_time_shift(sigmas)
|
||||
|
||||
if not self.config.reverse:
|
||||
sigmas = 1 - sigmas
|
||||
|
||||
self.sigmas = sigmas
|
||||
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
|
||||
dtype=torch.float32, device=device)
|
||||
|
||||
# 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):
|
||||
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
) -> 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
|
||||
@@ -0,0 +1,415 @@
|
||||
# ruff: noqa: F405, F403
|
||||
import argparse
|
||||
import re
|
||||
|
||||
from .constants import *
|
||||
from .modules.models import HUNYUAN_VIDEO_CONFIG
|
||||
|
||||
|
||||
def parse_args(namespace=None):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="HunyuanVideo inference script")
|
||||
|
||||
parser = add_network_args(parser)
|
||||
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)
|
||||
args = sanity_check_args(args)
|
||||
|
||||
return args
|
||||
|
||||
|
||||
def add_network_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="HunyuanVideo network args")
|
||||
|
||||
# Main model
|
||||
group.add_argument(
|
||||
"--model",
|
||||
type=str,
|
||||
choices=list(HUNYUAN_VIDEO_CONFIG.keys()),
|
||||
default="HYVideo-T/2-cfgdistill",
|
||||
)
|
||||
group.add_argument(
|
||||
"--latent-channels",
|
||||
type=str,
|
||||
default=16,
|
||||
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(
|
||||
"--precision",
|
||||
type=str,
|
||||
default="bf16",
|
||||
choices=PRECISIONS,
|
||||
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.")
|
||||
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)"
|
||||
)
|
||||
|
||||
# - VAE
|
||||
group.add_argument(
|
||||
"--vae",
|
||||
type=str,
|
||||
default="884-16c-hy",
|
||||
choices=list(VAE_PATH),
|
||||
help="Name of the VAE model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=PRECISIONS,
|
||||
help="Precision mode for the VAE model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--vae-tiling",
|
||||
action="store_true",
|
||||
help="Enable tiling for the VAE model to save GPU memory.",
|
||||
)
|
||||
group.set_defaults(vae_tiling=True)
|
||||
|
||||
group.add_argument(
|
||||
"--text-encoder",
|
||||
type=str,
|
||||
default="llm",
|
||||
choices=list(TEXT_ENCODER_PATH),
|
||||
help="Name of the text encoder model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=PRECISIONS,
|
||||
help="Precision mode for the text encoder model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--text-states-dim",
|
||||
type=int,
|
||||
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(
|
||||
"--tokenizer",
|
||||
type=str,
|
||||
default="llm",
|
||||
choices=list(TOKENIZER_PATH),
|
||||
help="Name of the tokenizer model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--prompt-template",
|
||||
type=str,
|
||||
default="dit-llm-encode",
|
||||
choices=PROMPT_TEMPLATE,
|
||||
help="Image prompt template for the decoder-only text encoder model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--prompt-template-video",
|
||||
type=str,
|
||||
default="dit-llm-encode-video",
|
||||
choices=PROMPT_TEMPLATE,
|
||||
help="Video prompt template for the decoder-only text encoder model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--hidden-state-skip-layer",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Skip layer for hidden states.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--apply-final-norm",
|
||||
action="store_true",
|
||||
help=
|
||||
"Apply final normalization to the used text encoder hidden states.",
|
||||
)
|
||||
|
||||
# - CLIP
|
||||
group.add_argument(
|
||||
"--text-encoder-2",
|
||||
type=str,
|
||||
default="clipL",
|
||||
choices=list(TEXT_ENCODER_PATH),
|
||||
help="Name of the second text encoder model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=PRECISIONS,
|
||||
help="Precision mode for the second text encoder model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--text-states-dim-2",
|
||||
type=int,
|
||||
default=768,
|
||||
help="Dimension of the second text encoder hidden states.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--tokenizer-2",
|
||||
type=str,
|
||||
default="clipL",
|
||||
choices=list(TOKENIZER_PATH),
|
||||
help="Name of the second tokenizer model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--text-len-2",
|
||||
type=int,
|
||||
default=77,
|
||||
help="Maximum length of the second text input.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="Denoise schedule args")
|
||||
|
||||
group.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default="flow",
|
||||
help="Denoise type for noised inputs.",
|
||||
)
|
||||
|
||||
# Flow Matching
|
||||
group.add_argument(
|
||||
"--flow-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.",
|
||||
)
|
||||
group.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)",
|
||||
)
|
||||
group.add_argument(
|
||||
"--linear-schedule-end",
|
||||
type=int,
|
||||
default=25,
|
||||
help="End step for linear quadratic schedule 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-base",
|
||||
type=str,
|
||||
default="ckpts",
|
||||
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."
|
||||
"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`"
|
||||
"2) named `*_model_states.pt`, where * can be `mp_rank_00`.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--model-resolution",
|
||||
type=str,
|
||||
default="540p",
|
||||
choices=["540p", "720p"],
|
||||
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.",
|
||||
)
|
||||
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(
|
||||
"--disable-autocast",
|
||||
action="store_true",
|
||||
help=
|
||||
"Disable autocast for denoising loop and vae decoding in pipeline sampling.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--save-path",
|
||||
type=str,
|
||||
default="./results",
|
||||
help="Path to save the generated samples.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--save-path-suffix",
|
||||
type=str,
|
||||
default="",
|
||||
help="Suffix for the directory of saved 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(
|
||||
"--video-size",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=(720, 1280),
|
||||
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.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--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",
|
||||
)
|
||||
# --- prompt ---
|
||||
group.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Prompt for sampling during evaluation.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--seed-type",
|
||||
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 "
|
||||
"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.")
|
||||
|
||||
# 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(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=6.0,
|
||||
help="Embedded classifier free guidance scale.",
|
||||
)
|
||||
|
||||
group.add_argument(
|
||||
"--reproduce",
|
||||
action="store_true",
|
||||
help=
|
||||
"Enable reproducibility by setting random seeds and deterministic algorithms.",
|
||||
)
|
||||
|
||||
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=1,
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--ring-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
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}'."
|
||||
)
|
||||
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})."
|
||||
)
|
||||
return args
|
||||
@@ -0,0 +1,534 @@
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
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.diffusion.pipelines import HunyuanVideoPipeline
|
||||
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
|
||||
from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
|
||||
class Inference(object):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
vae,
|
||||
vae_kwargs,
|
||||
text_encoder,
|
||||
model,
|
||||
text_encoder_2=None,
|
||||
pipeline=None,
|
||||
use_cpu_offload=False,
|
||||
device=None,
|
||||
logger=None,
|
||||
parallel_args=None,
|
||||
):
|
||||
self.vae = vae
|
||||
self.vae_kwargs = vae_kwargs
|
||||
|
||||
self.text_encoder = text_encoder
|
||||
self.text_encoder_2 = text_encoder_2
|
||||
|
||||
self.model = model
|
||||
self.pipeline = pipeline
|
||||
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.logger = logger
|
||||
self.parallel_args = parallel_args
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls,
|
||||
pretrained_model_path,
|
||||
args,
|
||||
device=None,
|
||||
**kwargs):
|
||||
"""
|
||||
Initialize the Inference pipeline.
|
||||
|
||||
Args:
|
||||
pretrained_model_path (str or pathlib.Path): The model path, including t2v, text encoder and vae checkpoints.
|
||||
args (argparse.Namespace): The arguments for the pipeline.
|
||||
device (int): The device for inference. Default is 0.
|
||||
"""
|
||||
# ========================================================================
|
||||
logger.info(
|
||||
f"Got text-to-video model root path: {pretrained_model_path}")
|
||||
|
||||
# ==================== Initialize Distributed Environment ================
|
||||
if nccl_info.sp_size > 1:
|
||||
device = torch.device(f"cuda:{os.environ['LOCAL_RANK']}")
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
parallel_args = None # {"ulysses_degree": args.ulysses_degree, "ring_degree": args.ring_degree}
|
||||
|
||||
# ======================== Get the args path =============================
|
||||
|
||||
# Disable gradient
|
||||
torch.set_grad_enabled(False)
|
||||
|
||||
# =========================== Build main model ===========================
|
||||
logger.info("Building model...")
|
||||
factor_kwargs = {
|
||||
"device": device,
|
||||
"dtype": PRECISION_TO_TYPE[args.precision]
|
||||
}
|
||||
in_channels = args.latent_channels
|
||||
out_channels = args.latent_channels
|
||||
|
||||
model = load_model(
|
||||
args,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
factor_kwargs=factor_kwargs,
|
||||
)
|
||||
model = model.to(device)
|
||||
model = Inference.load_state_dict(args, model, pretrained_model_path)
|
||||
model.eval()
|
||||
|
||||
# ============================= Build extra models ========================
|
||||
# VAE
|
||||
vae, _, s_ratio, t_ratio = load_vae(
|
||||
args.vae,
|
||||
args.vae_precision,
|
||||
logger=logger,
|
||||
device=device if not args.use_cpu_offload else "cpu",
|
||||
)
|
||||
vae_kwargs = {"s_ratio": s_ratio, "t_ratio": t_ratio}
|
||||
|
||||
# Text encoder
|
||||
if args.prompt_template_video is not None:
|
||||
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)
|
||||
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_video
|
||||
prompt_template_video = (PROMPT_TEMPLATE[args.prompt_template_video]
|
||||
if args.prompt_template_video is not None else
|
||||
None)
|
||||
|
||||
text_encoder = TextEncoder(
|
||||
text_encoder_type=args.text_encoder,
|
||||
max_length=max_length,
|
||||
text_encoder_precision=args.text_encoder_precision,
|
||||
tokenizer_type=args.tokenizer,
|
||||
prompt_template=prompt_template,
|
||||
prompt_template_video=prompt_template_video,
|
||||
hidden_state_skip_layer=args.hidden_state_skip_layer,
|
||||
apply_final_norm=args.apply_final_norm,
|
||||
reproduce=args.reproduce,
|
||||
logger=logger,
|
||||
device=device if not args.use_cpu_offload else "cpu",
|
||||
)
|
||||
text_encoder_2 = None
|
||||
if args.text_encoder_2 is not None:
|
||||
text_encoder_2 = TextEncoder(
|
||||
text_encoder_type=args.text_encoder_2,
|
||||
max_length=args.text_len_2,
|
||||
text_encoder_precision=args.text_encoder_precision_2,
|
||||
tokenizer_type=args.tokenizer_2,
|
||||
reproduce=args.reproduce,
|
||||
logger=logger,
|
||||
device=device if not args.use_cpu_offload else "cpu",
|
||||
)
|
||||
|
||||
return cls(
|
||||
args=args,
|
||||
vae=vae,
|
||||
vae_kwargs=vae_kwargs,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_2=text_encoder_2,
|
||||
model=model,
|
||||
use_cpu_offload=args.use_cpu_offload,
|
||||
device=device,
|
||||
logger=logger,
|
||||
parallel_args=parallel_args,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def load_state_dict(args, model, pretrained_model_path):
|
||||
load_key = args.load_key
|
||||
dit_weight = Path(args.dit_weight)
|
||||
|
||||
if dit_weight is None:
|
||||
model_dir = pretrained_model_path / f"t2v_{args.model_resolution}"
|
||||
files = list(model_dir.glob("*.pt"))
|
||||
if len(files) == 0:
|
||||
raise ValueError(f"No model weights found in {model_dir}")
|
||||
if str(files[0]).startswith("pytorch_model_"):
|
||||
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")
|
||||
]
|
||||
model_path = files[0]
|
||||
if len(files) > 1:
|
||||
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."
|
||||
)
|
||||
else:
|
||||
if dit_weight.is_dir():
|
||||
files = list(dit_weight.glob("*.pt"))
|
||||
if len(files) == 0:
|
||||
raise ValueError(f"No model weights found in {dit_weight}")
|
||||
if str(files[0]).startswith("pytorch_model_"):
|
||||
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")
|
||||
]
|
||||
model_path = files[0]
|
||||
if len(files) > 1:
|
||||
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."
|
||||
)
|
||||
elif dit_weight.is_file():
|
||||
model_path = dit_weight
|
||||
bare_model = "unknown"
|
||||
else:
|
||||
raise ValueError(f"Invalid model path: {dit_weight}")
|
||||
|
||||
if not model_path.exists():
|
||||
raise ValueError(f"model_path not exists: {model_path}")
|
||||
logger.info(f"Loading torch model {model_path}...")
|
||||
if model_path.suffix == ".safetensors":
|
||||
# Use safetensors library for .safetensors files
|
||||
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)
|
||||
else:
|
||||
raise ValueError(f"Unsupported file format: {model_path}")
|
||||
|
||||
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())}.")
|
||||
model.load_state_dict(state_dict, strict=True)
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
def parse_size(size):
|
||||
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}.")
|
||||
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}.")
|
||||
return size
|
||||
|
||||
|
||||
class HunyuanVideoSampler(Inference):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
vae,
|
||||
vae_kwargs,
|
||||
text_encoder,
|
||||
model,
|
||||
text_encoder_2=None,
|
||||
pipeline=None,
|
||||
use_cpu_offload=False,
|
||||
device=0,
|
||||
logger=None,
|
||||
parallel_args=None,
|
||||
):
|
||||
super().__init__(
|
||||
args,
|
||||
vae,
|
||||
vae_kwargs,
|
||||
text_encoder,
|
||||
model,
|
||||
text_encoder_2=text_encoder_2,
|
||||
pipeline=pipeline,
|
||||
use_cpu_offload=use_cpu_offload,
|
||||
device=device,
|
||||
logger=logger,
|
||||
parallel_args=parallel_args,
|
||||
)
|
||||
|
||||
self.pipeline = self.load_diffusion_pipeline(
|
||||
args=args,
|
||||
vae=self.vae,
|
||||
text_encoder=self.text_encoder,
|
||||
text_encoder_2=self.text_encoder_2,
|
||||
model=self.model,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
self.default_negative_prompt = NEGATIVE_PROMPT
|
||||
|
||||
def load_diffusion_pipeline(
|
||||
self,
|
||||
args,
|
||||
vae,
|
||||
text_encoder,
|
||||
text_encoder_2,
|
||||
model,
|
||||
scheduler=None,
|
||||
device=None,
|
||||
progress_bar_config=None,
|
||||
data_type="video",
|
||||
):
|
||||
"""Load the denoising scheduler for inference."""
|
||||
if scheduler is None:
|
||||
if args.denoise_type == "flow":
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=args.flow_shift,
|
||||
reverse=args.flow_reverse,
|
||||
solver=args.flow_solver,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid denoise type {args.denoise_type}")
|
||||
|
||||
pipeline = HunyuanVideoPipeline(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_2=text_encoder_2,
|
||||
transformer=model,
|
||||
scheduler=scheduler,
|
||||
progress_bar_config=progress_bar_config,
|
||||
args=args,
|
||||
)
|
||||
if self.use_cpu_offload:
|
||||
pipeline.enable_sequential_cpu_offload()
|
||||
else:
|
||||
pipeline = pipeline.to(device)
|
||||
|
||||
return pipeline
|
||||
|
||||
@torch.no_grad()
|
||||
def predict(
|
||||
self,
|
||||
prompt,
|
||||
height=192,
|
||||
width=336,
|
||||
video_length=129,
|
||||
seed=None,
|
||||
negative_prompt=None,
|
||||
infer_steps=50,
|
||||
guidance_scale=6,
|
||||
flow_shift=5.0,
|
||||
embedded_guidance_scale=None,
|
||||
batch_size=1,
|
||||
num_videos_per_prompt=1,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Predict the image/video from the given text.
|
||||
|
||||
Args:
|
||||
prompt (str or List[str]): The input text.
|
||||
kwargs:
|
||||
height (int): The height of the output video. Default is 192.
|
||||
width (int): The width of the output video. Default is 336.
|
||||
video_length (int): The frame number of the output video. Default is 129.
|
||||
seed (int or List[str]): The random seed for the generation. Default is a random integer.
|
||||
negative_prompt (str or List[str]): The negative text prompt. Default is an empty string.
|
||||
guidance_scale (float): The guidance scale for the generation. Default is 6.0.
|
||||
num_images_per_prompt (int): The number of images per prompt. Default is 1.
|
||||
infer_steps (int): The number of inference steps. Default is 100.
|
||||
"""
|
||||
|
||||
out_dict = dict()
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: seed
|
||||
# ========================================================================
|
||||
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)
|
||||
]
|
||||
elif isinstance(seed, int):
|
||||
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)
|
||||
]
|
||||
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}."
|
||||
)
|
||||
else:
|
||||
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
|
||||
]
|
||||
out_dict["seeds"] = seeds
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
# ========================================================================
|
||||
if width <= 0 or height <= 0 or video_length <= 0:
|
||||
raise ValueError(
|
||||
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}"
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"Input (height, width, video_length) = ({height}, {width}, {video_length})"
|
||||
)
|
||||
|
||||
target_height = align_to(height, 16)
|
||||
target_width = align_to(width, 16)
|
||||
target_video_length = video_length
|
||||
|
||||
out_dict["size"] = (target_height, target_width, target_video_length)
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: prompt, new_prompt, negative_prompt
|
||||
# ========================================================================
|
||||
if not isinstance(prompt, str):
|
||||
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)}"
|
||||
)
|
||||
negative_prompt = [negative_prompt.strip()]
|
||||
|
||||
# ========================================================================
|
||||
# Scheduler
|
||||
# ========================================================================
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=flow_shift,
|
||||
reverse=self.args.flow_reverse,
|
||||
solver=self.args.flow_solver,
|
||||
)
|
||||
self.pipeline.scheduler = scheduler
|
||||
|
||||
if "884" in self.args.vae:
|
||||
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]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
# ========================================================================
|
||||
# Print infer args
|
||||
# ========================================================================
|
||||
debug_str = f"""
|
||||
height: {target_height}
|
||||
width: {target_width}
|
||||
video_length: {target_video_length}
|
||||
prompt: {prompt}
|
||||
neg_prompt: {negative_prompt}
|
||||
seed: {seed}
|
||||
infer_steps: {infer_steps}
|
||||
num_videos_per_prompt: {num_videos_per_prompt}
|
||||
guidance_scale: {guidance_scale}
|
||||
n_tokens: {n_tokens}
|
||||
flow_shift: {flow_shift}
|
||||
embedded_guidance_scale: {embedded_guidance_scale}"""
|
||||
logger.debug(debug_str)
|
||||
|
||||
# ========================================================================
|
||||
# Pipeline inference
|
||||
# ========================================================================
|
||||
start_time = time.time()
|
||||
samples = self.pipeline(
|
||||
prompt=prompt,
|
||||
height=target_height,
|
||||
width=target_width,
|
||||
video_length=target_video_length,
|
||||
num_inference_steps=infer_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
negative_prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
generator=generator,
|
||||
output_type="pil",
|
||||
n_tokens=n_tokens,
|
||||
embedded_guidance_scale=embedded_guidance_scale,
|
||||
data_type="video" if target_video_length > 1 else "image",
|
||||
is_progress_bar=True,
|
||||
vae_ver=self.args.vae,
|
||||
enable_tiling=self.args.vae_tiling,
|
||||
enable_vae_sp=self.args.vae_sp,
|
||||
)[0]
|
||||
out_dict["samples"] = samples
|
||||
out_dict["prompts"] = prompt
|
||||
|
||||
gen_time = time.time() - start_time
|
||||
logger.info(f"Success, time: {gen_time}")
|
||||
|
||||
return out_dict
|
||||
@@ -0,0 +1,25 @@
|
||||
from .models import HUNYUAN_VIDEO_CONFIG, HYVideoDiffusionTransformer
|
||||
|
||||
|
||||
def load_model(args, in_channels, out_channels, factor_kwargs):
|
||||
"""load hunyuan video model
|
||||
|
||||
Args:
|
||||
args (dict): model args
|
||||
in_channels (int): input channels number
|
||||
out_channels (int): output channels number
|
||||
factor_kwargs (dict): factor kwargs
|
||||
|
||||
Returns:
|
||||
model (nn.Module): The hunyuan video model
|
||||
"""
|
||||
if args.model in HUNYUAN_VIDEO_CONFIG.keys():
|
||||
model = HYVideoDiffusionTransformer(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
**HUNYUAN_VIDEO_CONFIG[args.model],
|
||||
**factor_kwargs,
|
||||
)
|
||||
return model
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
@@ -0,0 +1,23 @@
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def get_activation_layer(act_type):
|
||||
"""get activation layer
|
||||
|
||||
Args:
|
||||
act_type (str): the activation type
|
||||
|
||||
Returns:
|
||||
torch.nn.functional: the activation layer
|
||||
"""
|
||||
if act_type == "gelu":
|
||||
return lambda: nn.GELU()
|
||||
elif act_type == "gelu_tanh":
|
||||
# Approximate `tanh` requires torch >= 1.13
|
||||
return lambda: nn.GELU(approximate="tanh")
|
||||
elif act_type == "relu":
|
||||
return nn.ReLU
|
||||
elif act_type == "silu":
|
||||
return nn.SiLU
|
||||
else:
|
||||
raise ValueError(f"Unknown activation type: {act_type}")
|
||||
@@ -0,0 +1,90 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
drop_rate=0,
|
||||
attn_mask=None,
|
||||
causal=False,
|
||||
):
|
||||
|
||||
qkv = torch.stack([q, k, v], dim=2)
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
query, encoder_query = q
|
||||
key, encoder_key = k
|
||||
value, encoder_value = v
|
||||
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)
|
||||
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
|
||||
value = all_to_all_4D(value, scatter_dim=2, gather_dim=1)
|
||||
|
||||
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)
|
||||
|
||||
encoder_query = shrink_head(encoder_query, dim=2)
|
||||
encoder_key = shrink_head(encoder_key, dim=2)
|
||||
encoder_value = shrink_head(encoder_value, dim=2)
|
||||
# [b, s, h, d]
|
||||
|
||||
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)
|
||||
|
||||
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 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 = hidden_states.to(query.dtype)
|
||||
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
|
||||
|
||||
attn = torch.cat([hidden_states, encoder_hidden_states], dim=1)
|
||||
|
||||
b, s, a, d = attn.shape
|
||||
attn = attn.reshape(b, s, -1)
|
||||
|
||||
return attn
|
||||
@@ -0,0 +1,163 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from ..utils.helpers import to_2tuple
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""2D Image to Patch Embedding
|
||||
|
||||
Image to Patch Embedding using Conv2d
|
||||
|
||||
A convolution based approach to patchifying a 2D image w/ embedding projection.
|
||||
|
||||
Based on the impl in https://github.com/google-research/vision_transformer
|
||||
|
||||
Hacked together by / Copyright 2020 Ross Wightman
|
||||
|
||||
Remove the _assert function in forward function to be compatible with multi-resolution images.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
dtype=None,
|
||||
device=None,
|
||||
):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
patch_size = to_2tuple(patch_size)
|
||||
self.patch_size = patch_size
|
||||
self.flatten = flatten
|
||||
|
||||
self.proj = nn.Conv3d(
|
||||
in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias,
|
||||
**factory_kwargs,
|
||||
)
|
||||
nn.init.xavier_uniform_(
|
||||
self.proj.weight.view(self.proj.weight.size(0), -1))
|
||||
if bias:
|
||||
nn.init.zeros_(self.proj.bias)
|
||||
|
||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
|
||||
class TextProjection(nn.Module):
|
||||
"""
|
||||
Projects text 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_channels,
|
||||
hidden_size,
|
||||
act_layer,
|
||||
dtype=None,
|
||||
device=None):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
self.linear_1 = nn.Linear(
|
||||
in_features=in_channels,
|
||||
out_features=hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
)
|
||||
self.act_1 = act_layer()
|
||||
self.linear_2 = nn.Linear(
|
||||
in_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
|
||||
Args:
|
||||
t (torch.Tensor): a 1-D Tensor of N indices, one per batch element. These may be fractional.
|
||||
dim (int): the dimension of the output.
|
||||
max_period (int): controls the minimum frequency of the embeddings.
|
||||
|
||||
Returns:
|
||||
embedding (torch.Tensor): An (N, D) Tensor of positional embeddings.
|
||||
|
||||
.. 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) /
|
||||
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)
|
||||
return embedding
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
act_layer,
|
||||
frequency_embedding_size=256,
|
||||
max_period=10000,
|
||||
out_size=None,
|
||||
dtype=None,
|
||||
device=None,
|
||||
):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.max_period = max_period
|
||||
if out_size is None:
|
||||
out_size = hidden_size
|
||||
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs),
|
||||
act_layer(),
|
||||
nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),
|
||||
)
|
||||
nn.init.normal_(self.mlp[0].weight, std=0.02)
|
||||
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_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
@@ -0,0 +1,133 @@
|
||||
# Modified from timm library:
|
||||
# https://github.com/huggingface/pytorch-image-models/blob/648aaa41233ba83eb38faf5ba9d415d574823241/timm/layers/mlp.py#L13
|
||||
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from ..utils.helpers import to_2tuple
|
||||
from .modulate_layers import modulate
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
"""MLP as used in Vision Transformer, MLP-Mixer and related networks"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
hidden_channels=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
norm_layer=None,
|
||||
bias=True,
|
||||
drop=0.0,
|
||||
use_conv=False,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
out_features = out_features or in_channels
|
||||
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
|
||||
|
||||
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.drop2 = nn.Dropout(drop_probs[1])
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop1(x)
|
||||
x = self.norm(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop2(x)
|
||||
return x
|
||||
|
||||
|
||||
#
|
||||
class MLPEmbedder(nn.Module):
|
||||
"""copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py"""
|
||||
|
||||
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.silu = nn.SiLU()
|
||||
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)))
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""The final layer of DiT."""
|
||||
|
||||
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)
|
||||
if isinstance(patch_size, int):
|
||||
self.linear = nn.Linear(
|
||||
hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
)
|
||||
else:
|
||||
self.linear = nn.Linear(
|
||||
hidden_size,
|
||||
patch_size[0] * patch_size[1] * patch_size[2] * out_channels,
|
||||
bias=True,
|
||||
)
|
||||
nn.init.zeros_(self.linear.weight)
|
||||
nn.init.zeros_(self.linear.bias)
|
||||
|
||||
# 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),
|
||||
)
|
||||
# Zero-initialize the modulation
|
||||
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
||||
nn.init.zeros_(self.adaLN_modulation[1].bias)
|
||||
|
||||
def forward(self, x, c):
|
||||
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift=shift, scale=scale)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
@@ -0,0 +1,750 @@
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
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.utils.parallel_states import nccl_info
|
||||
|
||||
from .activation_layers import get_activation_layer
|
||||
from .attenion import parallel_attention
|
||||
from .embed_layers import PatchEmbed, TextProjection, TimestepEmbedder
|
||||
from .mlp_layers import MLP, FinalLayer, MLPEmbedder
|
||||
from .modulate_layers import ModulateDiT, apply_gate, modulate
|
||||
from .norm_layers import get_norm_layer
|
||||
from .posemb_layers import apply_rotary_emb
|
||||
from .token_refiner import SingleTokenRefiner
|
||||
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
A multimodal dit block with separate modulation for
|
||||
text and image/video, see more details (SD3): https://arxiv.org/abs/2403.03206
|
||||
(Flux.1): https://github.com/black-forest-labs/flux
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
heads_num: int,
|
||||
mlp_width_ratio: float,
|
||||
mlp_act_type: str = "gelu_tanh",
|
||||
qk_norm: bool = True,
|
||||
qk_norm_type: str = "rms",
|
||||
qkv_bias: bool = False,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.heads_num = heads_num
|
||||
head_dim = hidden_size // heads_num
|
||||
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
|
||||
|
||||
self.img_mod = ModulateDiT(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
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_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)
|
||||
if qk_norm else nn.Identity())
|
||||
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_norm2 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
self.img_mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
act_layer=get_activation_layer(mlp_act_type),
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
)
|
||||
|
||||
self.txt_mod = ModulateDiT(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
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_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)
|
||||
if qk_norm else nn.Identity())
|
||||
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_mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
act_layer=get_activation_layer(mlp_act_type),
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
)
|
||||
self.hybrid_seq_parallel_attn = None
|
||||
|
||||
def enable_deterministic(self):
|
||||
self.deterministic = True
|
||||
|
||||
def disable_deterministic(self):
|
||||
self.deterministic = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
txt: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
(
|
||||
img_mod1_shift,
|
||||
img_mod1_scale,
|
||||
img_mod1_gate,
|
||||
img_mod2_shift,
|
||||
img_mod2_scale,
|
||||
img_mod2_gate,
|
||||
) = self.img_mod(vec).chunk(6, dim=-1)
|
||||
(
|
||||
txt_mod1_shift,
|
||||
txt_mod1_scale,
|
||||
txt_mod1_gate,
|
||||
txt_mod2_shift,
|
||||
txt_mod2_scale,
|
||||
txt_mod2_gate,
|
||||
) = self.txt_mod(vec).chunk(6, dim=-1)
|
||||
|
||||
# Prepare image for attention.
|
||||
img_modulated = self.img_norm1(img)
|
||||
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)
|
||||
# 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)
|
||||
|
||||
# Apply RoPE if needed.
|
||||
if freqs_cis is not None:
|
||||
|
||||
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)
|
||||
|
||||
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_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_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)
|
||||
# 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)
|
||||
|
||||
attn = parallel_attention(
|
||||
(img_q, txt_q),
|
||||
(img_k, txt_k),
|
||||
(img_v, txt_v),
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
)
|
||||
|
||||
# attention computation end
|
||||
|
||||
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_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_mlp(
|
||||
modulate(self.txt_norm2(txt),
|
||||
shift=txt_mod2_shift,
|
||||
scale=txt_mod2_scale)),
|
||||
gate=txt_mod2_gate,
|
||||
)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class MMSingleStreamBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with parallel linear layers as described in
|
||||
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
|
||||
Also refer to (SD3): https://arxiv.org/abs/2403.03206
|
||||
(Flux.1): https://github.com/black-forest-labs/flux
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
heads_num: int,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
mlp_act_type: str = "gelu_tanh",
|
||||
qk_norm: bool = True,
|
||||
qk_norm_type: str = "rms",
|
||||
qk_scale: float = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.hidden_size = hidden_size
|
||||
self.heads_num = heads_num
|
||||
head_dim = hidden_size // heads_num
|
||||
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
|
||||
self.mlp_hidden_dim = mlp_hidden_dim
|
||||
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)
|
||||
# proj and mlp_out
|
||||
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)
|
||||
if qk_norm else nn.Identity())
|
||||
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.mlp_act = get_activation_layer(mlp_act_type)()
|
||||
self.modulation = ModulateDiT(
|
||||
hidden_size,
|
||||
factor=3,
|
||||
act_layer=get_activation_layer("silu"),
|
||||
**factory_kwargs,
|
||||
)
|
||||
self.hybrid_seq_parallel_attn = None
|
||||
|
||||
def enable_deterministic(self):
|
||||
self.deterministic = True
|
||||
|
||||
def disable_deterministic(self):
|
||||
self.deterministic = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
txt_len: int,
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
text_mask: torch.Tensor = 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)
|
||||
|
||||
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)
|
||||
k = self.k_norm(k).to(v)
|
||||
|
||||
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)
|
||||
|
||||
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_q, img_k = img_qq, img_kk
|
||||
|
||||
attn = parallel_attention(
|
||||
(img_q, txt_q),
|
||||
(img_k, txt_k),
|
||||
(img_v, txt_v),
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
)
|
||||
|
||||
# attention computation end
|
||||
|
||||
# Compute activation in mlp stream, cat again and run second linear layer.
|
||||
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
||||
return x + apply_gate(output, gate=mod_gate)
|
||||
|
||||
|
||||
class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
"""
|
||||
HunyuanVideo Transformer backbone
|
||||
|
||||
Inherited from ModelMixin and ConfigMixin for compatibility with diffusers' sampler StableDiffusionPipeline.
|
||||
|
||||
Reference:
|
||||
[1] Flux.1: https://github.com/black-forest-labs/flux
|
||||
[2] MMDiT: http://arxiv.org/abs/2403.03206
|
||||
|
||||
Parameters
|
||||
----------
|
||||
args: argparse.Namespace
|
||||
The arguments parsed by argparse.
|
||||
patch_size: list
|
||||
The size of the patch.
|
||||
in_channels: int
|
||||
The number of input channels.
|
||||
out_channels: int
|
||||
The number of output channels.
|
||||
hidden_size: int
|
||||
The hidden size of the transformer backbone.
|
||||
heads_num: int
|
||||
The number of attention heads.
|
||||
mlp_width_ratio: float
|
||||
The ratio of the hidden size of the MLP in the transformer block.
|
||||
mlp_act_type: str
|
||||
The activation function of the MLP in the transformer block.
|
||||
depth_double_blocks: int
|
||||
The number of transformer blocks in the double blocks.
|
||||
depth_single_blocks: int
|
||||
The number of transformer blocks in the single blocks.
|
||||
rope_dim_list: list
|
||||
The dimension of the rotary embedding for t, h, w.
|
||||
qkv_bias: bool
|
||||
Whether to use bias in the qkv linear layer.
|
||||
qk_norm: bool
|
||||
Whether to use qk norm.
|
||||
qk_norm_type: str
|
||||
The type of qk norm.
|
||||
guidance_embed: bool
|
||||
Whether to use guidance embedding for distillation.
|
||||
text_projection: str
|
||||
The type of the text projection, default is single_refiner.
|
||||
use_attention_mask: bool
|
||||
Whether to use attention mask for text encoder.
|
||||
dtype: torch.dtype
|
||||
The dtype of the model.
|
||||
device: torch.device
|
||||
The device of the model.
|
||||
"""
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: list = [1, 2, 2],
|
||||
in_channels: int = 4, # Should be VAE.config.latent_channels.
|
||||
out_channels: int = None,
|
||||
hidden_size: int = 3072,
|
||||
heads_num: int = 24,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
mlp_act_type: str = "gelu_tanh",
|
||||
mm_double_blocks_depth: int = 20,
|
||||
mm_single_blocks_depth: int = 40,
|
||||
rope_dim_list: List[int] = [16, 56, 56],
|
||||
qkv_bias: bool = True,
|
||||
qk_norm: bool = True,
|
||||
qk_norm_type: str = "rms",
|
||||
guidance_embed: bool = False, # For modulation.
|
||||
text_projection: str = "single_refiner",
|
||||
use_attention_mask: bool = True,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
text_states_dim: int = 4096,
|
||||
text_states_dim_2: int = 768,
|
||||
rope_theta: int = 256,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels if out_channels is None else out_channels
|
||||
self.unpatchify_channels = self.out_channels
|
||||
self.guidance_embed = guidance_embed
|
||||
self.rope_dim_list = rope_dim_list
|
||||
self.rope_theta = rope_theta
|
||||
# Text projection. Default to linear projection.
|
||||
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
|
||||
self.use_attention_mask = use_attention_mask
|
||||
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}"
|
||||
)
|
||||
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}")
|
||||
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)
|
||||
|
||||
# text projection
|
||||
if self.text_projection == "linear":
|
||||
self.txt_in = TextProjection(
|
||||
self.config.text_states_dim,
|
||||
self.hidden_size,
|
||||
get_activation_layer("silu"),
|
||||
**factory_kwargs,
|
||||
)
|
||||
elif self.text_projection == "single_refiner":
|
||||
self.txt_in = SingleTokenRefiner(
|
||||
self.config.text_states_dim,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
depth=2,
|
||||
**factory_kwargs,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"Unsupported text_projection: {self.text_projection}")
|
||||
|
||||
# time modulation
|
||||
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)
|
||||
|
||||
# guidance modulation
|
||||
self.guidance_in = (TimestepEmbedder(
|
||||
self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
|
||||
if guidance_embed else None)
|
||||
|
||||
# double blocks
|
||||
self.double_blocks = nn.ModuleList([
|
||||
MMDoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_act_type=mlp_act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
) for _ in range(mm_double_blocks_depth)
|
||||
])
|
||||
|
||||
# single blocks
|
||||
self.single_blocks = nn.ModuleList([
|
||||
MMSingleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_act_type=mlp_act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
**factory_kwargs,
|
||||
) for _ in range(mm_single_blocks_depth)
|
||||
])
|
||||
|
||||
self.final_layer = FinalLayer(
|
||||
self.hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
get_activation_layer("silu"),
|
||||
**factory_kwargs,
|
||||
)
|
||||
|
||||
def enable_deterministic(self):
|
||||
for block in self.double_blocks:
|
||||
block.enable_deterministic()
|
||||
for block in self.single_blocks:
|
||||
block.enable_deterministic()
|
||||
|
||||
def disable_deterministic(self):
|
||||
for block in self.double_blocks:
|
||||
block.disable_deterministic()
|
||||
for block in self.single_blocks:
|
||||
block.disable_deterministic()
|
||||
|
||||
def get_rotary_pos_embed(self, rope_sizes):
|
||||
target_ndim = 3
|
||||
|
||||
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"
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
theta=self.rope_theta,
|
||||
use_real=True,
|
||||
theta_rescale_factor=1,
|
||||
)
|
||||
return freqs_cos, freqs_sin
|
||||
# x: torch.Tensor,
|
||||
# t: torch.Tensor, # Should be in range(0, 1000).
|
||||
# text_states: torch.Tensor = None,
|
||||
# text_mask: torch.Tensor = None, # Now we don't use it.
|
||||
# text_states_2: Optional[torch.Tensor] = None, # Text embedding for modulation.
|
||||
# guidance: torch.Tensor = None, # Guidance for modulation, should be cfg_scale x 1000.
|
||||
# return_dict: bool = True,
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
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
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
for _, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask]
|
||||
|
||||
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):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
]
|
||||
|
||||
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 unpatchify(self, x, t, h, w):
|
||||
"""
|
||||
x: (N, T, patch_size**2 * C)
|
||||
imgs: (N, H, W, C)
|
||||
"""
|
||||
c = self.unpatchify_channels
|
||||
pt, ph, pw = self.patch_size
|
||||
assert t * h * w == x.shape[1]
|
||||
|
||||
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
|
||||
x = torch.einsum("nthwcopq->nctohpwq", x)
|
||||
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
|
||||
|
||||
return imgs
|
||||
|
||||
def params_count(self):
|
||||
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())
|
||||
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())
|
||||
for block in self.single_blocks
|
||||
]),
|
||||
"total":
|
||||
sum(p.numel() for p in self.parameters()),
|
||||
}
|
||||
counts["attn+mlp"] = counts["double"] + counts["single"]
|
||||
return counts
|
||||
|
||||
|
||||
#################################################################################
|
||||
# HunyuanVideo Configs #
|
||||
#################################################################################
|
||||
|
||||
HUNYUAN_VIDEO_CONFIG = {
|
||||
"HYVideo-T/2": {
|
||||
"mm_double_blocks_depth": 20,
|
||||
"mm_single_blocks_depth": 40,
|
||||
"rope_dim_list": [16, 56, 56],
|
||||
"hidden_size": 3072,
|
||||
"heads_num": 24,
|
||||
"mlp_width_ratio": 4,
|
||||
},
|
||||
"HYVideo-T/2-cfgdistill": {
|
||||
"mm_double_blocks_depth": 20,
|
||||
"mm_single_blocks_depth": 40,
|
||||
"rope_dim_list": [16, 56, 56],
|
||||
"hidden_size": 3072,
|
||||
"heads_num": 24,
|
||||
"mlp_width_ratio": 4,
|
||||
"guidance_embed": True,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class ModulateDiT(nn.Module):
|
||||
"""Modulation layer for DiT."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
factor: int,
|
||||
act_layer: Callable,
|
||||
dtype=None,
|
||||
device=None,
|
||||
):
|
||||
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)
|
||||
# Zero-initialize the modulation
|
||||
nn.init.zeros_(self.linear.weight)
|
||||
nn.init.zeros_(self.linear.bias)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.linear(self.act(x))
|
||||
|
||||
|
||||
def modulate(x, shift=None, scale=None):
|
||||
"""modulate by shift and scale
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): input tensor.
|
||||
shift (torch.Tensor, optional): shift tensor. Defaults to None.
|
||||
scale (torch.Tensor, optional): scale tensor. Defaults to None.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: the output tensor after modulate.
|
||||
"""
|
||||
if scale is None and shift is None:
|
||||
return x
|
||||
elif shift is None:
|
||||
return x * (1 + scale.unsqueeze(1))
|
||||
elif scale is None:
|
||||
return x + shift.unsqueeze(1)
|
||||
else:
|
||||
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
|
||||
|
||||
def apply_gate(x, gate=None, tanh=False):
|
||||
"""AI is creating summary for apply_gate
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): input tensor.
|
||||
gate (torch.Tensor, optional): gate tensor. Defaults to None.
|
||||
tanh (bool, optional): whether to use tanh function. Defaults to False.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: the output tensor after apply gate.
|
||||
"""
|
||||
if gate is None:
|
||||
return x
|
||||
if tanh:
|
||||
return x * gate.unsqueeze(1).tanh()
|
||||
else:
|
||||
return x * gate.unsqueeze(1)
|
||||
|
||||
|
||||
def ckpt_wrapper(module):
|
||||
|
||||
def ckpt_forward(*inputs):
|
||||
outputs = module(*inputs)
|
||||
return outputs
|
||||
|
||||
return ckpt_forward
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def get_norm_layer(norm_layer):
|
||||
"""
|
||||
Get the normalization layer.
|
||||
|
||||
Args:
|
||||
norm_layer (str): The type of normalization layer.
|
||||
|
||||
Returns:
|
||||
norm_layer (nn.Module): The normalization layer.
|
||||
"""
|
||||
if norm_layer == "layer":
|
||||
return nn.LayerNorm
|
||||
elif norm_layer == "rms":
|
||||
return RMSNorm
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"Norm layer {norm_layer} is not implemented")
|
||||
@@ -0,0 +1,79 @@
|
||||
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
|
||||
|
||||
|
||||
def get_norm_layer(norm_layer):
|
||||
"""
|
||||
Get the normalization layer.
|
||||
|
||||
Args:
|
||||
norm_layer (str): The type of normalization layer.
|
||||
|
||||
Returns:
|
||||
norm_layer (nn.Module): The normalization layer.
|
||||
"""
|
||||
if norm_layer == "layer":
|
||||
return nn.LayerNorm
|
||||
elif norm_layer == "rms":
|
||||
return RMSNorm
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"Norm layer {norm_layer} is not implemented")
|
||||
@@ -0,0 +1,314 @@
|
||||
from typing import List, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _to_tuple(x, dim=2):
|
||||
if isinstance(x, int):
|
||||
return (x, ) * dim
|
||||
elif len(x) == dim:
|
||||
return x
|
||||
else:
|
||||
raise ValueError(f"Expected length {dim} or int, but got {x}")
|
||||
|
||||
|
||||
def get_meshgrid_nd(start, *args, dim=2):
|
||||
"""
|
||||
Get n-D meshgrid with start, stop and num.
|
||||
|
||||
Args:
|
||||
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
|
||||
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
|
||||
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
|
||||
n-tuples.
|
||||
*args: See above.
|
||||
dim (int): Dimension of the meshgrid. Defaults to 2.
|
||||
|
||||
Returns:
|
||||
grid (np.ndarray): [dim, ...]
|
||||
"""
|
||||
if len(args) == 0:
|
||||
# start is grid_size
|
||||
num = _to_tuple(start, dim=dim)
|
||||
start = (0, ) * dim
|
||||
stop = num
|
||||
elif len(args) == 1:
|
||||
# start is start, args[0] is stop, step is 1
|
||||
start = _to_tuple(start, dim=dim)
|
||||
stop = _to_tuple(args[0], dim=dim)
|
||||
num = [stop[i] - start[i] for i in range(dim)]
|
||||
elif len(args) == 2:
|
||||
# start is start, args[0] is stop, args[1] is num
|
||||
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
|
||||
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
|
||||
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
|
||||
else:
|
||||
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
|
||||
|
||||
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
|
||||
axis_grid = []
|
||||
for i in range(dim):
|
||||
a, b, n = start[i], stop[i], num[i]
|
||||
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
|
||||
axis_grid.append(g)
|
||||
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
|
||||
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
|
||||
|
||||
return grid
|
||||
|
||||
|
||||
#################################################################################
|
||||
# Rotary Positional Embedding Functions #
|
||||
#################################################################################
|
||||
# https://github.com/meta-llama/llama/blob/be327c427cc5e89cc1d3ab3d3fec4484df771245/llama/model.py#L80
|
||||
|
||||
|
||||
def reshape_for_broadcast(
|
||||
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
|
||||
x: torch.Tensor,
|
||||
head_first=False,
|
||||
):
|
||||
"""
|
||||
Reshape frequency tensor for broadcasting it with another tensor.
|
||||
|
||||
This function reshapes the frequency tensor to have the same shape as the target tensor 'x'
|
||||
for the purpose of broadcasting the frequency tensor during element-wise operations.
|
||||
|
||||
Notes:
|
||||
When using FlashMHAModified, head_first should be False.
|
||||
When using Attention, head_first should be True.
|
||||
|
||||
Args:
|
||||
freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped.
|
||||
x (torch.Tensor): Target tensor for broadcasting compatibility.
|
||||
head_first (bool): head dimension first (except batch dim) or not.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Reshaped frequency tensor.
|
||||
|
||||
Raises:
|
||||
AssertionError: If the frequency tensor doesn't match the expected shape.
|
||||
AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions.
|
||||
"""
|
||||
ndim = x.ndim
|
||||
assert 0 <= 1 < ndim
|
||||
|
||||
if isinstance(freqs_cis, tuple):
|
||||
# freqs_cis: (cos, sin) in real space
|
||||
if head_first:
|
||||
assert freqs_cis[0].shape == (
|
||||
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)
|
||||
]
|
||||
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)
|
||||
]
|
||||
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
|
||||
else:
|
||||
# freqs_cis: values in complex space
|
||||
if head_first:
|
||||
assert freqs_cis.shape == (
|
||||
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)
|
||||
]
|
||||
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)
|
||||
]
|
||||
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]
|
||||
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
||||
|
||||
|
||||
def apply_rotary_emb(
|
||||
xq: torch.Tensor,
|
||||
xk: torch.Tensor,
|
||||
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],
|
||||
head_first: bool = False,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply rotary embeddings to input tensors using the given frequency tensor.
|
||||
|
||||
This function applies rotary embeddings to the given query 'xq' and key 'xk' tensors using the provided
|
||||
frequency tensor 'freqs_cis'. The input tensors are reshaped as complex numbers, and the frequency tensor
|
||||
is reshaped for broadcasting compatibility. The resulting tensors contain rotary embeddings and are
|
||||
returned as real tensors.
|
||||
|
||||
Args:
|
||||
xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D]
|
||||
xk (torch.Tensor): Key tensor to apply rotary embeddings. [B, S, H, D]
|
||||
freqs_cis (torch.Tensor or tuple): Precomputed frequency tensor for complex exponential.
|
||||
head_first (bool): head dimension first (except batch dim) or not.
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
||||
|
||||
"""
|
||||
xk_out = None
|
||||
if isinstance(freqs_cis, tuple):
|
||||
cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D]
|
||||
cos, sin = cos.to(xq.device), sin.to(xq.device)
|
||||
# real * cos - imag * sin
|
||||
# imag * cos + real * sin
|
||||
xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq)
|
||||
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]
|
||||
# (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_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
|
||||
|
||||
return xq_out, xk_out
|
||||
|
||||
|
||||
def get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
start,
|
||||
*args,
|
||||
theta=10000.0,
|
||||
use_real=False,
|
||||
theta_rescale_factor: Union[float, List[float]] = 1.0,
|
||||
interpolation_factor: Union[float, List[float]] = 1.0,
|
||||
):
|
||||
"""
|
||||
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
||||
|
||||
Args:
|
||||
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
|
||||
sum(rope_dim_list) should equal to head_dim of attention layer.
|
||||
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
|
||||
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
|
||||
*args: See above.
|
||||
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
|
||||
use_real (bool): If True, return real part and imaginary part separately. Otherwise, return complex numbers.
|
||||
Some libraries such as TensorRT does not support complex64 data type. So it is useful to provide a real
|
||||
part and an imaginary part separately.
|
||||
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
|
||||
|
||||
Returns:
|
||||
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]
|
||||
|
||||
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:
|
||||
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)"
|
||||
|
||||
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:
|
||||
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)"
|
||||
|
||||
# use 1/ndim of dimensions to encode grid_axis
|
||||
embs = []
|
||||
for i in range(len(rope_dim_list)):
|
||||
emb = get_1d_rotary_pos_embed(
|
||||
rope_dim_list[i],
|
||||
grid[i].reshape(-1),
|
||||
theta,
|
||||
use_real=use_real,
|
||||
theta_rescale_factor=theta_rescale_factor[i],
|
||||
interpolation_factor=interpolation_factor[i],
|
||||
) # 2 x [WHD, rope_dim_list[i]]
|
||||
embs.append(emb)
|
||||
|
||||
if use_real:
|
||||
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
|
||||
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
|
||||
return cos, sin
|
||||
else:
|
||||
emb = torch.cat(embs, dim=1) # (WHD, D/2)
|
||||
return emb
|
||||
|
||||
|
||||
def get_1d_rotary_pos_embed(
|
||||
dim: int,
|
||||
pos: Union[torch.FloatTensor, int],
|
||||
theta: float = 10000.0,
|
||||
use_real: bool = False,
|
||||
theta_rescale_factor: float = 1.0,
|
||||
interpolation_factor: float = 1.0,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
|
||||
|
||||
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
|
||||
and the end index 'end'. The 'theta' parameter scales the frequencies.
|
||||
The returned tensor contains complex values in complex64 data type.
|
||||
|
||||
Args:
|
||||
dim (int): Dimension of the frequency tensor.
|
||||
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
|
||||
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
|
||||
use_real (bool, optional): If True, return real part and imaginary part separately.
|
||||
Otherwise, return complex numbers.
|
||||
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
|
||||
|
||||
Returns:
|
||||
freqs_cis: Precomputed frequency tensor with complex exponential. [S, D/2]
|
||||
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
|
||||
"""
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos).float()
|
||||
|
||||
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
||||
# has some connection to NTK literature
|
||||
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]
|
||||
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
if use_real:
|
||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
|
||||
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]
|
||||
return freqs_cis
|
||||
@@ -0,0 +1,230 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from .activation_layers import get_activation_layer
|
||||
from .attenion import attention
|
||||
from .embed_layers import TextProjection, TimestepEmbedder
|
||||
from .mlp_layers import MLP
|
||||
from .modulate_layers import apply_gate
|
||||
from .norm_layers import get_norm_layer
|
||||
|
||||
|
||||
class IndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
mlp_width_ratio: str = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
act_type: str = "silu",
|
||||
qk_norm: bool = False,
|
||||
qk_norm_type: str = "layer",
|
||||
qkv_bias: bool = True,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.heads_num = heads_num
|
||||
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)
|
||||
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)
|
||||
if qk_norm else nn.Identity())
|
||||
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.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,
|
||||
hidden_channels=mlp_hidden_dim,
|
||||
act_layer=act_layer,
|
||||
drop=mlp_drop_rate,
|
||||
**factory_kwargs,
|
||||
)
|
||||
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
act_layer(),
|
||||
nn.Linear(hidden_size,
|
||||
2 * hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs),
|
||||
)
|
||||
# Zero-initialize the modulation
|
||||
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
||||
nn.init.zeros_(self.adaLN_modulation[1].bias)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
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)
|
||||
# Apply QK-Norm if needed
|
||||
q = self.self_attn_q_norm(q).to(v)
|
||||
k = self.self_attn_k_norm(k).to(v)
|
||||
|
||||
# Self-Attention
|
||||
attn = attention(q, k, v, attn_mask=attn_mask)
|
||||
|
||||
x = x + apply_gate(self.self_attn_proj(attn), gate_msa)
|
||||
|
||||
# FFN Layer
|
||||
x = x + apply_gate(self.mlp(self.norm2(x)), gate_mlp)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class IndividualTokenRefiner(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
depth,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
act_type: str = "silu",
|
||||
qk_norm: bool = False,
|
||||
qk_norm_type: str = "layer",
|
||||
qkv_bias: bool = True,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.blocks = nn.ModuleList([
|
||||
IndividualTokenRefinerBlock(
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
act_type=act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
) for _ in range(depth)
|
||||
])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
c: torch.LongTensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
):
|
||||
mask = mask.clone().bool()
|
||||
# avoid attention weight become NaN
|
||||
mask[:, 0] = True
|
||||
for block in self.blocks:
|
||||
x = block(x, c, mask)
|
||||
return x
|
||||
|
||||
|
||||
class SingleTokenRefiner(nn.Module):
|
||||
"""
|
||||
A single token refiner block for llm text embedding refine.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
depth,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
act_type: str = "silu",
|
||||
qk_norm: bool = False,
|
||||
qk_norm_type: str = "layer",
|
||||
qkv_bias: bool = True,
|
||||
attn_mode: str = "torch",
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
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)
|
||||
|
||||
act_layer = get_activation_layer(act_type)
|
||||
# Build timestep embedding layer
|
||||
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.individual_token_refiner = IndividualTokenRefiner(
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
depth=depth,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
act_type=act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
t: torch.LongTensor,
|
||||
mask: Optional[torch.LongTensor] = None,
|
||||
):
|
||||
timestep_aware_representations = self.t_embedder(t)
|
||||
|
||||
if mask is None:
|
||||
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)
|
||||
c = timestep_aware_representations + context_aware_representations
|
||||
|
||||
x = self.input_embedder(x)
|
||||
|
||||
x = self.individual_token_refiner(x, c, mask)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,52 @@
|
||||
normal_mode_prompt = """Normal mode - Video Recaption Task:
|
||||
|
||||
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
|
||||
|
||||
0. Preserve ALL information, including style words and technical terms.
|
||||
|
||||
1. If the input is in Chinese, translate the entire description to English.
|
||||
|
||||
2. If the input is just one or two words describing an object or person, provide a brief, simple description focusing on basic visual characteristics. Limit the description to 1-2 short sentences.
|
||||
|
||||
3. If the input does not include style, lighting, atmosphere, you can make reasonable associations.
|
||||
|
||||
4. Output ALL must be in English.
|
||||
|
||||
Given Input:
|
||||
input: "{input}"
|
||||
"""
|
||||
|
||||
master_mode_prompt = """Master mode - Video Recaption Task:
|
||||
|
||||
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
|
||||
|
||||
0. Preserve ALL information, including style words and technical terms.
|
||||
|
||||
1. If the input is in Chinese, translate the entire description to English.
|
||||
|
||||
2. If the input is just one or two words describing an object or person, provide a brief, simple description focusing on basic visual characteristics. Limit the description to 1-2 short sentences.
|
||||
|
||||
3. If the input does not include style, lighting, atmosphere, you can make reasonable associations.
|
||||
|
||||
4. Output ALL must be in English.
|
||||
|
||||
Given Input:
|
||||
input: "{input}"
|
||||
"""
|
||||
|
||||
|
||||
def get_rewrite_prompt(ori_prompt, mode="Normal"):
|
||||
if mode == "Normal":
|
||||
prompt = normal_mode_prompt.format(input=ori_prompt)
|
||||
elif mode == "Master":
|
||||
prompt = master_mode_prompt.format(input=ori_prompt)
|
||||
else:
|
||||
raise Exception("Only supports Normal and Normal", mode)
|
||||
return prompt
|
||||
|
||||
|
||||
ori_prompt = "一只小狗在草地上奔跑。"
|
||||
normal_prompt = get_rewrite_prompt(ori_prompt, mode="Normal")
|
||||
master_prompt = get_rewrite_prompt(ori_prompt, mode="Master")
|
||||
|
||||
# Then you can use the normal_prompt or master_prompt to access the hunyuan-large rewrite model to get the final prompt.
|
||||
@@ -0,0 +1,353 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import AutoModel, AutoTokenizer, CLIPTextModel, CLIPTokenizer
|
||||
from transformers.utils import ModelOutput
|
||||
|
||||
from ..constants import PRECISION_TO_TYPE, TEXT_ENCODER_PATH, TOKENIZER_PATH
|
||||
|
||||
|
||||
def use_default(value, default):
|
||||
return value if value is not None else default
|
||||
|
||||
|
||||
def load_text_encoder(
|
||||
text_encoder_type,
|
||||
text_encoder_precision=None,
|
||||
text_encoder_path=None,
|
||||
logger=None,
|
||||
device=None,
|
||||
):
|
||||
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}"
|
||||
)
|
||||
|
||||
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.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.requires_grad_(False)
|
||||
|
||||
if logger is not None:
|
||||
logger.info(f"Text encoder to dtype: {text_encoder.dtype}")
|
||||
|
||||
if device is not None:
|
||||
text_encoder = text_encoder.to(device)
|
||||
|
||||
return text_encoder, text_encoder_path
|
||||
|
||||
|
||||
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}")
|
||||
|
||||
if tokenizer_type == "clipL":
|
||||
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path,
|
||||
max_length=77)
|
||||
elif tokenizer_type == "llm":
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
padding_side=padding_side)
|
||||
else:
|
||||
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
|
||||
|
||||
return tokenizer, tokenizer_path
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextEncoderModelOutput(ModelOutput):
|
||||
"""
|
||||
Base class for model's outputs that also contains a pooling of the last hidden states.
|
||||
|
||||
Args:
|
||||
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
|
||||
Sequence of hidden-states at the output of the last layer of the model.
|
||||
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
|
||||
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
|
||||
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
|
||||
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
|
||||
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
|
||||
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
|
||||
List of decoded texts.
|
||||
"""
|
||||
|
||||
hidden_state: torch.FloatTensor = None
|
||||
attention_mask: Optional[torch.LongTensor] = None
|
||||
hidden_states_list: Optional[Tuple[torch.FloatTensor, ...]] = None
|
||||
text_outputs: Optional[list] = None
|
||||
|
||||
|
||||
class TextEncoder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder_type: str,
|
||||
max_length: int,
|
||||
text_encoder_precision: Optional[str] = None,
|
||||
text_encoder_path: Optional[str] = None,
|
||||
tokenizer_type: Optional[str] = None,
|
||||
tokenizer_path: Optional[str] = None,
|
||||
output_key: Optional[str] = None,
|
||||
use_attention_mask: bool = True,
|
||||
input_max_length: Optional[int] = None,
|
||||
prompt_template: Optional[dict] = None,
|
||||
prompt_template_video: Optional[dict] = None,
|
||||
hidden_state_skip_layer: Optional[int] = None,
|
||||
apply_final_norm: bool = False,
|
||||
reproduce: bool = False,
|
||||
logger=None,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.text_encoder_type = text_encoder_type
|
||||
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.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)
|
||||
self.prompt_template = prompt_template
|
||||
self.prompt_template_video = prompt_template_video
|
||||
self.hidden_state_skip_layer = hidden_state_skip_layer
|
||||
self.apply_final_norm = apply_final_norm
|
||||
self.reproduce = reproduce
|
||||
self.logger = logger
|
||||
|
||||
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 "{}" in str(self.prompt_template["template"]), (
|
||||
"`prompt_template['template']` must contain a placeholder `{}` for the input text, "
|
||||
f"got {self.prompt_template['template']}")
|
||||
|
||||
self.use_video_template = self.prompt_template_video is not None
|
||||
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
|
||||
), 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, "
|
||||
f"got {self.prompt_template_video['template']}")
|
||||
|
||||
if "t5" in text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
elif "clip" in text_encoder_type:
|
||||
self.output_key = output_key or "pooler_output"
|
||||
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}")
|
||||
|
||||
self.model, self.model_path = load_text_encoder(
|
||||
text_encoder_type=self.text_encoder_type,
|
||||
text_encoder_precision=self.precision,
|
||||
text_encoder_path=self.model_path,
|
||||
logger=self.logger,
|
||||
device=device,
|
||||
)
|
||||
self.dtype = self.model.dtype
|
||||
self.device = self.model.device
|
||||
|
||||
self.tokenizer, self.tokenizer_path = load_tokenizer(
|
||||
tokenizer_type=self.tokenizer_type,
|
||||
tokenizer_path=self.tokenizer_path,
|
||||
padding_side="right",
|
||||
logger=self.logger,
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
|
||||
|
||||
@staticmethod
|
||||
def apply_text_to_template(text, template, prevent_empty_text=True):
|
||||
"""
|
||||
Apply text to template.
|
||||
|
||||
Args:
|
||||
text (str): Input text.
|
||||
template (str or list): Template string or list of chat conversation.
|
||||
prevent_empty_text (bool): If True, we will prevent the user text from being empty
|
||||
by adding a space. Defaults to True.
|
||||
"""
|
||||
if isinstance(template, str):
|
||||
# Will send string to tokenizer. Used for llm
|
||||
return template.format(text)
|
||||
else:
|
||||
raise TypeError(f"Unsupported template type: {type(template)}")
|
||||
|
||||
def text2tokens(self, text, data_type="image"):
|
||||
"""
|
||||
Tokenize the input text.
|
||||
|
||||
Args:
|
||||
text (str or list): Input text.
|
||||
"""
|
||||
tokenize_input_type = "str"
|
||||
if self.use_template:
|
||||
if data_type == "image":
|
||||
prompt_template = self.prompt_template["template"]
|
||||
elif data_type == "video":
|
||||
prompt_template = self.prompt_template_video["template"]
|
||||
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
|
||||
]
|
||||
if isinstance(text[0], list):
|
||||
tokenize_input_type = "list"
|
||||
elif isinstance(text, str):
|
||||
text = self.apply_text_to_template(text, prompt_template)
|
||||
if isinstance(text, list):
|
||||
tokenize_input_type = "list"
|
||||
else:
|
||||
raise TypeError(f"Unsupported text type: {type(text)}")
|
||||
|
||||
kwargs = dict(
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
if tokenize_input_type == "str":
|
||||
return self.tokenizer(
|
||||
text,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=True,
|
||||
**kwargs,
|
||||
)
|
||||
elif tokenize_input_type == "list":
|
||||
return self.tokenizer.apply_chat_template(
|
||||
text,
|
||||
add_generation_prompt=True,
|
||||
tokenize=True,
|
||||
return_dict=True,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported tokenize_input_type: {tokenize_input_type}")
|
||||
|
||||
def encode(
|
||||
self,
|
||||
batch_encoding,
|
||||
use_attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
do_sample=None,
|
||||
hidden_state_skip_layer=None,
|
||||
return_texts=False,
|
||||
data_type="image",
|
||||
device=None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
batch_encoding (dict): Batch encoding from tokenizer.
|
||||
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
|
||||
Defaults to None.
|
||||
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
|
||||
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
|
||||
output_hidden_states will be set True. Defaults to False.
|
||||
do_sample (bool): Whether to sample from the model. Used for Decoder-Only LLMs. Defaults to None.
|
||||
When self.produce is False, do_sample is set to True by default.
|
||||
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
|
||||
If None, self.output_key will be used. Defaults to None.
|
||||
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)
|
||||
do_sample = use_default(do_sample, not self.reproduce)
|
||||
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,
|
||||
)
|
||||
if hidden_state_skip_layer is not None:
|
||||
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)
|
||||
else:
|
||||
last_hidden_state = outputs[self.output_key]
|
||||
|
||||
# Remove hidden states of instruction tokens, only keep prompt tokens.
|
||||
if self.use_template:
|
||||
if data_type == "image":
|
||||
crop_start = self.prompt_template.get("crop_start", -1)
|
||||
elif data_type == "video":
|
||||
crop_start = self.prompt_template_video.get("crop_start", -1)
|
||||
else:
|
||||
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)
|
||||
|
||||
if output_hidden_states:
|
||||
return TextEncoderModelOutput(last_hidden_state, attention_mask,
|
||||
outputs.hidden_states)
|
||||
return TextEncoderModelOutput(last_hidden_state, attention_mask)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
text,
|
||||
use_attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
do_sample=False,
|
||||
hidden_state_skip_layer=None,
|
||||
return_texts=False,
|
||||
):
|
||||
batch_encoding = self.text2tokens(text)
|
||||
return self.encode(
|
||||
batch_encoding,
|
||||
use_attention_mask=use_attention_mask,
|
||||
output_hidden_states=output_hidden_states,
|
||||
do_sample=do_sample,
|
||||
hidden_state_skip_layer=hidden_state_skip_layer,
|
||||
return_texts=return_texts,
|
||||
)
|
||||
@@ -0,0 +1,14 @@
|
||||
import math
|
||||
|
||||
|
||||
def align_to(value, alignment):
|
||||
"""align height, width according to alignment
|
||||
|
||||
Args:
|
||||
value (int): height or width
|
||||
alignment (int): target alignment factor
|
||||
|
||||
Returns:
|
||||
int: the aligned value
|
||||
"""
|
||||
return int(math.ceil(value / alignment) * alignment)
|
||||
@@ -0,0 +1,75 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
CODE_SUFFIXES = {
|
||||
".py", # Python codes
|
||||
".sh", # Shell scripts
|
||||
".yaml",
|
||||
".yml", # Configuration files
|
||||
}
|
||||
|
||||
|
||||
def safe_dir(path):
|
||||
"""
|
||||
Create a directory (or the parent directory of a file) if it does not exist.
|
||||
|
||||
Args:
|
||||
path (str or Path): Path to the directory.
|
||||
|
||||
Returns:
|
||||
path (Path): Path object of the directory.
|
||||
"""
|
||||
path = Path(path)
|
||||
path.mkdir(exist_ok=True, parents=True)
|
||||
return path
|
||||
|
||||
|
||||
def safe_file(path):
|
||||
"""
|
||||
Create the parent directory of a file if it does not exist.
|
||||
|
||||
Args:
|
||||
path (str or Path): Path to the file.
|
||||
|
||||
Returns:
|
||||
path (Path): Path object of the file.
|
||||
"""
|
||||
path = Path(path)
|
||||
path.parent.mkdir(exist_ok=True, parents=True)
|
||||
return path
|
||||
|
||||
|
||||
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
|
||||
|
||||
Args:
|
||||
videos (torch.Tensor): video tensor predicted by the model
|
||||
path (str): path to save video
|
||||
rescale (bool, optional): rescale the video tensor from [-1, 1] to . Defaults to False.
|
||||
n_rows (int, optional): Defaults to 1.
|
||||
fps (int, optional): video save fps. Defaults to 8.
|
||||
"""
|
||||
videos = rearrange(videos, "b c t h w -> t b c h w")
|
||||
outputs = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=n_rows)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
if rescale:
|
||||
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
|
||||
x = torch.clamp(x, 0, 1)
|
||||
x = (x * 255).numpy().astype(np.uint8)
|
||||
outputs.append(x)
|
||||
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
imageio.mimsave(path, outputs, fps=fps)
|
||||
@@ -0,0 +1,41 @@
|
||||
import collections.abc
|
||||
from itertools import repeat
|
||||
|
||||
|
||||
def _ntuple(n):
|
||||
|
||||
def parse(x):
|
||||
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
|
||||
x = tuple(x)
|
||||
if len(x) == 1:
|
||||
x = tuple(repeat(x[0], n))
|
||||
return x
|
||||
return tuple(repeat(x, n))
|
||||
|
||||
return parse
|
||||
|
||||
|
||||
to_1tuple = _ntuple(1)
|
||||
to_2tuple = _ntuple(2)
|
||||
to_3tuple = _ntuple(3)
|
||||
to_4tuple = _ntuple(4)
|
||||
|
||||
|
||||
def as_tuple(x):
|
||||
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
|
||||
return tuple(x)
|
||||
if x is None or isinstance(x, (int, float, str)):
|
||||
return (x, )
|
||||
else:
|
||||
raise ValueError(f"Unknown type {type(x)}")
|
||||
|
||||
|
||||
def as_list_of_2tuple(x):
|
||||
x = as_tuple(x)
|
||||
if len(x) == 1:
|
||||
x = (x[0], x[0])
|
||||
assert len(x) % 2 == 0, f"Expect even length, got {len(x)}."
|
||||
lst = []
|
||||
for i in range(0, len(x), 2):
|
||||
lst.append((x[i], x[i + 1]))
|
||||
return lst
|
||||
@@ -0,0 +1,41 @@
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
from transformers import AutoProcessor, LlavaForConditionalGeneration
|
||||
|
||||
|
||||
def preprocess_text_encoder_tokenizer(args):
|
||||
|
||||
processor = AutoProcessor.from_pretrained(args.input_dir)
|
||||
model = LlavaForConditionalGeneration.from_pretrained(
|
||||
args.input_dir,
|
||||
torch_dtype=torch.float16,
|
||||
low_cpu_mem_usage=True,
|
||||
).to(0)
|
||||
|
||||
model.language_model.save_pretrained(f"{args.output_dir}")
|
||||
processor.tokenizer.save_pretrained(f"{args.output_dir}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--input_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="The path to the llava-llama-3-8b-v1_1-transformers.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default="",
|
||||
help="The output path of the llava-llama-3-8b-text-encoder-tokenizer."
|
||||
"if '', the parent dir of output will be the same as input dir.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if len(args.output_dir) == 0:
|
||||
args.output_dir = "/".join(args.input_dir.split("/")[:-1])
|
||||
|
||||
preprocess_text_encoder_tokenizer(args)
|
||||
@@ -0,0 +1,68 @@
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from ..constants import PRECISION_TO_TYPE, VAE_PATH
|
||||
from .autoencoder_kl_causal_3d import AutoencoderKLCausal3D
|
||||
|
||||
|
||||
def load_vae(
|
||||
vae_type: str = "884-16c-hy",
|
||||
vae_precision: str = None,
|
||||
sample_size: tuple = None,
|
||||
vae_path: str = None,
|
||||
logger=None,
|
||||
device=None,
|
||||
):
|
||||
"""the function to load the 3D VAE model
|
||||
|
||||
Args:
|
||||
vae_type (str): the type of the 3D VAE model. Defaults to "884-16c-hy".
|
||||
vae_precision (str, optional): the precision to load vae. Defaults to None.
|
||||
sample_size (tuple, optional): the tiling size. Defaults to None.
|
||||
vae_path (str, optional): the path to vae. Defaults to None.
|
||||
logger (_type_, optional): logger. Defaults to None.
|
||||
device (_type_, optional): device to load vae. Defaults to None.
|
||||
"""
|
||||
if vae_path is None:
|
||||
vae_path = VAE_PATH[vae_type]
|
||||
|
||||
if logger is not None:
|
||||
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)
|
||||
else:
|
||||
vae = AutoencoderKLCausal3D.from_config(config)
|
||||
|
||||
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
|
||||
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
|
||||
|
||||
ckpt = torch.load(vae_ckpt, map_location=vae.device)
|
||||
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.")
|
||||
}
|
||||
vae.load_state_dict(ckpt)
|
||||
|
||||
spatial_compression_ratio = vae.config.spatial_compression_ratio
|
||||
time_compression_ratio = vae.config.time_compression_ratio
|
||||
|
||||
if vae_precision is not None:
|
||||
vae = vae.to(dtype=PRECISION_TO_TYPE[vae_precision])
|
||||
|
||||
vae.requires_grad_(False)
|
||||
|
||||
if logger is not None:
|
||||
logger.info(f"VAE to dtype: {vae.dtype}")
|
||||
|
||||
if device is not None:
|
||||
vae = vae.to(device)
|
||||
|
||||
vae.eval()
|
||||
|
||||
return vae, vae_path, spatial_compression_ratio, time_compression_ratio
|
||||
@@ -0,0 +1,831 @@
|
||||
# Copyright 2024 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
#
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
from dataclasses import dataclass
|
||||
from math import prod
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
try:
|
||||
# This diffusers is modified and packed in the mirror.
|
||||
from diffusers.loaders import FromOriginalVAEMixin
|
||||
except ImportError:
|
||||
# Use this to be compatible with the original diffusers.
|
||||
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.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)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DecoderOutput2(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
posterior: Optional[DiagonalGaussianDistribution] = None
|
||||
|
||||
|
||||
class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
r"""
|
||||
A VAE model with KL loss for encoding images/videos into latents and decoding latent representations into images/videos.
|
||||
|
||||
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
|
||||
for all models (such as downloading or saving).
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: Tuple[str] = ("DownEncoderBlockCausal3D", ),
|
||||
up_block_types: Tuple[str] = ("UpDecoderBlockCausal3D", ),
|
||||
block_out_channels: Tuple[int] = (64, ),
|
||||
layers_per_block: int = 1,
|
||||
act_fn: str = "silu",
|
||||
latent_channels: int = 4,
|
||||
norm_num_groups: int = 32,
|
||||
sample_size: int = 32,
|
||||
sample_tsize: int = 64,
|
||||
scaling_factor: float = 0.18215,
|
||||
force_upcast: float = True,
|
||||
spatial_compression_ratio: int = 8,
|
||||
time_compression_ratio: int = 4,
|
||||
mid_block_add_attention: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_compression_ratio = time_compression_ratio
|
||||
|
||||
self.encoder = EncoderCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=latent_channels,
|
||||
down_block_types=down_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
act_fn=act_fn,
|
||||
norm_num_groups=norm_num_groups,
|
||||
double_z=True,
|
||||
time_compression_ratio=time_compression_ratio,
|
||||
spatial_compression_ratio=spatial_compression_ratio,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
self.decoder = DecoderCausal3D(
|
||||
in_channels=latent_channels,
|
||||
out_channels=out_channels,
|
||||
up_block_types=up_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
norm_num_groups=norm_num_groups,
|
||||
act_fn=act_fn,
|
||||
time_compression_ratio=time_compression_ratio,
|
||||
spatial_compression_ratio=spatial_compression_ratio,
|
||||
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.use_slicing = False
|
||||
self.use_spatial_tiling = False
|
||||
self.use_temporal_tiling = False
|
||||
self.use_parallel = False
|
||||
|
||||
# only relevant if vae tiling is enabled
|
||||
self.tile_sample_min_tsize = sample_tsize
|
||||
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)))
|
||||
self.tile_overlap_factor = 0.25
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if isinstance(module, (EncoderCausal3D, DecoderCausal3D)):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
def enable_temporal_tiling(self, use_tiling: bool = True):
|
||||
self.use_temporal_tiling = use_tiling
|
||||
|
||||
def disable_temporal_tiling(self):
|
||||
self.enable_temporal_tiling(False)
|
||||
|
||||
def enable_spatial_tiling(self, use_tiling: bool = True):
|
||||
self.use_spatial_tiling = use_tiling
|
||||
|
||||
def disable_spatial_tiling(self):
|
||||
self.enable_spatial_tiling(False)
|
||||
|
||||
def enable_tiling(self, use_tiling: bool = True):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger videos.
|
||||
"""
|
||||
self.enable_spatial_tiling(use_tiling)
|
||||
self.enable_temporal_tiling(use_tiling)
|
||||
|
||||
def disable_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
|
||||
decoding in one step.
|
||||
"""
|
||||
self.disable_spatial_tiling()
|
||||
self.disable_temporal_tiling()
|
||||
|
||||
def enable_parallel(self):
|
||||
r"""
|
||||
Enable sequence parallelism for the model. This will allow the vae to decode (with tiling) in parallel.
|
||||
"""
|
||||
self.use_parallel = True
|
||||
|
||||
def enable_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.use_slicing = True
|
||||
|
||||
def disable_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing
|
||||
decoding in one step.
|
||||
"""
|
||||
self.use_slicing = False
|
||||
|
||||
@property
|
||||
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.attn_processors
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
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(
|
||||
return_deprecated_lora=True)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child,
|
||||
processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_attn_processor
|
||||
def set_attn_processor(
|
||||
self,
|
||||
processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]],
|
||||
_remove_lora=False,
|
||||
):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child,
|
||||
processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
|
||||
def set_default_attn_processor(self):
|
||||
"""
|
||||
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()):
|
||||
processor = AttnAddedKVProcessor()
|
||||
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
|
||||
)
|
||||
|
||||
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]]:
|
||||
"""
|
||||
Encode a batch of images/videos into latents.
|
||||
|
||||
Args:
|
||||
x (`torch.FloatTensor`): Input batch of images/videos.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
The latent representations of the encoded images/videos. If `return_dict` is True, a
|
||||
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
|
||||
"""
|
||||
assert len(x.shape) == 5, "The input tensor should have 5 dimensions."
|
||||
|
||||
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):
|
||||
return self.spatial_tiled_encode(x, return_dict=return_dict)
|
||||
|
||||
if self.use_slicing and x.shape[0] > 1:
|
||||
encoded_slices = [self.encoder(x_slice) for x_slice in x.split(1)]
|
||||
h = torch.cat(encoded_slices)
|
||||
else:
|
||||
h = self.encoder(x)
|
||||
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
|
||||
if not return_dict:
|
||||
return (posterior, )
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
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:
|
||||
return self.parallel_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
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):
|
||||
return self.spatial_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z)
|
||||
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
@apply_forward_hook
|
||||
def decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
generator=None) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
"""
|
||||
Decode a batch of images/videos.
|
||||
|
||||
Args:
|
||||
z (`torch.FloatTensor`): Input batch of latent vectors.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.vae.DecoderOutput`] or `tuple`:
|
||||
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
|
||||
returned.
|
||||
|
||||
"""
|
||||
if self.use_slicing and z.shape[0] > 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
|
||||
|
||||
if not return_dict:
|
||||
return (decoded, )
|
||||
|
||||
return DecoderOutput(sample=decoded)
|
||||
|
||||
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)
|
||||
return b
|
||||
|
||||
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)
|
||||
return b
|
||||
|
||||
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)
|
||||
return b
|
||||
|
||||
def spatial_tiled_encode(
|
||||
self,
|
||||
x: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
return_moments: bool = False,
|
||||
) -> AutoencoderKLOutput:
|
||||
r"""Encode a batch of images/videos using a tiled encoder.
|
||||
|
||||
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
|
||||
steps. This is useful to keep memory use constant regardless of image/videos size. The end result of tiled encoding is
|
||||
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
|
||||
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
|
||||
output, but they should be much less noticeable.
|
||||
|
||||
Args:
|
||||
x (`torch.FloatTensor`): Input batch of images/videos.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
|
||||
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)
|
||||
row_limit = self.tile_latent_min_size - blend_extent
|
||||
|
||||
# Split video into tiles and encode them separately.
|
||||
rows = []
|
||||
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 = self.encoder(tile)
|
||||
tile = self.quant_conv(tile)
|
||||
row.append(tile)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(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])
|
||||
result_rows.append(torch.cat(result_row, dim=-1))
|
||||
|
||||
moments = torch.cat(result_rows, dim=-2)
|
||||
if return_moments:
|
||||
return moments
|
||||
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
if not return_dict:
|
||||
return (posterior, )
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def spatial_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
r"""
|
||||
Decode a batch of images/videos using a tiled decoder.
|
||||
|
||||
Args:
|
||||
z (`torch.FloatTensor`): Input batch of latent vectors.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.vae.DecoderOutput`] or `tuple`:
|
||||
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)
|
||||
row_limit = self.tile_sample_min_size - blend_extent
|
||||
|
||||
# Split z into overlapping tiles and decode them separately.
|
||||
# The tiles have an overlap to avoid seams between tiles.
|
||||
rows = []
|
||||
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 = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
row.append(decoded)
|
||||
rows.append(row)
|
||||
result_rows = []
|
||||
for i, row in enumerate(rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
# blend the above tile and the left tile
|
||||
# to the current tile and add the current tile to the result row
|
||||
if i > 0:
|
||||
tile = self.blend_v(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])
|
||||
result_rows.append(torch.cat(result_row, dim=-1))
|
||||
|
||||
dec = torch.cat(result_rows, dim=-2)
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
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)
|
||||
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):
|
||||
tile = self.spatial_tiled_encode(tile, return_moments=True)
|
||||
else:
|
||||
tile = self.encoder(tile)
|
||||
tile = self.quant_conv(tile)
|
||||
if i > 0:
|
||||
tile = tile[:, :, 1:, :, :]
|
||||
row.append(tile)
|
||||
result_row = []
|
||||
for i, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_t(row[i - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, :t_limit + 1, :, :])
|
||||
|
||||
moments = torch.cat(result_row, dim=2)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
|
||||
if not return_dict:
|
||||
return (posterior, )
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def temporal_tiled_decode(self,
|
||||
z: 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)
|
||||
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
|
||||
else:
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
if i > 0:
|
||||
decoded = decoded[:, :, 1:, :, :]
|
||||
row.append(decoded)
|
||||
result_row = []
|
||||
for i, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_t(row[i - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, :t_limit + 1, :, :])
|
||||
|
||||
dec = torch.cat(result_row, dim=2)
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
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)
|
||||
_start_shape += mul_shape
|
||||
global_idx += 1
|
||||
|
||||
def parallel_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
"""
|
||||
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
|
||||
"""
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
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_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_row_limit = self.tile_sample_min_size - s_blend_extent
|
||||
|
||||
# Calculate tile dimensions
|
||||
num_t_tiles = (T + t_overlap_size - 1) // t_overlap_size
|
||||
num_h_tiles = (H + s_overlap_size - 1) // s_overlap_size
|
||||
num_w_tiles = (W + s_overlap_size - 1) // s_overlap_size
|
||||
total_spatial_tiles = num_h_tiles * num_w_tiles
|
||||
total_tiles = num_t_tiles * total_spatial_tiles
|
||||
|
||||
# Calculate tiles per rank and padding
|
||||
tiles_per_rank = (total_tiles + world_size - 1) // world_size
|
||||
start_tile_idx = rank * tiles_per_rank
|
||||
end_tile_idx = min((rank + 1) * tiles_per_rank, total_tiles)
|
||||
|
||||
local_results = []
|
||||
local_dim_metadata = []
|
||||
# Process assigned tiles
|
||||
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
|
||||
h_idx = spatial_idx // num_w_tiles
|
||||
w_idx = spatial_idx % num_w_tiles
|
||||
|
||||
# Calculate positions
|
||||
t_start = t_idx * t_overlap_size
|
||||
h_start = h_idx * s_overlap_size
|
||||
w_start = w_idx * s_overlap_size
|
||||
|
||||
# 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]
|
||||
|
||||
# Process tile
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
|
||||
if t_start > 0:
|
||||
decoded = decoded[:, :, 1:, :, :]
|
||||
|
||||
# Store metadata
|
||||
shape = decoded.shape
|
||||
# Store decoded data (flattened)
|
||||
decoded_flat = decoded.reshape(-1)
|
||||
local_results.append(decoded_flat)
|
||||
local_dim_metadata.append(shape)
|
||||
|
||||
results = torch.cat(local_results, dim=0).contiguous()
|
||||
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)
|
||||
]
|
||||
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)
|
||||
padded_results[:results.size(0)] = results
|
||||
del results
|
||||
torch.cuda.empty_cache()
|
||||
# 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
|
||||
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):
|
||||
t_idx = global_idx // total_spatial_tiles
|
||||
spatial_idx = global_idx % total_spatial_tiles
|
||||
h_idx = spatial_idx // num_w_tiles
|
||||
w_idx = spatial_idx % num_w_tiles
|
||||
data[t_idx][h_idx][w_idx] = current_data
|
||||
# Merge results
|
||||
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)
|
||||
if i > 0:
|
||||
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, :, :])
|
||||
last_slice_data = slice_data
|
||||
dec = torch.cat(result_slices, dim=2)
|
||||
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def _merge_spatial_tiles(self, spatial_rows, blend_extent, row_limit):
|
||||
"""Helper function to merge spatial tiles with blending"""
|
||||
result_rows = []
|
||||
for i, row in enumerate(spatial_rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
if i > 0:
|
||||
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])
|
||||
result_rows.append(torch.cat(result_row, dim=-1))
|
||||
return torch.cat(result_rows, dim=-2)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
sample_posterior: bool = False,
|
||||
return_dict: bool = True,
|
||||
return_posterior: bool = False,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
) -> Union[DecoderOutput2, torch.FloatTensor]:
|
||||
r"""
|
||||
Args:
|
||||
sample (`torch.FloatTensor`): Input sample.
|
||||
sample_posterior (`bool`, *optional*, defaults to `False`):
|
||||
Whether to sample from the posterior.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
|
||||
"""
|
||||
x = sample
|
||||
posterior = self.encode(x).latent_dist
|
||||
if sample_posterior:
|
||||
z = posterior.sample(generator=generator)
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z).sample
|
||||
|
||||
if not return_dict:
|
||||
if return_posterior:
|
||||
return (dec, posterior)
|
||||
else:
|
||||
return (dec, )
|
||||
if return_posterior:
|
||||
return DecoderOutput2(sample=dec, posterior=posterior)
|
||||
else:
|
||||
return DecoderOutput2(sample=dec)
|
||||
|
||||
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections
|
||||
def fuse_qkv_projections(self):
|
||||
"""
|
||||
Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query,
|
||||
key, value) are fused. For cross-attention modules, key and value projection matrices are fused.
|
||||
|
||||
<Tip warning={true}>
|
||||
|
||||
This API is 🧪 experimental.
|
||||
|
||||
</Tip>
|
||||
"""
|
||||
self.original_attn_processors = None
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
self.original_attn_processors = self.attn_processors
|
||||
|
||||
for module in self.modules():
|
||||
if isinstance(module, Attention):
|
||||
module.fuse_projections(fuse=True)
|
||||
|
||||
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
|
||||
def unfuse_qkv_projections(self):
|
||||
"""Disables the fused QKV projection if enabled.
|
||||
|
||||
<Tip warning={true}>
|
||||
|
||||
This API is 🧪 experimental.
|
||||
|
||||
</Tip>
|
||||
|
||||
"""
|
||||
if self.original_attn_processors is not None:
|
||||
self.set_attn_processor(self.original_attn_processors)
|
||||
@@ -0,0 +1,829 @@
|
||||
# Copyright 2024 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
#
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.activations import get_activation
|
||||
from diffusers.models.attention_processor import Attention, SpatialNorm
|
||||
from diffusers.models.normalization import AdaGroupNorm, RMSNorm
|
||||
from diffusers.utils import logging
|
||||
from einops import rearrange
|
||||
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):
|
||||
seq_len = n_frame * n_hw
|
||||
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
|
||||
if batch_size is not None:
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
return mask
|
||||
|
||||
|
||||
class CausalConv3d(nn.Module):
|
||||
"""
|
||||
Implements a causal 3D convolution layer where each position only depends on previous timesteps and current spatial locations.
|
||||
This maintains temporal causality in video generation tasks.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chan_in,
|
||||
chan_out,
|
||||
kernel_size: Union[int, Tuple[int, int, int]],
|
||||
stride: Union[int, Tuple[int, int, int]] = 1,
|
||||
dilation: Union[int, Tuple[int, int, int]] = 1,
|
||||
pad_mode="replicate",
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.pad_mode = pad_mode
|
||||
padding = (
|
||||
kernel_size // 2,
|
||||
kernel_size // 2,
|
||||
kernel_size // 2,
|
||||
kernel_size // 2,
|
||||
kernel_size - 1,
|
||||
0,
|
||||
) # W, H, T
|
||||
self.time_causal_padding = padding
|
||||
|
||||
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)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class UpsampleCausal3D(nn.Module):
|
||||
"""
|
||||
A 3D upsampling layer with an optional convolution.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
use_conv: bool = False,
|
||||
use_conv_transpose: bool = False,
|
||||
out_channels: Optional[int] = None,
|
||||
name: str = "conv",
|
||||
kernel_size: Optional[int] = None,
|
||||
padding=1,
|
||||
norm_type=None,
|
||||
eps=None,
|
||||
elementwise_affine=None,
|
||||
bias=True,
|
||||
interpolate=True,
|
||||
upsample_factor=(2, 2, 2),
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.use_conv_transpose = use_conv_transpose
|
||||
self.name = name
|
||||
self.interpolate = interpolate
|
||||
self.upsample_factor = upsample_factor
|
||||
|
||||
if norm_type == "ln_norm":
|
||||
self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
|
||||
elif norm_type == "rms_norm":
|
||||
self.norm = RMSNorm(channels, eps, elementwise_affine)
|
||||
elif norm_type is None:
|
||||
self.norm = None
|
||||
else:
|
||||
raise ValueError(f"unknown norm_type: {norm_type}")
|
||||
|
||||
conv = None
|
||||
if use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
elif use_conv:
|
||||
if kernel_size is None:
|
||||
kernel_size = 3
|
||||
conv = CausalConv3d(self.channels,
|
||||
self.out_channels,
|
||||
kernel_size=kernel_size,
|
||||
bias=bias)
|
||||
|
||||
if name == "conv":
|
||||
self.conv = conv
|
||||
else:
|
||||
self.Conv2d_0 = conv
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
output_size: Optional[int] = None,
|
||||
scale: float = 1.0,
|
||||
) -> torch.FloatTensor:
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
if self.norm is not None:
|
||||
raise NotImplementedError
|
||||
|
||||
if self.use_conv_transpose:
|
||||
return self.conv(hidden_states)
|
||||
|
||||
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
||||
dtype = hidden_states.dtype
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
|
||||
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||
if hidden_states.shape[0] >= 64:
|
||||
hidden_states = hidden_states.contiguous()
|
||||
|
||||
# if `output_size` is passed we force the interpolation output
|
||||
# size and do not make use of `scale_factor=2`
|
||||
if self.interpolate:
|
||||
B, C, T, H, W = hidden_states.shape
|
||||
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")
|
||||
|
||||
first_h = first_h.squeeze(2)
|
||||
first_h = F.interpolate(first_h,
|
||||
scale_factor=self.upsample_factor[1:],
|
||||
mode="nearest")
|
||||
first_h = first_h.unsqueeze(2)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if T > 1:
|
||||
hidden_states = torch.cat((first_h, other_h), dim=2)
|
||||
else:
|
||||
hidden_states = first_h
|
||||
|
||||
# If the input is bfloat16, we cast back to bfloat16
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(dtype)
|
||||
|
||||
if self.use_conv:
|
||||
if self.name == "conv":
|
||||
hidden_states = self.conv(hidden_states)
|
||||
else:
|
||||
hidden_states = self.Conv2d_0(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DownsampleCausal3D(nn.Module):
|
||||
"""
|
||||
A 3D downsampling layer with an optional convolution.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
use_conv: bool = False,
|
||||
out_channels: Optional[int] = None,
|
||||
padding: int = 1,
|
||||
name: str = "conv",
|
||||
kernel_size=3,
|
||||
norm_type=None,
|
||||
eps=None,
|
||||
elementwise_affine=None,
|
||||
bias=True,
|
||||
stride=2,
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.padding = padding
|
||||
stride = stride
|
||||
self.name = name
|
||||
|
||||
if norm_type == "ln_norm":
|
||||
self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
|
||||
elif norm_type == "rms_norm":
|
||||
self.norm = RMSNorm(channels, eps, elementwise_affine)
|
||||
elif norm_type is None:
|
||||
self.norm = None
|
||||
else:
|
||||
raise ValueError(f"unknown norm_type: {norm_type}")
|
||||
|
||||
if use_conv:
|
||||
conv = CausalConv3d(
|
||||
self.channels,
|
||||
self.out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
bias=bias,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if name == "conv":
|
||||
self.Conv2d_0 = conv
|
||||
self.conv = conv
|
||||
elif name == "Conv2d_0":
|
||||
self.conv = conv
|
||||
else:
|
||||
self.conv = conv
|
||||
|
||||
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)
|
||||
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ResnetBlockCausal3D(nn.Module):
|
||||
r"""
|
||||
A Resnet block.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels: int,
|
||||
out_channels: Optional[int] = None,
|
||||
conv_shortcut: bool = False,
|
||||
dropout: float = 0.0,
|
||||
temb_channels: int = 512,
|
||||
groups: int = 32,
|
||||
groups_out: Optional[int] = None,
|
||||
pre_norm: bool = True,
|
||||
eps: float = 1e-6,
|
||||
non_linearity: str = "swish",
|
||||
skip_time_act: bool = False,
|
||||
# default, scale_shift, ada_group, spatial
|
||||
time_embedding_norm: str = "default",
|
||||
kernel: Optional[torch.FloatTensor] = None,
|
||||
output_scale_factor: float = 1.0,
|
||||
use_in_shortcut: Optional[bool] = None,
|
||||
up: bool = False,
|
||||
down: bool = False,
|
||||
conv_shortcut_bias: bool = True,
|
||||
conv_3d_out_channels: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.pre_norm = pre_norm
|
||||
self.pre_norm = True
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.up = up
|
||||
self.down = down
|
||||
self.output_scale_factor = output_scale_factor
|
||||
self.time_embedding_norm = time_embedding_norm
|
||||
self.skip_time_act = skip_time_act
|
||||
|
||||
linear_cls = nn.Linear
|
||||
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
|
||||
if self.time_embedding_norm == "ada_group":
|
||||
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.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 = None
|
||||
else:
|
||||
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)
|
||||
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.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.nonlinearity = get_activation(non_linearity)
|
||||
|
||||
self.upsample = self.downsample = None
|
||||
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.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:
|
||||
self.conv_shortcut = CausalConv3d(
|
||||
in_channels,
|
||||
conv_3d_out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
bias=conv_shortcut_bias,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_tensor: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
scale: float = 1.0,
|
||||
) -> torch.FloatTensor:
|
||||
hidden_states = input_tensor
|
||||
|
||||
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)
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
if self.upsample is not None:
|
||||
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||
if hidden_states.shape[0] >= 64:
|
||||
input_tensor = input_tensor.contiguous()
|
||||
hidden_states = hidden_states.contiguous()
|
||||
input_tensor = self.upsample(input_tensor, scale=scale)
|
||||
hidden_states = self.upsample(hidden_states, scale=scale)
|
||||
elif self.downsample is not None:
|
||||
input_tensor = self.downsample(input_tensor, scale=scale)
|
||||
hidden_states = self.downsample(hidden_states, scale=scale)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if self.time_emb_proj is not None:
|
||||
if not self.skip_time_act:
|
||||
temb = self.nonlinearity(temb)
|
||||
temb = self.time_emb_proj(temb, scale)[:, :, None, None]
|
||||
|
||||
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"):
|
||||
hidden_states = self.norm2(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "scale_shift":
|
||||
scale, shift = torch.chunk(temb, 2, dim=1)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
output_tensor = (input_tensor +
|
||||
hidden_states) / self.output_scale_factor
|
||||
|
||||
return output_tensor
|
||||
|
||||
|
||||
def get_down_block3d(
|
||||
down_block_type: str,
|
||||
num_layers: int,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
temb_channels: int,
|
||||
add_downsample: bool,
|
||||
downsample_stride: int,
|
||||
resnet_eps: float,
|
||||
resnet_act_fn: str,
|
||||
transformer_layers_per_block: int = 1,
|
||||
num_attention_heads: Optional[int] = None,
|
||||
resnet_groups: Optional[int] = None,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
downsample_padding: Optional[int] = None,
|
||||
dual_cross_attention: bool = False,
|
||||
use_linear_projection: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
attention_type: str = "default",
|
||||
resnet_skip_time_act: bool = False,
|
||||
resnet_out_scale_factor: float = 1.0,
|
||||
cross_attention_norm: Optional[str] = None,
|
||||
attention_head_dim: Optional[int] = None,
|
||||
downsample_type: Optional[str] = None,
|
||||
dropout: float = 0.0,
|
||||
):
|
||||
# If attn head dim is not defined, we default it to the number of heads
|
||||
if attention_head_dim is None:
|
||||
logger.warn(
|
||||
f"It is recommended to provide `attention_head_dim` when calling `get_down_block`. Defaulting `attention_head_dim` to {num_attention_heads}."
|
||||
)
|
||||
attention_head_dim = num_attention_heads
|
||||
|
||||
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,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
dropout=dropout,
|
||||
add_downsample=add_downsample,
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
downsample_padding=downsample_padding,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
)
|
||||
raise ValueError(f"{down_block_type} does not exist.")
|
||||
|
||||
|
||||
def get_up_block3d(
|
||||
up_block_type: str,
|
||||
num_layers: int,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
prev_output_channel: int,
|
||||
temb_channels: int,
|
||||
add_upsample: bool,
|
||||
upsample_scale_factor: Tuple,
|
||||
resnet_eps: float,
|
||||
resnet_act_fn: str,
|
||||
resolution_idx: Optional[int] = None,
|
||||
transformer_layers_per_block: int = 1,
|
||||
num_attention_heads: Optional[int] = None,
|
||||
resnet_groups: Optional[int] = None,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
dual_cross_attention: bool = False,
|
||||
use_linear_projection: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
attention_type: str = "default",
|
||||
resnet_skip_time_act: bool = False,
|
||||
resnet_out_scale_factor: float = 1.0,
|
||||
cross_attention_norm: Optional[str] = None,
|
||||
attention_head_dim: Optional[int] = None,
|
||||
upsample_type: Optional[str] = None,
|
||||
dropout: float = 0.0,
|
||||
) -> nn.Module:
|
||||
# If attn head dim is not defined, we default it to the number of heads
|
||||
if attention_head_dim is None:
|
||||
logger.warn(
|
||||
f"It is recommended to provide `attention_head_dim` when calling `get_up_block`. Defaulting `attention_head_dim` to {num_attention_heads}."
|
||||
)
|
||||
attention_head_dim = num_attention_heads
|
||||
|
||||
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,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
resolution_idx=resolution_idx,
|
||||
dropout=dropout,
|
||||
add_upsample=add_upsample,
|
||||
upsample_scale_factor=upsample_scale_factor,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
temb_channels=temb_channels,
|
||||
)
|
||||
raise ValueError(f"{up_block_type} does not exist.")
|
||||
|
||||
|
||||
class UNetMidBlockCausal3D(nn.Module):
|
||||
"""
|
||||
A 3D UNet mid-block [`UNetMidBlockCausal3D`] with multiple residual blocks and optional attention blocks.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
temb_channels: int,
|
||||
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,
|
||||
attn_groups: Optional[int] = None,
|
||||
resnet_pre_norm: bool = True,
|
||||
add_attention: bool = True,
|
||||
attention_head_dim: int = 1,
|
||||
output_scale_factor: float = 1.0,
|
||||
):
|
||||
super().__init__()
|
||||
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)
|
||||
|
||||
# there is always at least one resnet
|
||||
resnets = [
|
||||
ResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
]
|
||||
attentions = []
|
||||
|
||||
if attention_head_dim is None:
|
||||
logger.warn(
|
||||
f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {in_channels}."
|
||||
)
|
||||
attention_head_dim = in_channels
|
||||
|
||||
for _ in range(num_layers):
|
||||
if self.add_attention:
|
||||
attentions.append(
|
||||
Attention(
|
||||
in_channels,
|
||||
heads=in_channels // attention_head_dim,
|
||||
dim_head=attention_head_dim,
|
||||
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),
|
||||
residual_connection=True,
|
||||
bias=True,
|
||||
upcast_softmax=True,
|
||||
_from_deprecated_attn_block=True,
|
||||
))
|
||||
else:
|
||||
attentions.append(None)
|
||||
|
||||
resnets.append(
|
||||
ResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
))
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
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 = resnet(hidden_states, temb)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DownEncoderBlockCausal3D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
output_scale_factor: float = 1.0,
|
||||
add_downsample: bool = True,
|
||||
downsample_stride: int = 2,
|
||||
downsample_padding: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
|
||||
for i in range(num_layers):
|
||||
in_channels = in_channels if i == 0 else out_channels
|
||||
resnets.append(
|
||||
ResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=None,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
))
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_downsample:
|
||||
self.downsamplers = nn.ModuleList([
|
||||
DownsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
padding=downsample_padding,
|
||||
name="op",
|
||||
stride=downsample_stride,
|
||||
)
|
||||
])
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
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)
|
||||
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = downsampler(hidden_states, scale)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
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,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
|
||||
for i in range(num_layers):
|
||||
input_channels = in_channels if i == 0 else out_channels
|
||||
|
||||
resnets.append(
|
||||
ResnetBlockCausal3D(
|
||||
in_channels=input_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
))
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_upsample:
|
||||
self.upsamplers = nn.ModuleList([
|
||||
UpsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
upsample_factor=upsample_scale_factor,
|
||||
)
|
||||
])
|
||||
else:
|
||||
self.upsamplers = None
|
||||
|
||||
self.resolution_idx = resolution_idx
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
temb: Optional[torch.FloatTensor] = None,
|
||||
scale: float = 1.0,
|
||||
) -> torch.FloatTensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=temb, scale=scale)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = upsampler(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,385 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
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)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DecoderOutput(BaseOutput):
|
||||
r"""
|
||||
Output of decoding method.
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
|
||||
The decoded output sample from the last layer of the model.
|
||||
"""
|
||||
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
class EncoderCausal3D(nn.Module):
|
||||
r"""
|
||||
The `EncoderCausal3D` layer of a variational autoencoder that encodes its input into a latent representation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: Tuple[str, ...] = ("DownEncoderBlockCausal3D", ),
|
||||
block_out_channels: Tuple[int, ...] = (64, ),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
double_z: bool = True,
|
||||
mid_block_add_attention=True,
|
||||
time_compression_ratio: int = 4,
|
||||
spatial_compression_ratio: int = 8,
|
||||
):
|
||||
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.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
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_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)
|
||||
else:
|
||||
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)
|
||||
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),
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=1e-6,
|
||||
downsample_padding=0,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
attention_head_dim=output_channel,
|
||||
temb_channels=None,
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
# mid
|
||||
self.mid_block = UNetMidBlockCausal3D(
|
||||
in_channels=block_out_channels[-1],
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
output_scale_factor=1,
|
||||
resnet_time_scale_shift="default",
|
||||
attention_head_dim=block_out_channels[-1],
|
||||
resnet_groups=norm_num_groups,
|
||||
temb_channels=None,
|
||||
add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
# out
|
||||
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)
|
||||
|
||||
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"
|
||||
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
# down
|
||||
for down_block in self.down_blocks:
|
||||
sample = down_block(sample)
|
||||
|
||||
# middle
|
||||
sample = self.mid_block(sample)
|
||||
|
||||
# post-process
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
class DecoderCausal3D(nn.Module):
|
||||
r"""
|
||||
The `DecoderCausal3D` layer of a variational autoencoder that decodes its latent representation into an output sample.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
up_block_types: Tuple[str, ...] = ("UpDecoderBlockCausal3D", ),
|
||||
block_out_channels: Tuple[int, ...] = (64, ),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
norm_type: str = "group", # group, spatial
|
||||
mid_block_add_attention=True,
|
||||
time_compression_ratio: int = 4,
|
||||
spatial_compression_ratio: int = 8,
|
||||
):
|
||||
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.mid_block = None
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
temb_channels = in_channels if norm_type == "spatial" else None
|
||||
|
||||
# mid
|
||||
self.mid_block = UNetMidBlockCausal3D(
|
||||
in_channels=block_out_channels[-1],
|
||||
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,
|
||||
attention_head_dim=block_out_channels[-1],
|
||||
resnet_groups=norm_num_groups,
|
||||
temb_channels=temb_channels,
|
||||
add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
# up
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i, up_block_type in enumerate(up_block_types):
|
||||
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_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)
|
||||
else:
|
||||
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_T = (2, ) if add_time_upsample else (1, )
|
||||
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,
|
||||
in_channels=prev_output_channel,
|
||||
out_channels=output_channel,
|
||||
prev_output_channel=None,
|
||||
add_upsample=bool(add_spatial_upsample or add_time_upsample),
|
||||
upsample_scale_factor=upsample_scale_factor,
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
attention_head_dim=output_channel,
|
||||
temb_channels=temb_channels,
|
||||
resnet_time_scale_shift=norm_type,
|
||||
)
|
||||
self.up_blocks.append(up_block)
|
||||
prev_output_channel = output_channel
|
||||
|
||||
# out
|
||||
if norm_type == "spatial":
|
||||
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_act = nn.SiLU()
|
||||
self.conv_out = CausalConv3d(block_out_channels[0],
|
||||
out_channels,
|
||||
kernel_size=3)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
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."
|
||||
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
|
||||
if self.training and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
if is_torch_version(">=", "1.11.0"):
|
||||
# middle
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(self.mid_block),
|
||||
sample,
|
||||
latent_embeds,
|
||||
use_reentrant=False,
|
||||
)
|
||||
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,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
# middle
|
||||
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)
|
||||
else:
|
||||
# middle
|
||||
sample = self.mid_block(sample, latent_embeds)
|
||||
sample = sample.to(upscale_dtype)
|
||||
|
||||
# up
|
||||
for up_block in self.up_blocks:
|
||||
sample = up_block(sample, latent_embeds)
|
||||
|
||||
# post-process
|
||||
if latent_embeds is None:
|
||||
sample = self.conv_norm_out(sample)
|
||||
else:
|
||||
sample = self.conv_norm_out(sample, latent_embeds)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution(object):
|
||||
|
||||
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
|
||||
if parameters.ndim == 3:
|
||||
dim = 2 # (B, L, C)
|
||||
elif parameters.ndim == 5 or parameters.ndim == 4:
|
||||
dim = 1 # (B, C, T, H ,W) / (B, C, H, W)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=dim)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
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)
|
||||
|
||||
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,
|
||||
generator=generator,
|
||||
device=self.parameters.device,
|
||||
dtype=self.parameters.dtype,
|
||||
)
|
||||
x = self.mean + self.std * sample
|
||||
return x
|
||||
|
||||
def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor:
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
else:
|
||||
reduce_dim = list(range(1, self.mean.ndim))
|
||||
if other is None:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
|
||||
dim=reduce_dim,
|
||||
)
|
||||
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,
|
||||
dim=reduce_dim,
|
||||
)
|
||||
|
||||
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,
|
||||
dim=dims,
|
||||
)
|
||||
|
||||
def mode(self) -> torch.Tensor:
|
||||
return self.mean
|
||||
@@ -0,0 +1,952 @@
|
||||
# Copyright 2024 The Hunyuan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# 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.
|
||||
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
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.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 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)
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
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)
|
||||
|
||||
# 1. QKV projections
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(hidden_states)
|
||||
value = attn.to_v(hidden_states)
|
||||
|
||||
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
# 2. QK normalization
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query).to(value)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key).to(value)
|
||||
|
||||
image_rotary_emb = (
|
||||
shrink_head(image_rotary_emb[0], dim=0),
|
||||
shrink_head(image_rotary_emb[1], dim=0),
|
||||
)
|
||||
|
||||
# 3. Rotational positional embeddings applied to latent stream
|
||||
if image_rotary_emb is not None:
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
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),
|
||||
query[:, :, -encoder_hidden_states.shape[1]:],
|
||||
],
|
||||
dim=2,
|
||||
)
|
||||
key = torch.cat(
|
||||
[
|
||||
apply_rotary_emb(
|
||||
key[:, :, :-encoder_hidden_states.shape[1]],
|
||||
image_rotary_emb),
|
||||
key[:, :, -encoder_hidden_states.shape[1]:],
|
||||
],
|
||||
dim=2,
|
||||
)
|
||||
else:
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
# 4. Encoder condition QKV projection and normalization
|
||||
if attn.add_q_proj is not None and encoder_hidden_states is not None:
|
||||
encoder_query = attn.add_q_proj(encoder_hidden_states)
|
||||
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)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
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)
|
||||
|
||||
query = torch.cat([query, encoder_query], dim=2)
|
||||
key = torch.cat([key, encoder_key], dim=2)
|
||||
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) #
|
||||
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)
|
||||
|
||||
query_txt = shrink_head(query_txt, dim=1)
|
||||
key_txt = shrink_head(key_txt, dim=1)
|
||||
value_txt = shrink_head(value_txt, dim=1)
|
||||
query = torch.cat([query_img, query_txt], dim=2)
|
||||
key = torch.cat([key_img, key_txt], dim=2)
|
||||
value = torch.cat([value_img, value_txt], dim=2)
|
||||
|
||||
query = query.unsqueeze(2)
|
||||
key = key.unsqueeze(2)
|
||||
value = value.unsqueeze(2)
|
||||
qkv = torch.cat([query, key, value], dim=2)
|
||||
qkv = qkv.transpose(1, 3)
|
||||
|
||||
# 5. Attention
|
||||
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)
|
||||
|
||||
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()
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
encoder_hidden_states = encoder_hidden_states.flatten(2, 3)
|
||||
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
|
||||
else:
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# 6. Output projection
|
||||
if encoder_hidden_states is not None:
|
||||
hidden_states, encoder_hidden_states = (
|
||||
hidden_states[:, :-encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, -encoder_hidden_states.shape[1]:],
|
||||
)
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
if getattr(attn, "to_out", None) is not None:
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if getattr(attn, "to_add_out", None) is not None:
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoPatchEmbed(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Union[int, Tuple[int, int, int]] = 16,
|
||||
in_chans: int = 3,
|
||||
embed_dim: int = 768,
|
||||
) -> 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)
|
||||
|
||||
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
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoAdaNorm(nn.Module):
|
||||
|
||||
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]:
|
||||
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)
|
||||
return gate_msa, gate_mlp
|
||||
|
||||
|
||||
class HunyuanVideoIndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_width_ratio: str = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
attention_bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm1 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=True,
|
||||
eps=1e-6)
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
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.norm_out = HunyuanVideoAdaNorm(hidden_size, 2 * hidden_size)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
gate_msa, gate_mlp = self.norm_out(temb)
|
||||
hidden_states = hidden_states + attn_output * gate_msa
|
||||
|
||||
ff_output = self.ff(self.norm2(hidden_states))
|
||||
hidden_states = hidden_states + ff_output * gate_mlp
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoIndividualTokenRefiner(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_layers: int,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
attention_bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.refiner_blocks = nn.ModuleList([
|
||||
HunyuanVideoIndividualTokenRefinerBlock(
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
attention_bias=attention_bias,
|
||||
) for _ in range(num_layers)
|
||||
])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
self_attn_mask = None
|
||||
if attention_mask is not None:
|
||||
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_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
|
||||
|
||||
for block in self.refiner_blocks:
|
||||
hidden_states = block(hidden_states, temb, self_attn_mask)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoTokenRefiner(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_layers: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
attention_bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
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,
|
||||
attention_head_dim=attention_head_dim,
|
||||
num_layers=num_layers,
|
||||
mlp_width_ratio=mlp_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
attention_bias=attention_bias,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
attention_mask: Optional[torch.LongTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if attention_mask is None:
|
||||
pooled_projections = hidden_states.mean(dim=1)
|
||||
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 = pooled_projections.to(original_dtype)
|
||||
|
||||
temb = self.time_text_embed(timestep, pooled_projections)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
hidden_states = self.token_refiner(hidden_states, temb, attention_mask)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoRotaryPosEmbed(nn.Module):
|
||||
|
||||
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
|
||||
self.patch_size_t = patch_size_t
|
||||
self.rope_dim = rope_dim
|
||||
self.theta = theta
|
||||
|
||||
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
|
||||
]
|
||||
|
||||
axes_grids = []
|
||||
for i in range(3):
|
||||
# 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)
|
||||
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)
|
||||
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)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
class HunyuanVideoSingleTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
qk_norm: str = "rms_norm",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
mlp_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
dim_head=attention_head_dim,
|
||||
heads=num_attention_heads,
|
||||
out_dim=hidden_size,
|
||||
bias=True,
|
||||
processor=HunyuanVideoAttnProcessor2_0(),
|
||||
qk_norm=qk_norm,
|
||||
eps=1e-6,
|
||||
pre_only=True,
|
||||
)
|
||||
|
||||
self.norm = AdaLayerNormZeroSingle(hidden_size, norm_type="layer_norm")
|
||||
self.proj_mlp = nn.Linear(hidden_size, mlp_dim)
|
||||
self.act_mlp = nn.GELU(approximate="tanh")
|
||||
self.proj_out = nn.Linear(hidden_size + mlp_dim, hidden_size)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
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)
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
# 1. Input normalization
|
||||
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
||||
mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states))
|
||||
|
||||
norm_hidden_states, norm_encoder_hidden_states = (
|
||||
norm_hidden_states[:, :-text_seq_length, :],
|
||||
norm_hidden_states[:, -text_seq_length:, :],
|
||||
)
|
||||
|
||||
# 2. Attention
|
||||
attn_output, context_attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
attn_output = torch.cat([attn_output, context_attn_output], dim=1)
|
||||
|
||||
# 3. Modulation and residual connection
|
||||
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
||||
hidden_states = gate.unsqueeze(1) * self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states, encoder_hidden_states = (
|
||||
hidden_states[:, :-text_seq_length, :],
|
||||
hidden_states[:, -text_seq_length:, :],
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_ratio: float,
|
||||
qk_norm: str = "rms_norm",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
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.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
added_kv_proj_dim=hidden_size,
|
||||
dim_head=attention_head_dim,
|
||||
heads=num_attention_heads,
|
||||
out_dim=hidden_size,
|
||||
context_pre_only=False,
|
||||
bias=True,
|
||||
processor=HunyuanVideoAttnProcessor2_0(),
|
||||
qk_norm=qk_norm,
|
||||
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_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,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
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_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context(
|
||||
encoder_hidden_states, emb=temb)
|
||||
|
||||
# 2. Joint attention
|
||||
attn_output, context_attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=freqs_cis,
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
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]
|
||||
|
||||
# 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
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
FromOriginalModelMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data used in [HunyuanVideo](https://huggingface.co/tencent/HunyuanVideo).
|
||||
|
||||
Args:
|
||||
in_channels (`int`, defaults to `16`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, defaults to `16`):
|
||||
The number of channels in the output.
|
||||
num_attention_heads (`int`, defaults to `24`):
|
||||
The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, defaults to `128`):
|
||||
The number of channels in each head.
|
||||
num_layers (`int`, defaults to `20`):
|
||||
The number of layers of dual-stream blocks to use.
|
||||
num_single_layers (`int`, defaults to `40`):
|
||||
The number of layers of single-stream blocks to use.
|
||||
num_refiner_layers (`int`, defaults to `2`):
|
||||
The number of layers of refiner blocks to use.
|
||||
mlp_ratio (`float`, defaults to `4.0`):
|
||||
The ratio of the hidden layer size to the input size in the feedforward network.
|
||||
patch_size (`int`, defaults to `2`):
|
||||
The size of the spatial patches to use in the patch embedding layer.
|
||||
patch_size_t (`int`, defaults to `1`):
|
||||
The size of the tmeporal patches to use in the patch embedding layer.
|
||||
qk_norm (`str`, defaults to `rms_norm`):
|
||||
The normalization to use for the query and key projections in the attention layers.
|
||||
guidance_embeds (`bool`, defaults to `True`):
|
||||
Whether to use guidance embeddings in the model.
|
||||
text_embed_dim (`int`, defaults to `4096`):
|
||||
Input dimension of text embeddings from the text encoder.
|
||||
pooled_projection_dim (`int`, defaults to `768`):
|
||||
The dimension of the pooled projection of the text embeddings.
|
||||
rope_theta (`float`, defaults to `256.0`):
|
||||
The value of theta to use in the RoPE layer.
|
||||
rope_axes_dim (`Tuple[int]`, defaults to `(16, 56, 56)`):
|
||||
The dimensions of the axes to use in the RoPE layer.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
num_layers: int = 20,
|
||||
num_single_layers: int = 40,
|
||||
num_refiner_layers: int = 2,
|
||||
mlp_ratio: float = 4.0,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 1,
|
||||
qk_norm: str = "rms_norm",
|
||||
guidance_embeds: bool = True,
|
||||
text_embed_dim: int = 4096,
|
||||
pooled_projection_dim: int = 768,
|
||||
rope_theta: float = 256.0,
|
||||
rope_axes_dim: Tuple[int] = (16, 56, 56),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
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)
|
||||
|
||||
# 2. RoPE
|
||||
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)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
# 4. Single stream transformer blocks
|
||||
self.single_transformer_blocks = nn.ModuleList([
|
||||
HunyuanVideoSingleTransformerBlock(num_attention_heads,
|
||||
attention_head_dim,
|
||||
mlp_ratio=mlp_ratio,
|
||||
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.gradient_checkpointing = False
|
||||
|
||||
@property
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
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)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
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]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
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)
|
||||
else:
|
||||
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)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if hasattr(module, "gradient_checkpointing"):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
guidance: torch.Tensor = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
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)
|
||||
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# 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."
|
||||
)
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p, p_t = self.config.patch_size, self.config.patch_size_t
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
|
||||
pooled_projections = encoder_hidden_states[:, 0, :self.config.
|
||||
pooled_projection_dim]
|
||||
encoder_hidden_states = encoder_hidden_states[:, 1:]
|
||||
|
||||
# 1. RoPE
|
||||
image_rotary_emb = self.rope(hidden_states)
|
||||
|
||||
# 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)
|
||||
|
||||
# 3. Attention mask preparation
|
||||
latent_sequence_length = hidden_states.shape[1]
|
||||
condition_sequence_length = encoder_hidden_states.shape[1]
|
||||
sequence_length = latent_sequence_length + condition_sequence_length
|
||||
attention_mask = torch.zeros(batch_size,
|
||||
sequence_length,
|
||||
sequence_length,
|
||||
device=hidden_states.device,
|
||||
dtype=torch.bool) # [B, N, N]
|
||||
|
||||
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
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module, return_dict=None):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
if return_dict is not None:
|
||||
return module(*inputs, return_dict=return_dict)
|
||||
else:
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
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(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
temb,
|
||||
attention_mask,
|
||||
image_rotary_emb,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
|
||||
for block in self.single_transformer_blocks:
|
||||
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
temb,
|
||||
attention_mask,
|
||||
image_rotary_emb,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
|
||||
else:
|
||||
for block in self.transformer_blocks:
|
||||
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)
|
||||
|
||||
# 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.permute(0, 4, 1, 5, 2, 6, 3, 7)
|
||||
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return (hidden_states, )
|
||||
|
||||
return Transformer2DModelOutput(sample=hidden_states)
|
||||
@@ -0,0 +1,756 @@
|
||||
# Copyright 2024 The HunyuanVideo Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# 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.
|
||||
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
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.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 fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
|
||||
nccl_info)
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```python
|
||||
>>> import torch
|
||||
>>> from diffusers import HunyuanVideoPipeline, HunyuanVideoTransformer3DModel
|
||||
>>> from diffusers.utils import export_to_video
|
||||
|
||||
>>> model_id = "tencent/HunyuanVideo"
|
||||
>>> transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
... model_id, subfolder="transformer", torch_dtype=torch.bfloat16
|
||||
... )
|
||||
>>> pipe = HunyuanVideoPipeline.from_pretrained(model_id, transformer=transformer, torch_dtype=torch.float16)
|
||||
>>> pipe.vae.enable_tiling()
|
||||
>>> pipe.to("cuda")
|
||||
|
||||
>>> output = pipe(
|
||||
... prompt="A cat walks on the grass, realistic",
|
||||
... height=320,
|
||||
... width=512,
|
||||
... num_frames=61,
|
||||
... num_inference_steps=30,
|
||||
... ).frames[0]
|
||||
>>> export_to_video(output, "output.mp4", fps=15)
|
||||
```
|
||||
"""
|
||||
|
||||
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|>"),
|
||||
"crop_start":
|
||||
95,
|
||||
}
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
|
||||
def retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
timesteps: Optional[List[int]] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
|
||||
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
|
||||
|
||||
Args:
|
||||
scheduler (`SchedulerMixin`):
|
||||
The scheduler to get timesteps from.
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
|
||||
must be `None`.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
|
||||
`num_inference_steps` and `sigmas` must be `None`.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
|
||||
`num_inference_steps` and `timesteps` must be `None`.
|
||||
|
||||
Returns:
|
||||
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
|
||||
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"
|
||||
)
|
||||
if timesteps is not None:
|
||||
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."
|
||||
)
|
||||
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())
|
||||
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."
|
||||
)
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
else:
|
||||
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
r"""
|
||||
Pipeline for text-to-video generation using HunyuanVideo.
|
||||
|
||||
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:
|
||||
text_encoder ([`LlamaModel`]):
|
||||
[Llava Llama3-8B](https://huggingface.co/xtuner/llava-llama-3-8b-v1_1-transformers).
|
||||
tokenizer_2 (`LlamaTokenizer`):
|
||||
Tokenizer from [Llava Llama3-8B](https://huggingface.co/xtuner/llava-llama-3-8b-v1_1-transformers).
|
||||
transformer ([`HunyuanVideoTransformer3DModel`]):
|
||||
Conditional Transformer to denoise the encoded image latents.
|
||||
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKLHunyuanVideo`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
|
||||
text_encoder_2 ([`CLIPTextModel`]):
|
||||
[CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
|
||||
the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
|
||||
tokenizer_2 (`CLIPTokenizer`):
|
||||
Tokenizer of class
|
||||
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder: LlamaModel,
|
||||
tokenizer: LlamaTokenizerFast,
|
||||
transformer: HunyuanVideoTransformer3DModel,
|
||||
vae: AutoencoderKLHunyuanVideo,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
text_encoder_2: CLIPTextModel,
|
||||
tokenizer_2: CLIPTokenizer,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
text_encoder_2=text_encoder_2,
|
||||
tokenizer_2=tokenizer_2,
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = (self.vae.temporal_compression_ratio
|
||||
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)
|
||||
|
||||
def _get_llama_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_template: Dict[str, Any],
|
||||
num_videos_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 256,
|
||||
num_hidden_layers_to_skip: int = 2,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
prompt = [prompt_template["template"].format(p) for p in prompt]
|
||||
|
||||
crop_start = prompt_template.get("crop_start", None)
|
||||
if crop_start is None:
|
||||
prompt_template_input = self.tokenizer(
|
||||
prompt_template["template"],
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=False,
|
||||
)
|
||||
crop_start = prompt_template_input["input_ids"].shape[-1]
|
||||
# Remove <|eot_id|> token and placeholder {}
|
||||
crop_start -= 2
|
||||
|
||||
max_sequence_length += crop_start
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
max_length=max_sequence_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=True,
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids.to(device=device)
|
||||
prompt_attention_mask = text_inputs.attention_mask.to(device=device)
|
||||
|
||||
prompt_embeds = self.text_encoder(
|
||||
input_ids=text_input_ids,
|
||||
attention_mask=prompt_attention_mask,
|
||||
output_hidden_states=True,
|
||||
).hidden_states[-(num_hidden_layers_to_skip + 1)]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype)
|
||||
|
||||
if crop_start is not None and crop_start > 0:
|
||||
prompt_embeds = prompt_embeds[:, crop_start:]
|
||||
prompt_attention_mask = prompt_attention_mask[:, crop_start:]
|
||||
|
||||
# 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)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
def _get_clip_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
num_videos_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 77,
|
||||
) -> torch.Tensor:
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder_2.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer_2(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
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}")
|
||||
|
||||
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)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_2: Union[str, List[str]] = None,
|
||||
prompt_template: Dict[str, Any] = DEFAULT_PROMPT_TEMPLATE,
|
||||
num_videos_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 256,
|
||||
):
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_embeds, prompt_attention_mask = self._get_llama_prompt_embeds(
|
||||
prompt,
|
||||
prompt_template,
|
||||
num_videos_per_prompt,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
if pooled_prompt_embeds is None:
|
||||
if prompt_2 is None and pooled_prompt_embeds is None:
|
||||
prompt_2 = prompt
|
||||
pooled_prompt_embeds = self._get_clip_prompt_embeds(
|
||||
prompt,
|
||||
num_videos_per_prompt,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
max_sequence_length=77,
|
||||
)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, prompt_attention_mask
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
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}."
|
||||
)
|
||||
|
||||
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]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two.")
|
||||
elif prompt_2 is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" 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)}"
|
||||
)
|
||||
|
||||
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)}"
|
||||
)
|
||||
if "template" not in prompt_template:
|
||||
raise ValueError(
|
||||
f"`prompt_template` has to contain a key `template` but only found {prompt_template.keys()}"
|
||||
)
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_channels_latents: 32,
|
||||
height: int = 720,
|
||||
width: int = 1280,
|
||||
num_frames: int = 129,
|
||||
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)
|
||||
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
num_frames,
|
||||
int(height) // self.vae_scale_factor_spatial,
|
||||
int(width) // self.vae_scale_factor_spatial,
|
||||
)
|
||||
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."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
return latents
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
"""
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def disable_vae_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_tiling()
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def attention_kwargs(self):
|
||||
return self._attention_kwargs
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Union[str, List[str]] = None,
|
||||
height: int = 720,
|
||||
width: int = 1280,
|
||||
num_frames: int = 129,
|
||||
num_inference_steps: int = 50,
|
||||
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,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
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,
|
||||
MultiPipelineCallbacks]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
prompt_template: Dict[str, Any] = DEFAULT_PROMPT_TEMPLATE,
|
||||
max_sequence_length: int = 256,
|
||||
):
|
||||
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.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
will be used instead.
|
||||
height (`int`, defaults to `720`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, defaults to `1280`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, defaults to `129`):
|
||||
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.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
|
||||
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
|
||||
will be used.
|
||||
guidance_scale (`float`, defaults to `6.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. Note that the only available HunyuanVideo model is
|
||||
CFG-distilled, which means that traditional guidance between unconditional and conditional latent is
|
||||
not applied.
|
||||
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`.
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
|
||||
provided, text embeddings are generated from the `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`HunyuanVideoPipelineOutput`] instead of a plain tuple.
|
||||
attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
clip_skip (`int`, *optional*):
|
||||
Number of layers to be skipped from CLIP while computing the prompt embeddings. A value of 1 means that
|
||||
the output of the pre-final layer will be used for computing the prompt embeddings.
|
||||
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
|
||||
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
|
||||
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
|
||||
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
|
||||
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~HunyuanVideoPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`HunyuanVideoPipelineOutput`] 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.
|
||||
"""
|
||||
|
||||
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
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
prompt_template,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._interrupt = False
|
||||
|
||||
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]
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, pooled_prompt_embeds, prompt_attention_mask = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt,
|
||||
prompt_template=prompt_template,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
prompt_attention_mask=prompt_attention_mask,
|
||||
device=device,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
prompt_attention_mask = prompt_attention_mask.to(transformer_dtype)
|
||||
if pooled_prompt_embeds is not None:
|
||||
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
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
)
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
num_latent_frames = (num_frames -
|
||||
1) // self.vae_scale_factor_temporal + 1
|
||||
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_latent_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
# 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 = latents[:, :, rank, :, :, :]
|
||||
|
||||
# 6. Prepare guidance condition
|
||||
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
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
latent_model_input = latents.to(transformer_dtype)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
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]),
|
||||
value=0,
|
||||
).unsqueeze(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]
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
guidance=guidance,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
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)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
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):
|
||||
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
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(
|
||||
video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video, )
|
||||
|
||||
return HunyuanVideoPipelineOutput(frames=video)
|
||||
@@ -0,0 +1,500 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
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")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
def reverse_scale_shift(weight, dim):
|
||||
scale, shift = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([shift, scale], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
def reverse_proj_gate(weight):
|
||||
gate, proj = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([proj, gate], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
original_state_dict = state_dict.copy()
|
||||
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")
|
||||
|
||||
# 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")
|
||||
|
||||
# Convert transformer blocks
|
||||
num_layers = 48
|
||||
for i in range(num_layers):
|
||||
block_prefix = f"transformer_blocks.{i}."
|
||||
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")
|
||||
|
||||
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")
|
||||
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")
|
||||
|
||||
# Visual attention
|
||||
q = original_state_dict.pop(block_prefix + "attn1.to_q.weight")
|
||||
k = original_state_dict.pop(block_prefix + "attn1.to_k.weight")
|
||||
v = original_state_dict.pop(block_prefix + "attn1.to_v.weight")
|
||||
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")
|
||||
|
||||
# Context attention
|
||||
q = original_state_dict.pop(block_prefix + "attn1.add_q_proj.weight")
|
||||
k = original_state_dict.pop(block_prefix + "attn1.add_k_proj.weight")
|
||||
v = original_state_dict.pop(block_prefix + "attn1.add_v_proj.weight")
|
||||
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")
|
||||
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")
|
||||
|
||||
# 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")
|
||||
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")
|
||||
|
||||
# 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["pos_frequencies"] = original_state_dict.pop(
|
||||
"pos_frequencies")
|
||||
|
||||
print("Remaining Keys:", original_state_dict.keys())
|
||||
|
||||
return new_state_dict
|
||||
|
||||
|
||||
def convert_diffusers_vae_to_mochi(state_dict):
|
||||
original_state_dict = state_dict.copy()
|
||||
encoder_state_dict = {}
|
||||
decoder_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")
|
||||
|
||||
# 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")
|
||||
|
||||
# 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")
|
||||
|
||||
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")
|
||||
|
||||
# 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")
|
||||
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.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")
|
||||
|
||||
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.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")
|
||||
|
||||
# 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")
|
||||
|
||||
# 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")
|
||||
|
||||
# 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")
|
||||
|
||||
# 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")
|
||||
|
||||
# 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")
|
||||
|
||||
return encoder_state_dict, decoder_state_dict
|
||||
|
||||
|
||||
def ensure_safetensors_extension(path):
|
||||
if not path.endswith(".safetensors"):
|
||||
path = path + ".safetensors"
|
||||
return path
|
||||
|
||||
|
||||
def ensure_directory_exists(path):
|
||||
directory = os.path.dirname(path)
|
||||
if directory:
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
|
||||
|
||||
def main(args):
|
||||
from diffusers import MochiPipeline
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.diffusers_path)
|
||||
|
||||
if args.transformer_path:
|
||||
transformer_path = ensure_safetensors_extension(args.transformer_path)
|
||||
ensure_directory_exists(transformer_path)
|
||||
|
||||
print("Converting transformer model...")
|
||||
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}")
|
||||
|
||||
if args.vae_encoder_path and args.vae_decoder_path:
|
||||
encoder_path = ensure_safetensors_extension(args.vae_encoder_path)
|
||||
decoder_path = ensure_safetensors_extension(args.vae_decoder_path)
|
||||
|
||||
ensure_directory_exists(encoder_path)
|
||||
ensure_directory_exists(decoder_path)
|
||||
|
||||
print("Converting VAE models...")
|
||||
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}")
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main(args)
|
||||
+17
-11
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
|
||||
mochi_latents_mean = torch.tensor([
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
@@ -11,8 +12,8 @@ mochi_latents_mean = torch.tensor([
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
-0.0321993391887285,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor([
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
@@ -25,15 +26,20 @@ mochi_latents_std = torch.tensor([
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
0.9274757570805041
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
0.9274757570805041,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
def normalize_mochi_dit_input(latents):
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents):
|
||||
if model_type == "mochi":
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
elif model_type == "hunyuan_hf":
|
||||
return latents * 0.476986
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
@@ -16,30 +16,33 @@ from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import diffusers
|
||||
import torch.nn.functional as F
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import is_torch_version, logging
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
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.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.embeddings import (MochiCombinedTimestepCaptionEmbedding,
|
||||
PatchEmbed)
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from fastvideo.model.norm import MochiLayerNormContinuous, MochiRMSNormZero, MochiModulatedRMSNorm, MochiRMSNorm
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
import torch.nn.functional as F
|
||||
from diffusers.utils.torch_utils import is_torch_version, maybe_allow_in_graph
|
||||
from einops import rearrange
|
||||
|
||||
import numbers
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
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.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
|
||||
nccl_info)
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class FeedForward(HF_FeedForward):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
@@ -51,37 +54,19 @@ 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:
|
||||
hidden_states = self.net[0].proj(hidden_states)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
|
||||
return self.net[2](
|
||||
LigerSiLUMulFunction.apply(gate, hidden_states)
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
return self.net[2](LigerSiLUMulFunction.apply(gate, hidden_states))
|
||||
|
||||
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, max_s, dropout_p,
|
||||
softmax_scale=softmax_scale, causal=causal
|
||||
)
|
||||
output = rearrange(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)
|
||||
return output
|
||||
|
||||
class MochiAttention(nn.Module):
|
||||
class MochiAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
@@ -115,17 +100,26 @@ 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
|
||||
|
||||
@@ -143,7 +137,6 @@ class MochiAttention(nn.Module):
|
||||
attention_mask=attention_mask,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class MochiAttnProcessor2_0:
|
||||
@@ -151,7 +144,9 @@ 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,
|
||||
@@ -172,12 +167,11 @@ class MochiAttnProcessor2_0:
|
||||
key = key.unflatten(2, (attn.heads, -1))
|
||||
value = value.unflatten(2, (attn.heads, -1))
|
||||
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
# [b, 256, h * d]
|
||||
# [b, 256, h * d]
|
||||
encoder_query = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_key = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_value = attn.add_v_proj(encoder_hidden_states)
|
||||
@@ -186,37 +180,37 @@ class MochiAttnProcessor2_0:
|
||||
encoder_query = encoder_query.unflatten(2, (attn.heads, -1))
|
||||
encoder_key = encoder_key.unflatten(2, (attn.heads, -1))
|
||||
encoder_value = encoder_value.unflatten(2, (attn.heads, -1))
|
||||
|
||||
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_query = attn.norm_added_q(encoder_query)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_key = attn.norm_added_k(encoder_key)
|
||||
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
freqs_cos, freqs_sin = image_rotary_emb[0], image_rotary_emb[1]
|
||||
# shard the head dimension
|
||||
if get_sequence_parallel_state():
|
||||
# B, S, H, D to (S, B,) H, D
|
||||
# batch_size, seq_len, attn_heads, head_dim
|
||||
# batch_size, seq_len, attn_heads, head_dim
|
||||
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
|
||||
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
|
||||
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
|
||||
value = all_to_all_4D(value, scatter_dim=2, gather_dim=1)
|
||||
|
||||
|
||||
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)
|
||||
encoder_value = shrink_head(encoder_value, dim=2)
|
||||
if image_rotary_emb is not None:
|
||||
freqs_cos = shrink_head(freqs_cos, dim=1)
|
||||
freqs_sin = shrink_head(freqs_sin, dim=1)
|
||||
|
||||
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
|
||||
def apply_rotary_emb(x, freqs_cos, freqs_sin):
|
||||
x_even = x[..., 0::2].float()
|
||||
x_odd = x[..., 1::2].float()
|
||||
@@ -224,9 +218,10 @@ class MochiAttnProcessor2_0:
|
||||
sin = (x_even * freqs_sin + x_odd * freqs_cos).to(x.dtype)
|
||||
|
||||
return torch.stack([cos, sin], dim=-1).flatten(-2)
|
||||
|
||||
query = apply_rotary_emb(query, freqs_cos, freqs_sin)
|
||||
key = apply_rotary_emb(key, freqs_cos, freqs_sin)
|
||||
|
||||
|
||||
# query, key, value = query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2)
|
||||
# encoder_query, encoder_key, encoder_value = (
|
||||
# encoder_query.transpose(1, 2),
|
||||
@@ -246,22 +241,28 @@ 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)
|
||||
|
||||
|
||||
# valid_lengths = encoder_attention_mask.sum(dim=1) + sequence_length
|
||||
# def no_padding_mask(score, b, h, q_idx, kv_idx):
|
||||
# return torch.where(kv_idx < valid_lengths[b],score, -float("inf"))
|
||||
|
||||
|
||||
# hidden_states = flex_attention(query, key, value, score_mod=no_padding_mask)
|
||||
if get_sequence_parallel_state():
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
(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)
|
||||
@@ -271,10 +272,7 @@ class MochiAttnProcessor2_0:
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
|
||||
|
||||
(sequence_length, encoder_sequence_length), dim=1)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
@@ -286,6 +284,7 @@ class MochiAttnProcessor2_0:
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class MochiTransformerBlock(nn.Module):
|
||||
r"""
|
||||
@@ -325,10 +324,16 @@ 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,
|
||||
@@ -352,12 +357,17 @@ 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(
|
||||
@@ -377,40 +387,51 @@ class MochiTransformerBlock(nn.Module):
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
output_attn = False,
|
||||
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:
|
||||
norm_encoder_hidden_states, enc_gate_msa, enc_scale_mlp, enc_gate_mlp = self.norm1_context(
|
||||
encoder_hidden_states, temb
|
||||
)
|
||||
(
|
||||
norm_encoder_hidden_states,
|
||||
enc_gate_msa,
|
||||
enc_scale_mlp,
|
||||
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,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
encoder_attention_mask=encoder_attention_mask
|
||||
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)
|
||||
)
|
||||
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))
|
||||
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)
|
||||
)
|
||||
context_ff_output,
|
||||
torch.tanh(enc_gate_mlp).unsqueeze(1))
|
||||
|
||||
if not output_attn:
|
||||
attn_hidden_states = None
|
||||
@@ -434,7 +455,11 @@ 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(
|
||||
@@ -445,20 +470,28 @@ class MochiRoPE(nn.Module):
|
||||
device: Optional[torch.device] = None,
|
||||
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)
|
||||
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)
|
||||
|
||||
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("nd,dhf->nhf", pos.to(torch.float32), freqs.to(torch.float32))
|
||||
freqs = torch.einsum(
|
||||
"nd,dhf->nhf", # codespell:ignore
|
||||
pos.to(torch.float32), # codespell:ignore
|
||||
freqs.to(torch.float32))
|
||||
freqs_cos = torch.cos(freqs)
|
||||
freqs_sin = torch.sin(freqs)
|
||||
return freqs_cos, freqs_sin
|
||||
@@ -478,7 +511,7 @@ class MochiRoPE(nn.Module):
|
||||
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data introduced in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
|
||||
|
||||
@@ -545,28 +578,31 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
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(
|
||||
[
|
||||
MochiTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
pooled_projection_dim=pooled_projection_dim,
|
||||
qk_norm=qk_norm,
|
||||
activation_fn=activation_fn,
|
||||
context_pre_only=i == num_layers - 1,
|
||||
)
|
||||
for i in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
MochiTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
pooled_projection_dim=pooled_projection_dim,
|
||||
qk_norm=qk_norm,
|
||||
activation_fn=activation_fn,
|
||||
context_pre_only=i == num_layers - 1,
|
||||
) for i in range(num_layers)
|
||||
])
|
||||
|
||||
self.norm_out = AdaLayerNormContinuous(
|
||||
inner_dim, inner_dim, elementwise_affine=False, eps=1e-6, norm_type="layer_norm"
|
||||
inner_dim,
|
||||
inner_dim,
|
||||
elementwise_affine=False,
|
||||
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
|
||||
|
||||
@@ -580,24 +616,48 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
output_attn = False,
|
||||
output_features=False,
|
||||
output_features_stride=8,
|
||||
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()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# 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."
|
||||
)
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p = self.config.patch_size
|
||||
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
|
||||
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
|
||||
timestep = 1000 - timestep
|
||||
temb, encoder_hidden_states = self.time_embed(
|
||||
timestep, encoder_hidden_states, encoder_attention_mask, hidden_dtype=hidden_states.dtype
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
hidden_dtype=hidden_states.dtype,
|
||||
)
|
||||
|
||||
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,
|
||||
@@ -612,20 +672,27 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
if self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
|
||||
hidden_states, encoder_hidden_states, attn_outputs = torch.utils.checkpoint.checkpoint(
|
||||
ckpt_kwargs: Dict[str, Any] = ({
|
||||
"use_reentrant": False
|
||||
} if is_torch_version(">=", "1.11.0") else {})
|
||||
(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
attn_outputs,
|
||||
) = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
temb,
|
||||
image_rotary_emb,
|
||||
output_attn,
|
||||
output_features,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
else:
|
||||
@@ -635,20 +702,28 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
output_attn = output_attn,
|
||||
output_attn=output_features,
|
||||
)
|
||||
attn_outputs_list.append(attn_outputs)
|
||||
if i % output_features_stride == 0:
|
||||
attn_outputs_list.append(attn_outputs)
|
||||
|
||||
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)
|
||||
|
||||
if not output_attn :
|
||||
attn_outputs_list = None
|
||||
output = hidden_states.reshape(batch_size, -1, num_frames, height,
|
||||
width)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not output_features:
|
||||
attn_outputs_list = None
|
||||
else:
|
||||
attn_outputs_list = torch.stack(attn_outputs_list, dim=0)
|
||||
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
|
||||
return (-output, attn_outputs_list)
|
||||
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
|
||||
return (-output, attn_outputs_list)
|
||||
@@ -13,15 +13,14 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import numbers
|
||||
from typing import Dict, Optional, Tuple
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class MochiModulatedRMSNorm(nn.Module):
|
||||
|
||||
def __init__(self, eps: float):
|
||||
super().__init__()
|
||||
|
||||
@@ -38,9 +37,10 @@ class MochiModulatedRMSNorm(nn.Module):
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
|
||||
class MochiRMSNorm(nn.Module):
|
||||
|
||||
def __init__(self, dim, eps: float, elementwise_affine=True):
|
||||
super().__init__()
|
||||
|
||||
@@ -63,9 +63,10 @@ class MochiRMSNorm(nn.Module):
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
|
||||
class MochiLayerNormContinuous(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
@@ -77,7 +78,9 @@ 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(
|
||||
@@ -92,7 +95,7 @@ class MochiLayerNormContinuous(nn.Module):
|
||||
x = self.norm(x, (1 + scale.unsqueeze(1).to(torch.float32)))
|
||||
|
||||
return x.to(input_dtype)
|
||||
|
||||
|
||||
|
||||
class MochiRMSNormZero(nn.Module):
|
||||
r"""
|
||||
@@ -102,7 +105,11 @@ class MochiRMSNormZero(nn.Module):
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, embedding_dim: int, hidden_dim: int, eps: float = 1e-5, elementwise_affine: bool = False
|
||||
self,
|
||||
embedding_dim: int,
|
||||
hidden_dim: int,
|
||||
eps: float = 1e-5,
|
||||
elementwise_affine: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -118,7 +125,8 @@ class MochiRMSNormZero(nn.Module):
|
||||
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
|
||||
return hidden_states, gate_msa, scale_mlp, gate_mlp
|
||||
@@ -12,30 +12,29 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import inspect
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
import copy
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import T5EncoderModel, T5TokenizerFast
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.loaders import Mochi1LoraLoaderMixin
|
||||
from diffusers.models.autoencoders import AutoencoderKL
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
|
||||
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 diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.mochi.pipeline_output import MochiPipelineOutput
|
||||
from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather
|
||||
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)
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
@@ -44,7 +43,6 @@ if is_torch_xla_available():
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
@@ -80,14 +78,19 @@ 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]
|
||||
@@ -127,9 +130,12 @@ 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"
|
||||
@@ -139,7 +145,8 @@ def retrieve_timesteps(
|
||||
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"
|
||||
@@ -154,7 +161,7 @@ def retrieve_timesteps(
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class MochiPipeline(DiffusionPipeline):
|
||||
class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
r"""
|
||||
The mochi pipeline for text-to-video generation.
|
||||
|
||||
@@ -180,7 +187,9 @@ class MochiPipeline(DiffusionPipeline):
|
||||
|
||||
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,
|
||||
@@ -199,15 +208,15 @@ class MochiPipeline(DiffusionPipeline):
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
# TODO: determine these scaling factors from model parameters
|
||||
self.vae_spatial_scale_factor = 8
|
||||
self.vae_temporal_scale_factor = 6
|
||||
self.patch_size = 2
|
||||
|
||||
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
|
||||
)
|
||||
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)
|
||||
self.default_height = 480
|
||||
self.default_width = 848
|
||||
|
||||
@@ -238,25 +247,31 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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])
|
||||
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}"
|
||||
)
|
||||
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
|
||||
|
||||
@@ -320,21 +335,24 @@ class MochiPipeline(DiffusionPipeline):
|
||||
|
||||
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)}."
|
||||
)
|
||||
f" {type(prompt)}.")
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
" the batch size of `prompt`.")
|
||||
|
||||
negative_prompt_embeds, negative_prompt_attention_mask = self._get_t5_prompt_embeds(
|
||||
(
|
||||
negative_prompt_embeds,
|
||||
negative_prompt_attention_mask,
|
||||
) = self._get_t5_prompt_embeds(
|
||||
prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
@@ -342,7 +360,12 @@ class MochiPipeline(DiffusionPipeline):
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask
|
||||
return (
|
||||
prompt_embeds,
|
||||
prompt_attention_mask,
|
||||
negative_prompt_embeds,
|
||||
negative_prompt_attention_mask,
|
||||
)
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
@@ -356,11 +379,13 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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
|
||||
):
|
||||
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]}"
|
||||
)
|
||||
@@ -368,34 +393,39 @@ class MochiPipeline(DiffusionPipeline):
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
" 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 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:
|
||||
raise ValueError(
|
||||
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
f" {negative_prompt_embeds.shape}.")
|
||||
if prompt_attention_mask.shape != negative_prompt_attention_mask.shape:
|
||||
raise ValueError(
|
||||
"`prompt_attention_mask` and `negative_prompt_attention_mask` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_attention_mask` {prompt_attention_mask.shape} != `negative_prompt_attention_mask`"
|
||||
f" {negative_prompt_attention_mask.shape}."
|
||||
)
|
||||
f" {negative_prompt_attention_mask.shape}.")
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
@@ -452,7 +482,11 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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=torch.float32)
|
||||
latents = latents.to(dtype)
|
||||
return latents
|
||||
|
||||
@property
|
||||
@@ -467,6 +501,10 @@ class MochiPipeline(DiffusionPipeline):
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def attention_kwargs(self):
|
||||
return self._attention_kwargs
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
@@ -479,12 +517,13 @@ class MochiPipeline(DiffusionPipeline):
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_frames: int = 16,
|
||||
num_inference_steps: int = 28,
|
||||
num_frames: int = 19,
|
||||
num_inference_steps: int = 64,
|
||||
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,
|
||||
@@ -492,10 +531,12 @@ class MochiPipeline(DiffusionPipeline):
|
||||
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = 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,
|
||||
return_all_states=False,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
@@ -547,6 +588,10 @@ class MochiPipeline(DiffusionPipeline):
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.mochi.MochiPipelineOutput`] instead of a plain tuple.
|
||||
attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
@@ -567,7 +612,8 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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
|
||||
@@ -578,7 +624,8 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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,
|
||||
@@ -586,6 +633,7 @@ class MochiPipeline(DiffusionPipeline):
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._interrupt = False
|
||||
|
||||
# 2. Define call parameters
|
||||
@@ -617,8 +665,10 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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
|
||||
@@ -635,18 +685,20 @@ class MochiPipeline(DiffusionPipeline):
|
||||
)
|
||||
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):
|
||||
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
@@ -660,35 +712,48 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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
|
||||
}
|
||||
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,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
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:
|
||||
@@ -700,13 +765,17 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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:
|
||||
@@ -714,34 +783,38 @@ class MochiPipeline(DiffusionPipeline):
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
#latents_shape = list(latents.shape)
|
||||
#full_shape = [latents_shape[0] * world_size] + latents_shape[1:]
|
||||
#all_latents = torch.zeros(full_shape, dtype=latents.dtype, device=latents.device)
|
||||
#torch.distributed.all_gather_into_tensor(all_latents, latents)
|
||||
#latents_list = list(all_latents.chunk(world_size, dim=0))
|
||||
#latents = torch.cat(latents_list, dim=2)
|
||||
# latents_shape = list(latents.shape)
|
||||
# full_shape = [latents_shape[0] * world_size] + latents_shape[1:]
|
||||
# all_latents = torch.zeros(full_shape, dtype=latents.dtype, device=latents.device)
|
||||
# torch.distributed.all_gather_into_tensor(all_latents, latents)
|
||||
# latents_list = list(all_latents.chunk(world_size, dim=0))
|
||||
# latents = torch.cat(latents_list, dim=2)
|
||||
|
||||
if output_type == "latent":
|
||||
video = latents
|
||||
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()
|
||||
if return_all_states:
|
||||
@@ -751,6 +824,6 @@ class MochiPipeline(DiffusionPipeline):
|
||||
return original_noise, video, latents, prompt_embeds, prompt_attention_mask
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
return (video, )
|
||||
|
||||
return MochiPipelineOutput(frames=video)
|
||||
@@ -0,0 +1,102 @@
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
CONFIG_LIST = [
|
||||
triton.Config({"BLOCK_M": 256, "BLOCK_N": 32}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 128, "BLOCK_N": 64}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 128, "BLOCK_N": 32}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 64, "BLOCK_N": 64}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 64, "BLOCK_N": 32}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 32, "BLOCK_N": 64}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 32, "BLOCK_N": 128}, num_stages=2, num_warps=4),
|
||||
triton.Config({"BLOCK_M": 32, "BLOCK_N": 256}, num_stages=2, num_warps=4),
|
||||
]
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=CONFIG_LIST,
|
||||
key=["M", "N"],
|
||||
)
|
||||
@triton.jit
|
||||
def _modulate_fwd(
|
||||
x_ptr, # *Pointer* to first input vector.
|
||||
output_ptr, # *Pointer* to output vector.
|
||||
scale_ptr,
|
||||
shift_ptr,
|
||||
m_stride,
|
||||
s_stride,
|
||||
M,
|
||||
N,
|
||||
seq_len,
|
||||
BLOCK_M: tl.constexpr, # Number of elements each program should process.
|
||||
BLOCK_N: tl.constexpr,
|
||||
# NOTE: `constexpr` so it can be used as a shape value.
|
||||
):
|
||||
row_id = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
|
||||
rows = row_id * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
s_rows = (row_id // seq_len) * BLOCK_M
|
||||
col_id = tl.program_id(axis=1)
|
||||
cols = col_id * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
|
||||
x_ptrs = x_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
scale_ptrs = scale_ptr + s_rows * s_stride + cols[None, :]
|
||||
shift_ptrs = shift_ptr + s_rows * s_stride + cols[None, :]
|
||||
|
||||
col_mask = cols[None, :] < N
|
||||
block_mask = (rows[:, None] < M) & col_mask
|
||||
s_block_mask = col_mask
|
||||
x = tl.load(x_ptrs, mask=block_mask, other=0.0)
|
||||
scale = tl.load(scale_ptrs, mask=s_block_mask, other=0.0)
|
||||
shift = tl.load(shift_ptrs, mask=s_block_mask, other=0.0)
|
||||
|
||||
output = x * (1 + scale) + shift
|
||||
# Write x + y back to DRAM.
|
||||
tl.store(output_ptr + rows[:, None] * m_stride + cols[None, :], output, mask=block_mask)
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=CONFIG_LIST,
|
||||
key=["M", "N"],
|
||||
)
|
||||
@triton.jit
|
||||
def _modulate_bwd(
|
||||
dx_ptr, # *Pointer* to first input vector.
|
||||
x_ptr,
|
||||
dy_ptr, # *Pointer* to output vector.
|
||||
scale_ptr,
|
||||
dscale_ptr,
|
||||
m_stride,
|
||||
s_stride,
|
||||
M,
|
||||
N,
|
||||
seq_len,
|
||||
BLOCK_M: tl.constexpr, # Number of elements each program should process.
|
||||
BLOCK_N: tl.constexpr,
|
||||
# NOTE: `constexpr` so it can be used as a shape value.
|
||||
):
|
||||
row_id = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
|
||||
rows = row_id * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
s_rows = (row_id // seq_len) * BLOCK_M
|
||||
col_id = tl.program_id(axis=1)
|
||||
cols = col_id * BLOCK_N + tl.arange(0, BLOCK_N)
|
||||
|
||||
x_ptrs = x_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
dy_ptrs = dy_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
dx_ptrs = dx_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
dscale_ptrs = dscale_ptr + rows[:, None] * m_stride + cols[None, :]
|
||||
|
||||
scale_ptrs = scale_ptr + s_rows * s_stride + cols[None, :]
|
||||
|
||||
col_mask = cols[None, :] < N
|
||||
block_mask = (rows[:, None] < M) & col_mask
|
||||
s_block_mask = col_mask
|
||||
x = tl.load(x_ptrs, mask=block_mask, other=0.0)
|
||||
dy = tl.load(dy_ptrs, mask=block_mask, other=0.0)
|
||||
scale = tl.load(scale_ptrs, mask=s_block_mask, other=0.0)
|
||||
|
||||
dx = dy * (1 + scale)
|
||||
dscale = dy * x
|
||||
# Write x + y back to DRAM.
|
||||
tl.store(dx_ptrs, dx, mask=block_mask)
|
||||
tl.store(dscale_ptrs, dscale, mask=block_mask)
|
||||
@@ -0,0 +1,63 @@
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from fastvideo.ops.modulate.k_modulate import _modulate_fwd, _modulate_bwd
|
||||
|
||||
|
||||
class _FusedModulate(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x, scale, shift):
|
||||
y = torch.empty_like(x)
|
||||
batch, seq_len, dim = x.shape
|
||||
M = batch * seq_len
|
||||
N = dim
|
||||
x = x.view(-1, dim).contiguous()
|
||||
scale = scale.view(-1, dim).contiguous()
|
||||
shift = shift.view(-1, dim).contiguous()
|
||||
|
||||
def grid(meta):
|
||||
return (
|
||||
triton.cdiv(batch * seq_len, meta["BLOCK_M"]),
|
||||
triton.cdiv(dim, meta["BLOCK_N"]),
|
||||
)
|
||||
|
||||
_modulate_fwd[grid](x, y, scale, shift, x.stride(0), scale.stride(0), M, N, seq_len)
|
||||
|
||||
ctx.save_for_backward(x, scale)
|
||||
ctx.batch = batch
|
||||
ctx.seq_len = seq_len
|
||||
ctx.dim = dim
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, dy): # pragma: no cover # this is covered, but called directly from C++
|
||||
x, scale = ctx.saved_tensors
|
||||
|
||||
batch, seq_len, dim = ctx.batch, ctx.seq_len, ctx.dim
|
||||
M = batch * seq_len
|
||||
N = dim
|
||||
|
||||
# allocate output
|
||||
dy = dy.contiguous()
|
||||
dx = torch.empty_like(dy)
|
||||
dscale = torch.empty_like(dy)
|
||||
dshift = torch.sum(dy, dim=1)
|
||||
|
||||
def grid(meta):
|
||||
return (
|
||||
triton.cdiv(batch * seq_len, meta["BLOCK_M"]),
|
||||
triton.cdiv(dim, meta["BLOCK_N"]),
|
||||
)
|
||||
|
||||
_modulate_bwd[grid](dx, x, dy, scale, dscale, x.stride(0), scale.stride(0), M, N, seq_len)
|
||||
|
||||
dscale = torch.sum(dscale, dim=1)
|
||||
return dx, dscale, dshift
|
||||
|
||||
|
||||
def fused_modulate(
|
||||
x: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return _FusedModulate.apply(x, scale, shift)
|
||||
@@ -1,13 +1,16 @@
|
||||
import json
|
||||
|
||||
import torch.distributed as dist
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
import os
|
||||
from diffusers.utils import export_to_video
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
def generate_video_and_latent(pipe, prompt, height, width, num_frames, num_inference_steps, guidance_scale):
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
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):
|
||||
# Set the random seed for reproducibility
|
||||
generator = torch.Generator("cuda").manual_seed(12345)
|
||||
# Generate videos from the input prompt
|
||||
@@ -19,17 +22,17 @@ def generate_video_and_latent(pipe, prompt, height, width, num_frames, num_infer
|
||||
generator=generator,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
return_all_states=True,
|
||||
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)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
@@ -37,47 +40,75 @@ 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)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
if not isinstance(args.prompt_path, list):
|
||||
args.prompt_path = [args.prompt_path]
|
||||
if len(args.prompt_path) == 1 and args.prompt_path[0].endswith('txt'):
|
||||
text_prompt = open(args.prompt_path[0], 'r').readlines()
|
||||
if len(args.prompt_path) == 1 and args.prompt_path[0].endswith("txt"):
|
||||
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
|
||||
|
||||
# make dir if not exist
|
||||
|
||||
os.makedirs(args.dataset_output_dir, exist_ok=True)
|
||||
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:
|
||||
continue
|
||||
noise, video, latent, prompt_embed, prompt_attention_mask = generate_video_and_latent(pipe, prompt, args.height, args.width, args.num_frames, args.num_inference_steps, args.guidance_scale)
|
||||
(
|
||||
noise,
|
||||
video,
|
||||
latent,
|
||||
prompt_embed,
|
||||
prompt_attention_mask,
|
||||
) = generate_video_and_latent(
|
||||
pipe,
|
||||
prompt,
|
||||
args.height,
|
||||
args.width,
|
||||
args.num_frames,
|
||||
args.num_inference_steps,
|
||||
args.guidance_scale,
|
||||
)
|
||||
# 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)
|
||||
@@ -85,7 +116,7 @@ if __name__ == "__main__":
|
||||
torch.save(prompt_attention_mask, prompt_attention_mask_path)
|
||||
export_to_video(video, video_path, fps=30)
|
||||
item = {}
|
||||
|
||||
|
||||
item["cap"] = prompt
|
||||
item["video"] = video_name + ".mp4"
|
||||
item["noise"] = video_name + ".pt"
|
||||
@@ -97,11 +128,10 @@ if __name__ == "__main__":
|
||||
local_data = data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
|
||||
|
||||
# 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)
|
||||
|
||||
|
||||
@@ -0,0 +1,306 @@
|
||||
"""
|
||||
This script demonstrates how to generate a video using the CogVideoX model with the Hugging Face `diffusers` pipeline.
|
||||
The script supports different types of video generation, including text-to-video (t2v), image-to-video (i2v),
|
||||
and video-to-video (v2v), depending on the input data and different weight.
|
||||
|
||||
- text-to-video: THUDM/CogVideoX-5b, THUDM/CogVideoX-2b or THUDM/CogVideoX1.5-5b
|
||||
- video-to-video: THUDM/CogVideoX-5b, THUDM/CogVideoX-2b or THUDM/CogVideoX1.5-5b
|
||||
- image-to-video: THUDM/CogVideoX-5b-I2V or THUDM/CogVideoX1.5-5b-I2V
|
||||
|
||||
Running the Script:
|
||||
To run the script, use the following command with appropriate arguments:
|
||||
|
||||
```bash
|
||||
$ python cli_demo.py --prompt "A girl riding a bike." --model_path THUDM/CogVideoX1.5-5b --generate_type "t2v"
|
||||
```
|
||||
|
||||
Additional options are available to specify the model path, guidance scale, number of inference steps, video generation type, and output paths.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
from typing import Literal, Optional
|
||||
|
||||
import torch
|
||||
from diffusers import (
|
||||
CogVideoXDPMScheduler,
|
||||
CogVideoXImageToVideoPipeline,
|
||||
CogVideoXPipeline,
|
||||
CogVideoXVideoToVideoPipeline,
|
||||
)
|
||||
from diffusers.utils import export_to_video, load_image, load_video
|
||||
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
|
||||
# Recommended resolution for each model (width, height)
|
||||
RESOLUTION_MAP = {
|
||||
# cogvideox1.5-*
|
||||
"cogvideox1.5-5b-i2v": (1360, 768),
|
||||
"cogvideox1.5-5b": (1360, 768),
|
||||
# cogvideox-*
|
||||
"cogvideox-5b-i2v": (720, 480),
|
||||
"cogvideox-5b": (720, 480),
|
||||
"cogvideox-2b": (720, 480),
|
||||
}
|
||||
|
||||
|
||||
def generate_video(
|
||||
prompt: str,
|
||||
model_path: str,
|
||||
lora_path: str = None,
|
||||
lora_rank: int = 128,
|
||||
num_frames: int = 81,
|
||||
width: Optional[int] = None,
|
||||
height: Optional[int] = None,
|
||||
output_path: str = "./output.mp4",
|
||||
image_or_video_path: str = "",
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 6.0,
|
||||
num_videos_per_prompt: int = 1,
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
generate_type: str = Literal[
|
||||
"t2v", "i2v", "v2v"
|
||||
], # i2v: image to video, v2v: video to video
|
||||
seed: int = 42,
|
||||
fps: int = 16,
|
||||
):
|
||||
"""
|
||||
Generates a video based on the given prompt and saves it to the specified path.
|
||||
|
||||
Parameters:
|
||||
- prompt (str): The description of the video to be generated.
|
||||
- model_path (str): The path of the pre-trained model to be used.
|
||||
- lora_path (str): The path of the LoRA weights to be used.
|
||||
- lora_rank (int): The rank of the LoRA weights.
|
||||
- output_path (str): The path where the generated video will be saved.
|
||||
- num_inference_steps (int): Number of steps for the inference process. More steps can result in better quality.
|
||||
- num_frames (int): Number of frames to generate. CogVideoX1.0 generates 49 frames for 6 seconds at 8 fps, while CogVideoX1.5 produces either 81 or 161 frames, corresponding to 5 seconds or 10 seconds at 16 fps.
|
||||
- width (int): The width of the generated video, applicable only for CogVideoX1.5-5B-I2V
|
||||
- height (int): The height of the generated video, applicable only for CogVideoX1.5-5B-I2V
|
||||
- guidance_scale (float): The scale for classifier-free guidance. Higher values can lead to better alignment with the prompt.
|
||||
- num_videos_per_prompt (int): Number of videos to generate per prompt.
|
||||
- dtype (torch.dtype): The data type for computation (default is torch.bfloat16).
|
||||
- generate_type (str): The type of video generation (e.g., 't2v', 'i2v', 'v2v').·
|
||||
- seed (int): The seed for reproducibility.
|
||||
- fps (int): The frames per second for the generated video.
|
||||
"""
|
||||
|
||||
# 1. Load the pre-trained CogVideoX pipeline with the specified precision (bfloat16).
|
||||
# add device_map="balanced" in the from_pretrained function and remove the enable_model_cpu_offload()
|
||||
# function to use Multi GPUs.
|
||||
|
||||
image = None
|
||||
video = None
|
||||
|
||||
model_name = model_path.split("/")[-1].lower()
|
||||
desired_resolution = RESOLUTION_MAP[model_name]
|
||||
if width is None or height is None:
|
||||
width, height = desired_resolution
|
||||
logging.info(
|
||||
f"\033[1mUsing default resolution {desired_resolution} for {model_name}\033[0m"
|
||||
)
|
||||
elif (width, height) != desired_resolution:
|
||||
if generate_type == "i2v":
|
||||
# For i2v models, use user-defined width and height
|
||||
logging.warning(
|
||||
f"\033[1;31mThe width({width}) and height({height}) are not recommended for {model_name}. The best resolution is {desired_resolution}.\033[0m"
|
||||
)
|
||||
else:
|
||||
# Otherwise, use the recommended width and height
|
||||
logging.warning(
|
||||
f"\033[1;31m{model_name} is not supported for custom resolution. Setting back to default resolution {desired_resolution}.\033[0m"
|
||||
)
|
||||
width, height = desired_resolution
|
||||
|
||||
if generate_type == "i2v":
|
||||
pipe = CogVideoXImageToVideoPipeline.from_pretrained(
|
||||
model_path, torch_dtype=dtype
|
||||
)
|
||||
image = load_image(image=image_or_video_path)
|
||||
elif generate_type == "t2v":
|
||||
pipe = CogVideoXPipeline.from_pretrained(model_path, torch_dtype=dtype)
|
||||
else:
|
||||
pipe = CogVideoXVideoToVideoPipeline.from_pretrained(
|
||||
model_path, torch_dtype=dtype
|
||||
)
|
||||
video = load_video(image_or_video_path)
|
||||
|
||||
# If you're using with lora, add this code
|
||||
if lora_path:
|
||||
pipe.load_lora_weights(
|
||||
lora_path,
|
||||
weight_name="pytorch_lora_weights.safetensors",
|
||||
adapter_name="test_1",
|
||||
)
|
||||
pipe.fuse_lora(lora_scale=1 / lora_rank)
|
||||
|
||||
# 2. Set Scheduler.
|
||||
# Can be changed to `CogVideoXDPMScheduler` or `CogVideoXDDIMScheduler`.
|
||||
# We recommend using `CogVideoXDDIMScheduler` for CogVideoX-2B.
|
||||
# using `CogVideoXDPMScheduler` for CogVideoX-5B / CogVideoX-5B-I2V.
|
||||
|
||||
# pipe.scheduler = CogVideoXDDIMScheduler.from_config(pipe.scheduler.config, timestep_spacing="trailing")
|
||||
pipe.scheduler = CogVideoXDPMScheduler.from_config(
|
||||
pipe.scheduler.config, timestep_spacing="trailing"
|
||||
)
|
||||
|
||||
# 3. Enable CPU offload for the model.
|
||||
# turn off if you have multiple GPUs or enough GPU memory(such as H100) and it will cost less time in inference
|
||||
# and enable to("cuda")
|
||||
|
||||
# pipe.to("cuda")
|
||||
pipe.enable_sequential_cpu_offload()
|
||||
pipe.vae.enable_slicing()
|
||||
pipe.vae.enable_tiling()
|
||||
|
||||
# 4. Generate the video frames based on the prompt.
|
||||
# `num_frames` is the Number of frames to generate.
|
||||
if generate_type == "i2v":
|
||||
video_generate = pipe(
|
||||
height=height,
|
||||
width=width,
|
||||
prompt=prompt,
|
||||
image=image,
|
||||
# The path of the image, the resolution of video will be the same as the image for CogVideoX1.5-5B-I2V, otherwise it will be 720 * 480
|
||||
num_videos_per_prompt=num_videos_per_prompt, # Number of videos to generate per prompt
|
||||
num_inference_steps=num_inference_steps, # Number of inference steps
|
||||
num_frames=num_frames, # Number of frames to generate
|
||||
use_dynamic_cfg=True, # This id used for DPM scheduler, for DDIM scheduler, it should be False
|
||||
guidance_scale=guidance_scale,
|
||||
generator=torch.Generator().manual_seed(
|
||||
seed
|
||||
), # Set the seed for reproducibility
|
||||
).frames[0]
|
||||
elif generate_type == "t2v":
|
||||
video_generate = pipe(
|
||||
height=height,
|
||||
width=width,
|
||||
prompt=prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
num_inference_steps=num_inference_steps,
|
||||
num_frames=num_frames,
|
||||
use_dynamic_cfg=True,
|
||||
guidance_scale=guidance_scale,
|
||||
generator=torch.Generator().manual_seed(seed),
|
||||
).frames[0]
|
||||
else:
|
||||
video_generate = pipe(
|
||||
height=height,
|
||||
width=width,
|
||||
prompt=prompt,
|
||||
video=video, # The path of the video to be used as the background of the video
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
num_inference_steps=num_inference_steps,
|
||||
num_frames=num_frames,
|
||||
use_dynamic_cfg=True,
|
||||
guidance_scale=guidance_scale,
|
||||
generator=torch.Generator().manual_seed(
|
||||
seed
|
||||
), # Set the seed for reproducibility
|
||||
).frames[0]
|
||||
export_to_video(video_generate, output_path, fps=fps)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Generate a video from a text prompt using CogVideoX"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
required=True,
|
||||
help="The description of the video to be generated",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image_or_video_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The path of the image to be used as the background of the video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model_path",
|
||||
type=str,
|
||||
default="THUDM/CogVideoX1.5-5B",
|
||||
help="Path of the pre-trained model use",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The path of the LoRA weights to be used",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora_rank", type=int, default=128, help="The rank of the LoRA weights"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_path",
|
||||
type=str,
|
||||
default="./output.mp4",
|
||||
help="The path save generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance_scale",
|
||||
type=float,
|
||||
default=6.0,
|
||||
help="The scale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_inference_steps", type=int, default=50, help="Inference steps"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_frames",
|
||||
type=int,
|
||||
default=81,
|
||||
help="Number of steps for the inference process",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width", type=int, default=None, help="The width of the generated video"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--height", type=int, default=None, help="The height of the generated video"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=16,
|
||||
help="The frames per second for the generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_videos_per_prompt",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--generate_type", type=str, default="t2v", help="The type of video generation"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype", type=str, default="bfloat16", help="The data type for computation"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed", type=int, default=42, help="The seed for reproducibility"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
dtype = torch.float16 if args.dtype == "float16" else torch.bfloat16
|
||||
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
|
||||
generate_video(
|
||||
prompt=args.prompt,
|
||||
model_path=args.model_path,
|
||||
lora_path=args.lora_path,
|
||||
lora_rank=args.lora_rank,
|
||||
output_path=args.output_path,
|
||||
num_frames=args.num_frames,
|
||||
width=args.width,
|
||||
height=args.height,
|
||||
image_or_video_path=args.image_or_video_path,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
num_videos_per_prompt=args.num_videos_per_prompt,
|
||||
dtype=dtype,
|
||||
generate_type=args.generate_type,
|
||||
seed=args.seed,
|
||||
fps=args.fps,
|
||||
)
|
||||
@@ -0,0 +1,243 @@
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
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.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state, nccl_info)
|
||||
|
||||
|
||||
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
|
||||
|
||||
with open(args.prompt) as f:
|
||||
prompts = f.readlines()
|
||||
|
||||
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,
|
||||
)
|
||||
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(
|
||||
"--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)
|
||||
|
||||
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."
|
||||
)
|
||||
main(args)
|
||||
@@ -0,0 +1,372 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
|
||||
import torch
|
||||
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.pipeline_hunyuan import HunyuanVideoPipeline
|
||||
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 inference(args):
|
||||
initialize_distributed()
|
||||
print(nccl_info.sp_size)
|
||||
device = torch.cuda.current_device()
|
||||
# Peiyuan: GPU seed will cause A100 and H100 to produce different results .....
|
||||
weight_dtype = torch.bfloat16
|
||||
|
||||
if args.transformer_path is not None:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
args.transformer_path)
|
||||
else:
|
||||
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.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")
|
||||
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.set_adapters(["default"], [lora_scaling])
|
||||
print(
|
||||
f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}"
|
||||
)
|
||||
if args.cpu_offload:
|
||||
pipe.enable_model_cpu_offload(device)
|
||||
else:
|
||||
pipe.to(device)
|
||||
|
||||
# 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",
|
||||
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")]
|
||||
prompt_embeds = None
|
||||
encoder_attention_mask = None
|
||||
else:
|
||||
prompts = args.prompts
|
||||
prompt_embeds = None
|
||||
encoder_attention_mask = None
|
||||
|
||||
if prompts is not None:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
for prompt in prompts:
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
video = pipe(
|
||||
prompt=[prompt],
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
video[0],
|
||||
os.path.join(args.output_path, f"{suffix}.mp4"),
|
||||
fps=24,
|
||||
)
|
||||
else:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
videos = pipe(
|
||||
prompt_embeds=prompt_embeds,
|
||||
prompt_attention_mask=encoder_attention_mask,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=24)
|
||||
|
||||
|
||||
def inference_quantization(args):
|
||||
torch.manual_seed(args.seed)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
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)
|
||||
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)
|
||||
elif not args.quantization:
|
||||
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")
|
||||
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.transformer = transformer
|
||||
else:
|
||||
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()
|
||||
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
|
||||
torch.cuda.reset_max_memory_allocated(device)
|
||||
for prompt in prompts:
|
||||
start_time = time.perf_counter()
|
||||
output = pipe(
|
||||
prompt=prompt,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
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)
|
||||
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")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# Basic parameters
|
||||
parser.add_argument("--prompt", type=str, help="prompt file for inference")
|
||||
parser.add_argument("--prompt_embed_path", type=str, default=None)
|
||||
parser.add_argument("--prompt_path", type=str, default=None)
|
||||
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("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--output_path", type=str, default="./outputs/video")
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
parser.add_argument("--quantization", type=str, default=None)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
parser.add_argument(
|
||||
"--lora_checkpoint_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to the directory containing LoRA checkpoints",
|
||||
)
|
||||
# Additional parameters
|
||||
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(
|
||||
"--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", "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-tiling", action="store_true", default=True)
|
||||
|
||||
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)
|
||||
|
||||
args = parser.parse_args()
|
||||
if args.quantization:
|
||||
inference_quantization(args)
|
||||
else:
|
||||
inference(args)
|
||||
@@ -1,165 +1,125 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
import torch.distributed as dist
|
||||
|
||||
from diffusers.utils import export_to_video
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
|
||||
import argparse
|
||||
import os
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
import json
|
||||
from typing import Optional
|
||||
from safetensors.torch import save_file, load_file
|
||||
from peft import set_peft_model_state_dict, inject_adapter_in_model, load_peft_weights
|
||||
from peft import LoraConfig
|
||||
import sys
|
||||
import pdb
|
||||
import copy
|
||||
from typing import Dict
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
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)
|
||||
|
||||
|
||||
def initialize_distributed():
|
||||
local_rank = int(os.getenv('RANK', 0))
|
||||
world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
print('world_size', world_size)
|
||||
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)
|
||||
dist.init_process_group(backend="nccl",
|
||||
init_method="env://",
|
||||
world_size=world_size,
|
||||
rank=local_rank)
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
|
||||
def merge_lora_weights(
|
||||
base_model: torch.nn.Module,
|
||||
lora_weights: Dict[str, torch.Tensor],
|
||||
lora_config: LoraConfig,
|
||||
num_layers: Optional[int] = None
|
||||
) -> torch.nn.Module:
|
||||
merged_model = copy.deepcopy(base_model)
|
||||
if num_layers is None:
|
||||
num_layers = len(merged_model.transformer_blocks)
|
||||
scaling = lora_config.lora_alpha / lora_config.r
|
||||
|
||||
def merge_component(
|
||||
base_weight: torch.Tensor,
|
||||
lora_a: torch.Tensor,
|
||||
lora_b: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
device = base_weight.device
|
||||
lora_a = lora_a.to(device)
|
||||
lora_b = lora_b.to(device)
|
||||
lora_contribution = (lora_b @ lora_a) * scaling
|
||||
if lora_contribution.shape != base_weight.shape:
|
||||
raise ValueError(
|
||||
f"Shape mismatch: base={base_weight.shape}, "
|
||||
f"lora={lora_contribution.shape}"
|
||||
)
|
||||
return base_weight + lora_contribution
|
||||
|
||||
for layer_idx in range(num_layers):
|
||||
transformer_layer = merged_model.transformer_blocks[layer_idx].attn1
|
||||
for target_module in lora_config.target_modules:
|
||||
if target_module == "to_out.0":
|
||||
base_weight = transformer_layer.to_out[0].weight
|
||||
lora_a_key = f"transformer_blocks.{layer_idx}.attn1.to_out.0.lora_A.default.weight"
|
||||
lora_b_key = f"transformer_blocks.{layer_idx}.attn1.to_out.0.lora_B.default.weight"
|
||||
else:
|
||||
base_weight = getattr(transformer_layer, target_module).weight
|
||||
lora_a_key = f"transformer_blocks.{layer_idx}.attn1.{target_module}.lora_A.default.weight"
|
||||
lora_b_key = f"transformer_blocks.{layer_idx}.attn1.{target_module}.lora_B.default.weight"
|
||||
lora_a = lora_weights[lora_a_key]
|
||||
lora_b = lora_weights[lora_b_key]
|
||||
merged_weight = merge_component(base_weight, lora_a, lora_b)
|
||||
if target_module == "to_out.0":
|
||||
transformer_layer.to_out[0].weight.data.copy_(merged_weight)
|
||||
else:
|
||||
getattr(transformer_layer, target_module).weight.data.copy_(merged_weight)
|
||||
merged_model.transformer_blocks[layer_idx].attn1 = transformer_layer
|
||||
return merged_model
|
||||
|
||||
def load_lora_checkpoint(
|
||||
transformer: MochiTransformer3DModel,
|
||||
optimizer,
|
||||
lora_checkpoint_dir: str
|
||||
):
|
||||
config_path = os.path.join(lora_checkpoint_dir, "lora_config.json")
|
||||
with open(config_path, 'r') as f:
|
||||
lora_config_dict = json.load(f)
|
||||
|
||||
for key, value in lora_config['lora_params'].items():
|
||||
setattr(transformer.config, f"lora_{key}", value)
|
||||
|
||||
weight_path = os.path.join(lora_checkpoint_dir, "lora_weights.safetensors")
|
||||
lora_state_dict = load_file(weight_path)
|
||||
|
||||
lora_config = LoraConfig(
|
||||
r=lora_config_dict['lora_params']['lora_rank'],
|
||||
lora_alpha=lora_config_dict['lora_params']['lora_alpha'],
|
||||
target_modules=lora_config_dict['lora_params']['target_modules']
|
||||
)
|
||||
|
||||
transformer = merge_lora_weights(transformer, lora_state_dict, lora_config)
|
||||
step = lora_state_dict['step']
|
||||
print(f"--> Successfully loaded LoRA checkpoint from step {step}")
|
||||
return transformer
|
||||
|
||||
def main(args):
|
||||
initialize_distributed()
|
||||
print(nccl_info.sp_size)
|
||||
device = torch.cuda.current_device()
|
||||
generator = torch.Generator(device).manual_seed(args.seed)
|
||||
weight_dtype = torch.bfloat16
|
||||
# Peiyuan: GPU seed will cause A100 and H100 to produce different results .....
|
||||
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
|
||||
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, linear_quadratic,args.linear_threshold, args.linear_range)
|
||||
if args.transformer_path is not None:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder = 'transformer/')
|
||||
if args.lora_checkpoint_dir is not None:
|
||||
# Load and merge LoRA weights
|
||||
transformer = load_lora_checkpoint(
|
||||
transformer=transformer,
|
||||
optimizer=None, # No optimizer needed for inference
|
||||
output_dir=args.lora_checkpoint_dir
|
||||
scheduler = PCMFMScheduler(
|
||||
1000,
|
||||
args.shift,
|
||||
args.num_euler_timesteps,
|
||||
linear_quadratic,
|
||||
args.linear_threshold,
|
||||
args.linear_range,
|
||||
)
|
||||
print(f"Loaded and merged LoRA weights from {args.lora_checkpoint_dir}")
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer = transformer,scheduler=scheduler)
|
||||
|
||||
if args.transformer_path is not None:
|
||||
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)
|
||||
|
||||
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")
|
||||
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.set_adapters(["default"], [lora_scaling])
|
||||
print(
|
||||
f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}"
|
||||
)
|
||||
# pipe.to(device)
|
||||
|
||||
|
||||
pipe.enable_model_cpu_offload(device)
|
||||
|
||||
# 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", 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)
|
||||
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))
|
||||
prompts = None
|
||||
elif args.prompt_path is not None:
|
||||
prompts = [line.strip() for line in open(args.prompt_path, "r")]
|
||||
prompt_embeds = None
|
||||
encoder_attention_mask = None
|
||||
else:
|
||||
else:
|
||||
prompts = args.prompts
|
||||
prompt_embeds = None
|
||||
encoder_attention_mask = None
|
||||
|
||||
|
||||
if prompts is not None:
|
||||
videos = []
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
for prompt in prompts:
|
||||
video = pipe(
|
||||
prompt=[prompt],
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
videos.append(video[0])
|
||||
for prompt in prompts:
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
video = pipe(
|
||||
prompt=[prompt],
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
video[0],
|
||||
os.path.join(args.output_path, f"{suffix}.mp4"),
|
||||
fps=30,
|
||||
)
|
||||
else:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
videos = pipe(
|
||||
prompt_embeds=prompt_embeds,
|
||||
prompt_attention_mask=encoder_attention_mask,
|
||||
@@ -171,20 +131,14 @@ def main(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if prompts is not None:
|
||||
# mkdir
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
for video, prompt in zip(videos, prompts):
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(video, os.path.join(args.output_path, f"{suffix}.mp4"), fps=30)
|
||||
else:
|
||||
if nccl_info.global_rank <= 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# arg parse
|
||||
# arg parse
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs='+', default=[])
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
@@ -197,8 +151,15 @@ 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('--lora_checkpoint_dir', type=str, default=None, help='Path to the directory containing LoRA checkpoints')
|
||||
parser.add_argument("--encoder_attention_mask_path",
|
||||
type=str,
|
||||
default=None)
|
||||
parser.add_argument(
|
||||
"--lora_checkpoint_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to the directory containing LoRA checkpoints",
|
||||
)
|
||||
parser.add_argument("--shift", type=float, default=8.0)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.025)
|
||||
|
||||
@@ -1,19 +1,27 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import export_to_video, load_image, load_video
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
|
||||
|
||||
def main(args):
|
||||
# Set the random seed for reproducibility
|
||||
generator = torch.Generator("cuda").manual_seed(args.seed)
|
||||
# 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()
|
||||
@@ -29,14 +37,15 @@ def main(args):
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
).frames
|
||||
|
||||
for prompt,video in zip(args.prompts, videos):
|
||||
|
||||
for prompt, video in zip(args.prompts, videos):
|
||||
export_to_video(video, args.output_path + f"_{prompt}.mp4", fps=30)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# arg parse
|
||||
# arg parse
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs='+', default=[])
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
|
||||
+533
-256
@@ -1,56 +1,56 @@
|
||||
# !/bin/python3
|
||||
# isort: skip_file
|
||||
import argparse
|
||||
from email.policy import strict
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, \
|
||||
destroy_sequence_parallel_group, get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
|
||||
from fastvideo.model.mochi_latents_utils import normalize_mochi_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from torch.utils.data import DataLoader
|
||||
import torch
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
import json
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
|
||||
import diffusers
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import check_min_version
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import save_file, load_file
|
||||
from peft import LoraConfig, inject_adapter_in_model
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import save_checkpoint, save_lora_checkpoint, resume_lora_training
|
||||
from fastvideo.utils.logging import main_print
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version, convert_unet_state_dict_to_peft
|
||||
from peft import LoraConfig, set_peft_model_state_dict
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
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.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.dataset_utils import LengthGroupedSampler
|
||||
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.validation import log_validation
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
|
||||
|
||||
|
||||
|
||||
def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str, batch_size: int, generator, logit_mean: float = None, logit_std: float = None, mode_scale: float = None
|
||||
weighting_scheme: str,
|
||||
batch_size: int,
|
||||
generator,
|
||||
logit_mean: float = None,
|
||||
logit_std: float = None,
|
||||
mode_scale: float = None,
|
||||
):
|
||||
"""
|
||||
Compute the density for sampling the timesteps when doing SD3 training.
|
||||
@@ -61,20 +61,32 @@ def compute_density_for_timestep_sampling(
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device="cpu", generator=generator)
|
||||
u = torch.normal(
|
||||
mean=logit_mean,
|
||||
std=logit_std,
|
||||
size=(batch_size, ),
|
||||
device="cpu",
|
||||
generator=generator,
|
||||
)
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size,), device="cpu", generator=generator)
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u)
|
||||
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2)**2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size,), device="cpu", generator=generator)
|
||||
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
|
||||
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:
|
||||
@@ -82,16 +94,36 @@ def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32)
|
||||
return sigma
|
||||
|
||||
|
||||
def train_one_step_mochi(transformer, optimizer, lr_scheduler,loader, noise_scheduler, noise_random_generator, gradient_accumulation_steps, sp_size, precondition_outputs, max_grad_norm, weighting_scheme, logit_mean, logit_std, mode_scale):
|
||||
def train_one_step(
|
||||
transformer,
|
||||
model_type,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
noise_scheduler,
|
||||
noise_random_generator,
|
||||
gradient_accumulation_steps,
|
||||
sp_size,
|
||||
precondition_outputs,
|
||||
max_grad_norm,
|
||||
weighting_scheme,
|
||||
logit_mean,
|
||||
logit_std,
|
||||
mode_scale,
|
||||
):
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
latents, encoder_hidden_states, latents_attention_mask, encoder_attention_mask = next(loader)
|
||||
latents = normalize_mochi_dit_input(latents)
|
||||
|
||||
(
|
||||
latents,
|
||||
encoder_hidden_states,
|
||||
latents_attention_mask,
|
||||
encoder_attention_mask,
|
||||
) = next(loader)
|
||||
latents = normalize_dit_input(model_type, latents)
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn_like(latents)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=noise_random_generator,
|
||||
@@ -100,60 +132,62 @@ def train_one_step_mochi(transformer, optimizer, lr_scheduler,loader, noise_sche
|
||||
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)
|
||||
|
||||
sigmas = get_sigmas(noise_scheduler, latents.device, timesteps, n_dim=latents.ndim, dtype=latents.dtype)
|
||||
sigmas = get_sigmas(
|
||||
noise_scheduler,
|
||||
latents.device,
|
||||
timesteps,
|
||||
n_dim=latents.ndim,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
|
||||
with torch.autocast("cuda", torch.bfloat16):
|
||||
model_pred = transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
input_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if 'hunyuan' in model_type:
|
||||
input_kwargs["guidance"] = torch.tensor(
|
||||
[1000.0],
|
||||
device=noisy_model_input.device,
|
||||
dtype=torch.bfloat16)
|
||||
model_pred = transformer(**input_kwargs)[0]
|
||||
|
||||
if precondition_outputs:
|
||||
model_pred = noisy_model_input - model_pred * sigmas
|
||||
model_pred = noisy_model_input - model_pred * sigmas
|
||||
if precondition_outputs:
|
||||
target = latents
|
||||
else:
|
||||
target = noise - latents
|
||||
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()
|
||||
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
return total_loss, grad_norm.item()
|
||||
|
||||
def get_lora_model(transformer, lora_config):
|
||||
transformer.requires_grad_(False)
|
||||
transformer = inject_adapter_in_model(lora_config, transformer)
|
||||
return transformer
|
||||
|
||||
|
||||
|
||||
def main(args):
|
||||
# use LayerNorm, GeLu, SiLu always as fp32 mode
|
||||
# TODO:
|
||||
if args.enable_stable_fp32:
|
||||
raise NotImplementedError("enable_stable_fp32 is not supported now.")
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
local_rank = int(os.environ['LOCAL_RANK'])
|
||||
rank = int(os.environ['RANK'])
|
||||
world_size = int(os.environ['WORLD_SIZE'])
|
||||
|
||||
local_rank = int(os.environ["LOCAL_RANK"])
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
dist.init_process_group("nccl")
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.cuda.current_device()
|
||||
@@ -167,53 +201,94 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <=0 and args.output_dir is not None:
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weigths to half-precision
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
transformer = load_transformer(
|
||||
args.model_type,
|
||||
args.dit_model_name_or_path,
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
|
||||
)
|
||||
|
||||
|
||||
if args.use_lora:
|
||||
lora_config = LoraConfig(
|
||||
assert args.model_type != "hunyuan", "LoRA is only supported for huggingface model. Please use hunyuan_hf for lora finetuning"
|
||||
if args.model_type == "mochi":
|
||||
pipe = MochiPipeline
|
||||
elif args.model_type == "hunyuan_hf":
|
||||
pipe = HunyuanVideoPipeline
|
||||
transformer.requires_grad_(False)
|
||||
transformer_lora_config = LoraConfig(
|
||||
r=args.lora_rank,
|
||||
lora_alpha=args.lora_alpha,
|
||||
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
|
||||
init_lora_weights=True,
|
||||
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
|
||||
)
|
||||
transformer = get_lora_model(transformer, lora_config)
|
||||
transformer.add_adapter(transformer_lora_config)
|
||||
|
||||
if 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")
|
||||
if incompatible_keys is not None:
|
||||
# check only for unexpected keys
|
||||
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" 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,
|
||||
args.use_lora,
|
||||
args.use_cpu_offload,
|
||||
args.master_weight_type,
|
||||
)
|
||||
|
||||
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}")
|
||||
fsdp_kwargs = get_dit_fsdp_kwargs(args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload)
|
||||
|
||||
|
||||
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 = ["MochiTransformerBlock"]
|
||||
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,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
main_print(f"--> model loaded")
|
||||
main_print("--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(transformer, no_split_modules,
|
||||
args.selective_checkpointing)
|
||||
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
@@ -221,44 +296,45 @@ 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,
|
||||
lr=args.learning_rate,
|
||||
betas=(0.9,0.999),
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
init_steps = 0
|
||||
if args.resume_from_lora_checkpoint:
|
||||
transformer, optimizer, init_steps = resume_lora_training(
|
||||
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
|
||||
lr_scheduler = get_scheduler(
|
||||
args.lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=args.lr_warmup_steps * world_size,
|
||||
num_training_steps=args.max_train_steps * world_size,
|
||||
num_cycles=args.lr_num_cycles,
|
||||
power=args.lr_power,
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
|
||||
sampler = LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(train_dataset, rank=rank, num_replicas=world_size, shuffle=False)
|
||||
|
||||
args.lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=args.lr_warmup_steps,
|
||||
num_training_steps=args.max_train_steps,
|
||||
num_cycles=args.lr_num_cycles,
|
||||
power=args.lr_power,
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
|
||||
args.cfg)
|
||||
sampler = (LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(
|
||||
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
@@ -266,188 +342,360 @@ def main(args):
|
||||
pin_memory=True,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
drop_last=True,
|
||||
drop_last=True,
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = args.train_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")
|
||||
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.")
|
||||
# TODO
|
||||
|
||||
assert NotImplementedError(
|
||||
"resume_from_checkpoint is not supported now.")
|
||||
# TODO
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, args.max_train_steps),
|
||||
initial=init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable= local_rank > 0,
|
||||
disable=local_rank > 0,
|
||||
)
|
||||
|
||||
loader = sp_parallel_dataloader_wrapper(
|
||||
train_dataloader,
|
||||
device,
|
||||
args.train_batch_size,
|
||||
args.sp_size,
|
||||
args.train_sp_batch_size,
|
||||
)
|
||||
|
||||
loader = sp_parallel_dataloader_wrapper(train_dataloader, device, args.train_batch_size, args.sp_size, args.train_sp_batch_size)
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
|
||||
#todo future
|
||||
# todo future
|
||||
for i in range(init_steps):
|
||||
next(loader)
|
||||
for step in range(init_steps + 1, args.max_train_steps+1):
|
||||
for step in range(init_steps + 1, args.max_train_steps + 1):
|
||||
start_time = time.time()
|
||||
loss, grad_norm= train_one_step_mochi(transformer, optimizer, lr_scheduler, loader, noise_scheduler, noise_random_generator, args.gradient_accumulation_steps, args.sp_size, args.precondition_outputs, args.max_grad_norm, args.weighting_scheme, args.logit_mean, args.logit_std, args.mode_scale)
|
||||
loss, grad_norm = train_one_step(
|
||||
transformer,
|
||||
args.model_type,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
noise_scheduler,
|
||||
noise_random_generator,
|
||||
args.gradient_accumulation_steps,
|
||||
args.sp_size,
|
||||
args.precondition_outputs,
|
||||
args.max_grad_norm,
|
||||
args.weighting_scheme,
|
||||
args.logit_mean,
|
||||
args.logit_std,
|
||||
args.mode_scale,
|
||||
)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm
|
||||
})
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
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
|
||||
}, step=step)
|
||||
if step % args.checkpointing_steps == 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,
|
||||
},
|
||||
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, pipe)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
save_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
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)
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
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)
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
|
||||
args.max_train_steps, pipe)
|
||||
else:
|
||||
save_checkpoint(transformer, optimizer, 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()
|
||||
|
||||
|
||||
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]"
|
||||
)
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_height", type=int, default=480)
|
||||
parser.add_argument("--num_width", type=int, default=848)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--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.")
|
||||
parser.add_argument("--train_batch_size", type=int, 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("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument(
|
||||
"--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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
type=int,
|
||||
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("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--pretrained_model_name_or_path", type=str)
|
||||
parser.add_argument("--cache_dir", type=str, default='./cache_dir')
|
||||
parser.add_argument('--enable_stable_fp32', action='store_true') # TODO
|
||||
parser.add_argument("--dit_model_name_or_path", type=str, default=None)
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
|
||||
# diffusion setting
|
||||
parser.add_argument("--ema_decay", type=float, default=0.999)
|
||||
parser.add_argument("--ema_start_step", type=int, default=0)
|
||||
parser.add_argument('--cfg', type=float, default=0.1)
|
||||
parser.add_argument("--precondition_outputs", action="store_true", help="Whether to precondition the outputs of the model.")
|
||||
|
||||
parser.add_argument("--cfg", type=float, default=0.1)
|
||||
parser.add_argument(
|
||||
"--precondition_outputs",
|
||||
action="store_true",
|
||||
help="Whether to precondition the outputs of the model.",
|
||||
)
|
||||
|
||||
# validation & logs
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--uncond_prompt_dir", type=str)
|
||||
parser.add_argument("--validation_sampling_steps", type=int, default=64)
|
||||
parser.add_argument('--validation_guidance_scale', type=float, default=4.5)
|
||||
parser.add_argument('--validation_steps', type=float, default=4.5)
|
||||
parser.add_argument(
|
||||
"--validation_sampling_steps",
|
||||
type=str,
|
||||
default="64",
|
||||
help="use ',' to split multi sampling steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--validation_guidance_scale",
|
||||
type=str,
|
||||
default="4.5",
|
||||
help="use ',' to split multi scale",
|
||||
)
|
||||
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("--output_dir", type=str, default=None, help="The output directory where the model predictions and checkpoints will be written.")
|
||||
parser.add_argument("--checkpoints_total_limit", type=int, default=None, help=("Max number of checkpoints to store."))
|
||||
parser.add_argument("--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`."
|
||||
),
|
||||
)
|
||||
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.'
|
||||
),
|
||||
)
|
||||
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.'
|
||||
),
|
||||
)
|
||||
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***."
|
||||
),
|
||||
)
|
||||
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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoints_total_limit",
|
||||
type=int,
|
||||
default=None,
|
||||
help=("Max number of checkpoints to store."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--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`."),
|
||||
)
|
||||
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.'
|
||||
),
|
||||
)
|
||||
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.'
|
||||
),
|
||||
)
|
||||
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***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
parser.add_argument("--max_train_steps", type=int, default=None, 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.")
|
||||
parser.add_argument("--learning_rate", type=float, default=1e-4, 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.")
|
||||
parser.add_argument("--lr_warmup_steps", type=int, 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("--gradient_checkpointing", action="store_true", help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.")
|
||||
parser.add_argument(
|
||||
"--max_train_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--learning_rate",
|
||||
type=float,
|
||||
default=1e-4,
|
||||
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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr_warmup_steps",
|
||||
type=int,
|
||||
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(
|
||||
"--gradient_checkpointing",
|
||||
action="store_true",
|
||||
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"
|
||||
),
|
||||
)
|
||||
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."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--use_cpu_offload", action="store_true", help="Whether to use CPU offload for param & gradient & optimizer states.")
|
||||
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"
|
||||
),
|
||||
)
|
||||
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."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_cpu_offload",
|
||||
action="store_true",
|
||||
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("--train_sp_batch_size", type=int, default=1, help="Batch size for sequence parallel training")
|
||||
parser.add_argument("--sp_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="For sequence parallel")
|
||||
parser.add_argument(
|
||||
"--train_sp_batch_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Batch size for sequence parallel training",
|
||||
)
|
||||
|
||||
parser.add_argument("--use_lora", action="store_true", 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(
|
||||
"--use_lora",
|
||||
action="store_true",
|
||||
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("--fsdp_sharding_startegy", default="full")
|
||||
|
||||
parser.add_argument(
|
||||
@@ -457,26 +705,55 @@ if __name__ == "__main__":
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme."
|
||||
"--logit_mean",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="mean to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme."
|
||||
"--logit_std",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="std to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--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"]'
|
||||
),
|
||||
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"]'),
|
||||
)
|
||||
parser.add_argument("--lr_num_cycles", type=int, default=1, help="Number of cycles in the learning rate scheduler.")
|
||||
parser.add_argument("--lr_power", type=float, 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(
|
||||
"--lr_num_cycles",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of cycles in the learning rate scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr_power",
|
||||
type=float,
|
||||
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(
|
||||
"--master_weight_type",
|
||||
type=str,
|
||||
default="fp32",
|
||||
help="Weight type to use - fp32 or bf16.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
|
||||
+180
-111
@@ -1,33 +1,51 @@
|
||||
# import
|
||||
import os
|
||||
# import
|
||||
import json
|
||||
import torch
|
||||
from fastvideo.utils.logging import main_print
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, StateDictType, FullStateDictConfig
|
||||
from safetensors.torch import save_file, load_file
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
from torch.distributed.checkpoint.default_planner import DefaultSavePlanner, DefaultLoadPlanner
|
||||
from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_dict
|
||||
from torch.distributed.fsdp import FullOptimStateDictConfig
|
||||
def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=False):
|
||||
import os
|
||||
|
||||
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.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):
|
||||
with FSDP.state_dict_type(
|
||||
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
optim_state = FSDP.optim_state_dict(
|
||||
model,
|
||||
model,
|
||||
optimizer,
|
||||
)
|
||||
|
||||
#todo move to get_state_dict
|
||||
|
||||
# todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
# 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')
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
@@ -35,35 +53,72 @@ def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=Fals
|
||||
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)
|
||||
|
||||
|
||||
|
||||
def save_checkpoint_generator_discriminator(model, optimizer, discriminator, discriminator_optimizer, rank, output_dir, step,):
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
|
||||
|
||||
def save_checkpoint(transformer, rank, output_dir, step):
|
||||
main_print(f"--> saving checkpoint at step {step}")
|
||||
with FSDP.state_dict_type(
|
||||
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
# todo move to get_state_dict
|
||||
if rank <= 0:
|
||||
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")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(transformer.config)
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"] # TODO
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
|
||||
|
||||
def save_checkpoint_generator_discriminator(
|
||||
model,
|
||||
optimizer,
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
rank,
|
||||
output_dir,
|
||||
step,
|
||||
):
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
|
||||
#todo move to get_state_dict
|
||||
|
||||
# todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
hf_weight_dir = os.path.join(save_dir, "hf_weights")
|
||||
os.makedirs(hf_weight_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
# save using safetensors
|
||||
if rank <= 0:
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(hf_weight_dir, "config.json")
|
||||
# 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}")
|
||||
model_weight_dir = os.path.join(save_dir, "model_weights_state")
|
||||
os.makedirs(model_weight_dir, exist_ok=True)
|
||||
@@ -74,44 +129,53 @@ def save_checkpoint_generator_discriminator(model, optimizer, discriminator, dis
|
||||
model_state = model.state_dict()
|
||||
weight_state_dict = {"model": model_state}
|
||||
dist_cp.save_state_dict(
|
||||
state_dict=weight_state_dict,
|
||||
storage_writer=dist_cp.FileSystemWriter(model_weight_dir),
|
||||
planner=DefaultSavePlanner(),
|
||||
state_dict=weight_state_dict,
|
||||
storage_writer=dist_cp.FileSystemWriter(model_weight_dir),
|
||||
planner=DefaultSavePlanner(),
|
||||
)
|
||||
optimizer_state_dict = {"optimizer": optim_state}
|
||||
dist_cp.save_state_dict(
|
||||
state_dict=optimizer_state_dict,
|
||||
storage_writer=dist_cp.FileSystemWriter(model_optimizer_dir),
|
||||
planner=DefaultSavePlanner(),
|
||||
state_dict=optimizer_state_dict,
|
||||
storage_writer=dist_cp.FileSystemWriter(model_optimizer_dir),
|
||||
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, StateDictType.FULL_STATE_DICT, 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)
|
||||
with FSDP.state_dict_type(
|
||||
discriminator,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
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)
|
||||
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")
|
||||
if rank <= 0:
|
||||
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")
|
||||
|
||||
|
||||
|
||||
def load_sharded_model(model, optimizer, model_dir, optimizer_dir):
|
||||
with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
|
||||
weight_state_dict = {"model": model.state_dict()}
|
||||
|
||||
|
||||
optim_state = load_sharded_optimizer_state_dict(
|
||||
model_state_dict=weight_state_dict["model"],
|
||||
optimizer_key="optimizer",
|
||||
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,
|
||||
state_dict=weight_state_dict,
|
||||
storage_reader=dist_cp.FileSystemReader(model_dir),
|
||||
planner=DefaultLoadPlanner(),
|
||||
)
|
||||
@@ -120,38 +184,59 @@ def load_sharded_model(model, optimizer, model_dir, optimizer_dir):
|
||||
main_print(f"--> loaded model and optimizer from path {model_dir}")
|
||||
return model, optimizer
|
||||
|
||||
|
||||
def load_full_state_model(model, optimizer, checkpoint_file, rank):
|
||||
with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)):
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
discriminator_state = torch.load(checkpoint_file)
|
||||
model_state = discriminator_state["model"]
|
||||
if rank <= 0:
|
||||
if rank <= 0:
|
||||
optim_state = discriminator_state["optimizer"]
|
||||
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(
|
||||
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
current_state = model.state_dict()
|
||||
current_state.update(model_weights)
|
||||
@@ -162,83 +247,67 @@ def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
|
||||
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
|
||||
)
|
||||
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
|
||||
):
|
||||
main_print(f"--> saving LoRA checkpoint at step {step}")
|
||||
def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step,
|
||||
pipeline):
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
full_state_dict = transformer.state_dict()
|
||||
lora_state_dict = {
|
||||
k: v for k, v in full_state_dict.items()
|
||||
if 'lora' in k.lower()
|
||||
}
|
||||
lora_optim_state = FSDP.optim_state_dict(
|
||||
transformer,
|
||||
transformer,
|
||||
optimizer,
|
||||
)
|
||||
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
weight_path = os.path.join(save_dir, "lora_weights.safetensors")
|
||||
save_file(lora_state_dict, weight_path)
|
||||
|
||||
# save optimizer
|
||||
optim_path = os.path.join(save_dir, "lora_optimizer.pt")
|
||||
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)
|
||||
pipeline.save_lora_weights(
|
||||
save_directory=save_dir,
|
||||
transformer_lora_layers=transformer_lora_layers,
|
||||
is_main_process=True,
|
||||
)
|
||||
# save config
|
||||
lora_config = {
|
||||
'step': step,
|
||||
'lora_params': {
|
||||
'lora_rank': transformer.config.lora_rank,
|
||||
'lora_alpha': transformer.config.lora_alpha,
|
||||
'target_modules': transformer.config.lora_target_modules
|
||||
}
|
||||
"step": step,
|
||||
"lora_params": {
|
||||
"lora_rank": transformer.config.lora_rank,
|
||||
"lora_alpha": transformer.config.lora_alpha,
|
||||
"target_modules": transformer.config.lora_target_modules,
|
||||
},
|
||||
}
|
||||
config_path = os.path.join(save_dir, "lora_config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(lora_config, f, indent=4)
|
||||
main_print(f"--> LoRA checkpoint saved at step {step}")
|
||||
|
||||
def resume_lora_training(
|
||||
transformer,
|
||||
checkpoint_dir,
|
||||
optimizer
|
||||
):
|
||||
weight_path = os.path.join(checkpoint_dir, "lora_weights.safetensors")
|
||||
lora_weights = load_file(weight_path)
|
||||
|
||||
def resume_lora_optimizer(transformer, checkpoint_dir, optimizer):
|
||||
config_path = os.path.join(checkpoint_dir, "lora_config.json")
|
||||
with open(config_path, "r") as f:
|
||||
config_dict = json.load(f)
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
current_state = transformer.state_dict()
|
||||
current_state.update(lora_weights)
|
||||
transformer.load_state_dict(current_state, strict=False)
|
||||
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
|
||||
)
|
||||
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 training from step {step}")
|
||||
return transformer, optimizer, step
|
||||
step = config_dict["step"]
|
||||
main_print(f"--> Successfully resuming LoRA optimizer from step {step}")
|
||||
return transformer, optimizer, step
|
||||
|
||||
@@ -3,21 +3,24 @@
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from typing import Any, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
from typing import Any, Tuple
|
||||
from torch import Tensor
|
||||
from torch.nn import Module
|
||||
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
|
||||
def broadcast(input_: torch.Tensor):
|
||||
src = nccl_info.group_id * nccl_info.sp_size
|
||||
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
|
||||
|
||||
@@ -44,11 +47,8 @@ def _all_to_all_4D(
|
||||
|
||||
# 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,7 +62,8 @@ def _all_to_all_4D(
|
||||
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
|
||||
|
||||
@@ -75,13 +76,10 @@ def _all_to_all_4D(
|
||||
|
||||
# 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
|
||||
@@ -96,14 +94,17 @@ def _all_to_all_4D(
|
||||
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):
|
||||
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx: Any,
|
||||
@@ -112,7 +113,6 @@ class SeqAllToAll4D(torch.autograd.Function):
|
||||
scatter_idx: int,
|
||||
gather_idx: int,
|
||||
) -> Tensor:
|
||||
|
||||
ctx.group = group
|
||||
ctx.scatter_idx = scatter_idx
|
||||
ctx.gather_idx = gather_idx
|
||||
@@ -120,27 +120,26 @@ 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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def all_to_all_4D(
|
||||
input_: torch.Tensor,
|
||||
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(
|
||||
input_: torch.Tensor,
|
||||
world_size: int,
|
||||
@@ -148,7 +147,10 @@ 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()
|
||||
@@ -170,7 +172,8 @@ 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
|
||||
@@ -198,7 +201,6 @@ def all_to_all(
|
||||
return _AllToAll.apply(input_, nccl_info.group, scatter_dim, gather_dim)
|
||||
|
||||
|
||||
|
||||
class _AllGather(torch.autograd.Function):
|
||||
"""All-gather communication with autograd support.
|
||||
|
||||
@@ -237,6 +239,7 @@ class _AllGather(torch.autograd.Function):
|
||||
|
||||
return grad_input, None
|
||||
|
||||
|
||||
def all_gather(input_: torch.Tensor, dim: int = 1):
|
||||
"""Performs an all-gather operation on the input tensor along the specified dimension.
|
||||
|
||||
@@ -250,49 +253,82 @@ 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, encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
def prepare(hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask):
|
||||
|
||||
return (
|
||||
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)
|
||||
return hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
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,
|
||||
attention_mask,
|
||||
encoder_attention_mask,
|
||||
)
|
||||
|
||||
sp_size = nccl_info.sp_size
|
||||
frame = hidden_states.shape[2]
|
||||
assert frame % sp_size == 0, "frame should be a multiple of sp_size"
|
||||
|
||||
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask = prepare(hidden_states,
|
||||
encoder_hidden_states.repeat(1, sp_size, 1),
|
||||
attention_mask.repeat(1, sp_size, 1, 1),
|
||||
encoder_attention_mask.repeat(1, sp_size))
|
||||
|
||||
(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
attention_mask,
|
||||
encoder_attention_mask,
|
||||
) = prepare(
|
||||
hidden_states,
|
||||
encoder_hidden_states.repeat(1, sp_size, 1),
|
||||
attention_mask.repeat(1, sp_size, 1, 1),
|
||||
encoder_attention_mask.repeat(1, sp_size),
|
||||
)
|
||||
|
||||
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):
|
||||
while True:
|
||||
for data_item in dataloader:
|
||||
latents, cond,attn_mask, cond_mask = data_item
|
||||
latents = latents.to(device)
|
||||
cond = cond.to(device)
|
||||
attn_mask = attn_mask.to(device)
|
||||
cond_mask = cond_mask.to(device)
|
||||
frame = latents.shape[2]
|
||||
if frame == 1:
|
||||
yield latents, cond, attn_mask, cond_mask
|
||||
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):
|
||||
st_idx = iter * train_sp_batch_size
|
||||
ed_idx = (iter + 1) * train_sp_batch_size
|
||||
encoder_hidden_states=cond[st_idx: ed_idx]
|
||||
attention_mask=attn_mask[st_idx: ed_idx]
|
||||
encoder_attention_mask=cond_mask[st_idx: ed_idx]
|
||||
yield latents[st_idx: ed_idx], encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
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
|
||||
latents = latents.to(device)
|
||||
cond = cond.to(device)
|
||||
attn_mask = attn_mask.to(device)
|
||||
cond_mask = cond_mask.to(device)
|
||||
frame = latents.shape[2]
|
||||
if frame == 1:
|
||||
yield latents, cond, attn_mask, cond_mask
|
||||
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):
|
||||
st_idx = iter * train_sp_batch_size
|
||||
ed_idx = (iter + 1) * train_sp_batch_size
|
||||
encoder_hidden_states = cond[st_idx:ed_idx]
|
||||
attention_mask = attn_mask[st_idx:ed_idx]
|
||||
encoder_attention_mask = cond_mask[st_idx:ed_idx]
|
||||
yield (
|
||||
latents[st_idx:ed_idx],
|
||||
encoder_hidden_states,
|
||||
attention_mask,
|
||||
encoder_attention_mask,
|
||||
)
|
||||
|
||||
@@ -1,117 +0,0 @@
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate.logging import get_logger
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from diffusers.utils import export_to_video
|
||||
import json
|
||||
import os
|
||||
import torch.distributed as dist
|
||||
logger = get_logger(__name__)
|
||||
from torch.utils.data import Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
class T5dataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
vae_debug,
|
||||
):
|
||||
self.json_path = json_path
|
||||
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'])
|
||||
def __getitem__(self, idx):
|
||||
caption = self.train_dataset[idx]['caption']
|
||||
filename = self.train_dataset[idx]['latent_path'].split('.')[0]
|
||||
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']), map_location="cpu")
|
||||
else:
|
||||
latents = []
|
||||
|
||||
return dict(caption=caption, latents=latents, filename=filename, length=length)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.train_dataset)
|
||||
|
||||
def main(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)
|
||||
|
||||
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)
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path).to(device)
|
||||
pipe.vae.enable_tiling()
|
||||
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)
|
||||
|
||||
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
json_data = []
|
||||
for _, data in enumerate(train_dataloader):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
prompt_embeds, prompt_attention_mask, _, _ = pipe.encode_prompt(
|
||||
prompt=data['caption'],
|
||||
)
|
||||
if args.vae_debug:
|
||||
latents = data['latents']
|
||||
video = pipe.vae.decode(latents.to(device), return_dict=False)[0]
|
||||
video = pipe.video_processor.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")
|
||||
# save latent
|
||||
torch.save(prompt_embeds[idx], prompt_embed_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=30)
|
||||
item = {}
|
||||
item['length'] = int(data['length'][idx])
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["prompt_embed_path"] = video_name + ".pt"
|
||||
item["prompt_attention_mask"] = video_name + ".pt"
|
||||
item["caption"] = data['caption'][idx]
|
||||
json_data.append(item)
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
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:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--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.")
|
||||
parser.add_argument("--train_batch_size", type=int, 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("--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.")
|
||||
parser.add_argument("--vae_debug",action="store_true")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,106 +0,0 @@
|
||||
from fastvideo.dataset import getdataset
|
||||
from torch.utils.data import DataLoader
|
||||
from fastvideo.utils.dataset_utils import Collate
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate import Accelerator
|
||||
from accelerate.logging import get_logger
|
||||
from accelerate.utils import ProjectConfiguration
|
||||
import json
|
||||
import os
|
||||
from diffusers import AutoencoderKLMochi
|
||||
import torch.distributed as dist
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
def main(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)
|
||||
args.ae_stride_t, args.ae_stride_h, args.ae_stride_w = 4, 8, 8
|
||||
args.ae_stride = args.ae_stride_h
|
||||
patch_size_t, patch_size_h, patch_size_w = 1, 2, 2
|
||||
args.patch_size = patch_size_h
|
||||
args.patch_size_t, args.patch_size_h, args.patch_size_w = patch_size_t, patch_size_h, patch_size_w
|
||||
accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=args.logging_dir)
|
||||
accelerator = Accelerator(
|
||||
project_config=accelerator_project_config,
|
||||
)
|
||||
train_dataset = getdataset(args)
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
|
||||
encoder_device = torch.device(f"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)
|
||||
vae = AutoencoderKLMochi.from_pretrained(args.model_path, subfolder="vae").to("cuda")
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
json_data = []
|
||||
for _, data in enumerate(train_dataloader):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
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")
|
||||
torch.save(latents[idx].to(torch.bfloat16), latent_path)
|
||||
item = {}
|
||||
item["length"] = latents[idx].shape[1]
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["caption"] = data['text'][idx]
|
||||
json_data.append(item)
|
||||
print(f"{video_name} processed")
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
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:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--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.")
|
||||
parser.add_argument("--train_batch_size", type=int, 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("--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("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default='t2v')
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
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("--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.")
|
||||
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***."
|
||||
),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,23 +1,22 @@
|
||||
import math
|
||||
from einops import rearrange
|
||||
import random
|
||||
from collections import Counter
|
||||
from typing import List, Optional
|
||||
|
||||
import decord
|
||||
from torch.nn import functional as F
|
||||
import torch
|
||||
from typing import Optional
|
||||
import torch.utils
|
||||
import torch.utils.data
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
from torch.utils.data import Sampler
|
||||
from typing import List
|
||||
from collections import Counter
|
||||
import random
|
||||
|
||||
IMG_EXTENSIONS = [".jpg", ".JPG", ".jpeg", ".JPEG", ".png", ".PNG"]
|
||||
|
||||
IMG_EXTENSIONS = ['.jpg', '.JPG', '.jpeg', '.JPEG', '.png', '.PNG']
|
||||
|
||||
def is_image_file(filename):
|
||||
return any(filename.endswith(extension) for extension in IMG_EXTENSIONS)
|
||||
|
||||
|
||||
class DecordInit(object):
|
||||
"""Using Decord(https://github.com/dmlc/decord) to initialize the video_reader."""
|
||||
|
||||
@@ -37,11 +36,12 @@ class DecordInit(object):
|
||||
return reader
|
||||
|
||||
def __repr__(self):
|
||||
repr_str = (f'{self.__class__.__name__}('
|
||||
f'sr={self.sr},'
|
||||
f'num_threads={self.num_threads})')
|
||||
repr_str = (f"{self.__class__.__name__}("
|
||||
f"sr={self.sr},"
|
||||
f"num_threads={self.num_threads})")
|
||||
return repr_str
|
||||
|
||||
|
||||
def pad_to_multiple(number, ds_stride):
|
||||
remainder = number % ds_stride
|
||||
if remainder == 0:
|
||||
@@ -50,7 +50,10 @@ def pad_to_multiple(number, ds_stride):
|
||||
padding = ds_stride - remainder
|
||||
return number + padding
|
||||
|
||||
|
||||
# TODO
|
||||
class Collate:
|
||||
|
||||
def __init__(self, args):
|
||||
self.batch_size = args.train_batch_size
|
||||
self.group_frame = args.group_frame
|
||||
@@ -71,9 +74,9 @@ class Collate:
|
||||
self.max_thw = (self.num_frames, self.max_height, self.max_width)
|
||||
|
||||
def package(self, batch):
|
||||
batch_tubes = [i['pixel_values'] for i in batch] # b [c t h w]
|
||||
input_ids = [i['input_ids'] for i in batch] # b [1 l]
|
||||
cond_mask = [i['cond_mask'] for i in batch] # b [1 l]
|
||||
batch_tubes = [i["pixel_values"] for i in batch] # b [c t h w]
|
||||
input_ids = [i["input_ids"] for i in batch] # b [1 l]
|
||||
cond_mask = [i["cond_mask"] for i in batch] # b [1 l]
|
||||
return batch_tubes, input_ids, cond_mask
|
||||
|
||||
def __call__(self, batch):
|
||||
@@ -81,30 +84,65 @@ class Collate:
|
||||
|
||||
ds_stride = self.ae_stride * self.patch_size
|
||||
t_ds_stride = self.ae_stride_t * self.patch_size_t
|
||||
|
||||
pad_batch_tubes, attention_mask, input_ids, cond_mask = self.process(batch_tubes, input_ids, cond_mask, t_ds_stride, ds_stride, self.max_thw, self.ae_stride_thw)
|
||||
assert not torch.any(torch.isnan(pad_batch_tubes)), 'after pad_batch_tubes'
|
||||
|
||||
pad_batch_tubes, attention_mask, input_ids, cond_mask = self.process(
|
||||
batch_tubes,
|
||||
input_ids,
|
||||
cond_mask,
|
||||
t_ds_stride,
|
||||
ds_stride,
|
||||
self.max_thw,
|
||||
self.ae_stride_thw,
|
||||
)
|
||||
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(self, batch_tubes, input_ids, cond_mask, t_ds_stride, ds_stride, max_thw, ae_stride_thw):
|
||||
def process(
|
||||
self,
|
||||
batch_tubes,
|
||||
input_ids,
|
||||
cond_mask,
|
||||
t_ds_stride,
|
||||
ds_stride,
|
||||
max_thw,
|
||||
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]
|
||||
random_select_batch = [random.choice(candidate_batch) for _ in range(len(len_each_batch) - len(candidate_batch))]
|
||||
print(batch_input_size, idx_length_dict, count_dict, sorted_by_value, pick_length, candidate_batch, random_select_batch)
|
||||
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))
|
||||
]
|
||||
print(
|
||||
batch_input_size,
|
||||
idx_length_dict,
|
||||
count_dict,
|
||||
sorted_by_value,
|
||||
pick_length,
|
||||
candidate_batch,
|
||||
random_select_batch,
|
||||
)
|
||||
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]
|
||||
|
||||
@@ -115,50 +153,57 @@ class Collate:
|
||||
max_w = max([i[3] for i in batch_input_size])
|
||||
else:
|
||||
max_t, max_h, max_w = max_thw
|
||||
pad_max_t, pad_max_h, pad_max_w = pad_to_multiple(max_t-1+self.ae_stride_t, t_ds_stride), \
|
||||
pad_to_multiple(max_h, ds_stride), \
|
||||
pad_to_multiple(max_w, ds_stride)
|
||||
pad_max_t, pad_max_h, pad_max_w = (
|
||||
pad_to_multiple(max_t - 1 + self.ae_stride_t, t_ds_stride),
|
||||
pad_to_multiple(max_h, ds_stride),
|
||||
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)
|
||||
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)
|
||||
]
|
||||
]
|
||||
pad_batch_tubes = torch.stack(pad_batch_tubes, dim=0)
|
||||
|
||||
|
||||
max_tube_size = [pad_max_t, pad_max_h, pad_max_w]
|
||||
max_latent_size = [
|
||||
((max_tube_size[0]-1) // ae_stride_thw[0] + 1),
|
||||
((max_tube_size[0] - 1) // ae_stride_thw[0] + 1),
|
||||
max_tube_size[1] // ae_stride_thw[1],
|
||||
max_tube_size[2] // ae_stride_thw[2]
|
||||
]
|
||||
valid_latent_size = [
|
||||
[
|
||||
int(math.ceil((i[1]-1) / ae_stride_thw[0])) + 1,
|
||||
int(math.ceil(i[2] / ae_stride_thw[1])),
|
||||
int(math.ceil(i[3] / ae_stride_thw[2]))
|
||||
] for i in batch_input_size]
|
||||
max_tube_size[2] // ae_stride_thw[2],
|
||||
]
|
||||
valid_latent_size = [[
|
||||
int(math.ceil((i[1] - 1) / ae_stride_thw[0])) + 1,
|
||||
int(math.ceil(i[2] / ae_stride_thw[1])),
|
||||
int(math.ceil(i[3] / ae_stride_thw[2])),
|
||||
] for i in batch_input_size]
|
||||
attention_mask = [
|
||||
F.pad(torch.ones(i, dtype=pad_batch_tubes.dtype), (0, max_latent_size[2] - i[2],
|
||||
0, max_latent_size[1] - i[1],
|
||||
0, max_latent_size[0] - i[0]), value=0) for i in valid_latent_size]
|
||||
F.pad(
|
||||
torch.ones(i, dtype=pad_batch_tubes.dtype),
|
||||
(
|
||||
0,
|
||||
max_latent_size[2] - i[2],
|
||||
0,
|
||||
max_latent_size[1] - i[1],
|
||||
0,
|
||||
max_latent_size[0] - i[0],
|
||||
),
|
||||
value=0,
|
||||
) for i in valid_latent_size
|
||||
]
|
||||
attention_mask = torch.stack(attention_mask) # b t h w
|
||||
if self.batch_size == 1 or self.group_frame or self.group_resolution:
|
||||
assert torch.all(attention_mask.bool())
|
||||
|
||||
|
||||
input_ids = torch.stack(input_ids) # b 1 l
|
||||
cond_mask = torch.stack(cond_mask) # b 1 l
|
||||
|
||||
return pad_batch_tubes, attention_mask, input_ids, cond_mask
|
||||
|
||||
|
||||
|
||||
def split_to_even_chunks(indices, lengths, num_chunks, batch_size):
|
||||
"""
|
||||
Split a list of indices into `chunks` chunks of roughly equal lengths.
|
||||
@@ -184,13 +229,17 @@ 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)
|
||||
print(chunks[idx], "->", chunk)
|
||||
pad_chunks.append(chunk)
|
||||
return pad_chunks
|
||||
|
||||
|
||||
def group_frame_fun(indices, lengths):
|
||||
# sort by num_frames
|
||||
indices.sort(key=lambda i: lengths[i], reverse=True)
|
||||
@@ -204,48 +253,73 @@ def megabatch_frame_alignment(megabatches, lengths):
|
||||
len_each_megabatch = [lengths[i] for i in megabatch]
|
||||
idx_length_dict = dict([*zip(megabatch, len_each_megabatch)])
|
||||
count_dict = Counter(len_each_megabatch)
|
||||
|
||||
|
||||
# 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]
|
||||
random_select_batch = [random.choice(candidate_batch) for i in range(len(idx_length_dict) - len(candidate_batch))]
|
||||
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))
|
||||
]
|
||||
aligned_magabatch = candidate_batch + random_select_batch
|
||||
aligned_magabatches.append(aligned_magabatch)
|
||||
# already aligned megabatches
|
||||
else:
|
||||
aligned_magabatches.append(megabatch)
|
||||
|
||||
return aligned_magabatches
|
||||
|
||||
|
||||
def get_length_grouped_indices(lengths, batch_size, world_size, generator=None, group_frame=False, group_resolution=False, seed=42):
|
||||
return aligned_magabatches
|
||||
|
||||
|
||||
def get_length_grouped_indices(
|
||||
lengths,
|
||||
batch_size,
|
||||
world_size,
|
||||
generator=None,
|
||||
group_frame=False,
|
||||
group_resolution=False,
|
||||
seed=42,
|
||||
):
|
||||
# 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()
|
||||
|
||||
|
||||
# sort dataset according to frame
|
||||
indices = group_frame_fun(indices, lengths)
|
||||
|
||||
|
||||
# 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):
|
||||
@@ -259,9 +333,9 @@ class LengthGroupedSampler(Sampler):
|
||||
batch_size: int,
|
||||
rank: int,
|
||||
world_size: int,
|
||||
lengths: Optional[List[int]] = None,
|
||||
group_frame=False,
|
||||
group_resolution=False,
|
||||
lengths: Optional[List[int]] = None,
|
||||
group_frame=False,
|
||||
group_resolution=False,
|
||||
generator=None,
|
||||
):
|
||||
if lengths is None:
|
||||
@@ -279,8 +353,15 @@ class LengthGroupedSampler(Sampler):
|
||||
return len(self.lengths)
|
||||
|
||||
def __iter__(self):
|
||||
indices = get_length_grouped_indices(self.lengths, self.batch_size, self.world_size, group_frame=self.group_frame,
|
||||
group_resolution=self.group_resolution, generator=self.generator)
|
||||
indices = get_length_grouped_indices(
|
||||
self.lengths,
|
||||
self.batch_size,
|
||||
self.world_size,
|
||||
group_frame=self.group_frame,
|
||||
group_resolution=self.group_resolution,
|
||||
generator=self.generator,
|
||||
)
|
||||
|
||||
def distributed_sampler(lst, rank, batch_size, world_size):
|
||||
result = []
|
||||
index = rank * batch_size
|
||||
@@ -288,6 +369,7 @@ class LengthGroupedSampler(Sampler):
|
||||
result.extend(lst[index:index + batch_size])
|
||||
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,328 +0,0 @@
|
||||
import contextlib
|
||||
import copy
|
||||
import random
|
||||
from typing import Any, Dict, Iterable, List, Optional, Union
|
||||
|
||||
from diffusers.utils import (
|
||||
deprecate,
|
||||
is_torchvision_available,
|
||||
is_transformers_available,
|
||||
)
|
||||
|
||||
if is_transformers_available():
|
||||
import transformers
|
||||
|
||||
if is_torchvision_available():
|
||||
from torchvision import transforms
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
# Adapted from diffusers-style ema https://github.com/huggingface/diffusers/blob/main/src/diffusers/training_utils.py#L263
|
||||
class EMAModel:
|
||||
"""
|
||||
Exponential Moving Average of models weights
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parameters: Iterable[torch.nn.Parameter],
|
||||
decay: float = 0.9999,
|
||||
min_decay: float = 0.0,
|
||||
update_after_step: int = 0,
|
||||
use_ema_warmup: bool = False,
|
||||
inv_gamma: Union[float, int] = 1.0,
|
||||
power: Union[float, int] = 2 / 3,
|
||||
model_cls: Optional[Any] = None,
|
||||
model_config: Dict[str, Any] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
parameters (Iterable[torch.nn.Parameter]): The parameters to track.
|
||||
decay (float): The decay factor for the exponential moving average.
|
||||
min_decay (float): The minimum decay factor for the exponential moving average.
|
||||
update_after_step (int): The number of steps to wait before starting to update the EMA weights.
|
||||
use_ema_warmup (bool): Whether to use EMA warmup.
|
||||
inv_gamma (float):
|
||||
Inverse multiplicative factor of EMA warmup. Default: 1. Only used if `use_ema_warmup` is True.
|
||||
power (float): Exponential factor of EMA warmup. Default: 2/3. Only used if `use_ema_warmup` is True.
|
||||
device (Optional[Union[str, torch.device]]): The device to store the EMA weights on. If None, the EMA
|
||||
weights will be stored on CPU.
|
||||
|
||||
@crowsonkb's notes on EMA Warmup:
|
||||
If gamma=1 and power=1, implements a simple average. gamma=1, power=2/3 are good values for models you plan
|
||||
to train for a million or more steps (reaches decay factor 0.999 at 31.6K steps, 0.9999 at 1M steps),
|
||||
gamma=1, power=3/4 for models you plan to train for less (reaches decay factor 0.999 at 10K steps, 0.9999
|
||||
at 215.4k steps).
|
||||
"""
|
||||
|
||||
if isinstance(parameters, torch.nn.Module):
|
||||
deprecation_message = (
|
||||
"Passing a `torch.nn.Module` to `ExponentialMovingAverage` is deprecated. "
|
||||
"Please pass the parameters of the module instead."
|
||||
)
|
||||
deprecate(
|
||||
"passing a `torch.nn.Module` to `ExponentialMovingAverage`",
|
||||
"1.0.0",
|
||||
deprecation_message,
|
||||
standard_warn=False,
|
||||
)
|
||||
parameters = parameters.parameters()
|
||||
|
||||
# set use_ema_warmup to True if a torch.nn.Module is passed for backwards compatibility
|
||||
use_ema_warmup = True
|
||||
|
||||
if kwargs.get("max_value", None) is not None:
|
||||
deprecation_message = "The `max_value` argument is deprecated. Please use `decay` instead."
|
||||
deprecate("max_value", "1.0.0", deprecation_message, standard_warn=False)
|
||||
decay = kwargs["max_value"]
|
||||
|
||||
if kwargs.get("min_value", None) is not None:
|
||||
deprecation_message = "The `min_value` argument is deprecated. Please use `min_decay` instead."
|
||||
deprecate("min_value", "1.0.0", deprecation_message, standard_warn=False)
|
||||
min_decay = kwargs["min_value"]
|
||||
|
||||
parameters = list(parameters)
|
||||
self.shadow_params = [p.clone().detach() for p in parameters]
|
||||
|
||||
if kwargs.get("device", None) is not None:
|
||||
deprecation_message = "The `device` argument is deprecated. Please use `to` instead."
|
||||
deprecate("device", "1.0.0", deprecation_message, standard_warn=False)
|
||||
self.to(device=kwargs["device"])
|
||||
|
||||
self.temp_stored_params = None
|
||||
|
||||
self.decay = decay
|
||||
self.min_decay = min_decay
|
||||
self.update_after_step = update_after_step
|
||||
self.use_ema_warmup = use_ema_warmup
|
||||
self.inv_gamma = inv_gamma
|
||||
self.power = power
|
||||
self.optimization_step = 0
|
||||
self.cur_decay_value = None # set in `step()`
|
||||
|
||||
self.model_cls = model_cls
|
||||
self.model_config = model_config
|
||||
|
||||
@classmethod
|
||||
def extract_ema_kwargs(cls, kwargs):
|
||||
"""
|
||||
Extracts the EMA kwargs from the kwargs of a class method.
|
||||
"""
|
||||
ema_kwargs = {}
|
||||
for key in [
|
||||
"decay",
|
||||
"min_decay",
|
||||
"optimization_step",
|
||||
"update_after_step",
|
||||
"use_ema_warmup",
|
||||
"inv_gamma",
|
||||
"power",
|
||||
]:
|
||||
if kwargs.get(key, None) is not None:
|
||||
ema_kwargs[key] = kwargs.pop(key)
|
||||
return ema_kwargs
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, path, model_cls) -> "EMAModel":
|
||||
config = model_cls.load_config(path)
|
||||
ema_kwargs = cls.extract_ema_kwargs(config)
|
||||
model = model_cls.from_pretrained(path)
|
||||
|
||||
ema_model = cls(model.parameters(), model_cls=model_cls, model_config=config)
|
||||
|
||||
ema_model.load_state_dict(ema_kwargs)
|
||||
return ema_model
|
||||
|
||||
def save_pretrained(self, path):
|
||||
if self.model_cls is None:
|
||||
raise ValueError("`save_pretrained` can only be used if `model_cls` was defined at __init__.")
|
||||
|
||||
if self.model_config is None:
|
||||
raise ValueError("`save_pretrained` can only be used if `model_config` was defined at __init__.")
|
||||
|
||||
model = self.model_cls.from_config(self.model_config)
|
||||
state_dict = self.state_dict()
|
||||
state_dict.pop("shadow_params", None)
|
||||
|
||||
model.register_to_config(**state_dict)
|
||||
self.copy_to(model.parameters())
|
||||
model.save_pretrained(path)
|
||||
|
||||
def get_decay(self, optimization_step: int) -> float:
|
||||
"""
|
||||
Compute the decay factor for the exponential moving average.
|
||||
"""
|
||||
step = max(0, optimization_step - self.update_after_step - 1)
|
||||
|
||||
if step <= 0:
|
||||
return 0.0
|
||||
|
||||
if self.use_ema_warmup:
|
||||
cur_decay_value = 1 - (1 + step / self.inv_gamma) ** -self.power
|
||||
else:
|
||||
cur_decay_value = (1 + step) / (10 + step)
|
||||
|
||||
cur_decay_value = min(cur_decay_value, self.decay)
|
||||
# make sure decay is not smaller than min_decay
|
||||
cur_decay_value = max(cur_decay_value, self.min_decay)
|
||||
return cur_decay_value
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, parameters: Iterable[torch.nn.Parameter]):
|
||||
if isinstance(parameters, torch.nn.Module):
|
||||
deprecation_message = (
|
||||
"Passing a `torch.nn.Module` to `ExponentialMovingAverage.step` is deprecated. "
|
||||
"Please pass the parameters of the module instead."
|
||||
)
|
||||
deprecate(
|
||||
"passing a `torch.nn.Module` to `ExponentialMovingAverage.step`",
|
||||
"1.0.0",
|
||||
deprecation_message,
|
||||
standard_warn=False,
|
||||
)
|
||||
parameters = parameters.parameters()
|
||||
|
||||
parameters = list(parameters)
|
||||
|
||||
self.optimization_step += 1
|
||||
|
||||
# Compute the decay factor for the exponential moving average.
|
||||
decay = self.get_decay(self.optimization_step)
|
||||
self.cur_decay_value = decay
|
||||
one_minus_decay = 1 - decay
|
||||
|
||||
context_manager = contextlib.nullcontext
|
||||
if is_transformers_available() and transformers.integrations.is_deepspeed_zero3_enabled():
|
||||
import deepspeed
|
||||
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
if is_transformers_available() and transformers.integrations.is_deepspeed_zero3_enabled():
|
||||
context_manager = deepspeed.zero.GatheredParameters(param, modifier_rank=None)
|
||||
|
||||
with context_manager():
|
||||
if param.requires_grad:
|
||||
s_param.sub_(one_minus_decay * (s_param - param))
|
||||
else:
|
||||
s_param.copy_(param)
|
||||
|
||||
def copy_to(self, parameters: Iterable[torch.nn.Parameter]) -> None:
|
||||
"""
|
||||
Copy current averaged parameters into given collection of parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored moving averages. If `None`, the parameters with which this
|
||||
`ExponentialMovingAverage` was initialized will be used.
|
||||
"""
|
||||
parameters = list(parameters)
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
param.data.copy_(s_param.to(param.device).data)
|
||||
|
||||
|
||||
def to(self, device=None, dtype=None) -> None:
|
||||
r"""Move internal buffers of the ExponentialMovingAverage to `device`.
|
||||
|
||||
Args:
|
||||
device: like `device` argument to `torch.Tensor.to`
|
||||
"""
|
||||
# .to() on the tensors handles None correctly
|
||||
self.shadow_params = [
|
||||
p.to(device=device, dtype=dtype) if p.is_floating_point() else p.to(device=device)
|
||||
for p in self.shadow_params
|
||||
]
|
||||
|
||||
def state_dict(self) -> dict:
|
||||
r"""
|
||||
Returns the state of the ExponentialMovingAverage as a dict. This method is used by accelerate during
|
||||
checkpointing to save the ema state dict.
|
||||
"""
|
||||
# Following PyTorch conventions, references to tensors are returned:
|
||||
# "returns a reference to the state and not its copy!" -
|
||||
# https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict
|
||||
return {
|
||||
"decay": self.decay,
|
||||
"min_decay": self.min_decay,
|
||||
"optimization_step": self.optimization_step,
|
||||
"update_after_step": self.update_after_step,
|
||||
"use_ema_warmup": self.use_ema_warmup,
|
||||
"inv_gamma": self.inv_gamma,
|
||||
"power": self.power,
|
||||
"shadow_params": self.shadow_params,
|
||||
}
|
||||
|
||||
def store(self, parameters: Iterable[torch.nn.Parameter]) -> None:
|
||||
r"""
|
||||
Args:
|
||||
Save the current parameters for restoring later.
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
temporarily stored.
|
||||
"""
|
||||
self.temp_stored_params = [param.detach().cpu().clone() for param in parameters]
|
||||
|
||||
def restore(self, parameters: Iterable[torch.nn.Parameter]) -> None:
|
||||
r"""
|
||||
Args:
|
||||
Restore the parameters stored with the `store` method. Useful to validate the model with EMA parameters without:
|
||||
affecting the original optimization process. Store the parameters before the `copy_to()` method. After
|
||||
validation (or model saving), use this to restore the former parameters.
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters. If `None`, the parameters with which this
|
||||
`ExponentialMovingAverage` was initialized will be used.
|
||||
"""
|
||||
if self.temp_stored_params is None:
|
||||
raise RuntimeError("This ExponentialMovingAverage has no `store()`ed weights " "to `restore()`")
|
||||
for c_param, param in zip(self.temp_stored_params, parameters):
|
||||
param.data.copy_(c_param.data)
|
||||
|
||||
# Better memory-wise.
|
||||
self.temp_stored_params = None
|
||||
|
||||
def load_state_dict(self, state_dict: dict) -> None:
|
||||
r"""
|
||||
Args:
|
||||
Loads the ExponentialMovingAverage state. This method is used by accelerate during checkpointing to save the
|
||||
ema state dict.
|
||||
state_dict (dict): EMA state. Should be an object returned
|
||||
from a call to :meth:`state_dict`.
|
||||
"""
|
||||
# deepcopy, to be consistent with module API
|
||||
state_dict = copy.deepcopy(state_dict)
|
||||
|
||||
self.decay = state_dict.get("decay", self.decay)
|
||||
if self.decay < 0.0 or self.decay > 1.0:
|
||||
raise ValueError("Decay must be between 0 and 1")
|
||||
|
||||
self.min_decay = state_dict.get("min_decay", self.min_decay)
|
||||
if not isinstance(self.min_decay, float):
|
||||
raise ValueError("Invalid min_decay")
|
||||
|
||||
self.optimization_step = state_dict.get("optimization_step", self.optimization_step)
|
||||
if not isinstance(self.optimization_step, int):
|
||||
raise ValueError("Invalid optimization_step")
|
||||
|
||||
self.update_after_step = state_dict.get("update_after_step", self.update_after_step)
|
||||
if not isinstance(self.update_after_step, int):
|
||||
raise ValueError("Invalid update_after_step")
|
||||
|
||||
self.use_ema_warmup = state_dict.get("use_ema_warmup", self.use_ema_warmup)
|
||||
if not isinstance(self.use_ema_warmup, bool):
|
||||
raise ValueError("Invalid use_ema_warmup")
|
||||
|
||||
self.inv_gamma = state_dict.get("inv_gamma", self.inv_gamma)
|
||||
if not isinstance(self.inv_gamma, (float, int)):
|
||||
raise ValueError("Invalid inv_gamma")
|
||||
|
||||
self.power = state_dict.get("power", self.power)
|
||||
if not isinstance(self.power, (float, int)):
|
||||
raise ValueError("Invalid power")
|
||||
|
||||
shadow_params = state_dict.get("shadow_params", None)
|
||||
if shadow_params is not None:
|
||||
self.shadow_params = shadow_params
|
||||
if not isinstance(self.shadow_params, list):
|
||||
raise ValueError("shadow_params must be a list")
|
||||
if not all(isinstance(p, torch.Tensor) for p in self.shadow_params):
|
||||
raise ValueError("shadow_params must all be Tensors")
|
||||
@@ -0,0 +1,40 @@
|
||||
import platform
|
||||
|
||||
import accelerate
|
||||
import peft
|
||||
import torch
|
||||
import transformers
|
||||
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
|
||||
|
||||
VERSION = "1.2.0"
|
||||
|
||||
if __name__ == "__main__":
|
||||
info = {
|
||||
"FastVideo version": VERSION,
|
||||
"Platform": platform.platform(),
|
||||
"Python version": platform.python_version(),
|
||||
"PyTorch version": torch.__version__,
|
||||
"Transformers version": transformers.__version__,
|
||||
"Accelerate version": accelerate.__version__,
|
||||
"PEFT version": peft.__version__,
|
||||
}
|
||||
|
||||
if is_torch_cuda_available():
|
||||
info["PyTorch version"] += " (GPU)"
|
||||
info["GPU type"] = torch.cuda.get_device_name()
|
||||
|
||||
if is_torch_npu_available():
|
||||
info["PyTorch version"] += " (NPU)"
|
||||
info["NPU type"] = torch.npu.get_device_name()
|
||||
info["CANN version"] = torch.version.cann # codespell:ignore
|
||||
|
||||
try:
|
||||
import bitsandbytes
|
||||
|
||||
info["Bitsandbytes version"] = bitsandbytes.__version__
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
print("\n" +
|
||||
"\n".join([f"- {key}: {value}"
|
||||
for key, value in info.items()]) + "\n")
|
||||
@@ -1,31 +1,16 @@
|
||||
from sympy import use
|
||||
import torch
|
||||
import os
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
|
||||
checkpoint_wrapper,
|
||||
CheckpointImpl,
|
||||
apply_activation_checkpointing,
|
||||
)
|
||||
from peft.utils.other import fsdp_auto_wrap_policy
|
||||
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig, # general model non-sharded, non-flattened params
|
||||
LocalStateDictConfig, # flattened params, usable only by FSDP
|
||||
# ShardedStateDictConfig, # un-flattened param but shards, usable by other parallel schemes.
|
||||
)
|
||||
|
||||
from fastvideo.model.modeling_mochi import MochiTransformerBlock
|
||||
|
||||
# ruff: noqa: E731
|
||||
import functools
|
||||
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.fsdp import MixedPrecision, ShardingStrategy
|
||||
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
|
||||
|
||||
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
|
||||
import functools
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformerBlock
|
||||
from fastvideo.utils.load import get_no_split_modules
|
||||
|
||||
non_reentrant_wrapper = partial(
|
||||
checkpoint_wrapper,
|
||||
@@ -35,13 +20,12 @@ non_reentrant_wrapper = partial(
|
||||
check_fn = lambda submodule: isinstance(submodule, MochiTransformerBlock)
|
||||
|
||||
|
||||
def apply_fsdp_checkpointing(model, p=1):
|
||||
def apply_fsdp_checkpointing(model, no_split_modules, p=1):
|
||||
# https://github.com/foundation-model-stack/fms-fsdp/blob/408c7516d69ea9b6bcd4c0f5efab26c0f64b3c2d/fms_fsdp/policies/ac_handler.py#L16
|
||||
"""apply activation checkpointing to model
|
||||
returns None as model is updated directly
|
||||
"""
|
||||
print(f"--> applying fdsp activation checkpointing...")
|
||||
|
||||
print("--> applying fdsp activation checkpointing...")
|
||||
block_idx = 0
|
||||
cut_off = 1 / 2
|
||||
# when passing p as a fraction number (e.g. 1/3), it will be interpreted
|
||||
@@ -52,44 +36,52 @@ def apply_fsdp_checkpointing(model, p=1):
|
||||
nonlocal block_idx
|
||||
nonlocal cut_off
|
||||
|
||||
if isinstance(submodule, MochiTransformerBlock):
|
||||
if isinstance(submodule, no_split_modules):
|
||||
block_idx += 1
|
||||
if block_idx * p >= cut_off:
|
||||
cut_off += 1
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
|
||||
apply_activation_checkpointing(
|
||||
model, checkpoint_wrapper_fn=non_reentrant_wrapper, check_fn=selective_checkpointing
|
||||
model,
|
||||
checkpoint_wrapper_fn=non_reentrant_wrapper,
|
||||
check_fn=selective_checkpointing,
|
||||
)
|
||||
|
||||
|
||||
float32 = MixedPrecision(
|
||||
param_dtype=torch.float32,
|
||||
# Gradient communication precision.
|
||||
reduce_dtype=torch.float32,
|
||||
# Buffer precision.
|
||||
buffer_dtype=torch.float32,
|
||||
cast_forward_inputs=False
|
||||
)
|
||||
def get_mixed_precision(master_weight_type="fp32"):
|
||||
weight_type = torch.float32 if master_weight_type == "fp32" else torch.bfloat16
|
||||
mixed_precision = MixedPrecision(
|
||||
param_dtype=weight_type,
|
||||
# Gradient communication precision.
|
||||
reduce_dtype=weight_type,
|
||||
# Buffer precision.
|
||||
buffer_dtype=weight_type,
|
||||
cast_forward_inputs=False,
|
||||
)
|
||||
return mixed_precision
|
||||
|
||||
|
||||
|
||||
def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
|
||||
def get_dit_fsdp_kwargs(
|
||||
transformer,
|
||||
sharding_strategy,
|
||||
use_lora=False,
|
||||
cpu_offload=False,
|
||||
master_weight_type="fp32",
|
||||
):
|
||||
no_split_modules = get_no_split_modules(transformer)
|
||||
if use_lora:
|
||||
auto_wrap_policy = fsdp_auto_wrap_policy
|
||||
else:
|
||||
auto_wrap_policy = functools.partial(
|
||||
transformer_auto_wrap_policy,
|
||||
transformer_layer_cls={
|
||||
MochiTransformerBlock,
|
||||
},
|
||||
transformer_layer_cls=no_split_modules,
|
||||
)
|
||||
|
||||
|
||||
# we use float32 for fsdp but autocast during training
|
||||
mixed_precision = float32
|
||||
|
||||
mixed_precision = get_mixed_precision(master_weight_type)
|
||||
|
||||
if sharding_strategy == "full":
|
||||
sharding_strategy = ShardingStrategy.FULL_SHARD
|
||||
elif sharding_strategy == "hybrid_full":
|
||||
@@ -98,10 +90,11 @@ def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
|
||||
sharding_strategy = ShardingStrategy.NO_SHARD
|
||||
auto_wrap_policy = None
|
||||
elif sharding_strategy == "hybrid_zero2":
|
||||
sharding_strategy = ShardingStrategy._HYBRID_SHARD_ZERO2
|
||||
|
||||
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,
|
||||
@@ -110,29 +103,24 @@ def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
|
||||
"limit_all_gathers": True,
|
||||
"cpu_offload": cpu_offload,
|
||||
}
|
||||
|
||||
|
||||
# Add LoRA-specific settings when LoRA is enabled
|
||||
if use_lora:
|
||||
fsdp_kwargs.update({
|
||||
"use_orig_params": False, # Required for LoRA memory savings
|
||||
"sync_module_states": True,
|
||||
})
|
||||
|
||||
return fsdp_kwargs
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def get_discriminator_fsdp_kwargs():
|
||||
return fsdp_kwargs, no_split_modules
|
||||
|
||||
|
||||
def get_discriminator_fsdp_kwargs(master_weight_type="fp32"):
|
||||
auto_wrap_policy = None
|
||||
|
||||
|
||||
# Use existing mixed precision settings
|
||||
|
||||
mixed_precision = float32
|
||||
sharding_strategy = ShardingStrategy.NO_SHARD
|
||||
mixed_precision = get_mixed_precision(master_weight_type)
|
||||
sharding_strategy = ShardingStrategy.NO_SHARD
|
||||
device_id = torch.cuda.current_device()
|
||||
fsdp_kwargs = {
|
||||
"auto_wrap_policy": auto_wrap_policy,
|
||||
@@ -141,8 +129,5 @@ def get_discriminator_fsdp_kwargs():
|
||||
"device_id": device_id,
|
||||
"limit_all_gathers": True,
|
||||
}
|
||||
|
||||
|
||||
return fsdp_kwargs
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,377 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
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.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.utils.logging_ import main_print
|
||||
|
||||
hunyuan_config = {
|
||||
"mm_double_blocks_depth": 20,
|
||||
"mm_single_blocks_depth": 40,
|
||||
"rope_dim_list": [16, 56, 56],
|
||||
"hidden_size": 3072,
|
||||
"heads_num": 24,
|
||||
"mlp_width_ratio": 4,
|
||||
"guidance_embed": True,
|
||||
}
|
||||
|
||||
PROMPT_TEMPLATE_ENCODE = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|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|>")
|
||||
|
||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"dit-llm-encode": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE,
|
||||
"crop_start": 36,
|
||||
},
|
||||
"dit-llm-encode-video": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class HunyuanTextEncoderWrapper(nn.Module):
|
||||
|
||||
def __init__(self, pretrained_model_name_or_path, device):
|
||||
super().__init__()
|
||||
|
||||
text_len = 256
|
||||
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get(
|
||||
"crop_start", 0)
|
||||
|
||||
max_length = text_len + crop_start
|
||||
|
||||
# prompt_template
|
||||
prompt_template = PROMPT_TEMPLATE["dit-llm-encode"]
|
||||
|
||||
# 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")
|
||||
self.text_encoder = TextEncoder(
|
||||
text_encoder_type="llm",
|
||||
text_encoder_path=text_encoder_path,
|
||||
max_length=max_length,
|
||||
text_encoder_precision="fp16",
|
||||
tokenizer_type="llm",
|
||||
prompt_template=prompt_template,
|
||||
prompt_template_video=prompt_template_video,
|
||||
hidden_state_skip_layer=2,
|
||||
apply_final_norm=False,
|
||||
reproduce=False,
|
||||
logger=None,
|
||||
device=device,
|
||||
)
|
||||
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,
|
||||
max_length=77,
|
||||
text_encoder_precision="fp16",
|
||||
tokenizer_type="clipL",
|
||||
reproduce=False,
|
||||
logger=None,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def encode_(self, prompt, text_encoder, clip_skip=None):
|
||||
# TODO
|
||||
device = self.text_encoder.device
|
||||
data_type = "video"
|
||||
num_videos_per_prompt = 1
|
||||
|
||||
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_embeds = prompt_outputs.hidden_state
|
||||
else:
|
||||
prompt_outputs = text_encoder.encode(
|
||||
text_inputs,
|
||||
output_hidden_states=True,
|
||||
data_type=data_type,
|
||||
device=device,
|
||||
)
|
||||
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip + 1)]
|
||||
|
||||
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)
|
||||
|
||||
if text_encoder is not None:
|
||||
prompt_embeds_dtype = text_encoder.dtype
|
||||
elif self.transformer is not None:
|
||||
prompt_embeds_dtype = self.transformer.dtype
|
||||
else:
|
||||
prompt_embeds_dtype = prompt_embeds.dtype
|
||||
|
||||
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)
|
||||
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)
|
||||
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 = F.pad(
|
||||
prompt_embeds_2,
|
||||
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
prompt_embeds = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
|
||||
return prompt_embeds, attention_mask
|
||||
|
||||
|
||||
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.max_sequence_length = 256
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
device = self.text_encoder.device
|
||||
dtype = self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
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
|
||||
|
||||
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
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
|
||||
def load_hunyuan_state_dict(model, dit_model_name_or_path):
|
||||
load_key = "module"
|
||||
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)
|
||||
|
||||
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())}.")
|
||||
model.load_state_dict(state_dict, strict=True)
|
||||
return model
|
||||
|
||||
|
||||
def load_transformer(
|
||||
model_type,
|
||||
dit_model_name_or_path,
|
||||
pretrained_model_name_or_path,
|
||||
master_weight_type,
|
||||
):
|
||||
if model_type == "mochi":
|
||||
if dit_model_name_or_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
dit_model_name_or_path,
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "hunyuan_hf":
|
||||
if dit_model_name_or_path:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
dit_model_name_or_path,
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "hunyuan":
|
||||
transformer = HYVideoDiffusionTransformer(
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
**hunyuan_config,
|
||||
dtype=master_weight_type,
|
||||
)
|
||||
transformer = load_hunyuan_state_dict(transformer,
|
||||
dit_model_name_or_path)
|
||||
if master_weight_type == torch.bfloat16:
|
||||
transformer = transformer.bfloat16()
|
||||
else:
|
||||
raise ValueError(f"Unsupported model type: {model_type}")
|
||||
return 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")
|
||||
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")
|
||||
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")
|
||||
|
||||
config = AutoencoderKLCausal3D.load_config(vae_path)
|
||||
vae = AutoencoderKLCausal3D.from_config(config)
|
||||
|
||||
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
|
||||
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
|
||||
|
||||
ckpt = torch.load(vae_ckpt, map_location=vae.device, weights_only=True)
|
||||
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.")
|
||||
}
|
||||
vae.load_state_dict(ckpt)
|
||||
vae = vae.to(dtype=vae_precision)
|
||||
vae.requires_grad_(False)
|
||||
vae = vae.to("cuda")
|
||||
vae.eval()
|
||||
autocast_type = torch.float32
|
||||
fps = 24
|
||||
return vae, autocast_type, fps
|
||||
|
||||
|
||||
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)
|
||||
elif model_type == "hunyuan" or "hunyuan_hf":
|
||||
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path,
|
||||
device)
|
||||
else:
|
||||
raise ValueError(f"Unsupported model type: {model_type}")
|
||||
return text_encoder
|
||||
|
||||
|
||||
def get_no_split_modules(transformer):
|
||||
# if of type MochiTransformer3DModel
|
||||
if isinstance(transformer, MochiTransformer3DModel):
|
||||
return (MochiTransformerBlock, )
|
||||
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
|
||||
return (HunyuanVideoSingleTransformerBlock,
|
||||
HunyuanVideoTransformerBlock)
|
||||
elif isinstance(transformer, HYVideoDiffusionTransformer):
|
||||
return (MMDoubleStreamBlock, MMSingleStreamBlock)
|
||||
else:
|
||||
raise ValueError(f"Unsupported transformer type: {type(transformer)}")
|
||||
|
||||
|
||||
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)
|
||||
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)
|
||||
@@ -1,23 +1,24 @@
|
||||
|
||||
import sys
|
||||
import pdb
|
||||
import os
|
||||
import pdb
|
||||
import sys
|
||||
|
||||
|
||||
def main_print(content):
|
||||
if int(os.environ['LOCAL_RANK']) <= 0:
|
||||
if int(os.environ["LOCAL_RANK"]) <= 0:
|
||||
print(content)
|
||||
|
||||
#ForkedPdb().set_trace()
|
||||
|
||||
# ForkedPdb().set_trace()
|
||||
class ForkedPdb(pdb.Pdb):
|
||||
"""A Pdb subclass that may be used
|
||||
from a forked multiprocessing child
|
||||
|
||||
"""
|
||||
|
||||
def interaction(self, *args, **kwargs):
|
||||
_stdin = sys.stdin
|
||||
try:
|
||||
sys.stdin = open('/dev/stdin')
|
||||
sys.stdin = open("/dev/stdin")
|
||||
pdb.Pdb.interaction(self, *args, **kwargs)
|
||||
finally:
|
||||
sys.stdin = _stdin
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
import torch
|
||||
from accelerate.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
|
||||
# Optimizer creation
|
||||
supported_optimizers = ["adam", "adamw", "prodigy"]
|
||||
if args.optimizer not in supported_optimizers:
|
||||
logger.warning(
|
||||
f"Unsupported choice of optimizer: {args.optimizer}. Supported optimizers include {supported_optimizers}. Defaulting to AdamW"
|
||||
)
|
||||
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:
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except ImportError:
|
||||
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 = optimizer_class(
|
||||
params_to_optimize,
|
||||
betas=(args.adam_beta1, args.adam_beta2),
|
||||
eps=args.adam_epsilon,
|
||||
weight_decay=args.adam_weight_decay,
|
||||
)
|
||||
elif args.optimizer.lower() == "adam":
|
||||
optimizer_class = bnb.optim.Adam8bit if args.use_8bit_adam else torch.optim.Adam
|
||||
|
||||
optimizer = optimizer_class(
|
||||
params_to_optimize,
|
||||
betas=(args.adam_beta1, args.adam_beta2),
|
||||
eps=args.adam_epsilon,
|
||||
weight_decay=args.adam_weight_decay,
|
||||
)
|
||||
elif args.optimizer.lower() == "prodigy":
|
||||
try:
|
||||
import prodigyopt
|
||||
except ImportError:
|
||||
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"
|
||||
)
|
||||
|
||||
optimizer = optimizer_class(
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(args.adam_beta1, args.adam_beta2),
|
||||
beta3=args.prodigy_beta3,
|
||||
weight_decay=args.adam_weight_decay,
|
||||
eps=args.adam_epsilon,
|
||||
decouple=args.prodigy_decouple,
|
||||
use_bias_correction=args.prodigy_use_bias_correction,
|
||||
safeguard_warmup=args.prodigy_safeguard_warmup,
|
||||
)
|
||||
|
||||
return optimizer
|
||||
@@ -1,8 +1,10 @@
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import os
|
||||
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
class COMM_INFO:
|
||||
|
||||
def __init__(self):
|
||||
self.group = None
|
||||
self.sp_size = 1
|
||||
@@ -10,8 +12,11 @@ class COMM_INFO:
|
||||
self.rank_within_group = 0
|
||||
self.group_id = 0
|
||||
|
||||
|
||||
nccl_info = COMM_INFO()
|
||||
_SEQUENCE_PARALLEL_STATE = False
|
||||
|
||||
|
||||
def initialize_sequence_parallel_state(sequence_parallel_size):
|
||||
global _SEQUENCE_PARALLEL_STATE
|
||||
if sequence_parallel_size > 1:
|
||||
@@ -19,27 +24,34 @@ def initialize_sequence_parallel_state(sequence_parallel_size):
|
||||
initialize_sequence_parallel_group(sequence_parallel_size)
|
||||
else:
|
||||
nccl_info.sp_size = 1
|
||||
nccl_info.global_rank = int(os.getenv('RANK', '0'))
|
||||
nccl_info.global_rank = int(os.getenv("RANK", "0"))
|
||||
nccl_info.rank_within_group = 0
|
||||
nccl_info.group_id = int(os.getenv('RANK', '0'))
|
||||
nccl_info.group_id = int(os.getenv("RANK", "0"))
|
||||
|
||||
|
||||
def set_sequence_parallel_state(state):
|
||||
global _SEQUENCE_PARALLEL_STATE
|
||||
_SEQUENCE_PARALLEL_STATE = state
|
||||
|
||||
|
||||
def get_sequence_parallel_state():
|
||||
return _SEQUENCE_PARALLEL_STATE
|
||||
|
||||
|
||||
def initialize_sequence_parallel_group(sequence_parallel_size):
|
||||
"""Initialize the sequence parallel group."""
|
||||
rank = int(os.getenv('RANK', '0'))
|
||||
world_size = int(os.getenv("WORLD_SIZE", '1'))
|
||||
assert world_size % sequence_parallel_size == 0, "world_size must be divisible by sequence_parallel_size, but got world_size: {}, sequence_parallel_size: {}".format(world_size, sequence_parallel_size)
|
||||
rank = int(os.getenv("RANK", "0"))
|
||||
world_size = int(os.getenv("WORLD_SIZE", "1"))
|
||||
assert (
|
||||
world_size % sequence_parallel_size == 0
|
||||
), "world_size must be divisible by sequence_parallel_size, but got world_size: {}, sequence_parallel_size: {}".format(
|
||||
world_size, sequence_parallel_size)
|
||||
nccl_info.sp_size = 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
|
||||
|
||||
@@ -1,471 +0,0 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
import os
|
||||
import math
|
||||
import torch
|
||||
import logging
|
||||
import random
|
||||
import subprocess
|
||||
import numpy as np
|
||||
import torch.distributed as dist
|
||||
|
||||
# from torch._six import inf
|
||||
from torch import inf
|
||||
from PIL import Image
|
||||
from typing import Union, Iterable
|
||||
import collections
|
||||
from collections import OrderedDict
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from diffusers.utils import is_bs4_available, is_ftfy_available
|
||||
|
||||
import html
|
||||
import re
|
||||
import urllib.parse as ul
|
||||
|
||||
if is_bs4_available():
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
_tensor_or_tensors = Union[torch.Tensor, Iterable[torch.Tensor]]
|
||||
|
||||
def to_2tuple(x):
|
||||
if isinstance(x, collections.abc.Iterable):
|
||||
return x
|
||||
return (x, x)
|
||||
|
||||
def find_model(model_name):
|
||||
"""
|
||||
Finds a pre-trained Latte model, downloading it if necessary. Alternatively, loads a model from a local path.
|
||||
"""
|
||||
assert os.path.isfile(model_name), f'Could not find Latte checkpoint at {model_name}'
|
||||
checkpoint = torch.load(model_name, map_location=lambda storage, loc: storage)
|
||||
|
||||
# if "ema" in checkpoint: # supports checkpoints from train.py
|
||||
# print('Using Ema!')
|
||||
# checkpoint = checkpoint["ema"]
|
||||
# else:
|
||||
print('Using model!')
|
||||
checkpoint = checkpoint['model']
|
||||
return checkpoint
|
||||
|
||||
#################################################################################
|
||||
# Training Clip Gradients #
|
||||
#################################################################################
|
||||
|
||||
def get_grad_norm(
|
||||
parameters: _tensor_or_tensors, norm_type: float = 2.0) -> torch.Tensor:
|
||||
r"""
|
||||
Copy from torch.nn.utils.clip_grad_norm_
|
||||
|
||||
Clips gradient norm of an iterable of parameters.
|
||||
|
||||
The norm is computed over all gradients together, as if they were
|
||||
concatenated into a single vector. Gradients are modified in-place.
|
||||
|
||||
Args:
|
||||
parameters (Iterable[Tensor] or Tensor): an iterable of Tensors or a
|
||||
single Tensor that will have gradients normalized
|
||||
max_norm (float or int): max norm of the gradients
|
||||
norm_type (float or int): type of the used p-norm. Can be ``'inf'`` for
|
||||
infinity norm.
|
||||
error_if_nonfinite (bool): if True, an error is thrown if the total
|
||||
norm of the gradients from :attr:`parameters` is ``nan``,
|
||||
``inf``, or ``-inf``. Default: False (will switch to True in the future)
|
||||
|
||||
Returns:
|
||||
Total norm of the parameter gradients (viewed as a single vector).
|
||||
"""
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
norm_type = float(norm_type)
|
||||
if len(grads) == 0:
|
||||
return torch.tensor(0.)
|
||||
device = grads[0].device
|
||||
if norm_type == inf:
|
||||
norms = [g.detach().abs().max().to(device) for g in grads]
|
||||
total_norm = norms[0] if len(norms) == 1 else torch.max(torch.stack(norms))
|
||||
else:
|
||||
total_norm = torch.norm(torch.stack([torch.norm(g.detach(), norm_type).to(device) for g in grads]), norm_type)
|
||||
return total_norm
|
||||
|
||||
|
||||
def clip_grad_norm_(
|
||||
parameters: _tensor_or_tensors, max_norm: float, norm_type: float = 2.0,
|
||||
error_if_nonfinite: bool = False, clip_grad=True) -> torch.Tensor:
|
||||
r"""
|
||||
Copy from torch.nn.utils.clip_grad_norm_
|
||||
|
||||
Clips gradient norm of an iterable of parameters.
|
||||
|
||||
The norm is computed over all gradients together, as if they were
|
||||
concatenated into a single vector. Gradients are modified in-place.
|
||||
|
||||
Args:
|
||||
parameters (Iterable[Tensor] or Tensor): an iterable of Tensors or a
|
||||
single Tensor that will have gradients normalized
|
||||
max_norm (float or int): max norm of the gradients
|
||||
norm_type (float or int): type of the used p-norm. Can be ``'inf'`` for
|
||||
infinity norm.
|
||||
error_if_nonfinite (bool): if True, an error is thrown if the total
|
||||
norm of the gradients from :attr:`parameters` is ``nan``,
|
||||
``inf``, or ``-inf``. Default: False (will switch to True in the future)
|
||||
|
||||
Returns:
|
||||
Total norm of the parameter gradients (viewed as a single vector).
|
||||
"""
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
max_norm = float(max_norm)
|
||||
norm_type = float(norm_type)
|
||||
if len(grads) == 0:
|
||||
return torch.tensor(0.)
|
||||
device = grads[0].device
|
||||
if norm_type == inf:
|
||||
norms = [g.detach().abs().max().to(device) for g in grads]
|
||||
total_norm = norms[0] if len(norms) == 1 else torch.max(torch.stack(norms))
|
||||
else:
|
||||
total_norm = torch.norm(torch.stack([torch.norm(g.detach(), norm_type).to(device) for g in grads]), norm_type)
|
||||
|
||||
if clip_grad:
|
||||
if error_if_nonfinite and torch.logical_or(total_norm.isnan(), total_norm.isinf()):
|
||||
raise RuntimeError(
|
||||
f'The total norm of order {norm_type} for gradients from '
|
||||
'`parameters` is non-finite, so it cannot be clipped. To disable '
|
||||
'this error and scale the gradients by the non-finite norm anyway, '
|
||||
'set `error_if_nonfinite=False`')
|
||||
clip_coef = max_norm / (total_norm + 1e-6)
|
||||
# Note: multiplying by the clamped coef is redundant when the coef is clamped to 1, but doing so
|
||||
# avoids a `if clip_coef < 1:` conditional which can require a CPU <=> device synchronization
|
||||
# when the gradients do not reside in CPU memory.
|
||||
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||
for g in grads:
|
||||
g.detach().mul_(clip_coef_clamped.to(g.device))
|
||||
# gradient_cliped = torch.norm(torch.stack([torch.norm(g.detach(), norm_type).to(device) for g in grads]), norm_type)
|
||||
# print(gradient_cliped)
|
||||
return total_norm
|
||||
|
||||
|
||||
def get_experiment_dir(root_dir, args):
|
||||
# if args.pretrained is not None and 'Latte-XL-2-256x256.pt' not in args.pretrained:
|
||||
# root_dir += '-WOPRE'
|
||||
if args.use_compile:
|
||||
root_dir += '-Compile' # speedup by torch compile
|
||||
if args.attention_mode:
|
||||
root_dir += f'-{args.attention_mode.upper()}'
|
||||
# if args.enable_xformers_memory_efficient_attention:
|
||||
# root_dir += '-Xfor'
|
||||
if args.gradient_checkpointing:
|
||||
root_dir += '-Gc'
|
||||
if args.mixed_precision:
|
||||
root_dir += f'-{args.mixed_precision.upper()}'
|
||||
root_dir += f'-{args.max_image_size}'
|
||||
return root_dir
|
||||
|
||||
def get_precision(args):
|
||||
if args.mixed_precision == "bf16":
|
||||
dtype = torch.bfloat16
|
||||
elif args.mixed_precision == "fp16":
|
||||
dtype = torch.float16
|
||||
else:
|
||||
dtype = torch.float32
|
||||
return dtype
|
||||
|
||||
#################################################################################
|
||||
# Training Logger #
|
||||
#################################################################################
|
||||
|
||||
def create_logger(logging_dir):
|
||||
"""
|
||||
Create a logger that writes to a log file and stdout.
|
||||
"""
|
||||
if dist.get_rank() == 0: # real logger
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
# format='[\033[34m%(asctime)s\033[0m] %(message)s',
|
||||
format='[%(asctime)s] %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S',
|
||||
handlers=[logging.StreamHandler(), logging.FileHandler(f"{logging_dir}/log.txt")]
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
else: # dummy logger (does nothing)
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.addHandler(logging.NullHandler())
|
||||
return logger
|
||||
|
||||
|
||||
def create_tensorboard(tensorboard_dir):
|
||||
"""
|
||||
Create a tensorboard that saves losses.
|
||||
"""
|
||||
if dist.get_rank() == 0: # real tensorboard
|
||||
# tensorboard
|
||||
writer = SummaryWriter(tensorboard_dir)
|
||||
|
||||
return writer
|
||||
|
||||
|
||||
def write_tensorboard(writer, *args):
|
||||
'''
|
||||
write the loss information to a tensorboard file.
|
||||
Only for pytorch DDP mode.
|
||||
'''
|
||||
if dist.get_rank() == 0: # real tensorboard
|
||||
writer.add_scalar(args[0], args[1], args[2])
|
||||
|
||||
|
||||
#################################################################################
|
||||
# EMA Update/ DDP Training Utils #
|
||||
#################################################################################
|
||||
|
||||
@torch.no_grad()
|
||||
def update_ema(ema_model, model, decay=0.9999):
|
||||
"""
|
||||
Step the EMA model towards the current model.
|
||||
"""
|
||||
ema_params = OrderedDict(ema_model.named_parameters())
|
||||
model_params = OrderedDict(model.named_parameters())
|
||||
|
||||
for name, param in model_params.items():
|
||||
# TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed
|
||||
ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
|
||||
|
||||
|
||||
def requires_grad(model, flag=True):
|
||||
"""
|
||||
Set requires_grad flag for all parameters in a model.
|
||||
"""
|
||||
for p in model.parameters():
|
||||
p.requires_grad = flag
|
||||
|
||||
|
||||
def cleanup():
|
||||
"""
|
||||
End DDP training.
|
||||
"""
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
def setup_distributed(backend="nccl", port=None):
|
||||
"""Initialize distributed training environment.
|
||||
support both slurm and torch.distributed.launch
|
||||
see torch.distributed.init_process_group() for more details
|
||||
"""
|
||||
num_gpus = torch.cuda.device_count()
|
||||
|
||||
if "SLURM_JOB_ID" in os.environ:
|
||||
rank = int(os.environ["SLURM_PROCID"])
|
||||
world_size = int(os.environ["SLURM_NTASKS"])
|
||||
node_list = os.environ["SLURM_NODELIST"]
|
||||
addr = subprocess.getoutput(f"scontrol show hostname {node_list} | head -n1")
|
||||
# specify master port
|
||||
if port is not None:
|
||||
os.environ["MASTER_PORT"] = str(port)
|
||||
elif "MASTER_PORT" not in os.environ:
|
||||
# os.environ["MASTER_PORT"] = "29566"
|
||||
os.environ["MASTER_PORT"] = str(29567 + num_gpus)
|
||||
if "MASTER_ADDR" not in os.environ:
|
||||
os.environ["MASTER_ADDR"] = addr
|
||||
os.environ["WORLD_SIZE"] = str(world_size)
|
||||
os.environ["LOCAL_RANK"] = str(rank % num_gpus)
|
||||
os.environ["RANK"] = str(rank)
|
||||
else:
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
|
||||
# torch.cuda.set_device(rank % num_gpus)
|
||||
|
||||
dist.init_process_group(
|
||||
backend=backend,
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
)
|
||||
|
||||
|
||||
#################################################################################
|
||||
# Testing Utils #
|
||||
#################################################################################
|
||||
|
||||
def save_video_grid(video, nrow=None):
|
||||
b, t, h, w, c = video.shape
|
||||
|
||||
if nrow is None:
|
||||
nrow = math.ceil(math.sqrt(b))
|
||||
ncol = math.ceil(b / nrow)
|
||||
padding = 1
|
||||
video_grid = torch.zeros((t, (padding + h) * nrow + padding,
|
||||
(padding + w) * ncol + padding, c), dtype=torch.uint8)
|
||||
|
||||
print(video_grid.shape)
|
||||
for i in range(b):
|
||||
r = i // ncol
|
||||
c = i % ncol
|
||||
start_r = (padding + h) * r
|
||||
start_c = (padding + w) * c
|
||||
video_grid[:, start_r:start_r + h, start_c:start_c + w] = video[i]
|
||||
|
||||
return video_grid
|
||||
|
||||
|
||||
#################################################################################
|
||||
# MMCV Utils #
|
||||
#################################################################################
|
||||
|
||||
|
||||
def collect_env():
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from mmcv.utils import collect_env as collect_base_env
|
||||
from mmcv.utils import get_git_hash
|
||||
"""Collect the information of the running environments."""
|
||||
|
||||
env_info = collect_base_env()
|
||||
env_info['MMClassification'] = get_git_hash()[:7]
|
||||
|
||||
for name, val in env_info.items():
|
||||
print(f'{name}: {val}')
|
||||
|
||||
print(torch.cuda.get_arch_list())
|
||||
print(torch.version.cuda)
|
||||
|
||||
|
||||
#################################################################################
|
||||
# Pixart-alpha Utils #
|
||||
#################################################################################
|
||||
|
||||
|
||||
bad_punct_regex = re.compile(r'['+'#®•©™&@·º½¾¿¡§~'+'\)'+'\('+'\]'+'\['+'\}'+'\{'+'\|'+'\\'+'\/'+'\*' + r']{1,}') # noqa
|
||||
|
||||
def text_preprocessing(text, support_Chinese=True):
|
||||
# The exact text cleaning as was in the training stage:
|
||||
text = clean_caption(text, support_Chinese=support_Chinese)
|
||||
text = clean_caption(text, support_Chinese=support_Chinese)
|
||||
return text
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
def clean_caption(caption, support_Chinese=True):
|
||||
caption = str(caption)
|
||||
caption = ul.unquote_plus(caption)
|
||||
caption = caption.strip().lower()
|
||||
caption = re.sub('<person>', 'person', caption)
|
||||
# urls:
|
||||
caption = re.sub(
|
||||
r'\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))', # noqa
|
||||
'', caption) # regex for urls
|
||||
caption = re.sub(
|
||||
r'\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))', # noqa
|
||||
'', caption) # regex for urls
|
||||
# html:
|
||||
caption = BeautifulSoup(caption, features='html.parser').text
|
||||
|
||||
# @<nickname>
|
||||
caption = re.sub(r'@[\w\d]+\b', '', caption)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
caption = re.sub(r'[\u31c0-\u31ef]+', '', caption)
|
||||
caption = re.sub(r'[\u31f0-\u31ff]+', '', caption)
|
||||
caption = re.sub(r'[\u3200-\u32ff]+', '', caption)
|
||||
caption = re.sub(r'[\u3300-\u33ff]+', '', caption)
|
||||
caption = re.sub(r'[\u3400-\u4dbf]+', '', caption)
|
||||
caption = re.sub(r'[\u4dc0-\u4dff]+', '', caption)
|
||||
if not support_Chinese:
|
||||
caption = re.sub(r'[\u4e00-\u9fff]+', '', caption) # Chinese
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
caption = re.sub(
|
||||
r'[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+', # noqa
|
||||
'-', caption)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
caption = re.sub(r'[`´«»“”¨]', '"', caption)
|
||||
caption = re.sub(r'[‘’]', "'", caption)
|
||||
|
||||
# "
|
||||
caption = re.sub(r'"?', '', caption)
|
||||
# &
|
||||
caption = re.sub(r'&', '', caption)
|
||||
|
||||
# ip adresses:
|
||||
caption = re.sub(r'\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}', ' ', caption)
|
||||
|
||||
# article ids:
|
||||
caption = re.sub(r'\d:\d\d\s+$', '', caption)
|
||||
|
||||
# \n
|
||||
caption = re.sub(r'\\n', ' ', caption)
|
||||
|
||||
# "#123"
|
||||
caption = re.sub(r'#\d{1,3}\b', '', caption)
|
||||
# "#12345.."
|
||||
caption = re.sub(r'#\d{5,}\b', '', caption)
|
||||
# "123456.."
|
||||
caption = re.sub(r'\b\d{6,}\b', '', caption)
|
||||
# filenames:
|
||||
caption = re.sub(r'[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)', '', caption)
|
||||
|
||||
#
|
||||
caption = re.sub(r'[\"\']{2,}', r'"', caption) # """AUSVERKAUFT"""
|
||||
caption = re.sub(r'[\.]{2,}', r' ', caption) # """AUSVERKAUFT"""
|
||||
|
||||
caption = re.sub(bad_punct_regex, r' ', caption) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
caption = re.sub(r'\s+\.\s+', r' ', caption) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r'(?:\-|\_)')
|
||||
if len(re.findall(regex2, caption)) > 3:
|
||||
caption = re.sub(regex2, ' ', caption)
|
||||
|
||||
caption = basic_clean(caption)
|
||||
|
||||
caption = re.sub(r'\b[a-zA-Z]{1,3}\d{3,15}\b', '', caption) # jc6640
|
||||
caption = re.sub(r'\b[a-zA-Z]+\d+[a-zA-Z]+\b', '', caption) # jc6640vc
|
||||
caption = re.sub(r'\b\d+[a-zA-Z]+\d+\b', '', caption) # 6640vc231
|
||||
|
||||
caption = re.sub(r'(worldwide\s+)?(free\s+)?shipping', '', caption)
|
||||
caption = re.sub(r'(free\s)?download(\sfree)?', '', caption)
|
||||
caption = re.sub(r'\bclick\b\s(?:for|on)\s\w+', '', caption)
|
||||
caption = re.sub(r'\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?', '', caption)
|
||||
caption = re.sub(r'\bpage\s+\d+\b', '', caption)
|
||||
|
||||
caption = re.sub(r'\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b', r' ', caption) # j2d1a2a...
|
||||
|
||||
caption = re.sub(r'\b\d+\.?\d*[xх×]\d+\.?\d*\b', '', caption)
|
||||
|
||||
caption = re.sub(r'\b\s+\:\s+', r': ', caption)
|
||||
caption = re.sub(r'(\D[,\./])\b', r'\1 ', caption)
|
||||
caption = re.sub(r'\s+', ' ', caption)
|
||||
|
||||
caption.strip()
|
||||
|
||||
caption = re.sub(r'^[\"\']([\w\W]+)[\"\']$', r'\1', caption)
|
||||
caption = re.sub(r'^[\'\_,\-\:;]', r'', caption)
|
||||
caption = re.sub(r'[\'\_,\-\:\-\+]$', r'', caption)
|
||||
caption = re.sub(r'^\.\S+$', '', caption)
|
||||
|
||||
return caption.strip()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
# caption = re.sub(r'[\u4e00-\u9fff]+', '', caption)
|
||||
a = "امرأة مسنة بشعر أبيض ووجه مليء بالتجاعيد تجلس داخل سيارة قديمة الطراز، تنظر من خلال النافذة الجانبية بتعبير تأملي أو حزين قليلاً."
|
||||
print(a)
|
||||
print(text_preprocessing(a))
|
||||
|
||||
+180
-109
@@ -1,25 +1,26 @@
|
||||
import gc
|
||||
import os
|
||||
from typing import List, Optional, Union
|
||||
|
||||
|
||||
from typing import Optional, Union, List
|
||||
import numpy as np
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule, retrieve_timesteps
|
||||
from tqdm import tqdm
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
AutoencoderKLMochi,
|
||||
)
|
||||
from fastvideo.utils.logging import main_print
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
import os
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from einops import rearrange
|
||||
from tqdm import tqdm
|
||||
|
||||
import wandb
|
||||
import gc
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
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)
|
||||
|
||||
|
||||
def prepare_latents(
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
@@ -38,11 +39,15 @@ 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
|
||||
|
||||
|
||||
def sample_validation_video(
|
||||
model_type,
|
||||
transformer,
|
||||
vae,
|
||||
scheduler,
|
||||
@@ -60,8 +65,9 @@ def sample_validation_video(
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
vae_spatial_scale_factor = 8,
|
||||
vae_temporal_scale_factor = 6,
|
||||
vae_spatial_scale_factor=8,
|
||||
vae_temporal_scale_factor=6,
|
||||
num_channels_latents=12,
|
||||
):
|
||||
device = vae.device
|
||||
|
||||
@@ -69,12 +75,13 @@ 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
|
||||
num_channels_latents = 12
|
||||
latents = prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
@@ -85,20 +92,21 @@ def sample_validation_video(
|
||||
device,
|
||||
generator,
|
||||
vae_spatial_scale_factor,
|
||||
vae_temporal_scale_factor
|
||||
vae_temporal_scale_factor,
|
||||
)
|
||||
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
|
||||
# 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 = np.array(sigmas)
|
||||
if scheduler_type == "euler":
|
||||
if scheduler_type == "euler" and model_type == "mochi": #todo
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps,
|
||||
@@ -112,37 +120,46 @@ 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:
|
||||
# write with tqdm instead
|
||||
# only enable if nccl_info.global_rank == 0
|
||||
|
||||
with tqdm(total=num_inference_steps, disable= nccl_info.rank_within_group != 0, desc="Validation sampling...") as progress_bar:
|
||||
|
||||
with tqdm(
|
||||
total=num_inference_steps,
|
||||
disable=nccl_info.rank_within_group != 0,
|
||||
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]).to(latents.dtype)
|
||||
noise_pred = transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
noise_pred = transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
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:
|
||||
@@ -150,125 +167,179 @@ 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():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
|
||||
if output_type == "latent":
|
||||
video = latents
|
||||
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 = 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)
|
||||
|
||||
|
||||
|
||||
return (video,)
|
||||
|
||||
return (video, )
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast("cuda", dtype=torch.bfloat16)
|
||||
def log_validation(args, transformer, device, weight_dtype, global_step, scheduler_type="euler",shift=1.0, num_euler_timesteps=100, linear_quadratic_threshold=0.025, linear_range=0.5, ema=False):
|
||||
#TODO
|
||||
print(f"Running validation....\n")
|
||||
vae = AutoencoderKLMochi.from_pretrained(args.pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype).to("cuda")
|
||||
def log_validation(
|
||||
args,
|
||||
transformer,
|
||||
device,
|
||||
weight_dtype, # TODO
|
||||
global_step,
|
||||
scheduler_type="euler",
|
||||
shift=1.0,
|
||||
num_euler_timesteps=100,
|
||||
linear_quadratic_threshold=0.025,
|
||||
linear_range=0.5,
|
||||
ema=False,
|
||||
):
|
||||
# TODO
|
||||
print("Running validation....\n")
|
||||
if args.model_type == "mochi":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 6
|
||||
num_channels_latents = 12
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 4
|
||||
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.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
|
||||
else:
|
||||
linear_quadraic = True if scheduler_type == "pcm_linear_quadratic" else False
|
||||
scheduler = PCMFMScheduler(1000, shift, num_euler_timesteps, linear_quadraic, linear_quadratic_threshold, linear_range)
|
||||
scheduler = PCMFMScheduler(
|
||||
1000,
|
||||
shift,
|
||||
num_euler_timesteps,
|
||||
linear_quadraic,
|
||||
linear_quadratic_threshold,
|
||||
linear_range,
|
||||
)
|
||||
# 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:
|
||||
|
||||
videos = []
|
||||
# prompt_embed are named embed0 to embedN
|
||||
# check how many embeds are there
|
||||
num_embeds = len([f for f in os.listdir(args.validation_prompt_dir) if "embed" in f])
|
||||
embe_dir = os.path.join(args.validation_prompt_dir, "prompt_embed")
|
||||
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
|
||||
validation_prompt_ids += [0] * (num_sp_groups - num_embeds % num_sp_groups)
|
||||
if num_embeds % num_sp_groups != 0:
|
||||
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]
|
||||
|
||||
local_prompt_ids = validation_prompt_ids[nccl_info.group_id *
|
||||
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(args.validation_prompt_dir, f"embed{i}.pt")
|
||||
prompt_mask_path = os.path.join(args.validation_prompt_dir, f"mask{i}.pt")
|
||||
prompt_embeds = torch.load(prompt_embed_path, map_location="cpu", weights_only=True).to(device).to(weight_dtype).unsqueeze(0)
|
||||
prompt_attention_mask = torch.load(prompt_mask_path, map_location="cpu", weights_only=True).to(device).to(weight_dtype).unsqueeze(0)
|
||||
negative_prompt_embeds = torch.zeros(256, 4096).to(device).to(weight_dtype).unsqueeze(0)
|
||||
negative_prompt_attention_mask = torch.zeros(256).bool().to(device).unsqueeze(0)
|
||||
generator = torch.Generator(device="cuda").manual_seed(12345)
|
||||
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))
|
||||
generator = torch.Generator(device="cpu").manual_seed(12345)
|
||||
video = sample_validation_video(
|
||||
transformer,
|
||||
vae,
|
||||
scheduler,
|
||||
scheduler_type=scheduler_type,
|
||||
num_frames=args.num_frames,
|
||||
# Peiyuan TODO: remove hardcode
|
||||
height=480,
|
||||
width=848,
|
||||
num_inference_steps=validation_sampling_step,
|
||||
guidance_scale=validation_guidance_scale,
|
||||
generator=generator,
|
||||
prompt_embeds = prompt_embeds,
|
||||
prompt_attention_mask = prompt_attention_mask,
|
||||
negative_prompt_embeds = negative_prompt_embeds,
|
||||
negative_prompt_attention_mask = negative_prompt_attention_mask,
|
||||
)[0]
|
||||
args.model_type,
|
||||
transformer,
|
||||
vae,
|
||||
scheduler,
|
||||
scheduler_type=scheduler_type,
|
||||
num_frames=args.num_frames,
|
||||
height=args.num_height,
|
||||
width=args.num_width,
|
||||
num_inference_steps=validation_sampling_step,
|
||||
guidance_scale=validation_guidance_scale,
|
||||
generator=generator,
|
||||
prompt_embeds=prompt_embeds,
|
||||
prompt_attention_mask=prompt_attention_mask,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
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,
|
||||
)[0]
|
||||
if nccl_info.rank_within_group == 0:
|
||||
videos.append(video[0])
|
||||
# collect videos from all process to process zero
|
||||
|
||||
|
||||
gc.collect()
|
||||
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
|
||||
videos = [video for videos in all_videos for video in videos]
|
||||
videos = videos[:num_embeds]
|
||||
# linearize all videos
|
||||
# linearize all videos
|
||||
video_filenames = []
|
||||
for i, video in enumerate(videos):
|
||||
filename = os.path.join(args.output_dir, f"validation_step_{global_step}_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}_video_{i}.mp4")
|
||||
export_to_video(video, filename, fps=30)
|
||||
filename = os.path.join(
|
||||
args.output_dir,
|
||||
f"validation_step_{global_step}_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}_video_{i}.mp4",
|
||||
)
|
||||
export_to_video(video, filename, fps=fps)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
f"{'ema_' if ema else ''}validation_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}": [
|
||||
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.log(logs, step=global_step)
|
||||
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
#!/usr/bin/env bash
|
||||
# YAPF formatter, adapted from fastvideo.
|
||||
#
|
||||
# Usage:
|
||||
# # Do work and commit your work.
|
||||
|
||||
# # Format files that differ from origin/main.
|
||||
# bash format.sh
|
||||
|
||||
# # Commit changed files with message 'Run yapf and ruff'
|
||||
#
|
||||
#
|
||||
# This script formats all changed files from the last mergebase.
|
||||
# You are encouraged to run this locally before pushing changes for review.
|
||||
|
||||
# Cause the script to exit if a single command fails
|
||||
set -eo pipefail
|
||||
|
||||
# this stops git rev-parse from failing if we run this from the .git directory
|
||||
builtin cd "$(dirname "${BASH_SOURCE:-$0}")"
|
||||
ROOT="$(git rev-parse --show-toplevel)"
|
||||
builtin cd "$ROOT" || exit 1
|
||||
|
||||
check_command() {
|
||||
if ! command -v "$1" &> /dev/null; then
|
||||
echo "❓❓$1 is not installed, please run \`bash env_setup.sh\`"
|
||||
exit 1
|
||||
fi
|
||||
}
|
||||
|
||||
check_command yapf
|
||||
check_command ruff
|
||||
check_command codespell
|
||||
check_command isort
|
||||
|
||||
YAPF_VERSION=$(yapf --version | awk '{print $2}')
|
||||
RUFF_VERSION=$(ruff --version | awk '{print $2}')
|
||||
CODESPELL_VERSION=$(codespell --version)
|
||||
ISORT_VERSION=$(isort --vn)
|
||||
SPHINX_LINT_VERSION=$(sphinx-lint --version | awk '{print $2}')
|
||||
|
||||
|
||||
# # params: tool name, tool version, required version
|
||||
tool_version_check() {
|
||||
expected=$(grep "$1" requirements-lint.txt | cut -d'=' -f3)
|
||||
if [[ "$2" != "$expected" ]]; then
|
||||
echo "❓❓Wrong $1 version installed: $expected is required, not $2."
|
||||
exit 1
|
||||
fi
|
||||
}
|
||||
|
||||
tool_version_check "yapf" "$YAPF_VERSION"
|
||||
tool_version_check "ruff" "$RUFF_VERSION"
|
||||
tool_version_check "isort" "$ISORT_VERSION"
|
||||
tool_version_check "codespell" "$CODESPELL_VERSION"
|
||||
tool_version_check "sphinx-lint" "$SPHINX_LINT_VERSION"
|
||||
|
||||
YAPF_FLAGS=(
|
||||
'--recursive'
|
||||
'--parallel'
|
||||
)
|
||||
|
||||
YAPF_EXCLUDES=(
|
||||
'--exclude' 'data/**'
|
||||
)
|
||||
|
||||
# Format specified files
|
||||
format() {
|
||||
yapf --in-place "${YAPF_FLAGS[@]}" "$@"
|
||||
}
|
||||
|
||||
# Format files that differ from main branch. Ignores dirs that are not slated
|
||||
# for autoformat yet.
|
||||
format_changed() {
|
||||
# The `if` guard ensures that the list of filenames is not empty, which
|
||||
# could cause yapf to receive 0 positional arguments, making it hang
|
||||
# waiting for STDIN.
|
||||
#
|
||||
# `diff-filter=ACM` and $MERGEBASE is to ensure we only format files that
|
||||
# exist on both branches.
|
||||
MERGEBASE="$(git merge-base origin/main HEAD)"
|
||||
|
||||
if ! git diff --diff-filter=ACM --quiet --exit-code "$MERGEBASE" -- '*.py' '*.pyi' &>/dev/null; then
|
||||
git diff --name-only --diff-filter=ACM "$MERGEBASE" -- '*.py' '*.pyi' | xargs -P 5 \
|
||||
yapf --in-place "${YAPF_EXCLUDES[@]}" "${YAPF_FLAGS[@]}"
|
||||
fi
|
||||
|
||||
}
|
||||
|
||||
# Format all files
|
||||
format_all() {
|
||||
yapf --in-place "${YAPF_FLAGS[@]}" "${YAPF_EXCLUDES[@]}" .
|
||||
}
|
||||
|
||||
## This flag formats individual files. --files *must* be the first command line
|
||||
## arg to use this option.
|
||||
if [[ "$1" == '--files' ]]; then
|
||||
format "${@:2}"
|
||||
# If `--all` is passed, then any further arguments are ignored and the
|
||||
# entire python directory is formatted.
|
||||
elif [[ "$1" == '--all' ]]; then
|
||||
format_all
|
||||
else
|
||||
# Format only the files that changed in last commit.
|
||||
format_changed
|
||||
fi
|
||||
echo 'FastVideo yapf: Done'
|
||||
|
||||
|
||||
# If git diff returns a file that is in the skip list, the file may be checked anyway:
|
||||
# https://github.com/codespell-project/codespell/issues/1915
|
||||
# Avoiding the "./" prefix and using "/**" globs for directories appears to solve the problem
|
||||
CODESPELL_EXCLUDES=(
|
||||
'--skip' 'data/**,
|
||||
fastvideo/distill.py,
|
||||
fastvideo/models/hunyuan/modules/models.py,
|
||||
fastvideo/models/mochi_hf/modeling_mochi.py,
|
||||
fastvideo/utils/env_utils.py'
|
||||
)
|
||||
|
||||
# check spelling of specified files
|
||||
spell_check() {
|
||||
codespell "$@"
|
||||
}
|
||||
|
||||
spell_check_all(){
|
||||
codespell --toml pyproject.toml "${CODESPELL_EXCLUDES[@]}"
|
||||
}
|
||||
|
||||
# Spelling check of files that differ from main branch.
|
||||
spell_check_changed() {
|
||||
# The `if` guard ensures that the list of filenames is not empty, which
|
||||
# could cause ruff to receive 0 positional arguments, making it hang
|
||||
# waiting for STDIN.
|
||||
#
|
||||
# `diff-filter=ACM` and $MERGEBASE is to ensure we only lint files that
|
||||
# exist on both branches.
|
||||
MERGEBASE="$(git merge-base origin/main HEAD)"
|
||||
if ! git diff --diff-filter=ACM --quiet --exit-code "$MERGEBASE" -- '*.py' '*.pyi' &>/dev/null; then
|
||||
git diff --name-only --diff-filter=ACM "$MERGEBASE" -- '*.py' '*.pyi' | xargs \
|
||||
codespell "${CODESPELL_EXCLUDES[@]}"
|
||||
fi
|
||||
}
|
||||
|
||||
# Run Codespell
|
||||
## This flag runs spell check of individual files. --files *must* be the first command line
|
||||
## arg to use this option.
|
||||
if [[ "$1" == '--files' ]]; then
|
||||
spell_check "${@:2}"
|
||||
# If `--all` is passed, then any further arguments are ignored and the
|
||||
# entire python directory is linted.
|
||||
elif [[ "$1" == '--all' ]]; then
|
||||
spell_check_all
|
||||
else
|
||||
# Check spelling only of the files that changed in last commit.
|
||||
spell_check_changed
|
||||
fi
|
||||
echo 'FastVideo codespell: Done'
|
||||
|
||||
|
||||
# Lint specified files
|
||||
lint() {
|
||||
ruff check "$@"
|
||||
}
|
||||
|
||||
# Lint files that differ from main branch. Ignores dirs that are not slated
|
||||
# for autolint yet.
|
||||
lint_changed() {
|
||||
# The `if` guard ensures that the list of filenames is not empty, which
|
||||
# could cause ruff to receive 0 positional arguments, making it hang
|
||||
# waiting for STDIN.
|
||||
#
|
||||
# `diff-filter=ACM` and $MERGEBASE is to ensure we only lint files that
|
||||
# exist on both branches.
|
||||
MERGEBASE="$(git merge-base origin/main HEAD)"
|
||||
|
||||
if ! git diff --diff-filter=ACM --quiet --exit-code "$MERGEBASE" -- '*.py' '*.pyi' &>/dev/null; then
|
||||
git diff --name-only --diff-filter=ACM "$MERGEBASE" -- '*.py' '*.pyi' | xargs \
|
||||
ruff check
|
||||
fi
|
||||
|
||||
}
|
||||
|
||||
# Run Ruff
|
||||
### This flag lints individual files. --files *must* be the first command line
|
||||
### arg to use this option.
|
||||
if [[ "$1" == '--files' ]]; then
|
||||
lint "${@:2}"
|
||||
# If `--all` is passed, then any further arguments are ignored and the
|
||||
# entire python directory is linted.
|
||||
elif [[ "$1" == '--all' ]]; then
|
||||
lint fastvideo scripts
|
||||
else
|
||||
# Format only the files that changed in last commit.
|
||||
lint_changed
|
||||
fi
|
||||
echo 'FastVideo ruff: Done'
|
||||
|
||||
# check spelling of specified files
|
||||
isort_check() {
|
||||
isort "$@"
|
||||
}
|
||||
|
||||
isort_check_all(){
|
||||
isort .
|
||||
}
|
||||
|
||||
# Spelling check of files that differ from main branch.
|
||||
isort_check_changed() {
|
||||
# The `if` guard ensures that the list of filenames is not empty, which
|
||||
# could cause ruff to receive 0 positional arguments, making it hang
|
||||
# waiting for STDIN.
|
||||
#
|
||||
# `diff-filter=ACM` and $MERGEBASE is to ensure we only lint files that
|
||||
# exist on both branches.
|
||||
MERGEBASE="$(git merge-base origin/main HEAD)"
|
||||
|
||||
if ! git diff --diff-filter=ACM --quiet --exit-code "$MERGEBASE" -- '*.py' '*.pyi' &>/dev/null; then
|
||||
git diff --name-only --diff-filter=ACM "$MERGEBASE" -- '*.py' '*.pyi' | xargs \
|
||||
isort
|
||||
fi
|
||||
}
|
||||
|
||||
# Run Isort
|
||||
# This flag runs spell check of individual files. --files *must* be the first command line
|
||||
# arg to use this option.
|
||||
if [[ "$1" == '--files' ]]; then
|
||||
isort_check "${@:2}"
|
||||
# If `--all` is passed, then any further arguments are ignored and the
|
||||
# entire python directory is linted.
|
||||
elif [[ "$1" == '--all' ]]; then
|
||||
isort_check_all
|
||||
else
|
||||
# Check spelling only of the files that changed in last commit.
|
||||
isort_check_changed
|
||||
fi
|
||||
echo 'FastVideo isort: Done'
|
||||
@@ -1,39 +0,0 @@
|
||||
|
||||
|
||||
num_gpus=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_mochi.py \
|
||||
--model_path data/mochi \
|
||||
--prompt_path data/prompt.txt \
|
||||
--transformer_path data/outputs/video_distill_synthetic/checkpoint-1500 \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 8 \
|
||||
--guidance_scale 4.5 \
|
||||
--output_path outputs_video/distill_lq_163_1500_precision_stochastic_0.7 \
|
||||
--shift 8 \
|
||||
--seed 12345 \
|
||||
--scheduler_type "pcm_linear_quadratic"
|
||||
|
||||
|
||||
|
||||
|
||||
num_gpus=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_mochi.py \
|
||||
--model_path data/mochi \
|
||||
--prompt_embed_path "data/synthetic_debug2/prompt_embed/2.pt" \
|
||||
--encoder_attention_mask_path "data/synthetic_debug2/prompt_attention_mask/1.pt" \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 32 \
|
||||
--guidance_scale 4.5 \
|
||||
--output_path outputs_video/debug \
|
||||
--shift 8 \
|
||||
--seed 12345 \
|
||||
--scheduler_type "pcm_linear_quadratic"
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
python fastvideo/sample/sample_t2v_mochi_no_sp.py \
|
||||
--model_path data/mochi \
|
||||
--prompts "A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough." \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 64 \
|
||||
--guidance_scale 0.0 \
|
||||
--seed 12346 \
|
||||
--transformer_path data/outputs/debug/checkpoint-100/transformer \
|
||||
--output_path outputs_video/single_no_guidance
|
||||
+167
@@ -0,0 +1,167 @@
|
||||
# Prediction interface for Cog ⚙️
|
||||
# https://cog.run/python
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import time
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from cog import BasePredictor, Input, Path
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
|
||||
|
||||
MODEL_CACHE = 'FastHunyuan'
|
||||
os.environ['MODEL_BASE'] = './' + MODEL_CACHE
|
||||
|
||||
MODEL_URL = "https://weights.replicate.delivery/default/FastVideo/FastHunyuan/model.tar"
|
||||
|
||||
|
||||
def download_weights(url, dest):
|
||||
start = time.time()
|
||||
print("downloading url: ", url)
|
||||
print("downloading to: ", dest)
|
||||
subprocess.check_call(["pget", "-xf", url, dest], close_fds=False)
|
||||
print("downloading took: ", time.time() - start)
|
||||
|
||||
|
||||
class Predictor(BasePredictor):
|
||||
|
||||
def setup(self):
|
||||
"""Load the model into memory"""
|
||||
print("Model Base: " + os.environ['MODEL_BASE'])
|
||||
# Download weights
|
||||
if not os.path.exists(MODEL_CACHE):
|
||||
download_weights(MODEL_URL, MODEL_CACHE)
|
||||
|
||||
self.device = torch.device(
|
||||
"cuda" if torch.cuda.is_available() else "cpu")
|
||||
args = argparse.Namespace(
|
||||
num_frames=125,
|
||||
height=720,
|
||||
width=1280,
|
||||
num_inference_steps=6,
|
||||
fps=24,
|
||||
denoise_type='flow',
|
||||
seed=1024,
|
||||
neg_prompt=None,
|
||||
guidance_scale=1.0,
|
||||
embedded_cfg_scale=6.0,
|
||||
flow_shift=17,
|
||||
batch_size=1,
|
||||
num_videos=1,
|
||||
load_key='module',
|
||||
use_cpu_offload=False,
|
||||
dit_weight=
|
||||
'FastHunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt',
|
||||
reproduce=True,
|
||||
disable_autocast=False,
|
||||
flow_reverse=True,
|
||||
flow_solver='euler',
|
||||
use_linear_quadratic_schedule=False,
|
||||
linear_schedule_end=25,
|
||||
model='HYVideo-T/2-cfgdistill',
|
||||
latent_channels=16,
|
||||
precision='bf16',
|
||||
rope_theta=256,
|
||||
vae='884-16c-hy',
|
||||
vae_precision='fp16',
|
||||
vae_tiling=True,
|
||||
text_encoder='llm',
|
||||
text_encoder_precision='fp16',
|
||||
text_states_dim=4096,
|
||||
text_len=256,
|
||||
tokenizer='llm',
|
||||
prompt_template='dit-llm-encode',
|
||||
prompt_template_video='dit-llm-encode-video',
|
||||
hidden_state_skip_layer=2,
|
||||
apply_final_norm=False,
|
||||
text_encoder_2='clipL',
|
||||
text_encoder_precision_2='fp16',
|
||||
text_states_dim_2=768,
|
||||
tokenizer_2='clipL',
|
||||
text_len_2=77,
|
||||
model_path=MODEL_CACHE,
|
||||
)
|
||||
self.model = HunyuanVideoSampler.from_pretrained(MODEL_CACHE,
|
||||
args=args)
|
||||
|
||||
def predict(
|
||||
self,
|
||||
prompt: str = Input(
|
||||
description="Text prompt for video generation",
|
||||
default="A cat walks on the grass, realistic style."),
|
||||
negative_prompt: str = Input(
|
||||
description=
|
||||
"Text prompt to specify what you don't want in the video.",
|
||||
default=""),
|
||||
width: int = Input(description="Width of output video",
|
||||
default=1280,
|
||||
ge=256),
|
||||
height: int = Input(description="Height of output video",
|
||||
default=720,
|
||||
ge=256),
|
||||
num_frames: int = Input(description="Number of frames to generate",
|
||||
default=125,
|
||||
ge=16),
|
||||
num_inference_steps: int = Input(
|
||||
description="Number of denoising steps", default=6, ge=1, le=50),
|
||||
guidance_scale: float = Input(
|
||||
description="Classifier free guidance scale",
|
||||
default=1.0,
|
||||
ge=0.1,
|
||||
le=10.0),
|
||||
embedded_cfg_scale: float = Input(
|
||||
description="Embedded classifier free guidance scale",
|
||||
default=6.0,
|
||||
ge=0.1,
|
||||
le=10.0),
|
||||
flow_shift: int = Input(description="Flow shift parameter",
|
||||
default=17,
|
||||
ge=1,
|
||||
le=20),
|
||||
fps: int = Input(description="Frames per second of output video",
|
||||
default=24,
|
||||
ge=1,
|
||||
le=60),
|
||||
seed: int = Input(
|
||||
description="0 for Random seed. Set for reproducible generation",
|
||||
default=0),
|
||||
) -> Path:
|
||||
"""Run video generation"""
|
||||
if seed <= 0:
|
||||
seed = int.from_bytes(os.urandom(2), "big")
|
||||
print(f"Using seed: {seed}")
|
||||
|
||||
outputs = self.model.predict(
|
||||
prompt=prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
video_length=num_frames,
|
||||
seed=seed,
|
||||
negative_prompt=negative_prompt,
|
||||
infer_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
embedded_guidance_scale=embedded_cfg_scale,
|
||||
flow_shift=flow_shift,
|
||||
flow_reverse=True,
|
||||
batch_size=1,
|
||||
num_videos_per_prompt=1,
|
||||
)
|
||||
|
||||
# Process output video
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video
|
||||
output_path = Path("/tmp/output.mp4")
|
||||
imageio.mimsave(str(output_path), frames, fps=fps)
|
||||
return Path(output_path)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user