Compare commits
68
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d7b07adc08 | ||
|
|
8a41e61cec | ||
|
|
60ce6a62a6 | ||
|
|
7b2ca9abec | ||
|
|
b537e01a88 | ||
|
|
a233b58a6c | ||
|
|
d24b25a3e1 | ||
|
|
2a7c147c4e | ||
|
|
0c1c939d59 | ||
|
|
881e1f130a | ||
|
|
00e899cd90 | ||
|
|
de1e8d868e | ||
|
|
9f4151526b | ||
|
|
98b92be25e | ||
|
|
1ce3983d68 | ||
|
|
ee241cfa4d | ||
|
|
8106ee3f3d | ||
|
|
4dde52be9c | ||
|
|
58abda5c09 | ||
|
|
3e6019415a | ||
|
|
8d41d505fe | ||
|
|
8cfdf58a17 | ||
|
|
ccaf43c195 | ||
|
|
351538db29 | ||
|
|
361b24612d | ||
|
|
6ad03bea79 | ||
|
|
86e1f88877 | ||
|
|
e212a9c6b9 | ||
|
|
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 |
@@ -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,162 @@
|
||||
# Fast Video
|
||||
This is currently based on Open-Sora-1.2.0: https://github.com/PKU-YuanGroup/Open-Sora-Plan/tree/294993ca78bf65dec1c3b6fb25541432c545eda9
|
||||
# FastVideo
|
||||
|
||||
<div align="center">
|
||||
<a href=""><img src="https://img.shields.io/static/v1?label=API:H100&message=Replicate&color=pink"></a>  
|
||||
<a href=""><img src="https://img.shields.io/static/v1?label=Discuss&message=Discord&color=purple&logo=discord"></a>  
|
||||
</div>
|
||||
<br>
|
||||
<div align="center">
|
||||
<img src=assets/logo.png width="50%"/>
|
||||
</div>
|
||||
|
||||
FastVideo is a scalable framework for post-training video diffusion models, addressing the growing challenges of fine-tuning, distillation, and inference as model sizes and sequence lengths increase. As a first step, it provides an efficient script for distilling and fine-tuning the 10B Mochi model, with plans to expand features and support for more models.
|
||||
|
||||
### Features
|
||||
|
||||
- FastMochi, a distilled Mochi model that can generate videos with merely 8 sampling steps.
|
||||
- Finetuning with FSDP (both master weight and ema weight), sequence parallelism, and selective gradient checkpointing.
|
||||
- LoRA coupled with pecomputed the latents and text embedding for minumum memory consumption.
|
||||
- Finetuning with both image and videos.
|
||||
|
||||
## Change Log
|
||||
|
||||
|
||||
- ```2024/12/06```: `FastMochi` v0.0.1 is released.
|
||||
|
||||
|
||||
## Fast and High-Quality Text-to-video Generation
|
||||
|
||||
### 8-Step Results of FastMochi
|
||||
|
||||
<table class="center">
|
||||
<td><img src=assets/8steps/1.gif width="320"></td></td>
|
||||
<td><img src=assets/8steps/2.gif width="320"></td></td></td>
|
||||
<tr>
|
||||
<td style="text-align:center;" width="320">tmp</td>
|
||||
<td style="text-align:center;" width="320">tmp</td>
|
||||
<tr>
|
||||
</table >
|
||||
|
||||
|
||||
## Table of Contents
|
||||
|
||||
Jump to a specific section:
|
||||
|
||||
- [🔧 Installation](#-installation)
|
||||
- [🚀 Inference](#-inference)
|
||||
- [🎯 Distill](#-distill)
|
||||
- [⚡ Finetune](#-lora-finetune)
|
||||
|
||||
|
||||
## 🔧 Installation
|
||||
|
||||
## Envrironment
|
||||
Change the index-url cuda version according to your system.
|
||||
```
|
||||
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
|
||||
conda create -n fastmochi python=3.10.0 -y && conda activate fastmochi
|
||||
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
pip install "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo && pip install -e .
|
||||
```
|
||||
|
||||
```
|
||||
pip install -e . && pip install -e ".[train]"
|
||||
sudo apt-get update && apt install screen && pip install watch gpustat
|
||||
|
||||
|
||||
|
||||
## 🚀 Inference
|
||||
|
||||
Use [scripts/download_hf.py](scripts/download_hf.py) to download the hugging-face style model to a local directory. Use it like this:
|
||||
```bash
|
||||
python scripts/download_hf.py --repo_id=FastVideo/FastMochi --local_dir=data/FastMochi --repo_type=model
|
||||
```
|
||||
|
||||
## 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)
|
||||
|
||||
Start the gradio UI with
|
||||
```
|
||||
python fastvideo/demo/gradio_web_demo.py --model_path data/FastMochi
|
||||
```
|
||||
|
||||
We also provide CLI inference script featured with sequence parallelism.
|
||||
|
||||
```
|
||||
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
|
||||
export NUM_GPUS=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$NUM_GPUS \
|
||||
fastvideo/sample/sample_t2v_mochi.py \
|
||||
--model_path data/FastMochi \
|
||||
--prompt_path assets/prompt.txt \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 8 \
|
||||
--guidance_scale 1.5 \
|
||||
--output_path outputs_video/demo_video \
|
||||
--seed 12345 \
|
||||
--scheduler_type "pcm_linear_quadratic" \
|
||||
--linear_threshold 0.1 \
|
||||
--linear_range 0.75
|
||||
```
|
||||
|
||||
For the mochi style, simply following the scripts list in mochi repo.
|
||||
|
||||
```
|
||||
git clone https://github.com/genmoai/mochi.git
|
||||
cd mochi
|
||||
|
||||
# install env
|
||||
...
|
||||
|
||||
python3 ./demos/cli.py --model_dir weights/ --cpu_offload
|
||||
```
|
||||
|
||||
|
||||
## 🎯 Distill
|
||||
|
||||
## 💰Hardware requirement
|
||||
|
||||
- VRAM is required for both distill 10B mochi model
|
||||
|
||||
To launch distillation, you will first need to prepare data in the following formats
|
||||
|
||||
```bash
|
||||
asset/example_data
|
||||
├── AAA.txt
|
||||
├── AAA.png
|
||||
├── BCC.txt
|
||||
├── BCC.png
|
||||
├── ......
|
||||
├── CCC.txt
|
||||
└── CCC.png
|
||||
```
|
||||
|
||||
We provide a dataset example here. First download testing data. Use [scripts/download_hf.py](scripts/download_hf.py) to download the data to a local directory. Use it like this:
|
||||
```bash
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/Merge-425-Data --local_dir=data/Merge-425-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 ../..
|
||||
```
|
||||
|
||||
## 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 the distillation can be launched by:
|
||||
|
||||
## 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
|
||||
```
|
||||
bash scripts/distill_t2v.sh
|
||||
```
|
||||
|
||||
|
||||
17. no cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
18. shift16, euler_steps 50
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
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
|
||||
|
||||
## 💰Hardware requirement
|
||||
|
||||
- VRAM is required for both distill 10B mochi model
|
||||
|
||||
To launch finetuning, you will first need to prepare data in the following formats.
|
||||
|
||||
|
||||
|
||||
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
|
||||
Then the finetuning can be launched by:
|
||||
|
||||
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
|
||||
```
|
||||
bash scripts/lora_finetune.sh
|
||||
```
|
||||
|
||||
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
|
||||
|
||||
|
||||
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
|
||||
|
||||
49.
|
||||
## Acknowledgement
|
||||
We learned from and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), and [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan).
|
||||
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.
|
After Width: | Height: | Size: 380 KiB |
@@ -0,0 +1,9 @@
|
||||
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.
|
||||
En "The Matrix", Neo, interpretado por Keanu Reeves, personifica la lucha contra un sistema opresor a través de su icónica imagen, que incluye unos anteojos oscuros. Estos lentes no son solo un accesorio de moda; representan una barrera entre la realidad y la percepción. Al usar estos anteojos, Neo se sumerge en un mundo donde la verdad se oculta detrás de ilusiones y engaños. La oscuridad de los lentes simboliza la ignorancia y el control que las máquinas tienen sobre la humanidad, mientras que su propia búsqueda de la verdad lo lleva a descubrir sus auténticos poderes. La escena en que se los pone se convierte en un momento crucial, marcando su transformación de un simple programador a "El Elegido". Esta imagen se ha convertido en un ícono cultural, encapsulando el mensaje de que, al enfrentar la oscuridad, podemos encontrar la luz que nos guía hacia la libertad. Así, los anteojos de Neo se convierten en un símbolo de resistencia y autoconocimiento en un mundo manipulado.
|
||||
Medium close up. Low-angle shot. A woman in a 1950s retro dress sits in a diner bathed in neon light, surrounded by classic decor and lively chatter. The camera starts with a medium shot of her sitting at the counter, then slowly zooms in as she blows a shiny pink bubblegum bubble. The bubble swells dramatically before popping with a soft, playful burst. The scene is vibrant and nostalgic, evoking the fun and carefree spirit of the 1950s.
|
||||
Will Smith eats noodles.
|
||||
A short clip of the blonde woman taking a sip from her whiskey glass, her eyes locking with the camera as she smirks playfully. The background shows a group of people laughing and enjoying the party, with vibrant neon signs illuminating the space. The shot is taken in a way that conveys the feeling of a tipsy, carefree night out. The camera then zooms in on her face as she winks, creating a cheeky, flirtatious vibe.
|
||||
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 robot's immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
|
||||
A chimpanzee lead vocalist singing into a microphone on stage. The camera zooms in to show him singing. There is a spotlight on him.
|
||||
@@ -4,31 +4,51 @@ from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
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 (
|
||||
Normalize255,
|
||||
TemporalRandomCrop,
|
||||
CenterCropResizeVideo,
|
||||
)
|
||||
|
||||
|
||||
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)), ]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
# norm_fun
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
norm_fun
|
||||
])
|
||||
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
|
||||
]
|
||||
)
|
||||
transform_topcrop = transforms.Compose(
|
||||
[
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
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)
|
||||
|
||||
|
||||
@@ -37,32 +57,34 @@ if __name__ == "__main__":
|
||||
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,
|
||||
|
||||
}
|
||||
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 +92,9 @@ 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 +105,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")
|
||||
|
||||
@@ -4,13 +4,14 @@ import json
|
||||
import os
|
||||
import random
|
||||
|
||||
|
||||
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
|
||||
@@ -18,8 +19,10 @@ class LatentDataset(Dataset):
|
||||
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_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 +31,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
|
||||
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
# 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,25 +79,48 @@ 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
|
||||
for i, latent in enumerate(latents):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
|
||||
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)
|
||||
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()
|
||||
|
||||
@@ -13,15 +13,12 @@ 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__)
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
"""
|
||||
这是一个元类,用于创建单例类。
|
||||
"""
|
||||
_instances = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
@@ -53,7 +50,7 @@ class DataSetProg(metaclass=SingletonMeta):
|
||||
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,18 +58,20 @@ 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):
|
||||
@@ -95,11 +94,11 @@ 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 not ("mt5" 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
|
||||
@@ -121,35 +120,39 @@ class T2V_dataset(Dataset):
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.info(f'Error with {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))
|
||||
|
||||
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,70 @@ 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 = 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_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,78 +235,100 @@ 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]
|
||||
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 extention {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)}')
|
||||
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)}"
|
||||
)
|
||||
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()
|
||||
@@ -294,19 +338,20 @@ class T2V_dataset(Dataset):
|
||||
|
||||
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}...')
|
||||
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
|
||||
|
||||
+120
-66
@@ -32,7 +32,9 @@ def center_crop_arr(pil_image, image_size):
|
||||
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 +44,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 +125,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,15 +138,16 @@ 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)
|
||||
@@ -159,7 +177,9 @@ 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
|
||||
|
||||
@@ -219,7 +239,9 @@ 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
|
||||
@@ -235,7 +257,7 @@ class RandomCropVideo:
|
||||
|
||||
class SpatialStrideCropVideo:
|
||||
def __init__(self, stride):
|
||||
self.stride = stride
|
||||
self.stride = stride
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
@@ -258,17 +280,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 +314,31 @@ 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 +352,15 @@ 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 +368,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 +397,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 +408,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)
|
||||
@@ -398,13 +436,15 @@ class KineticsRandomCropResizeVideo:
|
||||
|
||||
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,13 +571,16 @@ 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__':
|
||||
|
||||
if __name__ == "__main__":
|
||||
from torchvision import transforms
|
||||
import torchvision.io as io
|
||||
import numpy as np
|
||||
@@ -544,18 +588,20 @@ if __name__ == '__main__':
|
||||
import os
|
||||
|
||||
vframes, aframes, info = io.read_video(
|
||||
filename='./v_Archery_g01_c03.avi',
|
||||
pts_unit='sec',
|
||||
output_format='TCHW'
|
||||
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)
|
||||
])
|
||||
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
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
target_video_len = 32
|
||||
frame_interval = 1
|
||||
@@ -569,7 +615,9 @@ 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 +628,18 @@ 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),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
import gradio as gr
|
||||
import torch
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
import tempfile
|
||||
import os
|
||||
import argparse
|
||||
from safetensors.torch import load_file
|
||||
|
||||
def init_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
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)
|
||||
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("--seed", type=int, default=42)
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--scheduler_type", type=str, default="euler")
|
||||
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=100)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.025)
|
||||
parser.add_argument("--linear_range", type=float, default=0.5)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
def load_model(args):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, False, args.linear_threshold, args.linear_range)
|
||||
|
||||
mochi_genmo = True
|
||||
if mochi_genmo:
|
||||
model_path = "/root/fastmochi_genmo/dit.safetensors"
|
||||
state_dcit = load_file(model_path)
|
||||
transformer = AsymmDiTJoint()
|
||||
transformer.load_state_dict(state_dcit)
|
||||
# from IPython import embed
|
||||
# embed()
|
||||
transformer.config.in_channels = 12
|
||||
print("load gennmo mochi successfully")
|
||||
else:
|
||||
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_model_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()
|
||||
|
||||
pipe = load_model(args)
|
||||
print("load model successfully")
|
||||
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()
|
||||
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# 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=8, maximum=256, value=args.num_frames)
|
||||
guidance_scale = gr.Slider(label="Guidance Scale", minimum=1, maximum=20, value=args.guidance_scale)
|
||||
num_inference_steps = gr.Slider(label="Inference Steps", minimum=10, 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)
|
||||
+505
-323
File diff suppressed because it is too large
Load Diff
@@ -23,7 +23,6 @@ from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
|
||||
class DiscriminatorHead(nn.Module):
|
||||
def __init__(self, input_channel, output_channel=1):
|
||||
super().__init__()
|
||||
@@ -48,9 +47,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)
|
||||
@@ -58,10 +57,9 @@ class DiscriminatorHead(nn.Module):
|
||||
|
||||
|
||||
class Discriminator(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stride = 8,
|
||||
stride=8,
|
||||
num_h_per_head=1,
|
||||
adapter_channel_dims=[3072],
|
||||
):
|
||||
@@ -82,24 +80,23 @@ class Discriminator(nn.Module):
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
|
||||
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]:
|
||||
for h in self.heads[i // self.stride]:
|
||||
# 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
|
||||
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ 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 fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -17,13 +17,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)))
|
||||
|
||||
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@@ -34,13 +35,14 @@ 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(
|
||||
@@ -238,6 +240,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
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
|
||||
@@ -279,7 +282,6 @@ class EulerSolver:
|
||||
multiphase,
|
||||
is_target=False,
|
||||
):
|
||||
|
||||
inference_indices = np.linspace(
|
||||
0, len(self.euler_timesteps), num=multiphase, endpoint=False
|
||||
)
|
||||
@@ -305,4 +307,3 @@ class EulerSolver:
|
||||
x_prev = sample + (sigma_prev - sigma) * model_pred
|
||||
|
||||
return x_prev, timestep_index_end
|
||||
|
||||
|
||||
+508
-259
File diff suppressed because it is too large
Load Diff
+41
-40
@@ -17,8 +17,8 @@ from torch.distributed.fsdp import (
|
||||
# ShardedStateDictConfig, # un-flattened param but shards, usable by other parallel schemes.
|
||||
)
|
||||
|
||||
from fastvideo.model.modeling_mochi import MochiTransformerBlock
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformerBlock
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmetricJointBlock
|
||||
from functools import partial
|
||||
|
||||
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
|
||||
@@ -58,38 +58,43 @@ def apply_fsdp_checkpointing(model, p=1):
|
||||
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(
|
||||
sharding_strategy, use_lora=False, cpu_offload=False, master_weight_type="fp32"
|
||||
):
|
||||
if use_lora:
|
||||
auto_wrap_policy = fsdp_auto_wrap_policy
|
||||
else:
|
||||
auto_wrap_policy = functools.partial(
|
||||
transformer_auto_wrap_policy,
|
||||
transformer_layer_cls={
|
||||
MochiTransformerBlock,
|
||||
MochiTransformerBlock, # AsymmetricJointBlock
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# 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 +103,12 @@ 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 +117,26 @@ 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,
|
||||
})
|
||||
|
||||
fsdp_kwargs.update(
|
||||
{
|
||||
"use_orig_params": False, # Required for LoRA memory savings
|
||||
"sync_module_states": True,
|
||||
}
|
||||
)
|
||||
|
||||
return fsdp_kwargs
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def get_discriminator_fsdp_kwargs():
|
||||
|
||||
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 +145,5 @@ def get_discriminator_fsdp_kwargs():
|
||||
"device_id": device_id,
|
||||
"limit_all_gathers": True,
|
||||
}
|
||||
|
||||
|
||||
return fsdp_kwargs
|
||||
|
||||
|
||||
|
||||
@@ -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,39 +0,0 @@
|
||||
import torch
|
||||
mochi_latents_mean = torch.tensor([
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor([
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
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
|
||||
|
||||
|
||||
@@ -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,29 @@
|
||||
from contextlib import contextmanager
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func as flash_varlen_attn
|
||||
except ImportError:
|
||||
flash_varlen_attn = None
|
||||
|
||||
try:
|
||||
from sageattention import sageattn as sage_attn
|
||||
except ImportError:
|
||||
sage_attn = None
|
||||
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
|
||||
training_backends = [SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION]
|
||||
eval_backends = list(training_backends)
|
||||
if torch.cuda.get_device_properties(0).major >= 9.0:
|
||||
# Enable fast CuDNN attention on Hopper.
|
||||
# This gives NaN on the backward pass for some reason,
|
||||
# so only use it for evaluation.
|
||||
eval_backends.append(SDPBackend.CUDNN_ATTENTION)
|
||||
|
||||
@contextmanager
|
||||
def sdpa_attn_ctx(training: bool = False):
|
||||
with sdpa_kernel(training_backends if training else eval_backends):
|
||||
yield
|
||||
@@ -0,0 +1,87 @@
|
||||
import contextlib
|
||||
from typing import Any, Iterable, Iterator, Optional
|
||||
|
||||
try:
|
||||
from tqdm import tqdm
|
||||
except ImportError:
|
||||
tqdm = None
|
||||
|
||||
try:
|
||||
from ray.experimental.tqdm_ray import tqdm as ray_tqdm
|
||||
except:
|
||||
ray_tqdm = None
|
||||
|
||||
# Global state
|
||||
_current_progress_type = "none"
|
||||
_is_progress_bar_active = False
|
||||
|
||||
|
||||
class DummyProgressBar:
|
||||
"""A no-op progress bar that mimics tqdm interface"""
|
||||
|
||||
def __init__(self, iterable=None, **kwargs):
|
||||
self.iterable = iterable
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.iterable)
|
||||
|
||||
def update(self, n=1):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
def set_description(self, desc):
|
||||
pass
|
||||
|
||||
|
||||
def get_new_progress_bar(iterable: Optional[Iterable] = None, **kwargs) -> Any:
|
||||
if not _is_progress_bar_active:
|
||||
return DummyProgressBar(iterable=iterable, **kwargs)
|
||||
|
||||
if _current_progress_type == "tqdm":
|
||||
if tqdm is None:
|
||||
raise ImportError("tqdm is required but not installed. Please install tqdm to use the tqdm progress bar.")
|
||||
return tqdm(iterable=iterable, **kwargs)
|
||||
elif _current_progress_type == "ray_tqdm":
|
||||
if ray_tqdm is None:
|
||||
raise ImportError("ray is required but not installed. Please install ray to use the ray_tqdm progress bar.")
|
||||
return ray_tqdm(iterable=iterable, **kwargs)
|
||||
return DummyProgressBar(iterable=iterable, **kwargs)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def progress_bar(type: str = "none", enabled=True):
|
||||
"""
|
||||
Context manager for setting progress bar type and options.
|
||||
|
||||
Args:
|
||||
type: Type of progress bar ("none" or "tqdm")
|
||||
**options: Options to pass to the progress bar (e.g., total, desc)
|
||||
|
||||
Raises:
|
||||
ValueError: If progress bar type is invalid
|
||||
RuntimeError: If progress bars are nested
|
||||
|
||||
Example:
|
||||
with progress_bar(type="tqdm", total=100):
|
||||
for i in get_new_progress_bar(range(100)):
|
||||
process(i)
|
||||
"""
|
||||
if type not in ("none", "tqdm", "ray_tqdm"):
|
||||
raise ValueError("Progress bar type must be 'none' or 'tqdm' or 'ray_tqdm'")
|
||||
if not enabled:
|
||||
type = "none"
|
||||
global _current_progress_type, _is_progress_bar_active
|
||||
|
||||
if _is_progress_bar_active:
|
||||
raise RuntimeError("Nested progress bars are not supported")
|
||||
|
||||
_is_progress_bar_active = True
|
||||
_current_progress_type = type
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_is_progress_bar_active = False
|
||||
_current_progress_type = "none"
|
||||
@@ -0,0 +1,67 @@
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
from moviepy.editor import ImageSequenceClip
|
||||
from PIL import Image
|
||||
|
||||
from genmo.lib.progress import get_new_progress_bar
|
||||
|
||||
|
||||
class Timer:
|
||||
def __init__(self):
|
||||
self.times = {} # Dictionary to store times per stage
|
||||
|
||||
def __call__(self, name):
|
||||
print(f"Timing {name}")
|
||||
return self.TimerContextManager(self, name)
|
||||
|
||||
def print_stats(self):
|
||||
total_time = sum(self.times.values())
|
||||
# Print table header
|
||||
print("{:<20} {:>10} {:>10}".format("Stage", "Time(s)", "Percent"))
|
||||
for name, t in self.times.items():
|
||||
percent = (t / total_time) * 100 if total_time > 0 else 0
|
||||
print("{:<20} {:>10.2f} {:>9.2f}%".format(name, t, percent))
|
||||
|
||||
class TimerContextManager:
|
||||
def __init__(self, outer, name):
|
||||
self.outer = outer # Reference to the Timer instance
|
||||
self.name = name
|
||||
self.start_time = None
|
||||
|
||||
def __enter__(self):
|
||||
self.start_time = time.perf_counter()
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback):
|
||||
end_time = time.perf_counter()
|
||||
elapsed = end_time - self.start_time
|
||||
self.outer.times[self.name] = self.outer.times.get(self.name, 0) + elapsed
|
||||
|
||||
|
||||
def save_video(final_frames, output_path, fps=30):
|
||||
assert final_frames.ndim == 4 and final_frames.shape[3] == 3, f"invalid shape: {final_frames} (need t h w c)"
|
||||
if final_frames.dtype != np.uint8:
|
||||
final_frames = (final_frames * 255).astype(np.uint8)
|
||||
ImageSequenceClip(list(final_frames), fps=fps).write_videofile(output_path)
|
||||
|
||||
|
||||
def create_memory_tracker():
|
||||
import torch
|
||||
|
||||
previous = [None] # Use list for mutable closure state
|
||||
|
||||
def track(label="all2all"):
|
||||
current = torch.cuda.memory_allocated() / 1e9
|
||||
if previous[0] is not None:
|
||||
diff = current - previous[0]
|
||||
sign = "+" if diff >= 0 else ""
|
||||
print(f"GPU memory ({label}): {current:.2f} GB ({sign}{diff:.2f} GB)")
|
||||
else:
|
||||
print(f"GPU memory ({label}): {current:.2f} GB")
|
||||
previous[0] = current # type: ignore
|
||||
|
||||
return track
|
||||
@@ -0,0 +1,721 @@
|
||||
import os
|
||||
from typing import Dict, List, Optional, Tuple, Any
|
||||
import warnings
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.attention import sdpa_kernel
|
||||
|
||||
from fastvideo.models.mochi_genmo.lib.attn_imports import flash_varlen_attn, sage_attn, sdpa_attn_ctx
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.layers import (
|
||||
FeedForward,
|
||||
PatchEmbed,
|
||||
RMSNorm,
|
||||
TimestepEmbedder,
|
||||
)
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.lora import LoraLinear
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.mod_rmsnorm import modulated_rmsnorm
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.residual_tanh_gated_rmsnorm import (
|
||||
residual_tanh_gated_rmsnorm,
|
||||
)
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.rope_mixed import (
|
||||
compute_mixed_rotation,
|
||||
create_position_matrix,
|
||||
)
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.temporal_rope import apply_rotary_emb_qk_real
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.utils import (
|
||||
AttentionPool,
|
||||
modulate,
|
||||
pad_and_split_xy,
|
||||
)
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.pipelines import compute_packed_indices
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
|
||||
COMPILE_FINAL_LAYER = os.environ.get("COMPILE_DIT") == "1"
|
||||
COMPILE_MMDIT_BLOCK = os.environ.get("COMPILE_DIT") == "1"
|
||||
|
||||
|
||||
def ck(fn, *args, enabled=True, **kwargs) -> torch.Tensor:
|
||||
if enabled:
|
||||
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs, use_reentrant=False)
|
||||
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
|
||||
class AsymmetricAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim_x: int,
|
||||
dim_y: int,
|
||||
num_heads: int = 8,
|
||||
qkv_bias: bool = False,
|
||||
qk_norm: bool = True,
|
||||
update_y: bool = True,
|
||||
out_bias: bool = True,
|
||||
attention_mode: str = "flash",
|
||||
softmax_scale: Optional[float] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
# Disable LoRA by default ...
|
||||
qkv_proj_lora_rank: int = 0,
|
||||
qkv_proj_lora_alpha: int = 0,
|
||||
qkv_proj_lora_dropout: float = 0.0,
|
||||
out_proj_lora_rank: int = 0,
|
||||
out_proj_lora_alpha: int = 0,
|
||||
out_proj_lora_dropout: float = 0.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.attention_mode = attention_mode
|
||||
self.dim_x = dim_x
|
||||
self.dim_y = dim_y
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim_x // num_heads
|
||||
self.update_y = update_y
|
||||
self.softmax_scale = softmax_scale
|
||||
if dim_x % num_heads != 0:
|
||||
raise ValueError(f"dim_x={dim_x} should be divisible by num_heads={num_heads}")
|
||||
|
||||
# Input layers.
|
||||
self.qkv_bias = qkv_bias
|
||||
qkv_lora_kwargs = dict(
|
||||
bias=qkv_bias,
|
||||
device=device,
|
||||
r=qkv_proj_lora_rank,
|
||||
lora_alpha=qkv_proj_lora_alpha,
|
||||
lora_dropout=qkv_proj_lora_dropout,
|
||||
)
|
||||
self.qkv_x = LoraLinear(dim_x, 3 * dim_x, **qkv_lora_kwargs)
|
||||
# Project text features to match visual features (dim_y -> dim_x)
|
||||
self.qkv_y = LoraLinear(dim_y, 3 * dim_x, **qkv_lora_kwargs)
|
||||
|
||||
# Query and key normalization for stability.
|
||||
assert qk_norm
|
||||
self.q_norm_x = RMSNorm(self.head_dim, device=device)
|
||||
self.k_norm_x = RMSNorm(self.head_dim, device=device)
|
||||
self.q_norm_y = RMSNorm(self.head_dim, device=device)
|
||||
self.k_norm_y = RMSNorm(self.head_dim, device=device)
|
||||
|
||||
# Output layers. y features go back down from dim_x -> dim_y.
|
||||
proj_lora_kwargs = dict(
|
||||
bias=out_bias,
|
||||
device=device,
|
||||
r=out_proj_lora_rank,
|
||||
lora_alpha=out_proj_lora_alpha,
|
||||
lora_dropout=out_proj_lora_dropout,
|
||||
)
|
||||
self.proj_x = LoraLinear(dim_x, dim_x, **proj_lora_kwargs)
|
||||
self.proj_y = LoraLinear(dim_x, dim_y, **proj_lora_kwargs) if update_y else nn.Identity()
|
||||
|
||||
def run_qkv_y(self, y):
|
||||
local_heads = self.num_heads
|
||||
|
||||
qkv_y = self.qkv_y(y) # (B, L, 3 * dim)
|
||||
|
||||
qkv_y = qkv_y.view(qkv_y.size(0), qkv_y.size(1), 3, local_heads, self.head_dim)
|
||||
q_y, k_y, v_y = qkv_y.unbind(2)
|
||||
|
||||
q_y = self.q_norm_y(q_y)
|
||||
k_y = self.k_norm_y(k_y)
|
||||
return q_y, k_y, v_y
|
||||
|
||||
def prepare_qkv(
|
||||
self,
|
||||
x: torch.Tensor, # (B, M, dim_x)
|
||||
y: torch.Tensor, # (B, L, dim_y)
|
||||
*,
|
||||
scale_x: torch.Tensor,
|
||||
scale_y: torch.Tensor,
|
||||
rope_cos: torch.Tensor,
|
||||
rope_sin: torch.Tensor,
|
||||
valid_token_indices: torch.Tensor,
|
||||
max_seqlen_in_batch: int,
|
||||
):
|
||||
# Process visual features
|
||||
x = modulated_rmsnorm(x, scale_x) # (B, M, dim_x) where M = N
|
||||
qkv_x = self.qkv_x(x) # (B, M, 3 * dim_x)
|
||||
assert qkv_x.dtype == torch.bfloat16
|
||||
B, M, _ = qkv_x.size()
|
||||
qkv_x = qkv_x.view(B, M, 3, self.num_heads, -1)
|
||||
qkv_x = qkv_x.permute(2, 0, 1, 3, 4)
|
||||
|
||||
# Split qkv_x into q, k, v
|
||||
q_x, k_x, v_x = qkv_x.unbind(0) # (B, N, local_h, head_dim)
|
||||
q_x = self.q_norm_x(q_x)
|
||||
q_x = apply_rotary_emb_qk_real(q_x, rope_cos, rope_sin)
|
||||
k_x = self.k_norm_x(k_x)
|
||||
k_x = apply_rotary_emb_qk_real(k_x, rope_cos, rope_sin)
|
||||
|
||||
# Concatenate streams
|
||||
B, N, num_heads, head_dim = q_x.size()
|
||||
D = num_heads * head_dim
|
||||
|
||||
# Process text features
|
||||
if B == 1:
|
||||
text_seqlen = max_seqlen_in_batch - N
|
||||
if text_seqlen > 0:
|
||||
y = y[:, :text_seqlen] # Remove padding tokens.
|
||||
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
|
||||
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
|
||||
|
||||
q = torch.cat([q_x, q_y], dim=1)
|
||||
k = torch.cat([k_x, k_y], dim=1)
|
||||
v = torch.cat([v_x, v_y], dim=1)
|
||||
else:
|
||||
q, k, v = q_x, k_x, v_x
|
||||
else:
|
||||
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
|
||||
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
|
||||
|
||||
indices = valid_token_indices[:, None].expand(-1, D)
|
||||
q = torch.cat([q_x, q_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
|
||||
k = torch.cat([k_x, k_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
|
||||
v = torch.cat([v_x, v_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
|
||||
|
||||
q = q.view(-1, num_heads, head_dim)
|
||||
k = k.view(-1, num_heads, head_dim)
|
||||
v = v.view(-1, num_heads, head_dim)
|
||||
return q, k, v
|
||||
|
||||
@torch.autocast("cuda", enabled=False)
|
||||
def flash_attention(self, q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim):
|
||||
out: torch.Tensor = flash_varlen_attn(
|
||||
q, k, v,
|
||||
cu_seqlens_q=cu_seqlens,
|
||||
cu_seqlens_k=cu_seqlens,
|
||||
max_seqlen_q=max_seqlen_in_batch,
|
||||
max_seqlen_k=max_seqlen_in_batch,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=self.softmax_scale,
|
||||
) # (total, local_heads, head_dim)
|
||||
return out.view(total, local_dim)
|
||||
|
||||
def sdpa_attention(self, q, k, v):
|
||||
with sdpa_attn_ctx(training=self.training):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
)
|
||||
return out
|
||||
|
||||
@torch.autocast("cuda", enabled=False)
|
||||
def sage_attention(self, q, k, v):
|
||||
return sage_attn(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
|
||||
|
||||
def run_attention(
|
||||
self,
|
||||
q: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
|
||||
k: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
|
||||
v: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
|
||||
*,
|
||||
B: int,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
max_seqlen_in_batch: Optional[int] = None,
|
||||
):
|
||||
local_heads = self.num_heads
|
||||
local_dim = local_heads * self.head_dim
|
||||
|
||||
# Check shapes
|
||||
assert q.ndim == 3 and k.ndim == 3 and v.ndim == 3
|
||||
total = q.size(0)
|
||||
assert k.size(0) == total and v.size(0) == total
|
||||
|
||||
if self.attention_mode == "flash":
|
||||
out = self.flash_attention(
|
||||
q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim) # (total, local_dim)
|
||||
else:
|
||||
assert B == 1, \
|
||||
f"Non-flash attention mode {self.attention_mode} only supports batch size 1, got {B}"
|
||||
|
||||
q = rearrange(q, "(b s) h d -> b h s d", b=B)
|
||||
k = rearrange(k, "(b s) h d -> b h s d", b=B)
|
||||
v = rearrange(v, "(b s) h d -> b h s d", b=B)
|
||||
|
||||
if self.attention_mode == "sdpa":
|
||||
out = self.sdpa_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
|
||||
elif self.attention_mode == "sage":
|
||||
out = self.sage_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
|
||||
else:
|
||||
raise ValueError(f"Unknown attention mode: {self.attention_mode}")
|
||||
|
||||
out = rearrange(out, "b h s d -> (b s) (h d)")
|
||||
|
||||
return out
|
||||
|
||||
def post_attention(
|
||||
self,
|
||||
out: torch.Tensor,
|
||||
B: int,
|
||||
M: int,
|
||||
L: int,
|
||||
dtype: torch.dtype,
|
||||
valid_token_indices: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
out: (total <= B * (N + L), local_dim)
|
||||
valid_token_indices: (total <= B * (N + L),)
|
||||
B: Batch size
|
||||
M: Number of visual tokens per context parallel rank
|
||||
L: Number of text tokens
|
||||
dtype: Data type of the input and output tensors
|
||||
|
||||
Returns:
|
||||
x: (B, N, dim_x) tensor of visual tokens where N = M
|
||||
y: (B, L, dim_y) tensor of text token features
|
||||
"""
|
||||
local_heads = self.num_heads
|
||||
local_dim = local_heads * self.head_dim
|
||||
N = M
|
||||
|
||||
# Split sequence into visual and text tokens, adding back padding.
|
||||
if B == 1:
|
||||
out = out.view(B, -1, local_dim)
|
||||
if out.size(1) > N:
|
||||
x, y = torch.tensor_split(out, (N,), dim=1) # (B, N, local_dim), (B, <= L, local_dim)
|
||||
y = F.pad(y, (0, 0, 0, L - y.size(1))) # (B, L, local_dim)
|
||||
else:
|
||||
# Empty prompt.
|
||||
x, y = out, out.new_zeros(B, L, local_dim)
|
||||
else:
|
||||
x, y = pad_and_split_xy(out, valid_token_indices, B, N, L, dtype)
|
||||
assert x.size() == (B, N, local_dim)
|
||||
assert y.size() == (B, L, local_dim)
|
||||
|
||||
# Communicate across context parallel ranks.
|
||||
x = x.view(B, N, local_heads, self.head_dim)
|
||||
x = x.view(x.size(0), x.size(1), x.size(2) * x.size(3)) # (B, M, dim_x = num_heads * head_dim)
|
||||
|
||||
x = self.proj_x(x)
|
||||
y = self.proj_y(y)
|
||||
return x, y
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # (B, M, dim_x)
|
||||
y: torch.Tensor, # (B, L, dim_y)
|
||||
*,
|
||||
scale_x: torch.Tensor, # (B, dim_x), modulation for pre-RMSNorm.
|
||||
scale_y: torch.Tensor, # (B, dim_y), modulation for pre-RMSNorm.
|
||||
packed_indices: Dict[str, torch.Tensor] = None,
|
||||
checkpoint_qkv: bool = False,
|
||||
checkpoint_post_attn: bool = False,
|
||||
**rope_rotation,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward pass of asymmetric multi-modal attention.
|
||||
|
||||
Args:
|
||||
x: (B, M, dim_x) tensor of visual tokens
|
||||
y: (B, L, dim_y) tensor of text token features
|
||||
packed_indices: Dict with keys for Flash Attention
|
||||
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
|
||||
|
||||
Returns:
|
||||
x: (B, M, dim_x) tensor of visual tokens after multi-modal attention
|
||||
y: (B, L, dim_y) tensor of text token features after multi-modal attention
|
||||
"""
|
||||
B, L, _ = y.shape
|
||||
_, M, _ = x.shape
|
||||
|
||||
# Predict a packed QKV tensor from visual and text features.
|
||||
q, k, v = ck(self.prepare_qkv,
|
||||
x=x,
|
||||
y=y,
|
||||
scale_x=scale_x,
|
||||
scale_y=scale_y,
|
||||
rope_cos=rope_rotation.get("rope_cos"),
|
||||
rope_sin=rope_rotation.get("rope_sin"),
|
||||
valid_token_indices=packed_indices["valid_token_indices_kv"],
|
||||
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
|
||||
enabled=checkpoint_qkv,
|
||||
) # (total <= B * (N + L), 3, local_heads, head_dim)
|
||||
|
||||
# Self-attention is expensive, so don't checkpoint it.
|
||||
out = self.run_attention(
|
||||
q, k, v, B=B,
|
||||
cu_seqlens=packed_indices["cu_seqlens_kv"],
|
||||
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
|
||||
)
|
||||
|
||||
x, y = ck(self.post_attention,
|
||||
out,
|
||||
B=B, M=M, L=L,
|
||||
dtype=v.dtype,
|
||||
valid_token_indices=packed_indices["valid_token_indices_kv"],
|
||||
enabled=checkpoint_post_attn,
|
||||
)
|
||||
|
||||
return x, y
|
||||
|
||||
|
||||
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
|
||||
class AsymmetricJointBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size_x: int,
|
||||
hidden_size_y: int,
|
||||
num_heads: int,
|
||||
*,
|
||||
mlp_ratio_x: float = 8.0, # Ratio of hidden size to d_model for MLP for visual tokens.
|
||||
mlp_ratio_y: float = 4.0, # Ratio of hidden size to d_model for MLP for text tokens.
|
||||
update_y: bool = True, # Whether to update text tokens in this block.
|
||||
device: Optional[torch.device] = None,
|
||||
**block_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.update_y = update_y
|
||||
self.hidden_size_x = hidden_size_x
|
||||
self.hidden_size_y = hidden_size_y
|
||||
self.mod_x = nn.Linear(hidden_size_x, 4 * hidden_size_x, device=device)
|
||||
if self.update_y:
|
||||
self.mod_y = nn.Linear(hidden_size_x, 4 * hidden_size_y, device=device)
|
||||
else:
|
||||
self.mod_y = nn.Linear(hidden_size_x, hidden_size_y, device=device)
|
||||
|
||||
# Self-attention:
|
||||
self.attn = AsymmetricAttention(
|
||||
hidden_size_x,
|
||||
hidden_size_y,
|
||||
num_heads=num_heads,
|
||||
update_y=update_y,
|
||||
device=device,
|
||||
**block_kwargs,
|
||||
)
|
||||
|
||||
# MLP.
|
||||
mlp_hidden_dim_x = int(hidden_size_x * mlp_ratio_x)
|
||||
assert mlp_hidden_dim_x == int(1536 * 8)
|
||||
self.mlp_x = FeedForward(
|
||||
in_features=hidden_size_x,
|
||||
hidden_size=mlp_hidden_dim_x,
|
||||
multiple_of=256,
|
||||
ffn_dim_multiplier=None,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# MLP for text not needed in last block.
|
||||
if self.update_y:
|
||||
mlp_hidden_dim_y = int(hidden_size_y * mlp_ratio_y)
|
||||
self.mlp_y = FeedForward(
|
||||
in_features=hidden_size_y,
|
||||
hidden_size=mlp_hidden_dim_y,
|
||||
multiple_of=256,
|
||||
ffn_dim_multiplier=None,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
c: torch.Tensor,
|
||||
y: torch.Tensor,
|
||||
# TODO: These could probably just go into attn_kwargs
|
||||
checkpoint_ff: bool = False,
|
||||
checkpoint_qkv: bool = False,
|
||||
checkpoint_post_attn: bool = False,
|
||||
**attn_kwargs,
|
||||
):
|
||||
"""Forward pass of a block.
|
||||
|
||||
Args:
|
||||
x: (B, N, dim) tensor of visual tokens
|
||||
c: (B, dim) tensor of conditioned features
|
||||
y: (B, L, dim) tensor of text tokens
|
||||
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
|
||||
|
||||
Returns:
|
||||
x: (B, N, dim) tensor of visual tokens after block
|
||||
y: (B, L, dim) tensor of text tokens after block
|
||||
"""
|
||||
N = x.size(1)
|
||||
|
||||
c = F.silu(c)
|
||||
mod_x = self.mod_x(c)
|
||||
scale_msa_x, gate_msa_x, scale_mlp_x, gate_mlp_x = mod_x.chunk(4, dim=1)
|
||||
mod_y = self.mod_y(c)
|
||||
|
||||
if self.update_y:
|
||||
scale_msa_y, gate_msa_y, scale_mlp_y, gate_mlp_y = mod_y.chunk(4, dim=1)
|
||||
else:
|
||||
scale_msa_y = mod_y
|
||||
|
||||
# Self-attention block.
|
||||
x_attn, y_attn = self.attn(
|
||||
x,
|
||||
y,
|
||||
scale_x=scale_msa_x,
|
||||
scale_y=scale_msa_y,
|
||||
checkpoint_qkv=checkpoint_qkv,
|
||||
checkpoint_post_attn=checkpoint_post_attn,
|
||||
**attn_kwargs,
|
||||
)
|
||||
|
||||
assert x_attn.size(1) == N
|
||||
x = residual_tanh_gated_rmsnorm(x, x_attn, gate_msa_x)
|
||||
|
||||
if self.update_y:
|
||||
y = residual_tanh_gated_rmsnorm(y, y_attn, gate_msa_y)
|
||||
|
||||
# MLP block.
|
||||
x = ck(self.ff_block_x, x, scale_mlp_x, gate_mlp_x, enabled=checkpoint_ff)
|
||||
if self.update_y:
|
||||
y = ck(self.ff_block_y, y, scale_mlp_y, gate_mlp_y, enabled=checkpoint_ff) # type: ignore
|
||||
return x, y
|
||||
|
||||
def ff_block_x(self, x, scale_x, gate_x):
|
||||
x_mod = modulated_rmsnorm(x, scale_x)
|
||||
x_res = self.mlp_x(x_mod)
|
||||
x = residual_tanh_gated_rmsnorm(x, x_res, gate_x) # Sandwich norm
|
||||
return x
|
||||
|
||||
def ff_block_y(self, y, scale_y, gate_y):
|
||||
y_mod = modulated_rmsnorm(y, scale_y)
|
||||
y_res = self.mlp_y(y_mod)
|
||||
y = residual_tanh_gated_rmsnorm(y, y_res, gate_y) # Sandwich norm
|
||||
return y
|
||||
|
||||
|
||||
@torch.compile(disable=not COMPILE_FINAL_LAYER)
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of DiT.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
patch_size,
|
||||
out_channels,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, device=device)
|
||||
self.mod = nn.Linear(hidden_size, 2 * hidden_size, device=device)
|
||||
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, device=device)
|
||||
|
||||
def forward(self, x, c):
|
||||
c = F.silu(c)
|
||||
shift, scale = self.mod(c).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class AsymmDiTJoint(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
"""
|
||||
Diffusion model with a Transformer backbone.
|
||||
|
||||
Ingests text embeddings instead of a label.
|
||||
"""
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
patch_size=2,
|
||||
in_channels=12,
|
||||
hidden_size_x=3072,
|
||||
hidden_size_y=1536,
|
||||
depth=48,
|
||||
num_heads=24,
|
||||
mlp_ratio_x=4.0,
|
||||
mlp_ratio_y=4.0,
|
||||
t5_feat_dim: int = 4096,
|
||||
t5_token_length: int = 256,
|
||||
patch_embed_bias: bool = True,
|
||||
timestep_mlp_bias: bool = True,
|
||||
timestep_scale: float = 1000.0,
|
||||
use_extended_posenc: bool = False,
|
||||
rope_theta: float = 10000.0,
|
||||
device: Optional[torch.device] = None,
|
||||
**block_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels
|
||||
self.patch_size = patch_size
|
||||
self.num_heads = num_heads
|
||||
self.hidden_size_x = hidden_size_x
|
||||
self.hidden_size_y = hidden_size_y
|
||||
self.head_dim = hidden_size_x // num_heads # Head dimension and count is determined by visual.
|
||||
self.use_extended_posenc = use_extended_posenc
|
||||
self.t5_token_length = t5_token_length
|
||||
self.t5_feat_dim = t5_feat_dim
|
||||
self.rope_theta = rope_theta # Scaling factor for frequency computation for temporal RoPE.
|
||||
self.timestep_scale = timestep_scale
|
||||
|
||||
self.x_embedder = PatchEmbed(
|
||||
patch_size=patch_size,
|
||||
in_chans=in_channels,
|
||||
embed_dim=hidden_size_x,
|
||||
bias=patch_embed_bias,
|
||||
device=device,
|
||||
)
|
||||
# Conditionings
|
||||
# Timestep
|
||||
self.t_embedder = TimestepEmbedder(hidden_size_x, bias=timestep_mlp_bias, timestep_scale=timestep_scale)
|
||||
|
||||
# Caption Pooling (T5)
|
||||
self.t5_y_embedder = AttentionPool(t5_feat_dim, num_heads=8, output_dim=hidden_size_x, device=device)
|
||||
|
||||
# Dense Embedding Projection (T5)
|
||||
self.t5_yproj = nn.Linear(t5_feat_dim, hidden_size_y, bias=True, device=device)
|
||||
|
||||
# Initialize pos_frequencies as an empty parameter.
|
||||
self.pos_frequencies = nn.Parameter(torch.empty(3, self.num_heads, self.head_dim // 2, device=device))
|
||||
|
||||
# for depth 48:
|
||||
# b = 0: AsymmetricJointBlock, update_y=True
|
||||
# b = 1: AsymmetricJointBlock, update_y=True
|
||||
# ...
|
||||
# b = 46: AsymmetricJointBlock, update_y=True
|
||||
# b = 47: AsymmetricJointBlock, update_y=False. No need to update text features.
|
||||
blocks = []
|
||||
for b in range(depth):
|
||||
# Joint multi-modal block
|
||||
update_y = b < depth - 1
|
||||
block = AsymmetricJointBlock(
|
||||
hidden_size_x,
|
||||
hidden_size_y,
|
||||
num_heads,
|
||||
mlp_ratio_x=mlp_ratio_x,
|
||||
mlp_ratio_y=mlp_ratio_y,
|
||||
update_y=update_y,
|
||||
device=device,
|
||||
**block_kwargs,
|
||||
)
|
||||
|
||||
blocks.append(block)
|
||||
self.blocks = nn.ModuleList(blocks)
|
||||
|
||||
self.final_layer = FinalLayer(hidden_size_x, patch_size, self.out_channels, device=device)
|
||||
|
||||
def embed_x(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: (B, C=12, T, H, W) tensor of visual tokens
|
||||
|
||||
Returns:
|
||||
x: (B, C=3072, N) tensor of visual tokens with positional embedding.
|
||||
"""
|
||||
return self.x_embedder(x) # Convert BcTHW to BCN
|
||||
|
||||
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
|
||||
def prepare(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
t5_feat: torch.Tensor,
|
||||
t5_mask: torch.Tensor,
|
||||
):
|
||||
"""Prepare input and conditioning embeddings."""
|
||||
|
||||
# Visual patch embeddings with positional encoding.
|
||||
T, H, W = x.shape[-3:]
|
||||
pH, pW = H // self.patch_size, W // self.patch_size
|
||||
x = self.embed_x(x) # (B, N, D), where N = T * H * W / patch_size ** 2
|
||||
assert x.ndim == 3
|
||||
B = x.size(0)
|
||||
|
||||
# Construct position array of size [N, 3].
|
||||
# pos[:, 0] is the frame index for each location,
|
||||
# pos[:, 1] is the row index for each location, and
|
||||
# pos[:, 2] is the column index for each location.
|
||||
N = T * pH * pW
|
||||
assert x.size(1) == N
|
||||
pos = create_position_matrix(T, pH=pH, pW=pW, device=x.device, dtype=torch.float32) # (N, 3)
|
||||
rope_cos, rope_sin = compute_mixed_rotation(
|
||||
freqs=self.pos_frequencies, pos=pos
|
||||
) # Each are (N, num_heads, dim // 2)
|
||||
|
||||
# Global vector embedding for conditionings.
|
||||
c_t = self.t_embedder(1 - sigma) # (B, D)
|
||||
|
||||
# Pool T5 tokens using attention pooler
|
||||
# Note encoder_hidden_states[1] contains T5 token features.
|
||||
assert (
|
||||
t5_feat.size(1) == self.t5_token_length
|
||||
), f"Expected L={self.t5_token_length}, got {t5_feat.shape} for encoder_hidden_states."
|
||||
t5_y_pool = self.t5_y_embedder(t5_feat, t5_mask) # (B, D)
|
||||
assert t5_y_pool.size(0) == B, f"Expected B={B}, got {t5_y_pool.shape} for t5_y_pool."
|
||||
|
||||
c = c_t + t5_y_pool
|
||||
|
||||
encoder_hidden_states = self.t5_yproj(t5_feat) # (B, L, t5_feat_dim) --> (B, L, D)
|
||||
|
||||
return x, c, encoder_hidden_states, rope_cos, rope_sin
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor, # [1, 12, 7, 60, 106]
|
||||
timestep: torch.Tensor, # [1]
|
||||
encoder_hidden_states: torch.Tensor, # [1, 256, 4096]
|
||||
encoder_attention_mask: torch.Tensor, # [1, 256]
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
rope_cos: torch.Tensor = None,
|
||||
rope_sin: torch.Tensor = None,
|
||||
num_ff_checkpoint: int = 48, # 48
|
||||
num_qkv_checkpoint: int = 48, # 48
|
||||
num_post_attn_checkpoint: int = 0, # 0
|
||||
):
|
||||
"""Forward pass of DiT.
|
||||
|
||||
Args:
|
||||
hidden_states: (B, C, T, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||
sigma: (B,) tensor of noise standard deviations
|
||||
encoder_hidden_states: List((B, L, encoder_hidden_states_dim) tensor of caption token features. For SDXL text encoders: L=77, encoder_hidden_states_dim=2048)
|
||||
encoder_attention_mask: List((B, L) boolean tensor indicating which tokens are not padding)
|
||||
packed_indices: Dict with keys for Flash Attention. Result of compute_packed_indices.
|
||||
{'cu_seqlens_kv': tensor([ 0, 11230], device='cuda:0', dtype=torch.int32),
|
||||
'max_seqlen_in_batch_kv': 11230,
|
||||
'valid_token_indices_kv': tensor([ 0, 1, 2, ..., 11227, 11228, 11229], device='cuda:0')}
|
||||
"""
|
||||
sigma = timestep / self.timestep_scale
|
||||
num_latent_toks = np.prod(hidden_states.shape[-3:])
|
||||
|
||||
packed_indices = compute_packed_indices(hidden_states.device, encoder_attention_mask, int(num_latent_toks))
|
||||
|
||||
_, _, T, H, W = hidden_states.shape
|
||||
|
||||
if self.pos_frequencies.dtype != torch.float32:
|
||||
warnings.warn(f"pos_frequencies dtype {self.pos_frequencies.dtype} != torch.float32")
|
||||
|
||||
# Use EFFICIENT_ATTENTION backend for T5 pooling, since we have a mask.
|
||||
# Have to call sdpa_kernel outside of a torch.compile region.
|
||||
with sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION):
|
||||
hidden_states, c, encoder_hidden_states, rope_cos, rope_sin = self.prepare(hidden_states, sigma, encoder_hidden_states, encoder_attention_mask) # [1, 11130, 3072], [1, 3072], [1, 256, 1536], [11130, 24, 64]
|
||||
del encoder_attention_mask
|
||||
|
||||
for i, block in enumerate(self.blocks):
|
||||
hidden_states, encoder_hidden_states = block( # [1, 11130, 3072], [1, 256, 1536]
|
||||
hidden_states,
|
||||
c,
|
||||
encoder_hidden_states,
|
||||
rope_cos=rope_cos,
|
||||
rope_sin=rope_sin,
|
||||
packed_indices=packed_indices,
|
||||
checkpoint_ff=i < num_ff_checkpoint,
|
||||
checkpoint_qkv=i < num_qkv_checkpoint,
|
||||
checkpoint_post_attn=i < num_post_attn_checkpoint,
|
||||
) # (B, M, D), (B, L, D)
|
||||
del encoder_hidden_states # Final layers don't use dense text features.
|
||||
|
||||
hidden_states = self.final_layer(hidden_states, c) # (B, M, patch_size ** 2 * out_channels) [1, 11130, 48]
|
||||
|
||||
hidden_states = rearrange( # [1, 12, 7, 60, 106]
|
||||
hidden_states,
|
||||
"B (T hp wp) (p1 p2 c) -> B c T (hp p1) (wp p2)",
|
||||
T=T,
|
||||
hp=H // self.patch_size,
|
||||
wp=W // self.patch_size,
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
c=self.out_channels,
|
||||
)
|
||||
|
||||
attn_outputs_list = None
|
||||
return (-hidden_states, attn_outputs_list)
|
||||
@@ -0,0 +1,737 @@
|
||||
import os
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.attention import sdpa_kernel
|
||||
|
||||
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
|
||||
from genmo.lib.attn_imports import flash_varlen_attn, sage_attn, sdpa_attn_ctx
|
||||
from genmo.mochi_preview.dit.joint_model.layers import (
|
||||
FeedForward,
|
||||
PatchEmbed,
|
||||
RMSNorm,
|
||||
TimestepEmbedder,
|
||||
)
|
||||
from genmo.mochi_preview.dit.joint_model.lora import LoraLinear
|
||||
from genmo.mochi_preview.dit.joint_model.mod_rmsnorm import modulated_rmsnorm
|
||||
from genmo.mochi_preview.dit.joint_model.residual_tanh_gated_rmsnorm import (
|
||||
residual_tanh_gated_rmsnorm,
|
||||
)
|
||||
from genmo.mochi_preview.dit.joint_model.rope_mixed import (
|
||||
compute_mixed_rotation,
|
||||
create_position_matrix,
|
||||
)
|
||||
from genmo.mochi_preview.dit.joint_model.temporal_rope import apply_rotary_emb_qk_real
|
||||
from genmo.mochi_preview.dit.joint_model.utils import (
|
||||
AttentionPool,
|
||||
modulate,
|
||||
pad_and_split_xy,
|
||||
)
|
||||
|
||||
COMPILE_FINAL_LAYER = os.environ.get("COMPILE_DIT") == "1"
|
||||
COMPILE_MMDIT_BLOCK = os.environ.get("COMPILE_DIT") == "1"
|
||||
|
||||
|
||||
def ck(fn, *args, enabled=True, **kwargs) -> torch.Tensor:
|
||||
if enabled:
|
||||
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs, use_reentrant=False)
|
||||
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
|
||||
class AsymmetricAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim_x: int,
|
||||
dim_y: int,
|
||||
num_heads: int = 8,
|
||||
qkv_bias: bool = True,
|
||||
qk_norm: bool = False,
|
||||
update_y: bool = True,
|
||||
out_bias: bool = True,
|
||||
attention_mode: str = "flash",
|
||||
softmax_scale: Optional[float] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
# Disable LoRA by default ...
|
||||
qkv_proj_lora_rank: int = 0,
|
||||
qkv_proj_lora_alpha: int = 0,
|
||||
qkv_proj_lora_dropout: float = 0.0,
|
||||
out_proj_lora_rank: int = 0,
|
||||
out_proj_lora_alpha: int = 0,
|
||||
out_proj_lora_dropout: float = 0.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.attention_mode = attention_mode
|
||||
self.dim_x = dim_x
|
||||
self.dim_y = dim_y
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim_x // num_heads
|
||||
self.update_y = update_y
|
||||
self.softmax_scale = softmax_scale
|
||||
if dim_x % num_heads != 0:
|
||||
raise ValueError(f"dim_x={dim_x} should be divisible by num_heads={num_heads}")
|
||||
|
||||
# Input layers.
|
||||
self.qkv_bias = qkv_bias
|
||||
qkv_lora_kwargs = dict(
|
||||
bias=qkv_bias,
|
||||
device=device,
|
||||
r=qkv_proj_lora_rank,
|
||||
lora_alpha=qkv_proj_lora_alpha,
|
||||
lora_dropout=qkv_proj_lora_dropout,
|
||||
)
|
||||
self.qkv_x = LoraLinear(dim_x, 3 * dim_x, **qkv_lora_kwargs)
|
||||
# Project text features to match visual features (dim_y -> dim_x)
|
||||
self.qkv_y = LoraLinear(dim_y, 3 * dim_x, **qkv_lora_kwargs)
|
||||
|
||||
# Query and key normalization for stability.
|
||||
assert qk_norm
|
||||
self.q_norm_x = RMSNorm(self.head_dim, device=device)
|
||||
self.k_norm_x = RMSNorm(self.head_dim, device=device)
|
||||
self.q_norm_y = RMSNorm(self.head_dim, device=device)
|
||||
self.k_norm_y = RMSNorm(self.head_dim, device=device)
|
||||
|
||||
# Output layers. y features go back down from dim_x -> dim_y.
|
||||
proj_lora_kwargs = dict(
|
||||
bias=out_bias,
|
||||
device=device,
|
||||
r=out_proj_lora_rank,
|
||||
lora_alpha=out_proj_lora_alpha,
|
||||
lora_dropout=out_proj_lora_dropout,
|
||||
)
|
||||
self.proj_x = LoraLinear(dim_x, dim_x, **proj_lora_kwargs)
|
||||
self.proj_y = LoraLinear(dim_x, dim_y, **proj_lora_kwargs) if update_y else nn.Identity()
|
||||
|
||||
def run_qkv_y(self, y):
|
||||
cp_rank, cp_size = cp.get_cp_rank_size()
|
||||
local_heads = self.num_heads // cp_size
|
||||
|
||||
if cp.is_cp_active():
|
||||
# Only predict local heads.
|
||||
assert not self.qkv_bias
|
||||
W_qkv_y = self.qkv_y.weight.view(3, self.num_heads, self.head_dim, self.dim_y)
|
||||
W_qkv_y = W_qkv_y.narrow(1, cp_rank * local_heads, local_heads)
|
||||
W_qkv_y = W_qkv_y.reshape(3 * local_heads * self.head_dim, self.dim_y)
|
||||
qkv_y = F.linear(y, W_qkv_y, None) # (B, L, 3 * local_h * head_dim)
|
||||
else:
|
||||
qkv_y = self.qkv_y(y) # (B, L, 3 * dim)
|
||||
|
||||
qkv_y = qkv_y.view(qkv_y.size(0), qkv_y.size(1), 3, local_heads, self.head_dim)
|
||||
q_y, k_y, v_y = qkv_y.unbind(2)
|
||||
|
||||
q_y = self.q_norm_y(q_y)
|
||||
k_y = self.k_norm_y(k_y)
|
||||
return q_y, k_y, v_y
|
||||
|
||||
def prepare_qkv(
|
||||
self,
|
||||
x: torch.Tensor, # (B, M, dim_x)
|
||||
y: torch.Tensor, # (B, L, dim_y)
|
||||
*,
|
||||
scale_x: torch.Tensor,
|
||||
scale_y: torch.Tensor,
|
||||
rope_cos: torch.Tensor,
|
||||
rope_sin: torch.Tensor,
|
||||
valid_token_indices: torch.Tensor,
|
||||
max_seqlen_in_batch: int,
|
||||
):
|
||||
# Process visual features
|
||||
x = modulated_rmsnorm(x, scale_x) # (B, M, dim_x) where M = N / cp_group_size
|
||||
qkv_x = self.qkv_x(x) # (B, M, 3 * dim_x)
|
||||
assert qkv_x.dtype == torch.bfloat16
|
||||
|
||||
qkv_x = cp.all_to_all_collect_tokens(qkv_x, self.num_heads) # (3, B, N, local_h, head_dim)
|
||||
|
||||
# Split qkv_x into q, k, v
|
||||
q_x, k_x, v_x = qkv_x.unbind(0) # (B, N, local_h, head_dim)
|
||||
q_x = self.q_norm_x(q_x)
|
||||
q_x = apply_rotary_emb_qk_real(q_x, rope_cos, rope_sin)
|
||||
k_x = self.k_norm_x(k_x)
|
||||
k_x = apply_rotary_emb_qk_real(k_x, rope_cos, rope_sin)
|
||||
|
||||
# Concatenate streams
|
||||
B, N, num_heads, head_dim = q_x.size()
|
||||
D = num_heads * head_dim
|
||||
|
||||
# Process text features
|
||||
if B == 1:
|
||||
text_seqlen = max_seqlen_in_batch - N
|
||||
if text_seqlen > 0:
|
||||
y = y[:, :text_seqlen] # Remove padding tokens.
|
||||
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
|
||||
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
|
||||
|
||||
q = torch.cat([q_x, q_y], dim=1)
|
||||
k = torch.cat([k_x, k_y], dim=1)
|
||||
v = torch.cat([v_x, v_y], dim=1)
|
||||
else:
|
||||
q, k, v = q_x, k_x, v_x
|
||||
else:
|
||||
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
|
||||
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
|
||||
|
||||
indices = valid_token_indices[:, None].expand(-1, D)
|
||||
q = torch.cat([q_x, q_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
|
||||
k = torch.cat([k_x, k_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
|
||||
v = torch.cat([v_x, v_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
|
||||
|
||||
q = q.view(-1, num_heads, head_dim)
|
||||
k = k.view(-1, num_heads, head_dim)
|
||||
v = v.view(-1, num_heads, head_dim)
|
||||
return q, k, v
|
||||
|
||||
@torch.autocast("cuda", enabled=False)
|
||||
def flash_attention(self, q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim):
|
||||
out: torch.Tensor = flash_varlen_attn(
|
||||
q, k, v,
|
||||
cu_seqlens_q=cu_seqlens,
|
||||
cu_seqlens_k=cu_seqlens,
|
||||
max_seqlen_q=max_seqlen_in_batch,
|
||||
max_seqlen_k=max_seqlen_in_batch,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=self.softmax_scale,
|
||||
) # (total, local_heads, head_dim)
|
||||
return out.view(total, local_dim)
|
||||
|
||||
def sdpa_attention(self, q, k, v):
|
||||
with sdpa_attn_ctx(training=self.training):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
)
|
||||
return out
|
||||
|
||||
@torch.autocast("cuda", enabled=False)
|
||||
def sage_attention(self, q, k, v):
|
||||
return sage_attn(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
|
||||
|
||||
def run_attention(
|
||||
self,
|
||||
q: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
|
||||
k: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
|
||||
v: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
|
||||
*,
|
||||
B: int,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
max_seqlen_in_batch: Optional[int] = None,
|
||||
):
|
||||
_, cp_size = cp.get_cp_rank_size()
|
||||
assert self.num_heads % cp_size == 0
|
||||
local_heads = self.num_heads // cp_size
|
||||
local_dim = local_heads * self.head_dim
|
||||
|
||||
# Check shapes
|
||||
assert q.ndim == 3 and k.ndim == 3 and v.ndim == 3
|
||||
total = q.size(0)
|
||||
assert k.size(0) == total and v.size(0) == total
|
||||
|
||||
if self.attention_mode == "flash":
|
||||
out = self.flash_attention(
|
||||
q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim) # (total, local_dim)
|
||||
else:
|
||||
assert B == 1, \
|
||||
f"Non-flash attention mode {self.attention_mode} only supports batch size 1, got {B}"
|
||||
|
||||
q = rearrange(q, "(b s) h d -> b h s d", b=B)
|
||||
k = rearrange(k, "(b s) h d -> b h s d", b=B)
|
||||
v = rearrange(v, "(b s) h d -> b h s d", b=B)
|
||||
|
||||
if self.attention_mode == "sdpa":
|
||||
out = self.sdpa_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
|
||||
elif self.attention_mode == "sage":
|
||||
out = self.sage_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
|
||||
else:
|
||||
raise ValueError(f"Unknown attention mode: {self.attention_mode}")
|
||||
|
||||
out = rearrange(out, "b h s d -> (b s) (h d)")
|
||||
|
||||
return out
|
||||
|
||||
def post_attention(
|
||||
self,
|
||||
out: torch.Tensor,
|
||||
B: int,
|
||||
M: int,
|
||||
L: int,
|
||||
dtype: torch.dtype,
|
||||
valid_token_indices: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
out: (total <= B * (N + L), local_dim)
|
||||
valid_token_indices: (total <= B * (N + L),)
|
||||
B: Batch size
|
||||
M: Number of visual tokens per context parallel rank
|
||||
L: Number of text tokens
|
||||
dtype: Data type of the input and output tensors
|
||||
|
||||
Returns:
|
||||
x: (B, N, dim_x) tensor of visual tokens where N = M * cp_size
|
||||
y: (B, L, dim_y) tensor of text token features
|
||||
"""
|
||||
_, cp_size = cp.get_cp_rank_size()
|
||||
local_heads = self.num_heads // cp_size
|
||||
local_dim = local_heads * self.head_dim
|
||||
N = M * cp_size
|
||||
|
||||
# Split sequence into visual and text tokens, adding back padding.
|
||||
if B == 1:
|
||||
out = out.view(B, -1, local_dim)
|
||||
if out.size(1) > N:
|
||||
x, y = torch.tensor_split(out, (N,), dim=1) # (B, N, local_dim), (B, <= L, local_dim)
|
||||
y = F.pad(y, (0, 0, 0, L - y.size(1))) # (B, L, local_dim)
|
||||
else:
|
||||
# Empty prompt.
|
||||
x, y = out, out.new_zeros(B, L, local_dim)
|
||||
else:
|
||||
x, y = pad_and_split_xy(out, valid_token_indices, B, N, L, dtype)
|
||||
assert x.size() == (B, N, local_dim)
|
||||
assert y.size() == (B, L, local_dim)
|
||||
|
||||
# Communicate across context parallel ranks.
|
||||
x = x.view(B, N, local_heads, self.head_dim)
|
||||
x = cp.all_to_all_collect_heads(x) # (B, M, dim_x = num_heads * head_dim)
|
||||
if cp.is_cp_active():
|
||||
y = cp.all_gather(y) # (cp_size * B, L, local_heads * head_dim)
|
||||
y = rearrange(y, "(G B) L D -> B L (G D)", G=cp_size, D=local_dim) # (B, L, dim_x)
|
||||
|
||||
x = self.proj_x(x)
|
||||
y = self.proj_y(y)
|
||||
return x, y
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # (B, M, dim_x)
|
||||
y: torch.Tensor, # (B, L, dim_y)
|
||||
*,
|
||||
scale_x: torch.Tensor, # (B, dim_x), modulation for pre-RMSNorm.
|
||||
scale_y: torch.Tensor, # (B, dim_y), modulation for pre-RMSNorm.
|
||||
packed_indices: Dict[str, torch.Tensor] = None,
|
||||
checkpoint_qkv: bool = False,
|
||||
checkpoint_post_attn: bool = False,
|
||||
**rope_rotation,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward pass of asymmetric multi-modal attention.
|
||||
|
||||
Args:
|
||||
x: (B, M, dim_x) tensor of visual tokens
|
||||
y: (B, L, dim_y) tensor of text token features
|
||||
packed_indices: Dict with keys for Flash Attention
|
||||
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
|
||||
|
||||
Returns:
|
||||
x: (B, M, dim_x) tensor of visual tokens after multi-modal attention
|
||||
y: (B, L, dim_y) tensor of text token features after multi-modal attention
|
||||
"""
|
||||
B, L, _ = y.shape
|
||||
_, M, _ = x.shape
|
||||
|
||||
# Predict a packed QKV tensor from visual and text features.
|
||||
q, k, v = ck(self.prepare_qkv,
|
||||
x=x,
|
||||
y=y,
|
||||
scale_x=scale_x,
|
||||
scale_y=scale_y,
|
||||
rope_cos=rope_rotation.get("rope_cos"),
|
||||
rope_sin=rope_rotation.get("rope_sin"),
|
||||
valid_token_indices=packed_indices["valid_token_indices_kv"],
|
||||
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
|
||||
enabled=checkpoint_qkv,
|
||||
) # (total <= B * (N + L), 3, local_heads, head_dim)
|
||||
|
||||
# Self-attention is expensive, so don't checkpoint it.
|
||||
out = self.run_attention(
|
||||
q, k, v, B=B,
|
||||
cu_seqlens=packed_indices["cu_seqlens_kv"],
|
||||
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
|
||||
)
|
||||
|
||||
x, y = ck(self.post_attention,
|
||||
out,
|
||||
B=B, M=M, L=L,
|
||||
dtype=v.dtype,
|
||||
valid_token_indices=packed_indices["valid_token_indices_kv"],
|
||||
enabled=checkpoint_post_attn,
|
||||
)
|
||||
|
||||
return x, y
|
||||
|
||||
|
||||
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
|
||||
class AsymmetricJointBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size_x: int,
|
||||
hidden_size_y: int,
|
||||
num_heads: int,
|
||||
*,
|
||||
mlp_ratio_x: float = 8.0, # Ratio of hidden size to d_model for MLP for visual tokens.
|
||||
mlp_ratio_y: float = 4.0, # Ratio of hidden size to d_model for MLP for text tokens.
|
||||
update_y: bool = True, # Whether to update text tokens in this block.
|
||||
device: Optional[torch.device] = None,
|
||||
**block_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.update_y = update_y
|
||||
self.hidden_size_x = hidden_size_x
|
||||
self.hidden_size_y = hidden_size_y
|
||||
self.mod_x = nn.Linear(hidden_size_x, 4 * hidden_size_x, device=device)
|
||||
if self.update_y:
|
||||
self.mod_y = nn.Linear(hidden_size_x, 4 * hidden_size_y, device=device)
|
||||
else:
|
||||
self.mod_y = nn.Linear(hidden_size_x, hidden_size_y, device=device)
|
||||
|
||||
# Self-attention:
|
||||
self.attn = AsymmetricAttention(
|
||||
hidden_size_x,
|
||||
hidden_size_y,
|
||||
num_heads=num_heads,
|
||||
update_y=update_y,
|
||||
device=device,
|
||||
**block_kwargs,
|
||||
)
|
||||
|
||||
# MLP.
|
||||
mlp_hidden_dim_x = int(hidden_size_x * mlp_ratio_x)
|
||||
assert mlp_hidden_dim_x == int(1536 * 8)
|
||||
self.mlp_x = FeedForward(
|
||||
in_features=hidden_size_x,
|
||||
hidden_size=mlp_hidden_dim_x,
|
||||
multiple_of=256,
|
||||
ffn_dim_multiplier=None,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# MLP for text not needed in last block.
|
||||
if self.update_y:
|
||||
mlp_hidden_dim_y = int(hidden_size_y * mlp_ratio_y)
|
||||
self.mlp_y = FeedForward(
|
||||
in_features=hidden_size_y,
|
||||
hidden_size=mlp_hidden_dim_y,
|
||||
multiple_of=256,
|
||||
ffn_dim_multiplier=None,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
c: torch.Tensor,
|
||||
y: torch.Tensor,
|
||||
# TODO: These could probably just go into attn_kwargs
|
||||
checkpoint_ff: bool = False,
|
||||
checkpoint_qkv: bool = False,
|
||||
checkpoint_post_attn: bool = False,
|
||||
**attn_kwargs,
|
||||
):
|
||||
"""Forward pass of a block.
|
||||
|
||||
Args:
|
||||
x: (B, N, dim) tensor of visual tokens
|
||||
c: (B, dim) tensor of conditioned features
|
||||
y: (B, L, dim) tensor of text tokens
|
||||
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
|
||||
|
||||
Returns:
|
||||
x: (B, N, dim) tensor of visual tokens after block
|
||||
y: (B, L, dim) tensor of text tokens after block
|
||||
"""
|
||||
N = x.size(1)
|
||||
|
||||
c = F.silu(c)
|
||||
mod_x = self.mod_x(c)
|
||||
scale_msa_x, gate_msa_x, scale_mlp_x, gate_mlp_x = mod_x.chunk(4, dim=1)
|
||||
mod_y = self.mod_y(c)
|
||||
|
||||
if self.update_y:
|
||||
scale_msa_y, gate_msa_y, scale_mlp_y, gate_mlp_y = mod_y.chunk(4, dim=1)
|
||||
else:
|
||||
scale_msa_y = mod_y
|
||||
|
||||
# Self-attention block.
|
||||
x_attn, y_attn = self.attn(
|
||||
x,
|
||||
y,
|
||||
scale_x=scale_msa_x,
|
||||
scale_y=scale_msa_y,
|
||||
checkpoint_qkv=checkpoint_qkv,
|
||||
checkpoint_post_attn=checkpoint_post_attn,
|
||||
**attn_kwargs,
|
||||
)
|
||||
|
||||
assert x_attn.size(1) == N
|
||||
x = residual_tanh_gated_rmsnorm(x, x_attn, gate_msa_x)
|
||||
|
||||
if self.update_y:
|
||||
y = residual_tanh_gated_rmsnorm(y, y_attn, gate_msa_y)
|
||||
|
||||
# MLP block.
|
||||
x = ck(self.ff_block_x, x, scale_mlp_x, gate_mlp_x, enabled=checkpoint_ff)
|
||||
if self.update_y:
|
||||
y = ck(self.ff_block_y, y, scale_mlp_y, gate_mlp_y, enabled=checkpoint_ff) # type: ignore
|
||||
return x, y
|
||||
|
||||
def ff_block_x(self, x, scale_x, gate_x):
|
||||
x_mod = modulated_rmsnorm(x, scale_x)
|
||||
x_res = self.mlp_x(x_mod)
|
||||
x = residual_tanh_gated_rmsnorm(x, x_res, gate_x) # Sandwich norm
|
||||
return x
|
||||
|
||||
def ff_block_y(self, y, scale_y, gate_y):
|
||||
y_mod = modulated_rmsnorm(y, scale_y)
|
||||
y_res = self.mlp_y(y_mod)
|
||||
y = residual_tanh_gated_rmsnorm(y, y_res, gate_y) # Sandwich norm
|
||||
return y
|
||||
|
||||
|
||||
@torch.compile(disable=not COMPILE_FINAL_LAYER)
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of DiT.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
patch_size,
|
||||
out_channels,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, device=device)
|
||||
self.mod = nn.Linear(hidden_size, 2 * hidden_size, device=device)
|
||||
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, device=device)
|
||||
|
||||
def forward(self, x, c):
|
||||
c = F.silu(c)
|
||||
shift, scale = self.mod(c).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class AsymmDiTJoint(nn.Module):
|
||||
"""
|
||||
Diffusion model with a Transformer backbone.
|
||||
|
||||
Ingests text embeddings instead of a label.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
patch_size=2,
|
||||
in_channels=4,
|
||||
hidden_size_x=1152,
|
||||
hidden_size_y=1152,
|
||||
depth=48,
|
||||
num_heads=16,
|
||||
mlp_ratio_x=8.0,
|
||||
mlp_ratio_y=4.0,
|
||||
t5_feat_dim: int = 4096,
|
||||
t5_token_length: int = 256,
|
||||
patch_embed_bias: bool = True,
|
||||
timestep_mlp_bias: bool = True,
|
||||
timestep_scale: Optional[float] = None,
|
||||
use_extended_posenc: bool = False,
|
||||
rope_theta: float = 10000.0,
|
||||
device: Optional[torch.device] = None,
|
||||
**block_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels
|
||||
self.patch_size = patch_size
|
||||
self.num_heads = num_heads
|
||||
self.hidden_size_x = hidden_size_x
|
||||
self.hidden_size_y = hidden_size_y
|
||||
self.head_dim = hidden_size_x // num_heads # Head dimension and count is determined by visual.
|
||||
self.use_extended_posenc = use_extended_posenc
|
||||
self.t5_token_length = t5_token_length
|
||||
self.t5_feat_dim = t5_feat_dim
|
||||
self.rope_theta = rope_theta # Scaling factor for frequency computation for temporal RoPE.
|
||||
|
||||
self.x_embedder = PatchEmbed(
|
||||
patch_size=patch_size,
|
||||
in_chans=in_channels,
|
||||
embed_dim=hidden_size_x,
|
||||
bias=patch_embed_bias,
|
||||
device=device,
|
||||
)
|
||||
# Conditionings
|
||||
# Timestep
|
||||
self.t_embedder = TimestepEmbedder(hidden_size_x, bias=timestep_mlp_bias, timestep_scale=timestep_scale)
|
||||
|
||||
# Caption Pooling (T5)
|
||||
self.t5_y_embedder = AttentionPool(t5_feat_dim, num_heads=8, output_dim=hidden_size_x, device=device)
|
||||
|
||||
# Dense Embedding Projection (T5)
|
||||
self.t5_yproj = nn.Linear(t5_feat_dim, hidden_size_y, bias=True, device=device)
|
||||
|
||||
# Initialize pos_frequencies as an empty parameter.
|
||||
self.pos_frequencies = nn.Parameter(torch.empty(3, self.num_heads, self.head_dim // 2, device=device))
|
||||
|
||||
# for depth 48:
|
||||
# b = 0: AsymmetricJointBlock, update_y=True
|
||||
# b = 1: AsymmetricJointBlock, update_y=True
|
||||
# ...
|
||||
# b = 46: AsymmetricJointBlock, update_y=True
|
||||
# b = 47: AsymmetricJointBlock, update_y=False. No need to update text features.
|
||||
blocks = []
|
||||
for b in range(depth):
|
||||
# Joint multi-modal block
|
||||
update_y = b < depth - 1
|
||||
block = AsymmetricJointBlock(
|
||||
hidden_size_x,
|
||||
hidden_size_y,
|
||||
num_heads,
|
||||
mlp_ratio_x=mlp_ratio_x,
|
||||
mlp_ratio_y=mlp_ratio_y,
|
||||
update_y=update_y,
|
||||
device=device,
|
||||
**block_kwargs,
|
||||
)
|
||||
|
||||
blocks.append(block)
|
||||
self.blocks = nn.ModuleList(blocks)
|
||||
|
||||
self.final_layer = FinalLayer(hidden_size_x, patch_size, self.out_channels, device=device)
|
||||
|
||||
def embed_x(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: (B, C=12, T, H, W) tensor of visual tokens
|
||||
|
||||
Returns:
|
||||
x: (B, C=3072, N) tensor of visual tokens with positional embedding.
|
||||
"""
|
||||
return self.x_embedder(x) # Convert BcTHW to BCN
|
||||
|
||||
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
|
||||
def prepare(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
t5_feat: torch.Tensor,
|
||||
t5_mask: torch.Tensor,
|
||||
):
|
||||
"""Prepare input and conditioning embeddings."""
|
||||
|
||||
# Visual patch embeddings with positional encoding.
|
||||
T, H, W = x.shape[-3:]
|
||||
pH, pW = H // self.patch_size, W // self.patch_size
|
||||
x = self.embed_x(x) # (B, N, D), where N = T * H * W / patch_size ** 2
|
||||
assert x.ndim == 3
|
||||
B = x.size(0)
|
||||
|
||||
# Construct position array of size [N, 3].
|
||||
# pos[:, 0] is the frame index for each location,
|
||||
# pos[:, 1] is the row index for each location, and
|
||||
# pos[:, 2] is the column index for each location.
|
||||
N = T * pH * pW
|
||||
assert x.size(1) == N
|
||||
pos = create_position_matrix(T, pH=pH, pW=pW, device=x.device, dtype=torch.float32) # (N, 3)
|
||||
rope_cos, rope_sin = compute_mixed_rotation(
|
||||
freqs=self.pos_frequencies, pos=pos
|
||||
) # Each are (N, num_heads, dim // 2)
|
||||
|
||||
# Global vector embedding for conditionings.
|
||||
c_t = self.t_embedder(1 - sigma) # (B, D)
|
||||
|
||||
# Pool T5 tokens using attention pooler
|
||||
# Note y_feat[1] contains T5 token features.
|
||||
assert (
|
||||
t5_feat.size(1) == self.t5_token_length
|
||||
), f"Expected L={self.t5_token_length}, got {t5_feat.shape} for y_feat."
|
||||
t5_y_pool = self.t5_y_embedder(t5_feat, t5_mask) # (B, D)
|
||||
assert t5_y_pool.size(0) == B, f"Expected B={B}, got {t5_y_pool.shape} for t5_y_pool."
|
||||
|
||||
c = c_t + t5_y_pool
|
||||
|
||||
y_feat = self.t5_yproj(t5_feat) # (B, L, t5_feat_dim) --> (B, L, D)
|
||||
|
||||
return x, c, y_feat, rope_cos, rope_sin
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor, # [1, 12, 7, 60, 106]
|
||||
sigma: torch.Tensor, # [1]
|
||||
y_feat: List[torch.Tensor], # [0][1, 256, 4096]
|
||||
y_mask: List[torch.Tensor], # [0][1, 256]
|
||||
packed_indices: Dict[str, torch.Tensor] = None,
|
||||
rope_cos: torch.Tensor = None,
|
||||
rope_sin: torch.Tensor = None,
|
||||
num_ff_checkpoint: int = 0, # 48
|
||||
num_qkv_checkpoint: int = 0, # 48
|
||||
num_post_attn_checkpoint: int = 0, # 0
|
||||
):
|
||||
"""Forward pass of DiT.
|
||||
|
||||
Args:
|
||||
x: (B, C, T, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||
sigma: (B,) tensor of noise standard deviations
|
||||
y_feat: List((B, L, y_feat_dim) tensor of caption token features. For SDXL text encoders: L=77, y_feat_dim=2048)
|
||||
y_mask: List((B, L) boolean tensor indicating which tokens are not padding)
|
||||
packed_indices: Dict with keys for Flash Attention. Result of compute_packed_indices.
|
||||
"""
|
||||
_, _, T, H, W = x.shape
|
||||
|
||||
if self.pos_frequencies.dtype != torch.float32:
|
||||
warnings.warn(f"pos_frequencies dtype {self.pos_frequencies.dtype} != torch.float32")
|
||||
|
||||
# Use EFFICIENT_ATTENTION backend for T5 pooling, since we have a mask.
|
||||
# Have to call sdpa_kernel outside of a torch.compile region.
|
||||
with sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION):
|
||||
x, c, y_feat, rope_cos, rope_sin = self.prepare(x, sigma, y_feat[0], y_mask[0]) # [1, 11130, 3072], [1, 3072], [1, 256, 1536], [11130, 24, 64]
|
||||
del y_mask
|
||||
|
||||
cp_rank, cp_size = cp.get_cp_rank_size()
|
||||
N = x.size(1)
|
||||
M = N // cp_size
|
||||
assert N % cp_size == 0, f"Visual sequence length ({x.shape[1]}) must be divisible by cp_size ({cp_size})."
|
||||
|
||||
if cp_size > 1:
|
||||
x = x.narrow(1, cp_rank * M, M)
|
||||
|
||||
assert self.num_heads % cp_size == 0
|
||||
local_heads = self.num_heads // cp_size
|
||||
rope_cos = rope_cos.narrow(1, cp_rank * local_heads, local_heads)
|
||||
rope_sin = rope_sin.narrow(1, cp_rank * local_heads, local_heads)
|
||||
|
||||
for i, block in enumerate(self.blocks):
|
||||
x, y_feat = block( # [1, 11130, 3072], [1, 256, 1536]
|
||||
x,
|
||||
c,
|
||||
y_feat,
|
||||
rope_cos=rope_cos,
|
||||
rope_sin=rope_sin,
|
||||
packed_indices=packed_indices,
|
||||
checkpoint_ff=i < num_ff_checkpoint,
|
||||
checkpoint_qkv=i < num_qkv_checkpoint,
|
||||
checkpoint_post_attn=i < num_post_attn_checkpoint,
|
||||
) # (B, M, D), (B, L, D)
|
||||
del y_feat # Final layers don't use dense text features.
|
||||
|
||||
x = self.final_layer(x, c) # (B, M, patch_size ** 2 * out_channels) [1, 11130, 48]
|
||||
|
||||
patch = x.size(2)
|
||||
x = cp.all_gather(x)
|
||||
x = rearrange(x, "(G B) M P -> B (G M) P", G=cp_size, P=patch) # [1, 11130, 48]
|
||||
x = rearrange( # [1, 12, 7, 60, 106]
|
||||
x,
|
||||
"B (T hp wp) (p1 p2 c) -> B c T (hp p1) (wp p2)",
|
||||
T=T,
|
||||
hp=H // self.patch_size,
|
||||
wp=W // self.patch_size,
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
c=self.out_channels,
|
||||
)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,158 @@
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from einops import rearrange
|
||||
|
||||
_CONTEXT_PARALLEL_GROUP = None
|
||||
_CONTEXT_PARALLEL_RANK = None
|
||||
_CONTEXT_PARALLEL_GROUP_SIZE = None
|
||||
_CONTEXT_PARALLEL_GROUP_RANKS = None
|
||||
|
||||
|
||||
def get_cp_rank_size() -> Tuple[int, int]:
|
||||
if _CONTEXT_PARALLEL_GROUP:
|
||||
assert isinstance(_CONTEXT_PARALLEL_RANK, int) and isinstance(_CONTEXT_PARALLEL_GROUP_SIZE, int)
|
||||
return _CONTEXT_PARALLEL_RANK, _CONTEXT_PARALLEL_GROUP_SIZE
|
||||
else:
|
||||
return 0, 1
|
||||
|
||||
|
||||
def local_shard(x: torch.Tensor, dim: int = 2) -> torch.Tensor:
|
||||
if not _CONTEXT_PARALLEL_GROUP:
|
||||
return x
|
||||
|
||||
cp_rank, cp_size = get_cp_rank_size()
|
||||
return x.tensor_split(cp_size, dim=dim)[cp_rank]
|
||||
|
||||
|
||||
def set_cp_group(cp_group, ranks, global_rank):
|
||||
global _CONTEXT_PARALLEL_GROUP, _CONTEXT_PARALLEL_RANK, _CONTEXT_PARALLEL_GROUP_SIZE, _CONTEXT_PARALLEL_GROUP_RANKS
|
||||
if _CONTEXT_PARALLEL_GROUP is not None:
|
||||
raise RuntimeError("CP group already initialized.")
|
||||
_CONTEXT_PARALLEL_GROUP = cp_group
|
||||
_CONTEXT_PARALLEL_RANK = dist.get_rank(cp_group)
|
||||
_CONTEXT_PARALLEL_GROUP_SIZE = dist.get_world_size(cp_group)
|
||||
_CONTEXT_PARALLEL_GROUP_RANKS = ranks
|
||||
|
||||
assert _CONTEXT_PARALLEL_RANK == ranks.index(
|
||||
global_rank
|
||||
), f"Rank mismatch: {global_rank} in {ranks} does not have position {_CONTEXT_PARALLEL_RANK} "
|
||||
assert _CONTEXT_PARALLEL_GROUP_SIZE == len(
|
||||
ranks
|
||||
), f"Group size mismatch: {_CONTEXT_PARALLEL_GROUP_SIZE} != len({ranks})"
|
||||
|
||||
|
||||
def get_cp_group():
|
||||
if _CONTEXT_PARALLEL_GROUP is None:
|
||||
raise RuntimeError("CP group not initialized")
|
||||
return _CONTEXT_PARALLEL_GROUP
|
||||
|
||||
|
||||
def is_cp_active():
|
||||
return _CONTEXT_PARALLEL_GROUP is not None
|
||||
|
||||
|
||||
class AllGatherIntoTensorFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x: torch.Tensor, reduce_dtype, group: dist.ProcessGroup):
|
||||
ctx.reduce_dtype = reduce_dtype
|
||||
ctx.group = group
|
||||
ctx.batch_size = x.size(0)
|
||||
group_size = dist.get_world_size(group)
|
||||
|
||||
x = x.contiguous()
|
||||
output = torch.empty(group_size * x.size(0), *x.shape[1:], dtype=x.dtype, device=x.device)
|
||||
dist.all_gather_into_tensor(output, x, group=group)
|
||||
return output
|
||||
|
||||
|
||||
def all_gather(tensor: torch.Tensor) -> torch.Tensor:
|
||||
if not _CONTEXT_PARALLEL_GROUP:
|
||||
return tensor
|
||||
|
||||
return AllGatherIntoTensorFunction.apply(tensor, torch.float32, _CONTEXT_PARALLEL_GROUP)
|
||||
|
||||
|
||||
@torch.compiler.disable()
|
||||
def _all_to_all_single(output, input, group):
|
||||
# Disable compilation since torch compile changes contiguity.
|
||||
assert input.is_contiguous(), "Input tensor must be contiguous."
|
||||
assert output.is_contiguous(), "Output tensor must be contiguous."
|
||||
return dist.all_to_all_single(output, input, group=group)
|
||||
|
||||
|
||||
class CollectTokens(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, qkv: torch.Tensor, group: dist.ProcessGroup, num_heads: int):
|
||||
"""Redistribute heads and receive tokens.
|
||||
|
||||
Args:
|
||||
qkv: query, key or value. Shape: [B, M, 3 * num_heads * head_dim]
|
||||
|
||||
Returns:
|
||||
qkv: shape: [3, B, N, local_heads, head_dim]
|
||||
|
||||
where M is the number of local tokens,
|
||||
N = cp_size * M is the number of global tokens,
|
||||
local_heads = num_heads // cp_size is the number of local heads.
|
||||
"""
|
||||
ctx.group = group
|
||||
ctx.num_heads = num_heads
|
||||
cp_size = dist.get_world_size(group)
|
||||
assert num_heads % cp_size == 0
|
||||
ctx.local_heads = num_heads // cp_size
|
||||
|
||||
qkv = rearrange(
|
||||
qkv,
|
||||
"B M (qkv G h d) -> G M h B (qkv d)",
|
||||
qkv=3,
|
||||
G=cp_size,
|
||||
h=ctx.local_heads,
|
||||
).contiguous()
|
||||
|
||||
output_chunks = torch.empty_like(qkv)
|
||||
_all_to_all_single(output_chunks, qkv, group=group)
|
||||
|
||||
return rearrange(output_chunks, "G M h B (qkv d) -> qkv B (G M) h d", qkv=3)
|
||||
|
||||
|
||||
def all_to_all_collect_tokens(x: torch.Tensor, num_heads: int) -> torch.Tensor:
|
||||
if not _CONTEXT_PARALLEL_GROUP:
|
||||
# Move QKV dimension to the front.
|
||||
# B M (3 H d) -> 3 B M H d
|
||||
B, M, _ = x.size()
|
||||
x = x.view(B, M, 3, num_heads, -1)
|
||||
return x.permute(2, 0, 1, 3, 4)
|
||||
|
||||
return CollectTokens.apply(x, _CONTEXT_PARALLEL_GROUP, num_heads)
|
||||
|
||||
|
||||
class CollectHeads(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x: torch.Tensor, group: dist.ProcessGroup):
|
||||
"""Redistribute tokens and receive heads.
|
||||
|
||||
Args:
|
||||
x: Output of attention. Shape: [B, N, local_heads, head_dim]
|
||||
|
||||
Returns:
|
||||
Shape: [B, M, num_heads * head_dim]
|
||||
"""
|
||||
ctx.group = group
|
||||
ctx.local_heads = x.size(2)
|
||||
ctx.head_dim = x.size(3)
|
||||
group_size = dist.get_world_size(group)
|
||||
x = rearrange(x, "B (G M) h D -> G h M B D", G=group_size).contiguous()
|
||||
output = torch.empty_like(x)
|
||||
_all_to_all_single(output, x, group=group)
|
||||
del x
|
||||
return rearrange(output, "G h M B D -> B M (G h D)")
|
||||
|
||||
|
||||
def all_to_all_collect_heads(x: torch.Tensor) -> torch.Tensor:
|
||||
if not _CONTEXT_PARALLEL_GROUP:
|
||||
# Merge heads.
|
||||
return x.view(x.size(0), x.size(1), x.size(2) * x.size(3))
|
||||
|
||||
return CollectHeads.apply(x, _CONTEXT_PARALLEL_GROUP)
|
||||
@@ -0,0 +1,179 @@
|
||||
import collections.abc
|
||||
import math
|
||||
from itertools import repeat
|
||||
from typing import Callable, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
# From PyTorch internals
|
||||
def _ntuple(n):
|
||||
def parse(x):
|
||||
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
|
||||
return tuple(x)
|
||||
return tuple(repeat(x, n))
|
||||
|
||||
return parse
|
||||
|
||||
|
||||
to_2tuple = _ntuple(2)
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
frequency_embedding_size: int = 256,
|
||||
*,
|
||||
bias: bool = True,
|
||||
timestep_scale: Optional[float] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=bias, device=device),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=bias, device=device),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.timestep_scale = timestep_scale
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
half = dim // 2
|
||||
freqs = torch.arange(start=0, end=half, dtype=torch.float32, device=t.device)
|
||||
freqs.mul_(-math.log(max_period) / half).exp_()
|
||||
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
|
||||
|
||||
def forward(self, t):
|
||||
if self.timestep_scale is not None:
|
||||
t = t * self.timestep_scale
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
class PooledCaptionEmbedder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
caption_feature_dim: int,
|
||||
hidden_size: int,
|
||||
*,
|
||||
bias: bool = True,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.caption_feature_dim = caption_feature_dim
|
||||
self.hidden_size = hidden_size
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(caption_feature_dim, hidden_size, bias=bias, device=device),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=bias, device=device),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.mlp(x)
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
hidden_size: int,
|
||||
multiple_of: int,
|
||||
ffn_dim_multiplier: Optional[float],
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
super().__init__()
|
||||
# keep parameter count and computation constant compared to standard FFN
|
||||
hidden_size = int(2 * hidden_size / 3)
|
||||
# custom dim factor multiplier
|
||||
if ffn_dim_multiplier is not None:
|
||||
hidden_size = int(ffn_dim_multiplier * hidden_size)
|
||||
hidden_size = multiple_of * ((hidden_size + multiple_of - 1) // multiple_of)
|
||||
|
||||
self.hidden_dim = hidden_size
|
||||
self.w1 = nn.Linear(in_features, 2 * hidden_size, bias=False, device=device)
|
||||
self.w2 = nn.Linear(hidden_size, in_features, bias=False, device=device)
|
||||
|
||||
def forward(self, x):
|
||||
# assert self.w1.weight.dtype == torch.bfloat16, f"FFN weight dtype {self.w1.weight.dtype} != bfloat16"
|
||||
x, gate = self.w1(x).chunk(2, dim=-1)
|
||||
x = self.w2(F.silu(x) * gate)
|
||||
return x
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 16,
|
||||
in_chans: int = 3,
|
||||
embed_dim: int = 768,
|
||||
norm_layer: Optional[Callable] = None,
|
||||
flatten: bool = True,
|
||||
bias: bool = True,
|
||||
dynamic_img_pad: bool = False,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = to_2tuple(patch_size)
|
||||
self.flatten = flatten
|
||||
self.dynamic_img_pad = dynamic_img_pad
|
||||
|
||||
self.proj = nn.Conv2d(
|
||||
in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias,
|
||||
device=device,
|
||||
)
|
||||
assert norm_layer is None
|
||||
self.norm = norm_layer(embed_dim, device=device) if norm_layer else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
B, _C, T, H, W = x.shape
|
||||
if not self.dynamic_img_pad:
|
||||
assert (
|
||||
H % self.patch_size[0] == 0
|
||||
), f"Input height ({H}) should be divisible by patch size ({self.patch_size[0]})."
|
||||
assert (
|
||||
W % self.patch_size[1] == 0
|
||||
), f"Input width ({W}) should be divisible by patch size ({self.patch_size[1]})."
|
||||
else:
|
||||
pad_h = (self.patch_size[0] - H % self.patch_size[0]) % self.patch_size[0]
|
||||
pad_w = (self.patch_size[1] - W % self.patch_size[1]) % self.patch_size[1]
|
||||
x = F.pad(x, (0, pad_w, 0, pad_h))
|
||||
|
||||
x = rearrange(x, "B C T H W -> (B T) C H W", B=B, T=T)
|
||||
x = self.proj(x)
|
||||
|
||||
# Flatten temporal and spatial dimensions.
|
||||
if not self.flatten:
|
||||
raise NotImplementedError("Must flatten output.")
|
||||
x = rearrange(x, "(B T) C H W -> B (T H W) C", B=B, T=T)
|
||||
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
def __init__(self, hidden_size, eps=1e-5, device=None):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = torch.nn.Parameter(torch.empty(hidden_size, device=device))
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def forward(self, x):
|
||||
# assert self.weight.dtype == torch.float32, f"RMSNorm weight dtype {self.weight.dtype} != float32"
|
||||
|
||||
x_fp32 = x.float()
|
||||
x_normed = x_fp32 * torch.rsqrt(x_fp32.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
return (x_normed * self.weight).type_as(x)
|
||||
@@ -0,0 +1,112 @@
|
||||
#! /usr/bin/env python3
|
||||
import math
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class LoRALayer:
|
||||
def __init__(
|
||||
self,
|
||||
r: int,
|
||||
lora_alpha: int,
|
||||
lora_dropout: float,
|
||||
merge_weights: bool,
|
||||
):
|
||||
self.r = r
|
||||
self.lora_alpha = lora_alpha
|
||||
if lora_dropout > 0.0:
|
||||
self.lora_dropout = nn.Dropout(p=lora_dropout)
|
||||
else:
|
||||
self.lora_dropout = lambda x: x
|
||||
self.merged = False
|
||||
self.merge_weights = merge_weights
|
||||
|
||||
|
||||
def mark_only_lora_as_trainable(model: nn.Module, bias: str = "none") -> None:
|
||||
assert bias == "none", f"Only bias='none' is supported"
|
||||
for n, p in model.named_parameters():
|
||||
if "lora_" not in n:
|
||||
p.requires_grad = False
|
||||
|
||||
|
||||
def lora_state_dict(model: nn.Module, bias: str = "none") -> Dict[str, torch.Tensor]:
|
||||
assert bias == "none", f"Only bias='none' is supported"
|
||||
my_state_dict = model.state_dict()
|
||||
return {k: my_state_dict[k] for k in my_state_dict if "lora_" in k}
|
||||
|
||||
|
||||
class LoraLinear(nn.Linear, LoRALayer):
|
||||
# LoRA implemented in a dense layer
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
r: int = 0,
|
||||
lora_alpha: int = 1,
|
||||
lora_dropout: float = 0.0,
|
||||
fan_in_fan_out: bool = False, # Set this to True if the layer to replace stores weight like (fan_in, fan_out)
|
||||
merge_weights: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
nn.Linear.__init__(self, in_features, out_features, **kwargs)
|
||||
LoRALayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=merge_weights)
|
||||
|
||||
self.fan_in_fan_out = fan_in_fan_out
|
||||
# Actual trainable parameters
|
||||
if r > 0:
|
||||
self.lora_A = nn.Parameter(self.weight.new_zeros((r, in_features)).to(torch.float32))
|
||||
self.lora_B = nn.Parameter(self.weight.new_zeros((out_features, r)).to(torch.float32))
|
||||
self.scaling = self.lora_alpha / self.r
|
||||
|
||||
# Freezing the pre-trained weight matrix
|
||||
self.weight.requires_grad = False
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
if fan_in_fan_out:
|
||||
self.weight.data = self.weight.data.transpose(0, 1)
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.Linear.reset_parameters(self)
|
||||
if hasattr(self, "lora_A"):
|
||||
# initialize B the same way as the default for nn.Linear and A to zero
|
||||
# this is different than what is described in the paper but should not affect performance
|
||||
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
|
||||
nn.init.zeros_(self.lora_B)
|
||||
|
||||
def train(self, mode: bool = True):
|
||||
def T(w):
|
||||
return w.transpose(0, 1) if self.fan_in_fan_out else w
|
||||
|
||||
nn.Linear.train(self, mode)
|
||||
if mode:
|
||||
if self.merge_weights and self.merged:
|
||||
# Make sure that the weights are not merged
|
||||
if self.r > 0:
|
||||
self.weight.data -= T(self.lora_B @ self.lora_A) * self.scaling
|
||||
self.merged = False
|
||||
else:
|
||||
if self.merge_weights and not self.merged:
|
||||
# Merge the weights and mark it
|
||||
if self.r > 0:
|
||||
self.weight.data += T(self.lora_B @ self.lora_A) * self.scaling
|
||||
self.merged = True
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
def T(w):
|
||||
return w.transpose(0, 1) if self.fan_in_fan_out else w
|
||||
|
||||
if self.r > 0 and not self.merged:
|
||||
result = F.linear(x, T(self.weight), bias=self.bias)
|
||||
|
||||
x = self.lora_dropout(x)
|
||||
x = x @ self.lora_A.transpose(0, 1)
|
||||
x = x @ self.lora_B.transpose(0, 1)
|
||||
x = x * self.scaling
|
||||
|
||||
return result + x
|
||||
else:
|
||||
return F.linear(x, T(self.weight), bias=self.bias)
|
||||
@@ -0,0 +1,15 @@
|
||||
import torch
|
||||
|
||||
|
||||
def modulated_rmsnorm(x, scale, eps=1e-6):
|
||||
dtype = x.dtype
|
||||
x = x.float()
|
||||
|
||||
# Compute RMS
|
||||
mean_square = x.pow(2).mean(-1, keepdim=True)
|
||||
inv_rms = torch.rsqrt(mean_square + eps)
|
||||
|
||||
# Normalize and modulate
|
||||
x_normed = x * inv_rms
|
||||
x_modulated = x_normed * (1 + scale.unsqueeze(1).float())
|
||||
return x_modulated.to(dtype)
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
import torch
|
||||
|
||||
|
||||
def residual_tanh_gated_rmsnorm(x, x_res, gate, eps=1e-6):
|
||||
# Convert to fp32 for precision
|
||||
x_res = x_res.float()
|
||||
|
||||
# Compute RMS
|
||||
mean_square = x_res.pow(2).mean(-1, keepdim=True)
|
||||
scale = torch.rsqrt(mean_square + eps)
|
||||
|
||||
# Apply tanh to gate
|
||||
tanh_gate = torch.tanh(gate).unsqueeze(1)
|
||||
|
||||
# Normalize and apply gated scaling
|
||||
x_normed = x_res * scale * tanh_gate
|
||||
|
||||
# Apply residual connection
|
||||
output = x + x_normed.type_as(x)
|
||||
return output
|
||||
@@ -0,0 +1,88 @@
|
||||
import functools
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def centers(start: float, stop, num, dtype=None, device=None):
|
||||
"""linspace through bin centers.
|
||||
|
||||
Args:
|
||||
start (float): Start of the range.
|
||||
stop (float): End of the range.
|
||||
num (int): Number of points.
|
||||
dtype (torch.dtype): Data type of the points.
|
||||
device (torch.device): Device of the points.
|
||||
|
||||
Returns:
|
||||
centers (Tensor): Centers of the bins. Shape: (num,).
|
||||
"""
|
||||
edges = torch.linspace(start, stop, num + 1, dtype=dtype, device=device)
|
||||
return (edges[:-1] + edges[1:]) / 2
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def create_position_matrix(
|
||||
T: int,
|
||||
pH: int,
|
||||
pW: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
*,
|
||||
target_area: float = 36864,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
T: int - Temporal dimension
|
||||
pH: int - Height dimension after patchify
|
||||
pW: int - Width dimension after patchify
|
||||
|
||||
Returns:
|
||||
pos: [T * pH * pW, 3] - position matrix
|
||||
"""
|
||||
with torch.no_grad():
|
||||
# Create 1D tensors for each dimension
|
||||
t = torch.arange(T, dtype=dtype)
|
||||
|
||||
# Positionally interpolate to area 36864.
|
||||
# (3072x3072 frame with 16x16 patches = 192x192 latents).
|
||||
# This automatically scales rope positions when the resolution changes.
|
||||
# We use a large target area so the model is more sensitive
|
||||
# to changes in the learned pos_frequencies matrix.
|
||||
scale = math.sqrt(target_area / (pW * pH))
|
||||
w = centers(-pW * scale / 2, pW * scale / 2, pW)
|
||||
h = centers(-pH * scale / 2, pH * scale / 2, pH)
|
||||
|
||||
# Use meshgrid to create 3D grids
|
||||
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
|
||||
|
||||
# Stack and reshape the grids.
|
||||
pos = torch.stack([grid_t, grid_h, grid_w], dim=-1) # [T, pH, pW, 3]
|
||||
pos = pos.view(-1, 3) # [T * pH * pW, 3]
|
||||
pos = pos.to(dtype=dtype, device=device)
|
||||
|
||||
return pos
|
||||
|
||||
|
||||
def compute_mixed_rotation(
|
||||
freqs: torch.Tensor,
|
||||
pos: torch.Tensor,
|
||||
):
|
||||
"""
|
||||
Project each 3-dim position into per-head, per-head-dim 1D frequencies.
|
||||
|
||||
Args:
|
||||
freqs: [3, num_heads, num_freqs] - learned rotation frequency (for t, row, col) for each head position
|
||||
pos: [N, 3] - position of each token
|
||||
num_heads: int
|
||||
|
||||
Returns:
|
||||
freqs_cos: [N, num_heads, num_freqs] - cosine components
|
||||
freqs_sin: [N, num_heads, num_freqs] - sine components
|
||||
"""
|
||||
with torch.autocast("cuda", enabled=False):
|
||||
assert freqs.ndim == 3
|
||||
freqs_sum = torch.einsum("Nd,dhf->Nhf", pos.to(freqs), freqs)
|
||||
freqs_cos = torch.cos(freqs_sum)
|
||||
freqs_sin = torch.sin(freqs_sum)
|
||||
return freqs_cos, freqs_sin
|
||||
@@ -0,0 +1,34 @@
|
||||
# Based on Llama3 Implementation.
|
||||
import torch
|
||||
|
||||
|
||||
def apply_rotary_emb_qk_real(
|
||||
xqk: torch.Tensor,
|
||||
freqs_cos: torch.Tensor,
|
||||
freqs_sin: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply rotary embeddings to input tensors using the given frequency tensor without complex numbers.
|
||||
|
||||
Args:
|
||||
xqk (torch.Tensor): Query and/or Key tensors to apply rotary embeddings. Shape: (B, S, *, num_heads, D)
|
||||
Can be either just query or just key, or both stacked along some batch or * dim.
|
||||
freqs_cos (torch.Tensor): Precomputed cosine frequency tensor.
|
||||
freqs_sin (torch.Tensor): Precomputed sine frequency tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The input tensor with rotary embeddings applied.
|
||||
"""
|
||||
assert xqk.dtype == torch.bfloat16
|
||||
# Split the last dimension into even and odd parts
|
||||
xqk_even = xqk[..., 0::2]
|
||||
xqk_odd = xqk[..., 1::2]
|
||||
|
||||
# Apply rotation
|
||||
cos_part = (xqk_even * freqs_cos - xqk_odd * freqs_sin).type_as(xqk)
|
||||
sin_part = (xqk_even * freqs_sin + xqk_odd * freqs_cos).type_as(xqk)
|
||||
|
||||
# Interleave the results back into the original shape
|
||||
out = torch.stack([cos_part, sin_part], dim=-1).flatten(-2)
|
||||
assert out.dtype == torch.bfloat16
|
||||
return out
|
||||
@@ -0,0 +1,109 @@
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def modulate(x, shift, scale):
|
||||
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
|
||||
|
||||
def pool_tokens(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False) -> torch.Tensor:
|
||||
"""
|
||||
Pool tokens in x using mask.
|
||||
|
||||
NOTE: We assume x does not require gradients.
|
||||
|
||||
Args:
|
||||
x: (B, L, D) tensor of tokens.
|
||||
mask: (B, L) boolean tensor indicating which tokens are not padding.
|
||||
|
||||
Returns:
|
||||
pooled: (B, D) tensor of pooled tokens.
|
||||
"""
|
||||
assert x.size(1) == mask.size(1) # Expected mask to have same length as tokens.
|
||||
assert x.size(0) == mask.size(0) # Expected mask to have same batch size as tokens.
|
||||
mask = mask[:, :, None].to(dtype=x.dtype)
|
||||
mask = mask / mask.sum(dim=1, keepdim=True).clamp(min=1)
|
||||
pooled = (x * mask).sum(dim=1, keepdim=keepdim)
|
||||
return pooled
|
||||
|
||||
|
||||
class AttentionPool(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
num_heads: int,
|
||||
output_dim: int = None,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
spatial_dim (int): Number of tokens in sequence length.
|
||||
embed_dim (int): Dimensionality of input tokens.
|
||||
num_heads (int): Number of attention heads.
|
||||
output_dim (int): Dimensionality of output tokens. Defaults to embed_dim.
|
||||
"""
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.to_kv = nn.Linear(embed_dim, 2 * embed_dim, device=device)
|
||||
self.to_q = nn.Linear(embed_dim, embed_dim, device=device)
|
||||
self.to_out = nn.Linear(embed_dim, output_dim or embed_dim, device=device)
|
||||
|
||||
def forward(self, x, mask):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): (B, L, D) tensor of input tokens.
|
||||
mask (torch.Tensor): (B, L) boolean tensor indicating which tokens are not padding.
|
||||
|
||||
NOTE: We assume x does not require gradients.
|
||||
|
||||
Returns:
|
||||
x (torch.Tensor): (B, D) tensor of pooled tokens.
|
||||
"""
|
||||
D = x.size(2)
|
||||
|
||||
# Construct attention mask, shape: (B, 1, num_queries=1, num_keys=1+L).
|
||||
attn_mask = mask[:, None, None, :].bool() # (B, 1, 1, L).
|
||||
attn_mask = F.pad(attn_mask, (1, 0), value=True) # (B, 1, 1, 1+L).
|
||||
|
||||
# Average non-padding token features. These will be used as the query.
|
||||
x_pool = pool_tokens(x, mask, keepdim=True) # (B, 1, D)
|
||||
|
||||
# Concat pooled features to input sequence.
|
||||
x = torch.cat([x_pool, x], dim=1) # (B, L+1, D)
|
||||
|
||||
# Compute queries, keys, values. Only the mean token is used to create a query.
|
||||
kv = self.to_kv(x) # (B, L+1, 2 * D)
|
||||
q = self.to_q(x[:, 0]) # (B, D)
|
||||
|
||||
# Extract heads.
|
||||
head_dim = D // self.num_heads
|
||||
kv = kv.unflatten(2, (2, self.num_heads, head_dim)) # (B, 1+L, 2, H, head_dim)
|
||||
kv = kv.transpose(1, 3) # (B, H, 2, 1+L, head_dim)
|
||||
k, v = kv.unbind(2) # (B, H, 1+L, head_dim)
|
||||
q = q.unflatten(1, (self.num_heads, head_dim)) # (B, H, head_dim)
|
||||
q = q.unsqueeze(2) # (B, H, 1, head_dim)
|
||||
|
||||
# Compute attention.
|
||||
x = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0) # (B, H, 1, head_dim)
|
||||
|
||||
# Concatenate heads and run output.
|
||||
x = x.squeeze(2).flatten(1, 2) # (B, D = H * head_dim)
|
||||
x = self.to_out(x)
|
||||
return x
|
||||
|
||||
|
||||
def pad_and_split_xy(xy, indices, B, N, L, dtype) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
D = xy.size(1)
|
||||
|
||||
# Pad sequences to (B, N + L, dim).
|
||||
assert indices.ndim == 1
|
||||
indices = indices.unsqueeze(1).expand(-1, D) # (total,) -> (total, num_heads * head_dim)
|
||||
output = torch.zeros(B * (N + L), D, device=xy.device, dtype=dtype)
|
||||
output = torch.scatter(output, 0, indices, xy)
|
||||
xy = output.view(B, N + L, D)
|
||||
|
||||
# Split visual and text tokens along the sequence length.
|
||||
return torch.tensor_split(xy, (N,), dim=1)
|
||||
@@ -0,0 +1,682 @@
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from functools import partial
|
||||
from typing import Any, Dict, List, Literal, Optional, Union, cast
|
||||
|
||||
import numpy as np
|
||||
import ray
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import repeat
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file
|
||||
from torch import nn
|
||||
from torch.distributed.fsdp import (
|
||||
BackwardPrefetch,
|
||||
MixedPrecision,
|
||||
ShardingStrategy,
|
||||
)
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp.wrap import (
|
||||
lambda_auto_wrap_policy,
|
||||
transformer_auto_wrap_policy,
|
||||
)
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
from transformers.models.t5.modeling_t5 import T5Block
|
||||
|
||||
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
|
||||
from genmo.lib.progress import get_new_progress_bar, progress_bar
|
||||
from genmo.lib.utils import Timer
|
||||
from genmo.mochi_preview.vae.models import (
|
||||
Decoder,
|
||||
Encoder,
|
||||
decode_latents,
|
||||
decode_latents_tiled_full,
|
||||
decode_latents_tiled_spatial,
|
||||
)
|
||||
from genmo.mochi_preview.vae.vae_stats import dit_latents_to_vae_latents
|
||||
|
||||
|
||||
def load_to_cpu(p, weights_only=True):
|
||||
if p.endswith(".safetensors"):
|
||||
return load_file(p)
|
||||
else:
|
||||
assert p.endswith(".pt")
|
||||
return torch.load(p, map_location="cpu", weights_only=weights_only)
|
||||
|
||||
|
||||
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)]
|
||||
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)
|
||||
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)
|
||||
]
|
||||
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule + [1.0]
|
||||
sigma_schedule = [1.0 - x for x in sigma_schedule]
|
||||
return sigma_schedule
|
||||
|
||||
|
||||
T5_MODEL = "google/t5-v1_1-xxl"
|
||||
MAX_T5_TOKEN_LENGTH = 256
|
||||
|
||||
|
||||
def setup_fsdp_sync(model, device_id, *, param_dtype, auto_wrap_policy) -> FSDP:
|
||||
model = FSDP(
|
||||
model,
|
||||
sharding_strategy=ShardingStrategy.FULL_SHARD,
|
||||
mixed_precision=MixedPrecision(
|
||||
param_dtype=param_dtype,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
),
|
||||
auto_wrap_policy=auto_wrap_policy,
|
||||
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
|
||||
limit_all_gathers=True,
|
||||
device_id=device_id,
|
||||
sync_module_states=True,
|
||||
use_orig_params=True,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
return model
|
||||
|
||||
|
||||
class ModelFactory(ABC):
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
@abstractmethod
|
||||
def get_model(self, *, local_rank: int, device_id: Union[int, Literal["cpu"]], world_size: int) -> Any:
|
||||
assert isinstance(device_id, int) or device_id == "cpu", "device_id must be an integer or 'cpu'"
|
||||
# FSDP does not work when the model is on the CPU
|
||||
if device_id == "cpu":
|
||||
assert world_size == 1, "CPU offload only supports single-GPU inference"
|
||||
|
||||
|
||||
class T5ModelFactory(ModelFactory):
|
||||
def __init__(self, model_dir=None):
|
||||
super().__init__()
|
||||
self.model_dir = model_dir or T5_MODEL
|
||||
|
||||
def get_model(self, *, local_rank, device_id, world_size):
|
||||
super().get_model(local_rank=local_rank, device_id=device_id, world_size=world_size)
|
||||
model = T5EncoderModel.from_pretrained(self.model_dir)
|
||||
if world_size > 1:
|
||||
model = setup_fsdp_sync(
|
||||
model,
|
||||
device_id=device_id,
|
||||
param_dtype=torch.float32,
|
||||
auto_wrap_policy=partial(
|
||||
transformer_auto_wrap_policy,
|
||||
transformer_layer_cls={
|
||||
T5Block,
|
||||
},
|
||||
),
|
||||
)
|
||||
elif isinstance(device_id, int):
|
||||
model = model.to(torch.device(f"cuda:{device_id}")) # type: ignore
|
||||
return model.eval()
|
||||
|
||||
|
||||
class DitModelFactory(ModelFactory):
|
||||
def __init__(
|
||||
self, *,
|
||||
model_path: str,
|
||||
model_dtype: str,
|
||||
lora_path: Optional[str] = None,
|
||||
attention_mode: Optional[str] = None
|
||||
):
|
||||
# Infer attention mode if not specified
|
||||
if attention_mode is None:
|
||||
from genmo.lib.attn_imports import flash_varlen_attn # type: ignore
|
||||
attention_mode = "sdpa" if flash_varlen_attn is None else "flash"
|
||||
print(f"Attention mode: {attention_mode}")
|
||||
|
||||
super().__init__(
|
||||
model_path=model_path,
|
||||
lora_path=lora_path,
|
||||
model_dtype=model_dtype,
|
||||
attention_mode=attention_mode
|
||||
)
|
||||
|
||||
def get_model(
|
||||
self,
|
||||
*,
|
||||
local_rank,
|
||||
device_id,
|
||||
world_size,
|
||||
model_kwargs=None,
|
||||
patch_model_fns=None,
|
||||
strict_load=True,
|
||||
load_checkpoint=True,
|
||||
fast_init=True,
|
||||
):
|
||||
from genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
|
||||
|
||||
if not model_kwargs:
|
||||
model_kwargs = {}
|
||||
|
||||
lora_sd = None
|
||||
lora_path = self.kwargs["lora_path"]
|
||||
if lora_path is not None:
|
||||
if lora_path.endswith(".safetensors"):
|
||||
lora_sd = {}
|
||||
with safe_open(lora_path, framework="pt") as f:
|
||||
for k in f.keys():
|
||||
lora_sd[k] = f.get_tensor(k)
|
||||
lora_kwargs = json.loads(f.metadata()["kwargs"])
|
||||
print(f"Loaded LoRA kwargs: {lora_kwargs}")
|
||||
else:
|
||||
lora = load_to_cpu(lora_path, weights_only=False)
|
||||
lora_sd, lora_kwargs = lora["state_dict"], lora["kwargs"]
|
||||
|
||||
model_kwargs.update(cast(dict, lora_kwargs))
|
||||
|
||||
model_args = dict(
|
||||
depth=48,
|
||||
patch_size=2,
|
||||
num_heads=24,
|
||||
hidden_size_x=3072,
|
||||
hidden_size_y=1536,
|
||||
mlp_ratio_x=4.0,
|
||||
mlp_ratio_y=4.0,
|
||||
in_channels=12,
|
||||
qk_norm=True,
|
||||
qkv_bias=False,
|
||||
out_bias=True,
|
||||
patch_embed_bias=True,
|
||||
timestep_mlp_bias=True,
|
||||
timestep_scale=1000.0,
|
||||
t5_feat_dim=4096,
|
||||
t5_token_length=256,
|
||||
rope_theta=10000.0,
|
||||
attention_mode=self.kwargs["attention_mode"],
|
||||
**model_kwargs,
|
||||
)
|
||||
|
||||
if fast_init:
|
||||
model: nn.Module = torch.nn.utils.skip_init(AsymmDiTJoint, **model_args)
|
||||
else:
|
||||
model: nn.Module = AsymmDiTJoint(**model_args)
|
||||
|
||||
for fn in patch_model_fns or []:
|
||||
model = fn(model)
|
||||
|
||||
# FSDP syncs weights from rank 0 to all other ranks
|
||||
if local_rank == 0 and load_checkpoint:
|
||||
model_path = self.kwargs["model_path"]
|
||||
sd = load_to_cpu(model_path)
|
||||
|
||||
# Load the state dictionary and capture the return value
|
||||
load_result = model.load_state_dict(sd, strict=strict_load)
|
||||
if not strict_load:
|
||||
# Print mismatched keys
|
||||
missing_keys = [k for k in load_result.missing_keys if ".lora_" not in k]
|
||||
if missing_keys:
|
||||
print(f"Missing keys from {model_path}: {missing_keys}")
|
||||
if load_result.unexpected_keys:
|
||||
print(f"Unexpected keys from {model_path}: {load_result.unexpected_keys}")
|
||||
|
||||
if lora_sd:
|
||||
model.load_state_dict(lora_sd, strict=strict_load) # type: ignore
|
||||
|
||||
if world_size > 1:
|
||||
assert self.kwargs["model_dtype"] == "bf16", "FP8 is not supported for multi-GPU inference"
|
||||
|
||||
model = setup_fsdp_sync(
|
||||
model,
|
||||
device_id=device_id,
|
||||
param_dtype=torch.float32,
|
||||
auto_wrap_policy=partial(
|
||||
lambda_auto_wrap_policy,
|
||||
lambda_fn=lambda m: m in model.blocks,
|
||||
),
|
||||
)
|
||||
elif isinstance(device_id, int):
|
||||
model = model.to(torch.device(f"cuda:{device_id}"))
|
||||
return model.eval()
|
||||
|
||||
|
||||
class DecoderModelFactory(ModelFactory):
|
||||
def __init__(self, *, model_path: str):
|
||||
super().__init__(model_path=model_path)
|
||||
|
||||
def get_model(self, *, local_rank=0, device_id=0, world_size=1):
|
||||
# TODO(ved): Set flag for torch.compile
|
||||
# TODO(ved): Use skip_init
|
||||
|
||||
decoder = Decoder(
|
||||
out_channels=3,
|
||||
base_channels=128,
|
||||
channel_multipliers=[1, 2, 4, 6],
|
||||
temporal_expansions=[1, 2, 3],
|
||||
spatial_expansions=[2, 2, 2],
|
||||
num_res_blocks=[3, 3, 4, 6, 3],
|
||||
latent_dim=12,
|
||||
has_attention=[False, False, False, False, False],
|
||||
output_norm=False,
|
||||
nonlinearity="silu",
|
||||
output_nonlinearity="silu",
|
||||
causal=True,
|
||||
)
|
||||
# VAE is not FSDP-wrapped
|
||||
state_dict = load_file(self.kwargs["model_path"])
|
||||
decoder.load_state_dict(state_dict, strict=True)
|
||||
device = torch.device(f"cuda:{device_id}") if isinstance(device_id, int) else "cpu"
|
||||
decoder.eval().to(device)
|
||||
return decoder
|
||||
|
||||
|
||||
class EncoderModelFactory(ModelFactory):
|
||||
def __init__(self, *, model_path: str):
|
||||
super().__init__(model_path=model_path)
|
||||
|
||||
def get_model(self, *, local_rank=0, device_id=0, world_size=1):
|
||||
# TODO(ved): Set flag for torch.compile
|
||||
# TODO(ved): Use skip_init
|
||||
|
||||
# We don't FSDP the encoder b/c it is small
|
||||
encoder = Encoder(
|
||||
in_channels=15,
|
||||
base_channels=64,
|
||||
channel_multipliers=[1, 2, 4, 6],
|
||||
num_res_blocks=[3, 3, 4, 6, 3],
|
||||
latent_dim=12,
|
||||
temporal_reductions=[1, 2, 3],
|
||||
spatial_reductions=[2, 2, 2],
|
||||
prune_bottlenecks=[False, False, False, False, False],
|
||||
has_attentions=[False, True, True, True, True],
|
||||
affine=True,
|
||||
bias=True,
|
||||
input_is_conv_1x1=True,
|
||||
padding_mode="replicate",
|
||||
)
|
||||
state_dict = load_file(self.kwargs["model_path"])
|
||||
encoder.load_state_dict(state_dict, strict=True)
|
||||
device = torch.device(f"cuda:{device_id}") if isinstance(device_id, int) else "cpu"
|
||||
encoder.eval().to(device)
|
||||
return encoder
|
||||
|
||||
|
||||
def get_conditioning(
|
||||
tokenizer: T5Tokenizer,
|
||||
encoder: Encoder,
|
||||
device: torch.device,
|
||||
batch_inputs: bool,
|
||||
*,
|
||||
prompt: str,
|
||||
negative_prompt: str,
|
||||
):
|
||||
if batch_inputs:
|
||||
return dict(
|
||||
batched=get_conditioning_for_prompts(
|
||||
tokenizer, encoder, device, [prompt, negative_prompt]
|
||||
)
|
||||
)
|
||||
else:
|
||||
cond_input = get_conditioning_for_prompts(tokenizer, encoder, device, [prompt])
|
||||
null_input = get_conditioning_for_prompts(tokenizer, encoder, device, [negative_prompt])
|
||||
return dict(cond=cond_input, null=null_input)
|
||||
|
||||
|
||||
def get_conditioning_for_prompts(tokenizer, encoder, device, prompts: List[str]):
|
||||
assert len(prompts) in [1, 2] # [neg] or [pos] or [pos, neg]
|
||||
B = len(prompts)
|
||||
t5_toks = tokenizer(
|
||||
prompts,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=MAX_T5_TOKEN_LENGTH,
|
||||
return_tensors="pt",
|
||||
return_attention_mask=True,
|
||||
)
|
||||
caption_input_ids_t5 = t5_toks["input_ids"]
|
||||
caption_attention_mask_t5 = t5_toks["attention_mask"].bool()
|
||||
del t5_toks
|
||||
|
||||
assert caption_input_ids_t5.shape == (B, MAX_T5_TOKEN_LENGTH)
|
||||
assert caption_attention_mask_t5.shape == (B, MAX_T5_TOKEN_LENGTH)
|
||||
|
||||
# Special-case empty negative prompt by zero-ing it
|
||||
if prompts[-1] == "":
|
||||
caption_input_ids_t5[-1] = 0
|
||||
caption_attention_mask_t5[-1] = False
|
||||
|
||||
caption_input_ids_t5 = caption_input_ids_t5.to(device, non_blocking=True)
|
||||
caption_attention_mask_t5 = caption_attention_mask_t5.to(device, non_blocking=True)
|
||||
|
||||
y_mask = [caption_attention_mask_t5]
|
||||
y_feat = [encoder(caption_input_ids_t5, caption_attention_mask_t5).last_hidden_state.detach()]
|
||||
# Sometimes returns a tensor, othertimes a tuple, not sure why
|
||||
# See: https://huggingface.co/genmo/mochi-1-preview/discussions/3
|
||||
assert tuple(y_feat[-1].shape) == (B, MAX_T5_TOKEN_LENGTH, 4096)
|
||||
assert y_feat[-1].dtype == torch.float32
|
||||
|
||||
return dict(y_mask=y_mask, y_feat=y_feat)
|
||||
|
||||
|
||||
def compute_packed_indices(
|
||||
device: torch.device, text_mask: torch.Tensor, num_latents: int
|
||||
) -> Dict[str, Union[torch.Tensor, int]]:
|
||||
"""
|
||||
Based on https://github.com/Dao-AILab/flash-attention/blob/765741c1eeb86c96ee71a3291ad6968cfbf4e4a1/flash_attn/bert_padding.py#L60-L80
|
||||
|
||||
Args:
|
||||
num_latents: Number of latent tokens
|
||||
text_mask: (B, L) List of boolean tensor indicating which text tokens are not padding.
|
||||
|
||||
Returns:
|
||||
packed_indices: Dict with keys for Flash Attention:
|
||||
- valid_token_indices_kv: up to (B * (N + L),) tensor of valid token indices (non-padding)
|
||||
in the packed sequence.
|
||||
- cu_seqlens_kv: (B + 1,) tensor of cumulative sequence lengths in the packed sequence.
|
||||
- max_seqlen_in_batch_kv: int of the maximum sequence length in the batch.
|
||||
"""
|
||||
# Create an expanded token mask saying which tokens are valid across both visual and text tokens.
|
||||
PATCH_SIZE = 2
|
||||
num_visual_tokens = num_latents // (PATCH_SIZE**2)
|
||||
assert num_visual_tokens > 0
|
||||
|
||||
mask = F.pad(text_mask, (num_visual_tokens, 0), value=True) # (B, N + L)
|
||||
seqlens_in_batch = mask.sum(dim=-1, dtype=torch.int32) # (B,)
|
||||
valid_token_indices = torch.nonzero(mask.flatten(), as_tuple=False).flatten() # up to (B * (N + L),)
|
||||
assert valid_token_indices.size(0) >= text_mask.size(0) * num_visual_tokens # At least (B * N,)
|
||||
cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
|
||||
max_seqlen_in_batch = seqlens_in_batch.max().item()
|
||||
|
||||
return {
|
||||
"cu_seqlens_kv": cu_seqlens.to(device, non_blocking=True),
|
||||
"max_seqlen_in_batch_kv": cast(int, max_seqlen_in_batch),
|
||||
"valid_token_indices_kv": valid_token_indices.to(device, non_blocking=True),
|
||||
}
|
||||
|
||||
|
||||
def assert_eq(x, y, msg=None):
|
||||
assert x == y, f"{msg or 'Assertion failed'}: {x} != {y}"
|
||||
|
||||
|
||||
def sample_model(device, dit, conditioning, **args):
|
||||
random.seed(args["seed"])
|
||||
np.random.seed(args["seed"])
|
||||
torch.manual_seed(args["seed"])
|
||||
|
||||
generator = torch.Generator(device=device)
|
||||
generator.manual_seed(args["seed"])
|
||||
|
||||
w, h, t = args["width"], args["height"], args["num_frames"]
|
||||
sample_steps = args["num_inference_steps"]
|
||||
cfg_schedule = args["cfg_schedule"]
|
||||
sigma_schedule = args["sigma_schedule"]
|
||||
|
||||
assert_eq(len(cfg_schedule), sample_steps, "cfg_schedule must have length sample_steps")
|
||||
assert_eq((t - 1) % 6, 0, "t - 1 must be divisible by 6")
|
||||
assert_eq(
|
||||
len(sigma_schedule),
|
||||
sample_steps + 1,
|
||||
"sigma_schedule must have length sample_steps + 1",
|
||||
)
|
||||
|
||||
B = 1
|
||||
SPATIAL_DOWNSAMPLE = 8
|
||||
TEMPORAL_DOWNSAMPLE = 6
|
||||
IN_CHANNELS = 12
|
||||
latent_t = ((t - 1) // TEMPORAL_DOWNSAMPLE) + 1
|
||||
latent_w, latent_h = w // SPATIAL_DOWNSAMPLE, h // SPATIAL_DOWNSAMPLE
|
||||
|
||||
z = torch.randn(
|
||||
(B, IN_CHANNELS, latent_t, latent_h, latent_w),
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
num_latents = latent_t * latent_h * latent_w
|
||||
cond_batched = cond_text = cond_null = None
|
||||
if "cond" in conditioning:
|
||||
cond_text = conditioning["cond"]
|
||||
cond_null = conditioning["null"]
|
||||
cond_text["packed_indices"] = compute_packed_indices(device, cond_text["y_mask"][0], num_latents)
|
||||
cond_null["packed_indices"] = compute_packed_indices(device, cond_null["y_mask"][0], num_latents)
|
||||
else:
|
||||
cond_batched = conditioning["batched"]
|
||||
cond_batched["packed_indices"] = compute_packed_indices(device, cond_batched["y_mask"][0], num_latents)
|
||||
z = repeat(z, "b ... -> (repeat b) ...", repeat=2)
|
||||
|
||||
def model_fn(*, z, sigma, cfg_scale):
|
||||
if cond_batched:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
out = dit(z, sigma, **cond_batched)
|
||||
out_cond, out_uncond = torch.chunk(out, chunks=2, dim=0)
|
||||
else:
|
||||
nonlocal cond_text, cond_null
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
out_cond = dit(z, sigma, **cond_text)
|
||||
out_uncond = dit(z, sigma, **cond_null)
|
||||
assert out_cond.shape == out_uncond.shape
|
||||
out_uncond = out_uncond.to(z)
|
||||
out_cond = out_cond.to(z)
|
||||
return out_uncond + cfg_scale * (out_cond - out_uncond)
|
||||
|
||||
# Euler sampler w/ customizable sigma schedule & cfg scale
|
||||
for i in get_new_progress_bar(range(0, sample_steps), desc="Sampling"):
|
||||
sigma = sigma_schedule[i]
|
||||
dsigma = sigma - sigma_schedule[i + 1]
|
||||
|
||||
# `pred` estimates `z_0 - eps`.
|
||||
pred = model_fn(
|
||||
z=z,
|
||||
sigma=torch.full([B] if cond_text else [B * 2], sigma, device=z.device),
|
||||
cfg_scale=cfg_schedule[i],
|
||||
)
|
||||
assert pred.dtype == torch.float32
|
||||
z = z + dsigma * pred
|
||||
|
||||
z = z[:B] if cond_batched else z
|
||||
return dit_latents_to_vae_latents(z)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def move_to_device(model: nn.Module, target_device, *, enabled=True):
|
||||
if not enabled:
|
||||
yield
|
||||
return
|
||||
|
||||
og_device = next(model.parameters()).device
|
||||
if og_device == target_device:
|
||||
print(f"move_to_device is a no-op model is already on {target_device}")
|
||||
else:
|
||||
print(f"moving model from {og_device} -> {target_device}")
|
||||
|
||||
model.to(target_device)
|
||||
yield
|
||||
if og_device != target_device:
|
||||
print(f"moving model from {target_device} -> {og_device}")
|
||||
model.to(og_device)
|
||||
|
||||
|
||||
def t5_tokenizer(model_dir=None):
|
||||
return T5Tokenizer.from_pretrained(model_dir or T5_MODEL, legacy=False)
|
||||
|
||||
|
||||
class MochiSingleGPUPipeline:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
text_encoder_factory: ModelFactory,
|
||||
dit_factory: ModelFactory,
|
||||
decoder_factory: ModelFactory,
|
||||
cpu_offload: Optional[bool] = False,
|
||||
decode_type: str = "full",
|
||||
decode_args: Optional[Dict[str, Any]] = None,
|
||||
fast_init=True,
|
||||
strict_load=True
|
||||
):
|
||||
self.device = torch.device("cuda:0")
|
||||
self.tokenizer = t5_tokenizer(text_encoder_factory.model_dir)
|
||||
t = Timer()
|
||||
self.cpu_offload = cpu_offload
|
||||
self.decode_args = decode_args or {}
|
||||
self.decode_type = decode_type
|
||||
init_id = "cpu" if cpu_offload else 0
|
||||
with t("load_text_encoder"):
|
||||
self.text_encoder = text_encoder_factory.get_model(
|
||||
local_rank=0,
|
||||
device_id=init_id,
|
||||
world_size=1,
|
||||
)
|
||||
with t("load_dit"):
|
||||
self.dit = dit_factory.get_model(local_rank=0, device_id=init_id, world_size=1, fast_init=fast_init, strict_load=strict_load) # type: ignore
|
||||
with t("load_vae"):
|
||||
self.decoder = decoder_factory.get_model(local_rank=0, device_id=init_id, world_size=1)
|
||||
t.print_stats()
|
||||
|
||||
def __call__(self, batch_cfg, prompt, negative_prompt, **kwargs):
|
||||
with torch.inference_mode():
|
||||
print_max_memory = lambda: print(
|
||||
f"Max memory reserved: {torch.cuda.max_memory_reserved() / 1024**3:.2f} GB"
|
||||
)
|
||||
print_max_memory()
|
||||
|
||||
with move_to_device(self.text_encoder, self.device):
|
||||
conditioning = get_conditioning(
|
||||
tokenizer=self.tokenizer,
|
||||
encoder=self.text_encoder,
|
||||
device=self.device,
|
||||
batch_inputs=batch_cfg,
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
)
|
||||
print_max_memory()
|
||||
|
||||
with move_to_device(self.dit, self.device):
|
||||
latents = sample_model(self.device, self.dit, conditioning, **kwargs)
|
||||
print_max_memory()
|
||||
|
||||
with move_to_device(self.decoder, self.device):
|
||||
if self.decode_type == "tiled_full":
|
||||
frames = decode_latents_tiled_full(
|
||||
self.decoder, latents, **self.decode_args)
|
||||
elif self.decode_type == "tiled_spatial":
|
||||
frames = decode_latents_tiled_spatial(
|
||||
self.decoder, latents, **self.decode_args,
|
||||
num_tiles_w=4, num_tiles_h=2)
|
||||
else:
|
||||
frames = decode_latents(self.decoder, latents)
|
||||
print_max_memory()
|
||||
return frames.cpu().numpy()
|
||||
|
||||
|
||||
def cast_dit(model, dtype):
|
||||
for name, module in model.named_modules():
|
||||
if isinstance(module, nn.Linear):
|
||||
assert any(
|
||||
n in name for n in ["mlp", "t5", "mod_", "attn.qkv_", "attn.proj_", "final_layer"]
|
||||
), f"Unexpected linear layer: {name}"
|
||||
module.to(dtype=dtype)
|
||||
elif isinstance(module, nn.Conv2d):
|
||||
assert "x_embedder.proj" in name, f"Unexpected conv2d layer: {name}"
|
||||
module.to(dtype=dtype)
|
||||
return model
|
||||
|
||||
|
||||
### ALL CODE BELOW HERE IS FOR MULTI-GPU MODE ###
|
||||
|
||||
|
||||
# In multi-gpu mode, all models must belong to a device which has a predefined context parallel group
|
||||
# So it doesn't make sense to work with models individually
|
||||
class MultiGPUContext:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
text_encoder_factory,
|
||||
dit_factory,
|
||||
decoder_factory,
|
||||
device_id,
|
||||
local_rank,
|
||||
world_size,
|
||||
):
|
||||
t = Timer()
|
||||
self.device = torch.device(f"cuda:{device_id}")
|
||||
print(f"Initializing rank {local_rank+1}/{world_size}")
|
||||
assert world_size > 1, f"Multi-GPU mode requires world_size > 1, got {world_size}"
|
||||
os.environ["MASTER_ADDR"] = "127.0.0.1"
|
||||
os.environ["MASTER_PORT"] = "29500"
|
||||
with t("init_process_group"):
|
||||
dist.init_process_group(
|
||||
"nccl",
|
||||
rank=local_rank,
|
||||
world_size=world_size,
|
||||
device_id=self.device, # force non-lazy init
|
||||
)
|
||||
pg = dist.group.WORLD
|
||||
cp.set_cp_group(pg, list(range(world_size)), local_rank)
|
||||
distributed_kwargs = dict(local_rank=local_rank, device_id=device_id, world_size=world_size)
|
||||
self.world_size = world_size
|
||||
self.tokenizer = t5_tokenizer(text_encoder_factory.model_dir)
|
||||
with t("load_text_encoder"):
|
||||
self.text_encoder = text_encoder_factory.get_model(**distributed_kwargs)
|
||||
with t("load_dit"):
|
||||
self.dit = dit_factory.get_model(**distributed_kwargs)
|
||||
with t("load_vae"):
|
||||
self.decoder = decoder_factory.get_model(**distributed_kwargs)
|
||||
self.local_rank = local_rank
|
||||
t.print_stats()
|
||||
|
||||
def run(self, *, fn, **kwargs):
|
||||
return fn(self, **kwargs)
|
||||
|
||||
|
||||
class MochiMultiGPUPipeline:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
text_encoder_factory: ModelFactory,
|
||||
dit_factory: ModelFactory,
|
||||
decoder_factory: ModelFactory,
|
||||
world_size: int,
|
||||
):
|
||||
ray.init()
|
||||
RemoteClass = ray.remote(MultiGPUContext)
|
||||
self.ctxs = [
|
||||
RemoteClass.options(num_gpus=1).remote(
|
||||
text_encoder_factory=text_encoder_factory,
|
||||
dit_factory=dit_factory,
|
||||
decoder_factory=decoder_factory,
|
||||
world_size=world_size,
|
||||
device_id=0,
|
||||
local_rank=i,
|
||||
)
|
||||
for i in range(world_size)
|
||||
]
|
||||
for ctx in self.ctxs:
|
||||
ray.get(ctx.__ray_ready__.remote())
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
def sample(ctx, *, batch_cfg, prompt, negative_prompt, **kwargs):
|
||||
with progress_bar(type="ray_tqdm", enabled=ctx.local_rank == 0), torch.inference_mode():
|
||||
conditioning = get_conditioning(
|
||||
ctx.tokenizer,
|
||||
ctx.text_encoder,
|
||||
ctx.device,
|
||||
batch_cfg,
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
)
|
||||
latents = sample_model(ctx.device, ctx.dit, conditioning=conditioning, **kwargs)
|
||||
if ctx.local_rank == 0:
|
||||
torch.save(latents, "latents.pt")
|
||||
frames = decode_latents(ctx.decoder, latents)
|
||||
return frames.cpu().numpy()
|
||||
|
||||
return ray.get([ctx.run.remote(fn=sample, **kwargs, show_progress=i == 0) for i, ctx in enumerate(self.ctxs)])[
|
||||
0
|
||||
]
|
||||
@@ -0,0 +1,155 @@
|
||||
from typing import Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
|
||||
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
|
||||
|
||||
|
||||
def cast_tuple(t, length=1):
|
||||
return t if isinstance(t, tuple) else ((t,) * length)
|
||||
|
||||
|
||||
def cp_pass_frames(x: torch.Tensor, frames_to_send: int) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass that handles communication between ranks for inference.
|
||||
Args:
|
||||
x: Tensor of shape (B, C, T, H, W)
|
||||
frames_to_send: int, number of frames to communicate between ranks
|
||||
Returns:
|
||||
output: Tensor of shape (B, C, T', H, W)
|
||||
"""
|
||||
cp_rank, cp_world_size = cp.get_cp_rank_size()
|
||||
if frames_to_send == 0 or cp_world_size == 1:
|
||||
return x
|
||||
|
||||
group = cp.get_cp_group()
|
||||
global_rank = dist.get_rank()
|
||||
|
||||
# Send to next rank
|
||||
if cp_rank < cp_world_size - 1:
|
||||
assert x.size(2) >= frames_to_send
|
||||
tail = x[:, :, -frames_to_send:].contiguous()
|
||||
dist.send(tail, global_rank + 1, group=group)
|
||||
|
||||
# Receive from previous rank
|
||||
if cp_rank > 0:
|
||||
B, C, _, H, W = x.shape
|
||||
recv_buffer = torch.empty(
|
||||
(B, C, frames_to_send, H, W),
|
||||
dtype=x.dtype,
|
||||
device=x.device,
|
||||
)
|
||||
dist.recv(recv_buffer, global_rank - 1, group=group)
|
||||
x = torch.cat([recv_buffer, x], dim=2)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def _pad_to_max(x: torch.Tensor, max_T: int) -> torch.Tensor:
|
||||
if max_T > x.size(2):
|
||||
pad_T = max_T - x.size(2)
|
||||
pad_dims = (0, 0, 0, 0, 0, pad_T)
|
||||
return F.pad(x, pad_dims)
|
||||
return x
|
||||
|
||||
|
||||
def gather_all_frames(x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Gathers all frames from all processes for inference.
|
||||
Args:
|
||||
x: Tensor of shape (B, C, T, H, W)
|
||||
Returns:
|
||||
output: Tensor of shape (B, C, T_total, H, W)
|
||||
"""
|
||||
cp_rank, cp_size = cp.get_cp_rank_size()
|
||||
if cp_size == 1:
|
||||
return x
|
||||
|
||||
cp_group = cp.get_cp_group()
|
||||
|
||||
# Ensure the tensor is contiguous for collective operations
|
||||
x = x.contiguous()
|
||||
|
||||
# Get the local time dimension size
|
||||
local_T = x.size(2)
|
||||
local_T_tensor = torch.tensor([local_T], device=x.device, dtype=torch.int64)
|
||||
|
||||
# Gather all T sizes from all processes
|
||||
all_T = [torch.zeros(1, dtype=torch.int64, device=x.device) for _ in range(cp_size)]
|
||||
dist.all_gather(all_T, local_T_tensor, group=cp_group)
|
||||
all_T = [t.item() for t in all_T]
|
||||
|
||||
# Pad the tensor at the end of the time dimension to match max_T
|
||||
max_T = max(all_T)
|
||||
x = _pad_to_max(x, max_T).contiguous()
|
||||
|
||||
# Prepare a list to hold the gathered tensors
|
||||
gathered_x = [torch.zeros_like(x).contiguous() for _ in range(cp_size)]
|
||||
|
||||
# Perform the all_gather operation
|
||||
dist.all_gather(gathered_x, x, group=cp_group)
|
||||
|
||||
# Slice each gathered tensor back to its original T size
|
||||
for idx, t_size in enumerate(all_T):
|
||||
gathered_x[idx] = gathered_x[idx][:, :, :t_size]
|
||||
|
||||
return torch.cat(gathered_x, dim=2)
|
||||
|
||||
|
||||
def excessive_memory_usage(input: torch.Tensor, max_gb: float = 2.0) -> bool:
|
||||
"""Estimate memory usage based on input tensor size and data type."""
|
||||
element_size = input.element_size() # Size in bytes of each element
|
||||
memory_bytes = input.numel() * element_size
|
||||
memory_gb = memory_bytes / 1024**3
|
||||
return memory_gb > max_gb
|
||||
|
||||
|
||||
class ContextParallelCausalConv3d(torch.nn.Conv3d):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size: Union[int, Tuple[int, int, int]],
|
||||
stride: Union[int, Tuple[int, int, int]],
|
||||
**kwargs,
|
||||
):
|
||||
kernel_size = cast_tuple(kernel_size, 3)
|
||||
stride = cast_tuple(stride, 3)
|
||||
height_pad = (kernel_size[1] - 1) // 2
|
||||
width_pad = (kernel_size[2] - 1) // 2
|
||||
|
||||
super().__init__(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
dilation=(1, 1, 1),
|
||||
padding=(0, height_pad, width_pad),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
cp_rank, cp_world_size = cp.get_cp_rank_size()
|
||||
|
||||
context_size = self.kernel_size[0] - 1
|
||||
if cp_rank == 0:
|
||||
mode = "constant" if self.padding_mode == "zeros" else self.padding_mode
|
||||
x = F.pad(x, (0, 0, 0, 0, context_size, 0), mode=mode)
|
||||
|
||||
if cp_world_size == 1:
|
||||
return super().forward(x)
|
||||
|
||||
if all(s == 1 for s in self.stride):
|
||||
# Receive some frames from previous rank.
|
||||
x = cp_pass_frames(x, context_size)
|
||||
return super().forward(x)
|
||||
|
||||
# Less efficient implementation for strided convs.
|
||||
# All gather x, infer and chunk.
|
||||
x = gather_all_frames(x) # [B, C, k - 1 + global_T, H, W]
|
||||
x = super().forward(x)
|
||||
x_chunks = x.tensor_split(cp_world_size, dim=2)
|
||||
assert len(x_chunks) == cp_world_size
|
||||
return x_chunks[cp_rank]
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Container for latent space posterior."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class LatentDistribution:
|
||||
def __init__(self, mean: torch.Tensor, logvar: torch.Tensor):
|
||||
"""Initialize latent distribution.
|
||||
|
||||
Args:
|
||||
mean: Mean of the distribution. Shape: [B, C, T, H, W].
|
||||
logvar: Logarithm of variance of the distribution. Shape: [B, C, T, H, W].
|
||||
"""
|
||||
assert mean.shape == logvar.shape
|
||||
self.mean = mean
|
||||
self.logvar = logvar
|
||||
|
||||
def sample(self, temperature=1.0, generator: torch.Generator = None, noise=None):
|
||||
if temperature == 0.0:
|
||||
return self.mean
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn(self.mean.shape, device=self.mean.device, dtype=self.mean.dtype, generator=generator)
|
||||
else:
|
||||
assert noise.device == self.mean.device
|
||||
noise = noise.to(self.mean.dtype)
|
||||
|
||||
if temperature != 1.0:
|
||||
raise NotImplementedError(f"Temperature {temperature} is not supported.")
|
||||
|
||||
# Just Gaussian sample with no scaling of variance.
|
||||
return noise * torch.exp(self.logvar * 0.5) + self.mean
|
||||
|
||||
def mode(self):
|
||||
return self.mean
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,67 @@
|
||||
import torch
|
||||
|
||||
# Channel-wise mean and standard deviation of VAE encoder latents
|
||||
STATS = {
|
||||
"mean": torch.Tensor(
|
||||
[
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285,
|
||||
]
|
||||
),
|
||||
"std": torch.Tensor(
|
||||
[
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
0.9274757570805041,
|
||||
]
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def dit_latents_to_vae_latents(dit_outputs: torch.Tensor) -> torch.Tensor:
|
||||
"""Unnormalize latents output by Mochi's DiT to be compatible with VAE.
|
||||
Run this on sampled latents before calling the VAE decoder.
|
||||
|
||||
Args:
|
||||
latents (torch.Tensor): [B, C_z, T_z, H_z, W_z], float
|
||||
|
||||
Returns:
|
||||
torch.Tensor: [B, C_z, T_z, H_z, W_z], float
|
||||
"""
|
||||
mean = STATS["mean"][:, None, None, None]
|
||||
std = STATS["std"][:, None, None, None]
|
||||
|
||||
assert dit_outputs.ndim == 5
|
||||
assert dit_outputs.size(1) == mean.size(0) == std.size(0)
|
||||
return dit_outputs * std.to(dit_outputs) + mean.to(dit_outputs)
|
||||
|
||||
|
||||
def vae_latents_to_dit_latents(vae_latents: torch.Tensor):
|
||||
"""Normalize latents output by the VAE encoder to be compatible with Mochi's DiT.
|
||||
E.g, for fine-tuning or video-to-video.
|
||||
"""
|
||||
mean = STATS["mean"][:, None, None, None]
|
||||
std = STATS["std"][:, None, None, None]
|
||||
|
||||
assert vae_latents.ndim == 5
|
||||
assert vae_latents.size(1) == mean.size(0) == std.size(0)
|
||||
return (vae_latents - mean.to(vae_latents)) / std.to(vae_latents)
|
||||
@@ -0,0 +1,431 @@
|
||||
import torch
|
||||
import argparse
|
||||
from safetensors.torch import save_file
|
||||
import os
|
||||
|
||||
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(f"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(f"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)
|
||||
@@ -0,0 +1,42 @@
|
||||
import torch
|
||||
|
||||
mochi_latents_mean = torch.tensor(
|
||||
[
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285,
|
||||
]
|
||||
).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor(
|
||||
[
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
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
|
||||
@@ -19,13 +19,29 @@ import torch.nn as nn
|
||||
import diffusers
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import is_torch_version, logging
|
||||
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 diffusers.models.attention import FeedForward as HF_FeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.embeddings import MochiCombinedTimestepCaptionEmbedding, PatchEmbed
|
||||
from diffusers.models.embeddings import (
|
||||
MochiCombinedTimestepCaptionEmbedding,
|
||||
PatchEmbed,
|
||||
)
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from fastvideo.model.norm import MochiLayerNormContinuous, MochiRMSNormZero, MochiModulatedRMSNorm, MochiRMSNorm
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
from fastvideo.models.mochi_hf.norm import (
|
||||
MochiLayerNormContinuous,
|
||||
MochiRMSNormZero,
|
||||
MochiModulatedRMSNorm,
|
||||
MochiRMSNorm,
|
||||
)
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
@@ -39,6 +55,9 @@ from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from liger_kernel.ops.swiglu import LigerSiLUMulFunction
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class FeedForward(HF_FeedForward):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -51,37 +70,50 @@ 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):
|
||||
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)
|
||||
|
||||
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)
|
||||
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
|
||||
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,
|
||||
)
|
||||
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 +147,25 @@ 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.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,15 +183,15 @@ class MochiAttention(nn.Module):
|
||||
attention_mask=attention_mask,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class MochiAttnProcessor2_0:
|
||||
"""Attention processor used in Mochi."""
|
||||
|
||||
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 +212,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 +225,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 +263,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),
|
||||
@@ -237,6 +277,7 @@ class MochiAttnProcessor2_0:
|
||||
sequence_length = query.size(1)
|
||||
encoder_sequence_length = encoder_query.size(1)
|
||||
|
||||
|
||||
# H
|
||||
query = torch.cat([query, encoder_query], dim=1).unsqueeze(2)
|
||||
key = torch.cat([key, encoder_key], dim=1).unsqueeze(2)
|
||||
@@ -246,14 +287,15 @@ 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 = 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(
|
||||
@@ -261,7 +303,9 @@ class MochiAttnProcessor2_0:
|
||||
)
|
||||
# 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()
|
||||
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)
|
||||
@@ -273,8 +317,6 @@ class MochiAttnProcessor2_0:
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
|
||||
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
@@ -286,6 +328,7 @@ class MochiAttnProcessor2_0:
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class MochiTransformerBlock(nn.Module):
|
||||
r"""
|
||||
@@ -328,7 +371,9 @@ class MochiTransformerBlock(nn.Module):
|
||||
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 +397,18 @@ 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,14 +428,19 @@ 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)
|
||||
|
||||
@@ -392,20 +448,27 @@ class MochiTransformerBlock(nn.Module):
|
||||
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)
|
||||
)
|
||||
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(
|
||||
@@ -447,18 +510,22 @@ class MochiRoPE(nn.Module):
|
||||
) -> 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)
|
||||
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:
|
||||
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", pos.to(torch.float32), freqs.to(torch.float32)
|
||||
)
|
||||
freqs_cos = torch.cos(freqs)
|
||||
freqs_sin = torch.sin(freqs)
|
||||
return freqs_cos, freqs_sin
|
||||
@@ -478,7 +545,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,7 +612,9 @@ 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(
|
||||
@@ -564,7 +633,11 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
@@ -576,32 +649,57 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
hidden_states: torch.Tensor, # [2, 12, 28, 60, 106]
|
||||
encoder_hidden_states: torch.Tensor, # [2, 256, 4096]
|
||||
timestep: torch.LongTensor, # [2]
|
||||
encoder_attention_mask: torch.Tensor, #[2, 256]
|
||||
output_attn = False,
|
||||
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"
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
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
|
||||
temb, encoder_hidden_states = self.time_embed( # [2, 3072], [2, 256, 1536]
|
||||
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.permute(0, 2, 1, 3, 4).flatten(0, 1) # [56, 12, 60, 106]
|
||||
hidden_states = self.patch_embed(hidden_states) # [56, 1590, 3072]
|
||||
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2) # [2, 44520, 3072]
|
||||
|
||||
image_rotary_emb = self.rope(
|
||||
self.pos_frequencies,
|
||||
num_frames,
|
||||
image_rotary_emb = self.rope( #[0][44520, 24, 64]
|
||||
self.pos_frequencies, #[3, 24, 64]
|
||||
num_frames, # 28
|
||||
post_patch_height,
|
||||
post_patch_width,
|
||||
device=hidden_states.device,
|
||||
@@ -617,8 +715,14 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
|
||||
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,
|
||||
@@ -629,26 +733,30 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
else:
|
||||
hidden_states, encoder_hidden_states, attn_outputs = block(
|
||||
hidden_states, encoder_hidden_states, attn_outputs = block( # [2, 44520, 3072], [2, 256, 1536],
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
output_attn = output_attn,
|
||||
output_attn=output_attn,
|
||||
)
|
||||
attn_outputs_list.append(attn_outputs)
|
||||
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = self.proj_out(hidden_states) #[2, 44520, 48]
|
||||
|
||||
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)
|
||||
hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1) # [2, 28, 30, 53, 2, 2, 12]
|
||||
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5) # [2, 12, 28, 30, 2, 53, 2]
|
||||
output = hidden_states.reshape(batch_size, -1, num_frames, height, width) # [2, 12, 28, 60, 106]
|
||||
|
||||
if not output_attn :
|
||||
attn_outputs_list = None
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not output_attn:
|
||||
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)
|
||||
@@ -38,7 +38,7 @@ 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):
|
||||
@@ -63,7 +63,7 @@ class MochiRMSNorm(nn.Module):
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
|
||||
class MochiLayerNormContinuous(nn.Module):
|
||||
def __init__(
|
||||
@@ -92,7 +92,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 +102,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 +122,9 @@ 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
|
||||
@@ -13,7 +13,7 @@
|
||||
# limitations under the License.
|
||||
|
||||
import inspect
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
from typing import Callable, Dict, List, Optional, Union, Any
|
||||
import copy
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -21,7 +21,7 @@ from transformers import T5EncoderModel, T5TokenizerFast
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.models.autoencoders import AutoencoderKL
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import (
|
||||
@@ -35,7 +35,8 @@ 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 fastvideo.utils.communications import all_gather
|
||||
from diffusers.loaders import Mochi1LoraLoaderMixin
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
@@ -80,14 +81,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)
|
||||
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 +133,13 @@ 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 +149,9 @@ 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 +166,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.
|
||||
|
||||
@@ -199,14 +211,17 @@ 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.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.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,22 +253,32 @@ 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}"
|
||||
)
|
||||
|
||||
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)
|
||||
@@ -320,7 +345,11 @@ 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):
|
||||
raise TypeError(
|
||||
@@ -334,7 +363,10 @@ class MochiPipeline(DiffusionPipeline):
|
||||
" 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 +374,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,10 +393,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]}"
|
||||
@@ -374,14 +414,25 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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:
|
||||
@@ -467,6 +518,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
|
||||
@@ -492,10 +547,11 @@ class MochiPipeline(DiffusionPipeline):
|
||||
negative_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[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 +603,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,
|
||||
@@ -586,6 +646,7 @@ class MochiPipeline(DiffusionPipeline):
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._interrupt = False
|
||||
|
||||
# 2. Define call parameters
|
||||
@@ -618,7 +679,9 @@ class MochiPipeline(DiffusionPipeline):
|
||||
)
|
||||
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_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,9 +698,10 @@ 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
|
||||
@@ -646,7 +710,7 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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,7 +724,9 @@ 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
|
||||
@@ -669,26 +735,36 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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)
|
||||
|
||||
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:
|
||||
@@ -706,7 +782,9 @@ class MochiPipeline(DiffusionPipeline):
|
||||
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 +792,49 @@ 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)
|
||||
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)
|
||||
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 = 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:
|
||||
@@ -2,12 +2,15 @@ import json
|
||||
|
||||
import torch.distributed as dist
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
import os
|
||||
from diffusers.utils import export_to_video
|
||||
import argparse
|
||||
|
||||
def generate_video_and_latent(pipe, prompt, height, width, num_frames, num_inference_steps, guidance_scale):
|
||||
|
||||
def generate_video_and_latent(
|
||||
pipe, prompt, height, width, num_frames, num_inference_steps, guidance_scale
|
||||
):
|
||||
# Set the random seed for reproducibility
|
||||
generator = torch.Generator("cuda").manual_seed(12345)
|
||||
# Generate videos from the input prompt
|
||||
@@ -19,17 +22,16 @@ 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 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 +39,71 @@ 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(
|
||||
"--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.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_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")
|
||||
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")
|
||||
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 +111,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 +123,11 @@ 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)
|
||||
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.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
|
||||
from fastvideo.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state,
|
||||
nccl_info,
|
||||
)
|
||||
import argparse
|
||||
import os
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
import json
|
||||
from typing import Optional
|
||||
from safetensors.torch import save_file, load_file
|
||||
@@ -17,88 +20,21 @@ import pdb
|
||||
import copy
|
||||
from typing import Dict
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import convert_unet_state_dict_to_peft
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
|
||||
from safetensors.torch import load_file
|
||||
|
||||
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()
|
||||
@@ -110,54 +46,92 @@ def main(args):
|
||||
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)
|
||||
|
||||
mochi_genmo = False
|
||||
if mochi_genmo:
|
||||
model_path = "/root/weights/dit.safetensors"
|
||||
state_dcit = load_file(model_path)
|
||||
transformer = AsymmDiTJoint()
|
||||
transformer.load_state_dict(state_dcit)
|
||||
# from IPython import embed
|
||||
# embed()
|
||||
transformer.config.in_channels = 12
|
||||
print("load gennmo mochi successfully")
|
||||
else:
|
||||
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()
|
||||
|
||||
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:
|
||||
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])
|
||||
else:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
videos = pipe(
|
||||
@@ -173,18 +147,21 @@ def main(args):
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if prompts is not None:
|
||||
# mkdir
|
||||
# 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)
|
||||
export_to_video(
|
||||
video, os.path.join(args.output_path, f"{suffix}.mp4"), fps=30
|
||||
)
|
||||
else:
|
||||
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)
|
||||
@@ -198,7 +175,12 @@ if __name__ == "__main__":
|
||||
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(
|
||||
"--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,9 +1,11 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import export_to_video, load_image, load_video
|
||||
import argparse
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
|
||||
def main(args):
|
||||
# Set the random seed for reproducibility
|
||||
generator = torch.Generator("cuda").manual_seed(args.seed)
|
||||
@@ -12,8 +14,12 @@ def main(args):
|
||||
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)
|
||||
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 +35,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)
|
||||
|
||||
+445
-187
@@ -5,10 +5,14 @@ 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.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.models.mochi_hf.mochi_latents_utils import normalize_mochi_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from torch.utils.data import DataLoader
|
||||
@@ -25,32 +29,40 @@ 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.utils import convert_unet_state_dict_to_peft
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.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 peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import save_checkpoint, save_lora_checkpoint, resume_lora_training
|
||||
from fastvideo.utils.checkpoint import (
|
||||
save_checkpoint,
|
||||
save_lora_checkpoint,
|
||||
resume_lora_optimizer,
|
||||
)
|
||||
from fastvideo.utils.logging import main_print
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
# 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
|
||||
|
||||
|
||||
|
||||
|
||||
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,7 +73,13 @@ 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)
|
||||
@@ -70,6 +88,7 @@ def compute_density_for_timestep_sampling(
|
||||
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):
|
||||
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
|
||||
schedule_timesteps = noise_scheduler.timesteps.to(device)
|
||||
@@ -82,16 +101,35 @@ 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_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,
|
||||
):
|
||||
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,
|
||||
encoder_hidden_states,
|
||||
latents_attention_mask,
|
||||
encoder_attention_mask,
|
||||
) = next(loader)
|
||||
latents = normalize_mochi_dit_input(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,
|
||||
@@ -104,56 +142,58 @@ def train_one_step_mochi(transformer, optimizer, lr_scheduler,loader, noise_sche
|
||||
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
|
||||
# if rank<=0:
|
||||
# print("2222222222222222222222222222222222222222222222")
|
||||
# print(type(latents_attention_mask))
|
||||
# print(latents_attention_mask)
|
||||
with torch.autocast("cuda", torch.bfloat16):
|
||||
model_pred = transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# if rank<=0:
|
||||
# print("333333333333333333333333333333333333333333333333")
|
||||
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,45 +207,77 @@ 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
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
|
||||
f
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
weight_type = torch.float32 if args.master_weight_type == 'fp32' else torch.bfloat16
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
torch_dtype=torch.float32
|
||||
if args.master_weight_type == "fp32"
|
||||
else torch.bfloat16,
|
||||
)
|
||||
|
||||
|
||||
if args.use_lora:
|
||||
lora_config = LoraConfig(
|
||||
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 = MochiPipeline.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 = get_dit_fsdp_kwargs(
|
||||
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)
|
||||
|
||||
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
@@ -226,39 +298,44 @@ def main(args):
|
||||
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, 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,
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
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,93 +343,136 @@ 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)
|
||||
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 = (
|
||||
args.train_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" 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}")
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
# 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_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,
|
||||
)
|
||||
|
||||
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
|
||||
})
|
||||
progress_bar.set_postfix(
|
||||
{
|
||||
"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
|
||||
)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
save_checkpoint(transformer, optimizer, 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)
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
save_lora_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, args.max_train_steps
|
||||
)
|
||||
else:
|
||||
save_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
save_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, args.max_train_steps
|
||||
)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
|
||||
@@ -363,91 +483,195 @@ if __name__ == "__main__":
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
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("--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(
|
||||
"--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(
|
||||
"--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,10 +681,16 @@ 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",
|
||||
@@ -469,14 +699,42 @@ if __name__ == "__main__":
|
||||
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",
|
||||
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.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--Mochi_type",
|
||||
type=str,
|
||||
default="hf",
|
||||
help="Choose Mochi model between hf and genmo(original mochi).",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
|
||||
+135
-97
@@ -1,29 +1,42 @@
|
||||
# import
|
||||
# import
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
from fastvideo.utils.logging import main_print
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, StateDictType, FullStateDictConfig
|
||||
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.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):
|
||||
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
|
||||
|
||||
def save_checkpoint(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")
|
||||
save_file(cpu_state, weight_path)
|
||||
@@ -39,21 +52,30 @@ def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=Fals
|
||||
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,):
|
||||
|
||||
|
||||
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)
|
||||
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")
|
||||
@@ -62,8 +84,7 @@ def save_checkpoint_generator_discriminator(model, optimizer, discriminator, dis
|
||||
json.dump(config_dict, f, indent=4)
|
||||
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 +95,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")
|
||||
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)):
|
||||
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 +150,62 @@ 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")
|
||||
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 +216,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):
|
||||
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
|
||||
)
|
||||
MochiPipeline.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
|
||||
|
||||
@@ -10,11 +10,12 @@ from typing import Any, Tuple
|
||||
from torch import Tensor
|
||||
from torch.nn import Module
|
||||
|
||||
|
||||
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:
|
||||
@@ -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
|
||||
@@ -129,18 +129,16 @@ class SeqAllToAll4D(torch.autograd.Function):
|
||||
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 +146,9 @@ 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 +170,9 @@ 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 +200,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 +238,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 +252,83 @@ 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)
|
||||
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_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,16 +1,18 @@
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate.logging import get_logger
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.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,
|
||||
@@ -21,31 +23,40 @@ class T5dataset(Dataset):
|
||||
self.vae_debug = vae_debug
|
||||
with open(self.json_path, "r") as f:
|
||||
train_dataset = json.load(f)
|
||||
self.train_dataset = sorted(train_dataset, key=lambda x: x['latent_path'])
|
||||
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
|
||||
|
||||
def __getitem__(self, idx):
|
||||
caption = self.train_dataset[idx]['caption']
|
||||
filename = self.train_dataset[idx]['latent_path'].split('.')[0]
|
||||
length = self.train_dataset[idx]['length']
|
||||
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")
|
||||
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)
|
||||
|
||||
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
@@ -53,32 +64,40 @@ def main(args):
|
||||
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)
|
||||
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'],
|
||||
prompt=data["caption"],
|
||||
)
|
||||
if args.vae_debug:
|
||||
latents = data['latents']
|
||||
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")
|
||||
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)
|
||||
@@ -86,11 +105,11 @@ def main(args):
|
||||
if args.vae_debug:
|
||||
export_to_video(video[idx], video_path, fps=30)
|
||||
item = {}
|
||||
item['length'] = int(data['length'][idx])
|
||||
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]
|
||||
item["caption"] = data["caption"][idx]
|
||||
json_data.append(item)
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
@@ -99,19 +118,35 @@ def main(args):
|
||||
if local_rank == 0:
|
||||
# os.remove(latents_json_path)
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption.json"), 'w') as f:
|
||||
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
|
||||
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")
|
||||
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)
|
||||
|
||||
@@ -14,21 +14,30 @@ 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)
|
||||
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)
|
||||
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)
|
||||
sampler = DistributedSampler(
|
||||
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
|
||||
)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
@@ -36,29 +45,36 @@ def main(args):
|
||||
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")
|
||||
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']):
|
||||
latents = vae.encode(data["pixel_values"].to(encoder_device))[
|
||||
"latent_dist"
|
||||
].sample()
|
||||
for idx, video_path in enumerate(data["path"]):
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
|
||||
latent_path = os.path.join(
|
||||
args.output_dir, "latent", video_name + ".pt"
|
||||
)
|
||||
torch.save(latents[idx].to(torch.bfloat16), latent_path)
|
||||
item = {}
|
||||
item["length"] = latents[idx].shape[1]
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["caption"] = data['text'][idx]
|
||||
item["caption"] = data["text"][idx]
|
||||
json_data.append(item)
|
||||
print(f"{video_name} processed")
|
||||
dist.barrier()
|
||||
@@ -67,40 +83,61 @@ def main(args):
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
if local_rank == 0:
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), 'w') as f:
|
||||
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
|
||||
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(
|
||||
"--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("--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***."
|
||||
),
|
||||
)
|
||||
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)
|
||||
|
||||
@@ -13,11 +13,13 @@ 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."""
|
||||
|
||||
@@ -31,17 +33,20 @@ class DecordInit(object):
|
||||
results (dict): The resulting dict to be modified and passed
|
||||
to the next transform in pipeline.
|
||||
"""
|
||||
reader = decord.VideoReader(filename,
|
||||
ctx=self.ctx,
|
||||
num_threads=self.num_threads)
|
||||
reader = decord.VideoReader(
|
||||
filename, ctx=self.ctx, num_threads=self.num_threads
|
||||
)
|
||||
return reader
|
||||
|
||||
def __repr__(self):
|
||||
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,6 +55,7 @@ def pad_to_multiple(number, ds_stride):
|
||||
padding = ds_stride - remainder
|
||||
return number + padding
|
||||
|
||||
|
||||
class Collate:
|
||||
def __init__(self, args):
|
||||
self.batch_size = args.train_batch_size
|
||||
@@ -71,9 +77,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,13 +87,29 @@ 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)]
|
||||
assert len(batch_input_size) == self.batch_size
|
||||
@@ -98,13 +120,30 @@ class Collate:
|
||||
if len(count_dict) != 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 +154,61 @@ 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
|
||||
]
|
||||
[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]
|
||||
]
|
||||
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[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]
|
||||
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 +234,16 @@ 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 +257,70 @@ 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])
|
||||
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 +334,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,15 +354,24 @@ 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
|
||||
while index < len(lst):
|
||||
result.extend(lst[index:index + batch_size])
|
||||
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")
|
||||
@@ -1,23 +1,24 @@
|
||||
|
||||
import sys
|
||||
import pdb
|
||||
import os
|
||||
|
||||
|
||||
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,77 @@
|
||||
from accelerate.logging import get_logger
|
||||
import torch
|
||||
|
||||
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
|
||||
@@ -2,6 +2,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import os
|
||||
|
||||
|
||||
class COMM_INFO:
|
||||
def __init__(self):
|
||||
self.group = None
|
||||
@@ -10,8 +11,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,22 +23,29 @@ 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
|
||||
|
||||
@@ -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))
|
||||
|
||||
+136
-69
@@ -1,13 +1,14 @@
|
||||
|
||||
|
||||
from typing import Optional, Union, List
|
||||
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 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 fastvideo.models.mochi_hf.pipeline_mochi import (
|
||||
linear_quadratic_schedule,
|
||||
retrieve_timesteps,
|
||||
)
|
||||
from tqdm import tqdm
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers import (
|
||||
@@ -20,6 +21,8 @@ from diffusers.utils import export_to_video
|
||||
import os
|
||||
import wandb
|
||||
import gc
|
||||
|
||||
|
||||
def prepare_latents(
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
@@ -38,10 +41,10 @@ def prepare_latents(
|
||||
|
||||
shape = (batch_size, num_channels_latents, num_frames, height, width)
|
||||
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
|
||||
def sample_validation_video(
|
||||
transformer,
|
||||
vae,
|
||||
@@ -60,8 +63,8 @@ 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,
|
||||
):
|
||||
device = vae.device
|
||||
|
||||
@@ -70,7 +73,9 @@ 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_attention_mask = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask], dim=0
|
||||
)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
# TODO: Remove hardcore
|
||||
@@ -85,13 +90,14 @@ 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
|
||||
@@ -118,12 +124,16 @@ def sample_validation_video(
|
||||
# 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(
|
||||
@@ -133,16 +143,20 @@ def sample_validation_video(
|
||||
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,28 +164,35 @@ 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)
|
||||
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)
|
||||
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:
|
||||
@@ -180,87 +201,134 @@ def sample_validation_video(
|
||||
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,)
|
||||
|
||||
|
||||
|
||||
@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
|
||||
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")
|
||||
vae = AutoencoderKLMochi.from_pretrained(
|
||||
args.pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype
|
||||
).to("cuda")
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
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])
|
||||
num_embeds = len(
|
||||
[f for f in os.listdir(args.validation_prompt_dir) if "embed" in f]
|
||||
)
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
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]
|
||||
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]
|
||||
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")
|
||||
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)
|
||||
video_filenames.append(filename)
|
||||
|
||||
@@ -271,4 +339,3 @@ def log_validation(args, transformer, device, weight_dtype, global_step, schedu
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
|
||||
|
||||
@@ -1,24 +1,5 @@
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
path = "data/outputs/BW_Testrun/checkpoint-0/config.json"
|
||||
|
||||
with open(path, 'r') as f:
|
||||
data = json.load(f)
|
||||
|
||||
# save with indent
|
||||
with open(path, 'w') as f:
|
||||
json.dump(data, f, indent=4)
|
||||
+1
-5
@@ -21,12 +21,8 @@ dependencies = [
|
||||
"timm==1.0.11", "torchdiffeq==0.2.4", "torchmetrics==1.5.1", "tqdm==4.66.5", "urllib3==2.2.0", "uvicorn==0.32.0",
|
||||
"scikit-video==1.1.11", "imageio-ffmpeg==0.5.1", "sentencepiece==0.2.0", "beautifulsoup4==4.12.3", "ftfy==6.3.0",
|
||||
"moviepy==1.0.3", "wandb==0.18.5", "tensorboard==2.18.0", "pydantic==2.9.2", "gradio==5.3.0", "huggingface_hub==0.26.1", "protobuf==5.28.3",
|
||||
"watch", "gpustat", "peft==0.13.2", "liger_kernel==0.4.1", "einops==0.8.0", "wheel==0.44.0"
|
||||
]
|
||||
"watch", "gpustat", "peft==0.13.2", "liger_kernel==0.4.1", "einops==0.8.0", "wheel==0.44.0"]
|
||||
|
||||
[project.optional-dependencies]
|
||||
train = ["deepspeed==0.15.3"]
|
||||
dev = ["mypy==1.8.0"]
|
||||
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
|
||||
+34
-14
@@ -1,19 +1,39 @@
|
||||
from huggingface_hub import snapshot_download, hf_hub_download
|
||||
import argparse
|
||||
# set args for repo_id, local_dir, repo_type,
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser(description='Download a dataset or model from the Hugging Face Hub')
|
||||
parser.add_argument('--repo_id', type=str, help='The ID of the repository to download')
|
||||
parser.add_argument('--local_dir', type=str, help='The local directory to download the repository to')
|
||||
parser.add_argument('--repo_type', type=str, help='The type of repository to download (dataset or model)')
|
||||
parser.add_argument('--file_name', type=str, help='The file name to download')
|
||||
# set args for repo_id, local_dir, repo_type,
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Download a dataset or model from the Hugging Face Hub"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--repo_id", type=str, help="The ID of the repository to download"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--local_dir",
|
||||
type=str,
|
||||
help="The local directory to download the repository to",
|
||||
)
|
||||
parser.add_argument(di
|
||||
"--repo_type",
|
||||
type=str,
|
||||
help="The type of repository to download (dataset or model)",
|
||||
)
|
||||
parser.add_argument("--file_name", type=str, help="The file name to download")
|
||||
args = parser.parse_args()
|
||||
if args.file_name:
|
||||
hf_hub_download(repo_id=args.repo_id, filename=args.file_name, repo_type=args.repo_type, local_dir=args.local_dir)
|
||||
else:
|
||||
snapshot_download(repo_id=args.repo_id,
|
||||
local_dir=args.local_dir,
|
||||
repo_type=args.repo_type,
|
||||
local_dir_use_symlinks=False,
|
||||
resume_download=True)
|
||||
hf_hub_download(
|
||||
repo_id=args.repo_id,
|
||||
filename=args.file_name,
|
||||
repo_type=args.repo_type,
|
||||
local_dir=args.local_dir,
|
||||
)
|
||||
else:
|
||||
snapshot_download(
|
||||
repo_id=args.repo_id,
|
||||
local_dir=args.local_dir,
|
||||
repo_type=args.repo_type,
|
||||
local_dir_use_symlinks=False,
|
||||
resume_download=True,
|
||||
)
|
||||
|
||||
@@ -1,56 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export NCCL_DEBUG=INFO
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path data/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "data/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--uncond_prompt_dir "data/validation_embeddings/uncond_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 250\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="data/outputs/lq_euler_50_thresh_0.025"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale 4.5 \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
gsutil cp data/outputs/lq_euler_50/checkpoint-4000 gs://vid_gen/runlong_temp_folder_for_pandas70m_debugging/fastvid/lq_euler_50/checkpoint-4000
|
||||
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 4\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.05_bs32"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "2.5,3.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.05
|
||||
@@ -1,45 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export NCCL_DEBUG=INFO
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path data/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "data/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="data/outputs/lq_euler_50_thresh0.05_lr_1e-7"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "2.5,3.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.05
|
||||
@@ -1,44 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export NCCL_DEBUG=INFO
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path data/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "data/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/shift1_euler_50"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--shift 1 \
|
||||
--validation_guidance_scale "2.5,3.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
@@ -1,51 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_cfg_4.5"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--distill_cfg 4.5
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.15_lrg_0.75"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.15 \
|
||||
--linear_range 0.75
|
||||
|
||||
@@ -1,54 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_ema_0.95_decay0.0"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--use_ema \
|
||||
--ema_decay 0.95 \
|
||||
--weight_decay 0.0
|
||||
|
||||
|
||||
@@ -1,57 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=lq_euler_50_thres0.1_lrg_0.75_cfg0.0
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--not_apply_cfg_solver
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=lq_euler_50_thres0.1_lrg_0.75_cfg0.0
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--not_apply_cfg_solver
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=shift16_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--shift 16 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=shift16_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--shift 16 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_shift16_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--shift 16 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_shift16_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--shift 16 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
@@ -1,48 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export NCCL_DEBUG=INFO
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path data/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "data/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--uncond_prompt_dir "data/validation_embeddings/uncond_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 250\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="data/outputs/lq_euler_50_thresh0.05"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale 4.5 \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.05
|
||||
|
||||
gsutil cp data/outputs/lq_euler_50_thresh0.05/checkpoint-4000 gs://vid_gen/runlong_temp_folder_for_pandas70m_debugging/fastvid/lq_euler_50_thresh0.05/checkpoint-4000
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_shift12_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--shift 12 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_shift12_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--shift 12 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50
|
||||
|
||||
@@ -1,54 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_lq_euler_50_thresh0.1_lrg_0.75
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75
|
||||
|
||||
@@ -1,54 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_lq_euler_50_thresh0.1_lrg_0.75
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=lq_euler_50_thresh0.1_lrg_0.75_phase1
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 8\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=lq_euler_50_thresh0.1_lrg_0.75_phase1
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=lq_euler_50_thres0.1_lrg_0.75_bs_64
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 8 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1
|
||||
@@ -1,52 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_shift16_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=5e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1
|
||||
@@ -1,51 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
CACHE_DIR=/data/.cache
|
||||
EXPERIMENT=4step_infer_shift16_euler_50
|
||||
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
|
||||
export WANDB_DIR=$DATA_DIR/wandb/
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir $CACHE_DIR \
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 8\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--learning_rate=5e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 139 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1
|
||||
@@ -1,49 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=data
|
||||
IP=172.23.30.16
|
||||
|
||||
torchrun --nnodes 4 --nproc_per_node 4\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/shift1_euler_50_0.75_phase1"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--shift 1 \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
|
||||
|
||||
DATA_DIR=data
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--use_ema \
|
||||
--ema_decay 0.95
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
|
||||
|
||||
DATA_DIR=data
|
||||
IP=172.23.30.16
|
||||
|
||||
torchrun --nnodes 4 --nproc_per_node 4\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95_cfg4.5"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--use_ema \
|
||||
--ema_decay 0.95 \
|
||||
--distill_cfg 4.5
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.95"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--use_ema \
|
||||
--ema_decay 0.95
|
||||
|
||||
@@ -1,54 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95_cfg6.0"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--use_ema \
|
||||
--ema_decay 0.95 \
|
||||
--distill_cfg 6.0
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export NCCL_DEBUG=INFO
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path data/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "data/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--uncond_prompt_dir "data/validation_embeddings/uncond_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="data/outputs/shift8_euler_100"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--shift 8.0 \
|
||||
--validation_guidance_scale 4.5
|
||||
|
||||
|
||||
gsutil cp data/outputs/shift8_euler_100/checkpoint-4000 gs://vid_gen/runlong_temp_folder_for_pandas70m_debugging/fastvid/shift8_euler_100/checkpoint-4000
|
||||
@@ -1,54 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.98_cfg4.5"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--use_ema \
|
||||
--ema_decay 0.98 \
|
||||
--distill_cfg 4.5
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=3e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.1_lrg_0.75_phase1_lr_3e-7"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
@@ -1,54 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.15_lrg_0.75_phase1_ema0.95_cfg4.5"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.15 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--use_ema \
|
||||
--ema_decay 0.95 \
|
||||
--distill_cfg 4.5
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.142.161
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=500\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--ema_decay 0.999\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_linear_range_0.75_cfg_7.0"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--distill_cfg 6.0 \
|
||||
--multi_phased_distill_schedule "4000-8" \
|
||||
|
||||
@@ -1,46 +0,0 @@
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
|
||||
|
||||
DATA_DIR=data
|
||||
IP=172.23.30.16
|
||||
|
||||
torchrun --nnodes 4 --nproc_per_node 4\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=125\
|
||||
--validation_steps 125\
|
||||
--validation_sampling_steps "8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_reproduce"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-8" \
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/mochi\
|
||||
--cache_dir "data/.cache"\
|
||||
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 28\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 2\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=4000\
|
||||
--learning_rate=5e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1_lr5e-6"\
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale "0.5,1.5,2.5" \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user