Add Animate-X (see details)
https://lucaria-academy.github.io/Animate-X/ https://github.com/antgroup/animate-x @ fdc80909f911d8e487cb4e6847f2e4c6501b62af
@@ -0,0 +1,18 @@
|
||||
*.pkl
|
||||
*.pt
|
||||
*.mov
|
||||
*.pth
|
||||
*.mov
|
||||
*.npz
|
||||
*.npy
|
||||
*.boj
|
||||
*.onnx
|
||||
*.tar
|
||||
*.bin
|
||||
cache*
|
||||
.DS_Store
|
||||
*DS_Store
|
||||
outputs/
|
||||
**/__pycache__
|
||||
***/__pycache__
|
||||
*/__pycache__
|
||||
@@ -0,0 +1,201 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"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 [yyyy] [name of copyright owner]
|
||||
|
||||
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.
|
||||
@@ -0,0 +1,181 @@
|
||||
|
||||
<p align="center">
|
||||
|
||||
<h2 align="center">Animate-X: Universal Character Image Animation with Enhanced Motion Representation</h2>
|
||||
<p align="center">
|
||||
<a href=""><strong>Shuai Tan</strong></a>
|
||||
·
|
||||
<a href="https://scholar.google.com/citations?user=BwdpTiQAAAAJ"><strong>Biao Gong</strong></a><sup>†</sup>
|
||||
·
|
||||
<a href="https://scholar.google.com/citations?user=cQbXvkcAAAAJ"><strong>Xiang Wang</strong></a>
|
||||
·
|
||||
<a href="https://scholar.google.com/citations?user=ZO3OQ-8AAAAJ"><strong>Shiwei Zhang</strong></a>
|
||||
<br>
|
||||
<a href="https://openreview.net/profile?id=~DanDan_Zheng1"><strong>Dandan Zheng</strong></a>
|
||||
·
|
||||
<a href="https://scholar.google.com.hk/citations?user=S8FmqTUAAAAJ"><strong>Ruobing Zheng</strong></a>
|
||||
·
|
||||
<a href="https://scholar.google.com/citations?user=hMDQifQAAAAJ"><strong>Kecheng Zheng</strong></a>
|
||||
·
|
||||
<a href="https://openreview.net/profile?id=~Jingdong_Chen1"><strong>Jingdong Chen</strong></a>
|
||||
·
|
||||
<a href="https://openreview.net/profile?id=~Ming_Yang2"><strong>Ming Yang</strong></a>
|
||||
<br>
|
||||
<br>
|
||||
<a href="https://arxiv.org/abs/2410.10306"><img src='https://img.shields.io/badge/arXiv-Animate--X-red' alt='Paper PDF'></a>
|
||||
<a href='https://lucaria-academy.github.io/Animate-X/'><img src='https://img.shields.io/badge/Project_Page-Animate--X-blue' alt='Project Page'></a>
|
||||
<a href='https://mp.weixin.qq.com/s/vDR4kPLqnCUwfPiBNKKV9A'><img src='https://badges.aleen42.com/src/wechat.svg'></a>
|
||||
<a href='https://huggingface.co/Shuaishuai0219/Animate-X'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20HuggingFace-Model-yellow'></a>
|
||||
<br>
|
||||
<b></a>Ant Group | </a>Tongyi Lab </b>
|
||||
<br>
|
||||
</p>
|
||||
</p>
|
||||
|
||||
This repository is the official implementation of paper "Animate-X: Universal Character Image Animation with Enhanced Motion Representation". Animate-X is a universal animation framework based on latent diffusion models for various character types (collectively named X), including anthropomorphic characters.
|
||||
<table align="center">
|
||||
<tr>
|
||||
<td>
|
||||
<img src="https://github.com/user-attachments/assets/fb2f4396-341f-4206-8d70-44d8b034f810">
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
|
||||
## 📌 Updates
|
||||
* [2024.12.20] 🔥 We release our [Animate-X](https://github.com/antgroup/animate-x) inference codes.
|
||||
* [2024.12.10] 🔥 We release our [Animate-X CKPT](https://huggingface.co/Shuaishuai0219/Animate-X) checkpoints.
|
||||
* [2024.10.14] 🔥 Our [paper](https://arxiv.org/abs/2410.10306) is in public on arxiv.
|
||||
|
||||
|
||||
|
||||
<!-- <video controls loop src="https://cloud.video.taobao.com/vod/vs4L24EAm6IQ5zM3SbN5AyHCSqZIXwmuobrzqNztMRM.mp4" muted="false"></video> -->
|
||||
|
||||
## 🌄 Gallery
|
||||
### Introduction
|
||||
<table class="center">
|
||||
<tr>
|
||||
<td width=47% style="border: none">
|
||||
<video controls loop src="https://github.com/user-attachments/assets/085b70c4-cb68-4ac1-b45f-ed7f1c75bd5c" muted="false"></video>
|
||||
</td>
|
||||
<td width=53% style="border: none">
|
||||
<video controls loop src="https://github.com/user-attachments/assets/f6275c0d-fbca-43b4-b6d6-cf095723729e" muted="false"></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### Animations produced by Animate-X
|
||||
<table class="center">
|
||||
<tr>
|
||||
<td width=50% style="border: none">
|
||||
<video controls loop src="https://github.com/user-attachments/assets/732a3445-2054-4e7b-9c2d-9db21c39771e" muted="false"></video>
|
||||
</td>
|
||||
<td width=50% style="border: none">
|
||||
<video controls loop src="https://github.com/user-attachments/assets/f25af02c-e5be-4cab-ae64-c9e0b392643a" muted="false"></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
|
||||
|
||||
|
||||
## 🚀 Installation
|
||||
Install with `conda`:
|
||||
```bash
|
||||
conda env create -f environment.yaml
|
||||
conda activate animate-x
|
||||
```
|
||||
|
||||
|
||||
## 🚀 Download Checkpoints
|
||||
Download Animate-X [checkpoints](https://huggingface.co/Shuaishuai0219/Animate-X) and put all files in `checkpoints` dir, which should be like:
|
||||
```
|
||||
./checkpoints/
|
||||
|---- animate-x.pth
|
||||
|---- dw-ll_ucoco_384.onnx
|
||||
|---- open_clip_pytorch_model.bin
|
||||
|---- v2-1_512-ema-pruned.ckpt
|
||||
└---- yolox_l.onnx
|
||||
```
|
||||
|
||||
## 💡 Inference
|
||||
|
||||
The default inputs are a image (.jpg) and a dance video (.mp4). The default output is a 32-frame video (.mp4) with 768x512 resolution, which will be saved in `./results` dir.
|
||||
|
||||
1. pre-process the video.
|
||||
```bash
|
||||
python process_data.py \
|
||||
--source_video_paths data/videos \
|
||||
--saved_pose_dir data/saved_pkl \
|
||||
--saved_pose data/saved_pose \
|
||||
--saved_frame_dir data/saved_frames
|
||||
```
|
||||
2. run Animate-X.
|
||||
```bash
|
||||
python inference.py --cfg configs/Animate_X_infer.yaml
|
||||
```
|
||||
|
||||
|
||||
Some key parameters in the `.yaml` configuration file are described as follows. For example, users can adjust the `max_frames` or `sampling interval` of the dance video to generate videos of varying durations or speeds.
|
||||
- `max_frames`: Number of frames (default as 32) in the generated video (fps: 8).
|
||||
- If you want to generage longer video with more frames, you should modify
|
||||
- `max_frames` as the number of frames
|
||||
- `seq_len` in `UNet` as the number of frames + 1
|
||||
- We take 96 frames as an example, and the config should be:
|
||||
```python
|
||||
{
|
||||
max_frames: 96 # 1. modify `max_frames` as the number of frames (e.g. 96)
|
||||
......
|
||||
UNet: {
|
||||
......
|
||||
'use_sim_mask': False,
|
||||
'seq_len': 97, # 2. modify `seq_len` in `UNet` as the number of frames + 1 (e.g. 97 = 96 + 1)
|
||||
}
|
||||
}
|
||||
```
|
||||
- `round`: The number of times each test case is generated.
|
||||
- `test_list_path`: The input paths for all test cases, for example:
|
||||
```python
|
||||
[
|
||||
[2, "data/images/1.jpg", "data/saved_pose/dance_1","data/saved_frames/dance_1","data/saved_pkl/dance_1.pkl", 14],
|
||||
[2, "data/images/4.png", "data/saved_pose/dance_1","data/saved_frames/dance_1","data/saved_pkl/dance_1.pkl", 14],
|
||||
......
|
||||
]
|
||||
```
|
||||
- `2` indicates that 1 frame is sampled from every 2 frames of the reference dance video to be used as input for the model.
|
||||
- `"data/images/1.jpg"` indicates the path to the reference image.
|
||||
- `"data/saved_pose/dance_1"` indicates the path to the saved pose images. (output by `process_data.py`, $I^p$, keypoints visualization)
|
||||
- `"data/saved_frames/dance_1"` indicates the path to the saved frames from the driven video. (output by `process_data.py`)
|
||||
- `"data/saved_pkl/dance_1.pkl"` indicates the path to the saved pose keypoints. (output by `process_data.py`, $p^d$, DWPose)
|
||||
- `14` indicates the random seed.
|
||||
- `log_dir`: path to the generated animation videos, e.g., `./results`.
|
||||
|
||||
|
||||
**✔ Some tips**:
|
||||
|
||||
> Although Animate-x does not rely on strict pose alignment and we did not perform any manual alignment operations for all the results in the paper, we cannot guarantee that all cases are perfect. Therefore, users can perform handmade pose alignment operations themselves, e.g, applying the overall x/y translation and scaling on the pose skeleton of each frame to align with the position of the subject in the reference image. (put in `data/saved_pose`)
|
||||
|
||||
|
||||
## 📧 Acknowledgement
|
||||
Our implementation is based on [UniAnimate](https://github.com/ali-vilab/UniAnimate), [MimicMotion](https://github.com/Tencent/MimicMotion), and [MusePose](https://github.com/TMElyralab/MusePose). Thanks for their remarkable contribution and released code! If we missed any open-source projects or related articles, we would like to complement the acknowledgement of this specific work immediately.
|
||||
|
||||
## ⚖ License
|
||||
This repository is released under the Apache-2.0 license as found in the [LICENSE](LICENSE) file.
|
||||
|
||||
## 📚 Citation
|
||||
If you find this codebase useful for your research, please use the following entry.
|
||||
```BibTeX
|
||||
@article{AnimateX2025,
|
||||
title={Animate-X: Universal Character Image Animation with Enhanced Motion Representation},
|
||||
author={Tan, Shuai and Gong, Biao and Wang, Xiang and Zhang, Shiwei and Zheng, Dandan and Zheng, Ruobing and Zheng, Kecheng and Chen, Jingdong and Yang, Ming},
|
||||
journal={arXiv preprint arXiv:2410.10306},
|
||||
year={2025}
|
||||
}
|
||||
|
||||
@article{Mimir2025,
|
||||
title={Mimir: Improving Video Diffusion Models for Precise Text Understanding},
|
||||
author={Tan, Shuai and Gong, Biao and Feng, Yutong and Zheng, Kecheng and Zheng, Dandan and Shi, Shuwei and Shen, Yujun and Chen, Jingdong and Yang, Ming},
|
||||
journal={arXiv preprint arXiv:2412.03085},
|
||||
year={2025}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1 @@
|
||||
from .inference_animate_x_entrance import *
|
||||
@@ -0,0 +1,207 @@
|
||||
import torch
|
||||
import logging
|
||||
import os.path as osp
|
||||
from datetime import datetime
|
||||
from easydict import EasyDict
|
||||
import os
|
||||
|
||||
cfg = EasyDict(__name__='Config: VideoLDM Decoder')
|
||||
|
||||
# -------------------------------distributed training--------------------------
|
||||
pmi_world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
gpus_per_machine = torch.cuda.device_count()
|
||||
world_size = pmi_world_size * gpus_per_machine
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
# ---------------------------Dataset Parameter---------------------------------
|
||||
cfg.mean = [0.5, 0.5, 0.5]
|
||||
cfg.std = [0.5, 0.5, 0.5]
|
||||
cfg.max_words = 1000
|
||||
cfg.num_workers = 8
|
||||
cfg.prefetch_factor = 2
|
||||
|
||||
# PlaceHolder
|
||||
cfg.resolution = [448, 256]
|
||||
cfg.vit_out_dim = 1024
|
||||
cfg.vit_resolution = 336
|
||||
cfg.depth_clamp = 10.0
|
||||
cfg.misc_size = 384
|
||||
cfg.depth_std = 20.0
|
||||
|
||||
cfg.save_fps = 8
|
||||
|
||||
cfg.frame_lens = [32, 32, 32, 1]
|
||||
cfg.sample_fps = [4, ]
|
||||
cfg.vid_dataset = {
|
||||
'type': 'VideoBaseDataset',
|
||||
'data_list': [],
|
||||
'max_words': cfg.max_words,
|
||||
'resolution': cfg.resolution}
|
||||
cfg.img_dataset = {
|
||||
'type': 'ImageBaseDataset',
|
||||
'data_list': ['laion_400m',],
|
||||
'max_words': cfg.max_words,
|
||||
'resolution': cfg.resolution}
|
||||
|
||||
cfg.batch_sizes = {
|
||||
str(1):256,
|
||||
str(4):4,
|
||||
str(8):4,
|
||||
str(16):4}
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
# ---------------------------Mode Parameters-----------------------------------
|
||||
# Diffusion
|
||||
cfg.Diffusion = {
|
||||
'type': 'DiffusionDDIM',
|
||||
'schedule': 'cosine', # cosine
|
||||
'schedule_param': {
|
||||
'num_timesteps': 1000,
|
||||
'cosine_s': 0.008,
|
||||
'zero_terminal_snr': True,
|
||||
},
|
||||
'mean_type': 'v', # [v, eps]
|
||||
'loss_type': 'mse',
|
||||
'var_type': 'fixed_small',
|
||||
'rescale_timesteps': False,
|
||||
'noise_strength': 0.1,
|
||||
'ddim_timesteps': 50
|
||||
}
|
||||
cfg.ddim_timesteps = 50 # official: 250
|
||||
cfg.use_div_loss = False
|
||||
# classifier-free guidance
|
||||
cfg.p_zero = 0.9
|
||||
cfg.guide_scale = 3.0
|
||||
|
||||
# clip vision encoder
|
||||
cfg.vit_mean = [0.48145466, 0.4578275, 0.40821073]
|
||||
cfg.vit_std = [0.26862954, 0.26130258, 0.27577711]
|
||||
|
||||
# sketch
|
||||
cfg.sketch_mean = [0.485, 0.456, 0.406]
|
||||
cfg.sketch_std = [0.229, 0.224, 0.225]
|
||||
# cfg.misc_size = 256
|
||||
cfg.depth_std = 20.0
|
||||
cfg.depth_clamp = 10.0
|
||||
cfg.hist_sigma = 10.0
|
||||
|
||||
# Model
|
||||
cfg.scale_factor = 0.18215
|
||||
cfg.use_checkpoint = True
|
||||
cfg.use_sharded_ddp = False
|
||||
cfg.use_fsdp = False
|
||||
cfg.use_fp16 = True
|
||||
cfg.temporal_attention = True
|
||||
|
||||
cfg.UNet = {
|
||||
'type': 'UNetSD',
|
||||
'in_dim': 4,
|
||||
'dim': 320,
|
||||
'y_dim': cfg.vit_out_dim,
|
||||
'context_dim': 1024,
|
||||
'out_dim': 8,
|
||||
'dim_mult': [1, 2, 4, 4],
|
||||
'num_heads': 8,
|
||||
'head_dim': 64,
|
||||
'num_res_blocks': 2,
|
||||
'attn_scales': [1 / 1, 1 / 2, 1 / 4],
|
||||
'dropout': 0.1,
|
||||
'temporal_attention': cfg.temporal_attention,
|
||||
'temporal_attn_times': 1,
|
||||
'use_checkpoint': cfg.use_checkpoint,
|
||||
'use_fps_condition': False,
|
||||
'use_sim_mask': False
|
||||
}
|
||||
|
||||
# auotoencoder from stabel diffusion
|
||||
cfg.guidances = []
|
||||
cfg.auto_encoder = {
|
||||
'type': 'AutoencoderKL',
|
||||
'ddconfig': {
|
||||
'double_z': True,
|
||||
'z_channels': 4,
|
||||
'resolution': 256,
|
||||
'in_channels': 3,
|
||||
'out_ch': 3,
|
||||
'ch': 128,
|
||||
'ch_mult': [1, 2, 4, 4],
|
||||
'num_res_blocks': 2,
|
||||
'attn_resolutions': [],
|
||||
'dropout': 0.0,
|
||||
'video_kernel_size': [3, 1, 1]
|
||||
},
|
||||
'embed_dim': 4,
|
||||
'pretrained': 'models/v2-1_512-ema-pruned.ckpt'
|
||||
}
|
||||
# clip embedder
|
||||
cfg.embedder = {
|
||||
'type': 'FrozenOpenCLIPEmbedder',
|
||||
'layer': 'penultimate',
|
||||
'pretrained': 'models/open_clip_pytorch_model.bin'
|
||||
}
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
# ---------------------------Training Settings---------------------------------
|
||||
# training and optimizer
|
||||
cfg.ema_decay = 0.9999
|
||||
cfg.num_steps = 600000
|
||||
cfg.lr = 5e-5
|
||||
cfg.weight_decay = 0.0
|
||||
cfg.betas = (0.9, 0.999)
|
||||
cfg.eps = 1.0e-8
|
||||
cfg.chunk_size = 16
|
||||
cfg.decoder_bs = 8
|
||||
cfg.alpha = 0.7
|
||||
cfg.save_ckp_interval = 1000
|
||||
|
||||
# scheduler
|
||||
cfg.warmup_steps = 10
|
||||
cfg.decay_mode = 'cosine'
|
||||
|
||||
# acceleration
|
||||
cfg.use_ema = True
|
||||
if world_size<2:
|
||||
cfg.use_ema = False
|
||||
cfg.load_from = None
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
# ----------------------------Pretrain Settings---------------------------------
|
||||
cfg.Pretrain = {
|
||||
'type': 'pretrain_specific_strategies',
|
||||
'fix_weight': False,
|
||||
'grad_scale': 0.2,
|
||||
'resume_checkpoint': 'models/jiuniu_0267000.pth',
|
||||
'sd_keys_path': 'models/stable_diffusion_image_key_temporal_attention_x1.json',
|
||||
}
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
# -----------------------------Visual-------------------------------------------
|
||||
# Visual videos
|
||||
cfg.viz_interval = 1000
|
||||
cfg.resume_checkpoint = ""
|
||||
cfg.visual_train = {
|
||||
'type': 'VisualTrainTextImageToVideo',
|
||||
}
|
||||
cfg.visual_inference = {
|
||||
'type': 'VisualGeneratedVideos',
|
||||
}
|
||||
cfg.inference_list_path = ''
|
||||
|
||||
# logging
|
||||
cfg.log_interval = 100
|
||||
|
||||
### Default log_dir
|
||||
cfg.log_dir = 'outputs/'
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
# ---------------------------Others--------------------------------------------
|
||||
# seed
|
||||
cfg.seed = 8888
|
||||
cfg.negative_prompt = 'Distorted, discontinuous, Ugly, blurry, low resolution, motionless, static, disfigured, disconnected limbs, Ugly faces, incomplete arms'
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .diffusion_ddim import *
|
||||
@@ -0,0 +1,880 @@
|
||||
import torch
|
||||
import math
|
||||
|
||||
from utils.registry_class import DIFFUSION
|
||||
from .schedules import beta_schedule, sigma_schedule
|
||||
from typing import Callable, List, Optional
|
||||
import numpy as np
|
||||
|
||||
def _i(tensor, t, x):
|
||||
r"""Index tensor using t and format the output according to x.
|
||||
"""
|
||||
if tensor.device != x.device:
|
||||
tensor = tensor.to(x.device)
|
||||
shape = (x.size(0), ) + (1, ) * (x.ndim - 1)
|
||||
return tensor[t].view(shape).to(x)
|
||||
|
||||
@DIFFUSION.register_class()
|
||||
class DiffusionDDIMSR(object):
|
||||
def __init__(self, reverse_diffusion, forward_diffusion, **kwargs):
|
||||
from .diffusion_gauss import GaussianDiffusion
|
||||
self.reverse_diffusion = GaussianDiffusion(sigmas=sigma_schedule(reverse_diffusion.schedule, **reverse_diffusion.schedule_param),
|
||||
prediction_type=reverse_diffusion.mean_type)
|
||||
self.forward_diffusion = GaussianDiffusion(sigmas=sigma_schedule(forward_diffusion.schedule, **forward_diffusion.schedule_param),
|
||||
prediction_type=forward_diffusion.mean_type)
|
||||
|
||||
|
||||
@DIFFUSION.register_class()
|
||||
class DiffusionDPM(object):
|
||||
def __init__(self, forward_diffusion, **kwargs):
|
||||
from .diffusion_gauss import GaussianDiffusion
|
||||
self.forward_diffusion = GaussianDiffusion(sigmas=sigma_schedule(forward_diffusion.schedule, **forward_diffusion.schedule_param),
|
||||
prediction_type=forward_diffusion.mean_type)
|
||||
|
||||
|
||||
@DIFFUSION.register_class()
|
||||
class DiffusionDDIM(object):
|
||||
def __init__(self,
|
||||
schedule='linear_sd',
|
||||
schedule_param={},
|
||||
mean_type='eps',
|
||||
var_type='learned_range',
|
||||
loss_type='mse',
|
||||
epsilon = 1e-12,
|
||||
rescale_timesteps=False,
|
||||
noise_strength=0.0,
|
||||
**kwargs):
|
||||
|
||||
assert mean_type in ['x0', 'x_{t-1}', 'eps', 'v']
|
||||
assert var_type in ['learned', 'learned_range', 'fixed_large', 'fixed_small']
|
||||
assert loss_type in ['mse', 'rescaled_mse', 'kl', 'rescaled_kl', 'l1', 'rescaled_l1','charbonnier']
|
||||
|
||||
betas = beta_schedule(schedule, **schedule_param)
|
||||
assert min(betas) > 0 and max(betas) <= 1
|
||||
|
||||
if not isinstance(betas, torch.DoubleTensor):
|
||||
betas = torch.tensor(betas, dtype=torch.float64)
|
||||
|
||||
self.betas = betas
|
||||
self.num_timesteps = len(betas)
|
||||
self.mean_type = mean_type # eps
|
||||
self.var_type = var_type # 'fixed_small'
|
||||
self.loss_type = loss_type # mse
|
||||
self.epsilon = epsilon # 1e-12
|
||||
self.rescale_timesteps = rescale_timesteps # False
|
||||
self.noise_strength = noise_strength # 0.0
|
||||
|
||||
# alphas
|
||||
alphas = 1 - self.betas
|
||||
self.alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
self.alphas_cumprod_prev = torch.cat([alphas.new_ones([1]), self.alphas_cumprod[:-1]])
|
||||
self.alphas_cumprod_next = torch.cat([self.alphas_cumprod[1:], alphas.new_zeros([1])])
|
||||
|
||||
# q(x_t | x_{t-1})
|
||||
self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)
|
||||
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)
|
||||
self.log_one_minus_alphas_cumprod = torch.log(1.0 - self.alphas_cumprod)
|
||||
self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod)
|
||||
self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod - 1)
|
||||
|
||||
# q(x_{t-1} | x_t, x_0)
|
||||
self.posterior_variance = betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
self.posterior_log_variance_clipped = torch.log(self.posterior_variance.clamp(1e-20))
|
||||
self.posterior_mean_coef1 = betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - self.alphas_cumprod)
|
||||
|
||||
|
||||
def sample_loss(self, x0, noise=None):
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x0)
|
||||
if self.noise_strength > 0:
|
||||
b, c, f, _, _= x0.shape
|
||||
offset_noise = torch.randn(b, c, f, 1, 1, device=x0.device)
|
||||
noise = noise + self.noise_strength * offset_noise
|
||||
return noise
|
||||
|
||||
|
||||
def q_sample(self, x0, t, noise=None):
|
||||
r"""Sample from q(x_t | x_0).
|
||||
"""
|
||||
# noise = torch.randn_like(x0) if noise is None else noise
|
||||
noise = self.sample_loss(x0, noise)
|
||||
return _i(self.sqrt_alphas_cumprod, t, x0) * x0 + \
|
||||
_i(self.sqrt_one_minus_alphas_cumprod, t, x0) * noise
|
||||
|
||||
def q_mean_variance(self, x0, t):
|
||||
r"""Distribution of q(x_t | x_0).
|
||||
"""
|
||||
mu = _i(self.sqrt_alphas_cumprod, t, x0) * x0
|
||||
var = _i(1.0 - self.alphas_cumprod, t, x0)
|
||||
log_var = _i(self.log_one_minus_alphas_cumprod, t, x0)
|
||||
return mu, var, log_var
|
||||
|
||||
def q_posterior_mean_variance(self, x0, xt, t):
|
||||
r"""Distribution of q(x_{t-1} | x_t, x_0).
|
||||
"""
|
||||
mu = _i(self.posterior_mean_coef1, t, xt) * x0 + _i(self.posterior_mean_coef2, t, xt) * xt
|
||||
var = _i(self.posterior_variance, t, xt)
|
||||
log_var = _i(self.posterior_log_variance_clipped, t, xt)
|
||||
return mu, var, log_var
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None):
|
||||
r"""Sample from p(x_{t-1} | x_t).
|
||||
- condition_fn: for classifier-based guidance (guided-diffusion).
|
||||
- guide_scale: for classifier-free guidance (glide/dalle-2).
|
||||
"""
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
mu, var, log_var, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
|
||||
# random sample (with optional conditional function)
|
||||
noise = torch.randn_like(xt)
|
||||
mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1))) # no noise when t == 0
|
||||
if condition_fn is not None:
|
||||
grad = condition_fn(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
mu = mu.float() + var * grad.float()
|
||||
xt_1 = mu + mask * torch.exp(0.5 * log_var) * noise
|
||||
return xt_1, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None):
|
||||
r"""Sample from p(x_{t-1} | x_t) p(x_{t-2} | x_{t-1}) ... p(x_0 | x_1).
|
||||
"""
|
||||
# prepare input
|
||||
b = noise.size(0)
|
||||
xt = noise
|
||||
|
||||
# diffusion process
|
||||
for step in torch.arange(self.num_timesteps).flip(0):
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.p_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale)
|
||||
return xt
|
||||
|
||||
def p_mean_variance(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None):
|
||||
r"""Distribution of p(x_{t-1} | x_t).
|
||||
"""
|
||||
# print("============")
|
||||
# print(model_kwargs) # 'pose_embedding'
|
||||
# print("============")
|
||||
# predict distribution
|
||||
# print("=============p_mean_variance============")
|
||||
# print("=============xt.shape============", xt.shape)
|
||||
# print("=============t.shape============", t.shape)
|
||||
if guide_scale is None:
|
||||
out = model(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
else:
|
||||
# classifier-free guidance
|
||||
# (model_kwargs[0]: conditional kwargs; model_kwargs[1]: non-conditional kwargs)
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
|
||||
|
||||
# print("===============model_kwargs[0]==============")
|
||||
|
||||
# print(model_kwargs[0])
|
||||
|
||||
# print("===============model_kwargs[0]==============")
|
||||
|
||||
y_out = model(xt, self._scale_timesteps(t), **model_kwargs[0])
|
||||
|
||||
# print("===============model_kwargs[1]==============")
|
||||
|
||||
# print(model_kwargs[1])
|
||||
|
||||
# print("===============model_kwargs[1]==============")
|
||||
|
||||
u_out = model(xt, self._scale_timesteps(t), **model_kwargs[1])
|
||||
dim = y_out.size(1) if self.var_type.startswith('fixed') else y_out.size(1) // 2
|
||||
out = torch.cat([
|
||||
u_out[:, :dim] + guide_scale * (y_out[:, :dim] - u_out[:, :dim]),
|
||||
y_out[:, dim:]], dim=1) # guide_scale=9.0
|
||||
|
||||
# compute variance
|
||||
if self.var_type == 'learned':
|
||||
out, log_var = out.chunk(2, dim=1)
|
||||
var = torch.exp(log_var)
|
||||
elif self.var_type == 'learned_range':
|
||||
out, fraction = out.chunk(2, dim=1)
|
||||
min_log_var = _i(self.posterior_log_variance_clipped, t, xt)
|
||||
max_log_var = _i(torch.log(self.betas), t, xt)
|
||||
fraction = (fraction + 1) / 2.0
|
||||
log_var = fraction * max_log_var + (1 - fraction) * min_log_var
|
||||
var = torch.exp(log_var)
|
||||
elif self.var_type == 'fixed_large':
|
||||
var = _i(torch.cat([self.posterior_variance[1:2], self.betas[1:]]), t, xt)
|
||||
log_var = torch.log(var)
|
||||
elif self.var_type == 'fixed_small':
|
||||
var = _i(self.posterior_variance, t, xt)
|
||||
log_var = _i(self.posterior_log_variance_clipped, t, xt)
|
||||
|
||||
# compute mean and x0
|
||||
if self.mean_type == 'x_{t-1}':
|
||||
mu = out # x_{t-1}
|
||||
x0 = _i(1.0 / self.posterior_mean_coef1, t, xt) * mu - \
|
||||
_i(self.posterior_mean_coef2 / self.posterior_mean_coef1, t, xt) * xt
|
||||
elif self.mean_type == 'x0':
|
||||
x0 = out
|
||||
mu, _, _ = self.q_posterior_mean_variance(x0, xt, t)
|
||||
elif self.mean_type == 'eps':
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * out
|
||||
mu, _, _ = self.q_posterior_mean_variance(x0, xt, t)
|
||||
elif self.mean_type == 'v':
|
||||
x0 = _i(self.sqrt_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_one_minus_alphas_cumprod, t, xt) * out
|
||||
mu, _, _ = self.q_posterior_mean_variance(x0, xt, t)
|
||||
|
||||
# restrict the range of x0
|
||||
if percentile is not None:
|
||||
assert percentile > 0 and percentile <= 1 # e.g., 0.995
|
||||
s = torch.quantile(x0.flatten(1).abs(), percentile, dim=1).clamp_(1.0).view(-1, 1, 1, 1)
|
||||
x0 = torch.min(s, torch.max(-s, x0)) / s
|
||||
elif clamp is not None:
|
||||
x0 = x0.clamp(-clamp, clamp)
|
||||
return mu, var, log_var, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, ddim_timesteps=20, eta=0.0):
|
||||
r"""Sample from p(x_{t-1} | x_t) using DDIM.
|
||||
- condition_fn: for classifier-based guidance (guided-diffusion).
|
||||
- guide_scale: for classifier-free guidance (glide/dalle-2).
|
||||
"""
|
||||
stride = self.num_timesteps // ddim_timesteps
|
||||
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
if condition_fn is not None:
|
||||
# x0 -> eps
|
||||
alpha = _i(self.alphas_cumprod, t, xt)
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
eps = eps - (1 - alpha).sqrt() * condition_fn(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
|
||||
# eps -> x0
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps
|
||||
|
||||
# derive variables
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
alphas = _i(self.alphas_cumprod, t, xt)
|
||||
alphas_prev = _i(self.alphas_cumprod, (t - stride).clamp(0), xt)
|
||||
sigmas = eta * torch.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev))
|
||||
|
||||
# random sample
|
||||
noise = torch.randn_like(xt)
|
||||
direction = torch.sqrt(1 - alphas_prev - sigmas ** 2) * eps
|
||||
mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1)))
|
||||
xt_1 = torch.sqrt(alphas_prev) * x0 + direction + mask * sigmas * noise
|
||||
return xt_1, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, ddim_timesteps=20, eta=0.0):
|
||||
# prepare input
|
||||
b = noise.size(0)
|
||||
xt = noise
|
||||
|
||||
# diffusion process (TODO: clamp is inaccurate! Consider replacing the stride by explicit prev/next steps)
|
||||
steps = (1 + torch.arange(0, self.num_timesteps, self.num_timesteps // ddim_timesteps)).clamp(0, self.num_timesteps - 1).flip(0)
|
||||
from tqdm import tqdm
|
||||
for step in tqdm(steps):
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.ddim_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale, ddim_timesteps, eta)
|
||||
# from ipdb import set_trace; set_trace()
|
||||
return xt
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_reverse_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, ddim_timesteps=20):
|
||||
r"""Sample from p(x_{t+1} | x_t) using DDIM reverse ODE (deterministic).
|
||||
"""
|
||||
stride = self.num_timesteps // ddim_timesteps
|
||||
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
|
||||
# derive variables
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
alphas_next = _i(
|
||||
torch.cat([self.alphas_cumprod, self.alphas_cumprod.new_zeros([1])]),
|
||||
(t + stride).clamp(0, self.num_timesteps), xt)
|
||||
|
||||
# reverse sample
|
||||
mu = torch.sqrt(alphas_next) * x0 + torch.sqrt(1 - alphas_next) * eps
|
||||
return mu, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_reverse_sample_loop(self, x0, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, ddim_timesteps=20):
|
||||
# prepare input
|
||||
b = x0.size(0)
|
||||
xt = x0
|
||||
|
||||
# reconstruction steps
|
||||
steps = torch.arange(0, self.num_timesteps, self.num_timesteps // ddim_timesteps)
|
||||
for step in steps:
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.ddim_reverse_sample(xt, t, model, model_kwargs, clamp, percentile, guide_scale, ddim_timesteps)
|
||||
return xt
|
||||
|
||||
@torch.no_grad()
|
||||
def plms_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, plms_timesteps=20):
|
||||
r"""Sample from p(x_{t-1} | x_t) using PLMS.
|
||||
- condition_fn: for classifier-based guidance (guided-diffusion).
|
||||
- guide_scale: for classifier-free guidance (glide/dalle-2).
|
||||
"""
|
||||
stride = self.num_timesteps // plms_timesteps
|
||||
|
||||
# function for compute eps
|
||||
def compute_eps(xt, t):
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
|
||||
# condition
|
||||
if condition_fn is not None:
|
||||
# x0 -> eps
|
||||
alpha = _i(self.alphas_cumprod, t, xt)
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
eps = eps - (1 - alpha).sqrt() * condition_fn(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
|
||||
# eps -> x0
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps
|
||||
|
||||
# derive eps
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
return eps
|
||||
|
||||
# function for compute x_0 and x_{t-1}
|
||||
def compute_x0(eps, t):
|
||||
# eps -> x0
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps
|
||||
|
||||
# deterministic sample
|
||||
alphas_prev = _i(self.alphas_cumprod, (t - stride).clamp(0), xt)
|
||||
direction = torch.sqrt(1 - alphas_prev) * eps
|
||||
mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1)))
|
||||
xt_1 = torch.sqrt(alphas_prev) * x0 + direction
|
||||
return xt_1, x0
|
||||
|
||||
# PLMS sample
|
||||
eps = compute_eps(xt, t)
|
||||
if len(eps_cache) == 0:
|
||||
# 2nd order pseudo improved Euler
|
||||
xt_1, x0 = compute_x0(eps, t)
|
||||
eps_next = compute_eps(xt_1, (t - stride).clamp(0))
|
||||
eps_prime = (eps + eps_next) / 2.0
|
||||
elif len(eps_cache) == 1:
|
||||
# 2nd order pseudo linear multistep (Adams-Bashforth)
|
||||
eps_prime = (3 * eps - eps_cache[-1]) / 2.0
|
||||
elif len(eps_cache) == 2:
|
||||
# 3nd order pseudo linear multistep (Adams-Bashforth)
|
||||
eps_prime = (23 * eps - 16 * eps_cache[-1] + 5 * eps_cache[-2]) / 12.0
|
||||
elif len(eps_cache) >= 3:
|
||||
# 4nd order pseudo linear multistep (Adams-Bashforth)
|
||||
eps_prime = (55 * eps - 59 * eps_cache[-1] + 37 * eps_cache[-2] - 9 * eps_cache[-3]) / 24.0
|
||||
xt_1, x0 = compute_x0(eps_prime, t)
|
||||
return xt_1, x0, eps
|
||||
|
||||
@torch.no_grad()
|
||||
def plms_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, plms_timesteps=20):
|
||||
# prepare input
|
||||
b = noise.size(0)
|
||||
xt = noise
|
||||
|
||||
# diffusion process
|
||||
steps = (1 + torch.arange(0, self.num_timesteps, self.num_timesteps // plms_timesteps)).clamp(0, self.num_timesteps - 1).flip(0)
|
||||
eps_cache = []
|
||||
for step in steps:
|
||||
# PLMS sampling step
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _, eps = self.plms_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale, plms_timesteps, eps_cache)
|
||||
|
||||
# update eps cache
|
||||
eps_cache.append(eps)
|
||||
if len(eps_cache) >= 4:
|
||||
eps_cache.pop(0)
|
||||
return xt
|
||||
|
||||
def _scale_timesteps(self, t):
|
||||
if self.rescale_timesteps:
|
||||
return t.float() * 1000.0 / self.num_timesteps
|
||||
return t
|
||||
#return t.float()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@DIFFUSION.register_class()
|
||||
class DiffusionDDIMLong(object):
|
||||
def __init__(self,
|
||||
schedule='linear_sd',
|
||||
schedule_param={},
|
||||
mean_type='eps',
|
||||
var_type='learned_range',
|
||||
loss_type='mse',
|
||||
epsilon = 1e-12,
|
||||
rescale_timesteps=False,
|
||||
noise_strength=0.0,
|
||||
**kwargs):
|
||||
|
||||
assert mean_type in ['x0', 'x_{t-1}', 'eps', 'v']
|
||||
assert var_type in ['learned', 'learned_range', 'fixed_large', 'fixed_small']
|
||||
assert loss_type in ['mse', 'rescaled_mse', 'kl', 'rescaled_kl', 'l1', 'rescaled_l1','charbonnier']
|
||||
|
||||
betas = beta_schedule(schedule, **schedule_param)
|
||||
assert min(betas) > 0 and max(betas) <= 1
|
||||
|
||||
if not isinstance(betas, torch.DoubleTensor):
|
||||
betas = torch.tensor(betas, dtype=torch.float64)
|
||||
|
||||
self.betas = betas
|
||||
self.num_timesteps = len(betas)
|
||||
self.mean_type = mean_type # v
|
||||
self.var_type = var_type # 'fixed_small'
|
||||
self.loss_type = loss_type # mse
|
||||
self.epsilon = epsilon # 1e-12
|
||||
self.rescale_timesteps = rescale_timesteps # False
|
||||
self.noise_strength = noise_strength
|
||||
|
||||
# alphas
|
||||
alphas = 1 - self.betas
|
||||
self.alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
self.alphas_cumprod_prev = torch.cat([alphas.new_ones([1]), self.alphas_cumprod[:-1]])
|
||||
self.alphas_cumprod_next = torch.cat([self.alphas_cumprod[1:], alphas.new_zeros([1])])
|
||||
|
||||
# q(x_t | x_{t-1})
|
||||
self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)
|
||||
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)
|
||||
self.log_one_minus_alphas_cumprod = torch.log(1.0 - self.alphas_cumprod)
|
||||
self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod)
|
||||
self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod - 1)
|
||||
|
||||
# q(x_{t-1} | x_t, x_0)
|
||||
self.posterior_variance = betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
self.posterior_log_variance_clipped = torch.log(self.posterior_variance.clamp(1e-20))
|
||||
self.posterior_mean_coef1 = betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - self.alphas_cumprod)
|
||||
|
||||
|
||||
def sample_loss(self, x0, noise=None):
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x0)
|
||||
if self.noise_strength > 0:
|
||||
b, c, f, _, _= x0.shape
|
||||
offset_noise = torch.randn(b, c, f, 1, 1, device=x0.device)
|
||||
noise = noise + self.noise_strength * offset_noise
|
||||
return noise
|
||||
|
||||
|
||||
def q_sample(self, x0, t, noise=None):
|
||||
r"""Sample from q(x_t | x_0).
|
||||
"""
|
||||
# noise = torch.randn_like(x0) if noise is None else noise
|
||||
noise = self.sample_loss(x0, noise)
|
||||
return _i(self.sqrt_alphas_cumprod, t, x0) * x0 + \
|
||||
_i(self.sqrt_one_minus_alphas_cumprod, t, x0) * noise
|
||||
|
||||
def q_mean_variance(self, x0, t):
|
||||
r"""Distribution of q(x_t | x_0).
|
||||
"""
|
||||
mu = _i(self.sqrt_alphas_cumprod, t, x0) * x0
|
||||
var = _i(1.0 - self.alphas_cumprod, t, x0)
|
||||
log_var = _i(self.log_one_minus_alphas_cumprod, t, x0)
|
||||
return mu, var, log_var
|
||||
|
||||
def q_posterior_mean_variance(self, x0, xt, t):
|
||||
r"""Distribution of q(x_{t-1} | x_t, x_0).
|
||||
"""
|
||||
mu = _i(self.posterior_mean_coef1, t, xt) * x0 + _i(self.posterior_mean_coef2, t, xt) * xt
|
||||
var = _i(self.posterior_variance, t, xt)
|
||||
log_var = _i(self.posterior_log_variance_clipped, t, xt)
|
||||
return mu, var, log_var
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None):
|
||||
r"""Sample from p(x_{t-1} | x_t).
|
||||
- condition_fn: for classifier-based guidance (guided-diffusion).
|
||||
- guide_scale: for classifier-free guidance (glide/dalle-2).
|
||||
"""
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
mu, var, log_var, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
|
||||
# random sample (with optional conditional function)
|
||||
noise = torch.randn_like(xt)
|
||||
mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1))) # no noise when t == 0
|
||||
if condition_fn is not None:
|
||||
grad = condition_fn(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
mu = mu.float() + var * grad.float()
|
||||
xt_1 = mu + mask * torch.exp(0.5 * log_var) * noise
|
||||
return xt_1, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None):
|
||||
r"""Sample from p(x_{t-1} | x_t) p(x_{t-2} | x_{t-1}) ... p(x_0 | x_1).
|
||||
"""
|
||||
# prepare input
|
||||
b = noise.size(0)
|
||||
xt = noise
|
||||
|
||||
# diffusion process
|
||||
for step in torch.arange(self.num_timesteps).flip(0):
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.p_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale)
|
||||
return xt
|
||||
|
||||
def p_mean_variance(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, context_size=32, context_stride=1, context_overlap=4, context_batch_size=1):
|
||||
r"""Distribution of p(x_{t-1} | x_t).
|
||||
"""
|
||||
|
||||
|
||||
|
||||
|
||||
noise = xt
|
||||
context_queue = list(
|
||||
context_scheduler(
|
||||
0,
|
||||
31,
|
||||
noise.shape[2],
|
||||
context_size=context_size,
|
||||
context_stride=1,
|
||||
context_overlap=4,
|
||||
)
|
||||
)
|
||||
context_step = min(
|
||||
context_stride, int(np.ceil(np.log2(noise.shape[2] / context_size))) + 1
|
||||
)
|
||||
# replace the final segment to improve temporal consistency
|
||||
num_frames = noise.shape[2]
|
||||
context_queue[-1] = [
|
||||
e % num_frames
|
||||
for e in range(num_frames - context_size * context_step, num_frames, context_step)
|
||||
]
|
||||
|
||||
import math
|
||||
# context_batch_size = 1
|
||||
num_context_batches = math.ceil(len(context_queue) / context_batch_size)
|
||||
global_context = []
|
||||
for i in range(num_context_batches):
|
||||
global_context.append(
|
||||
context_queue[
|
||||
i * context_batch_size : (i + 1) * context_batch_size
|
||||
]
|
||||
)
|
||||
noise_pred = torch.zeros_like(noise)
|
||||
noise_pred_uncond = torch.zeros_like(noise)
|
||||
counter = torch.zeros(
|
||||
(1, 1, xt.shape[2], 1, 1),
|
||||
device=xt.device,
|
||||
dtype=xt.dtype,
|
||||
)
|
||||
|
||||
for i_index, context in enumerate(global_context):
|
||||
|
||||
|
||||
latent_model_input = torch.cat([xt[:, :, c] for c in context])
|
||||
bs_context = len(context)
|
||||
|
||||
# print("model_kwargs[0][local_image].shape: ", model_kwargs[0]['local_image'].shape)
|
||||
# print("model_kwargs[0][dwpose].shape: ", model_kwargs[0]['dwpose'].shape)
|
||||
# print("model_kwargs[0][pose_embedding].shape: ", model_kwargs[0]['pose_embedding'].shape)
|
||||
|
||||
model_kwargs_new = [{
|
||||
'y': None,
|
||||
"local_image": None if not model_kwargs[0].__contains__('local_image') else torch.cat([model_kwargs[0]["local_image"][:, :, c] for c in context]),
|
||||
'image': None if not model_kwargs[0].__contains__('image') else model_kwargs[0]["image"].repeat(bs_context, 1, 1),
|
||||
'dwpose': None if not model_kwargs[0].__contains__('dwpose') else torch.cat([model_kwargs[0]["dwpose"][:, :, [0]+[ii+1 for ii in c]] for c in context]),
|
||||
'randomref': None if not model_kwargs[0].__contains__('randomref') else torch.cat([model_kwargs[0]["randomref"][:, :, c] for c in context]),
|
||||
},
|
||||
{
|
||||
'y': None,
|
||||
"local_image": None,
|
||||
'image': None,
|
||||
'randomref': None,
|
||||
'dwpose': None,
|
||||
}]
|
||||
|
||||
if model_kwargs[0].__contains__('pose_embedding'):
|
||||
model_kwargs_new[0]['pose_embedding'] = torch.cat([model_kwargs[0]["pose_embedding"][:, :, c] for c in context])
|
||||
|
||||
if guide_scale is None:
|
||||
out = model(latent_model_input, self._scale_timesteps(t), **model_kwargs)
|
||||
for j, c in enumerate(context):
|
||||
noise_pred[:, :, c] = noise_pred[:, :, c] + out
|
||||
counter[:, :, c] = counter[:, :, c] + 1
|
||||
else:
|
||||
# classifier-free guidance
|
||||
# (model_kwargs[0]: conditional kwargs; model_kwargs[1]: non-conditional kwargs)
|
||||
# assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
|
||||
y_out = model(latent_model_input, self._scale_timesteps(t).repeat(bs_context), **model_kwargs_new[0])
|
||||
u_out = model(latent_model_input, self._scale_timesteps(t).repeat(bs_context), **model_kwargs_new[1])
|
||||
dim = y_out.size(1) if self.var_type.startswith('fixed') else y_out.size(1) // 2
|
||||
for j, c in enumerate(context):
|
||||
noise_pred[:, :, c] = noise_pred[:, :, c] + y_out[j:j+1]
|
||||
noise_pred_uncond[:, :, c] = noise_pred_uncond[:, :, c] + u_out[j:j+1]
|
||||
counter[:, :, c] = counter[:, :, c] + 1
|
||||
|
||||
noise_pred = noise_pred / counter
|
||||
noise_pred_uncond = noise_pred_uncond / counter
|
||||
out = torch.cat([
|
||||
noise_pred_uncond[:, :dim] + guide_scale * (noise_pred[:, :dim] - noise_pred_uncond[:, :dim]),
|
||||
noise_pred[:, dim:]], dim=1) # guide_scale=2.5
|
||||
|
||||
|
||||
# compute variance
|
||||
if self.var_type == 'learned':
|
||||
out, log_var = out.chunk(2, dim=1)
|
||||
var = torch.exp(log_var)
|
||||
elif self.var_type == 'learned_range':
|
||||
out, fraction = out.chunk(2, dim=1)
|
||||
min_log_var = _i(self.posterior_log_variance_clipped, t, xt)
|
||||
max_log_var = _i(torch.log(self.betas), t, xt)
|
||||
fraction = (fraction + 1) / 2.0
|
||||
log_var = fraction * max_log_var + (1 - fraction) * min_log_var
|
||||
var = torch.exp(log_var)
|
||||
elif self.var_type == 'fixed_large':
|
||||
var = _i(torch.cat([self.posterior_variance[1:2], self.betas[1:]]), t, xt)
|
||||
log_var = torch.log(var)
|
||||
elif self.var_type == 'fixed_small':
|
||||
var = _i(self.posterior_variance, t, xt)
|
||||
log_var = _i(self.posterior_log_variance_clipped, t, xt)
|
||||
|
||||
# compute mean and x0
|
||||
if self.mean_type == 'x_{t-1}':
|
||||
mu = out # x_{t-1}
|
||||
x0 = _i(1.0 / self.posterior_mean_coef1, t, xt) * mu - \
|
||||
_i(self.posterior_mean_coef2 / self.posterior_mean_coef1, t, xt) * xt
|
||||
elif self.mean_type == 'x0':
|
||||
x0 = out
|
||||
mu, _, _ = self.q_posterior_mean_variance(x0, xt, t)
|
||||
elif self.mean_type == 'eps':
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * out
|
||||
mu, _, _ = self.q_posterior_mean_variance(x0, xt, t)
|
||||
elif self.mean_type == 'v':
|
||||
x0 = _i(self.sqrt_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_one_minus_alphas_cumprod, t, xt) * out
|
||||
mu, _, _ = self.q_posterior_mean_variance(x0, xt, t)
|
||||
|
||||
# restrict the range of x0
|
||||
if percentile is not None:
|
||||
assert percentile > 0 and percentile <= 1 # e.g., 0.995
|
||||
s = torch.quantile(x0.flatten(1).abs(), percentile, dim=1).clamp_(1.0).view(-1, 1, 1, 1)
|
||||
x0 = torch.min(s, torch.max(-s, x0)) / s
|
||||
elif clamp is not None:
|
||||
x0 = x0.clamp(-clamp, clamp)
|
||||
return mu, var, log_var, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, ddim_timesteps=20, eta=0.0, context_size=32, context_stride=1, context_overlap=4, context_batch_size=1):
|
||||
r"""Sample from p(x_{t-1} | x_t) using DDIM.
|
||||
- condition_fn: for classifier-based guidance (guided-diffusion).
|
||||
- guide_scale: for classifier-free guidance (glide/dalle-2).
|
||||
"""
|
||||
stride = self.num_timesteps // ddim_timesteps
|
||||
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale, context_size, context_stride, context_overlap, context_batch_size)
|
||||
if condition_fn is not None:
|
||||
# x0 -> eps
|
||||
alpha = _i(self.alphas_cumprod, t, xt)
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
eps = eps - (1 - alpha).sqrt() * condition_fn(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
|
||||
# eps -> x0
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps
|
||||
|
||||
# derive variables
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
alphas = _i(self.alphas_cumprod, t, xt)
|
||||
alphas_prev = _i(self.alphas_cumprod, (t - stride).clamp(0), xt)
|
||||
sigmas = eta * torch.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev))
|
||||
|
||||
# random sample
|
||||
noise = torch.randn_like(xt)
|
||||
direction = torch.sqrt(1 - alphas_prev - sigmas ** 2) * eps
|
||||
mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1)))
|
||||
xt_1 = torch.sqrt(alphas_prev) * x0 + direction + mask * sigmas * noise
|
||||
return xt_1, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sample_loop(self, noise, context_size, context_stride, context_overlap, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, ddim_timesteps=20, eta=0.0, context_batch_size=1):
|
||||
# prepare input
|
||||
b = noise.size(0)
|
||||
xt = noise
|
||||
|
||||
# diffusion process (TODO: clamp is inaccurate! Consider replacing the stride by explicit prev/next steps)
|
||||
steps = (1 + torch.arange(0, self.num_timesteps, self.num_timesteps // ddim_timesteps)).clamp(0, self.num_timesteps - 1).flip(0)
|
||||
from tqdm import tqdm
|
||||
|
||||
for step in tqdm(steps):
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.ddim_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale, ddim_timesteps, eta, context_size=context_size, context_stride=context_stride, context_overlap=context_overlap, context_batch_size=context_batch_size)
|
||||
return xt
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_reverse_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, ddim_timesteps=20):
|
||||
r"""Sample from p(x_{t+1} | x_t) using DDIM reverse ODE (deterministic).
|
||||
"""
|
||||
stride = self.num_timesteps // ddim_timesteps
|
||||
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
|
||||
# derive variables
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
alphas_next = _i(
|
||||
torch.cat([self.alphas_cumprod, self.alphas_cumprod.new_zeros([1])]),
|
||||
(t + stride).clamp(0, self.num_timesteps), xt)
|
||||
|
||||
# reverse sample
|
||||
mu = torch.sqrt(alphas_next) * x0 + torch.sqrt(1 - alphas_next) * eps
|
||||
return mu, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_reverse_sample_loop(self, x0, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, ddim_timesteps=20):
|
||||
# prepare input
|
||||
b = x0.size(0)
|
||||
xt = x0
|
||||
|
||||
# reconstruction steps
|
||||
steps = torch.arange(0, self.num_timesteps, self.num_timesteps // ddim_timesteps)
|
||||
for step in steps:
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.ddim_reverse_sample(xt, t, model, model_kwargs, clamp, percentile, guide_scale, ddim_timesteps)
|
||||
return xt
|
||||
|
||||
@torch.no_grad()
|
||||
def plms_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, plms_timesteps=20):
|
||||
r"""Sample from p(x_{t-1} | x_t) using PLMS.
|
||||
- condition_fn: for classifier-based guidance (guided-diffusion).
|
||||
- guide_scale: for classifier-free guidance (glide/dalle-2).
|
||||
"""
|
||||
stride = self.num_timesteps // plms_timesteps
|
||||
|
||||
# function for compute eps
|
||||
def compute_eps(xt, t):
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale)
|
||||
|
||||
# condition
|
||||
if condition_fn is not None:
|
||||
# x0 -> eps
|
||||
alpha = _i(self.alphas_cumprod, t, xt)
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
eps = eps - (1 - alpha).sqrt() * condition_fn(xt, self._scale_timesteps(t), **model_kwargs)
|
||||
|
||||
# eps -> x0
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps
|
||||
|
||||
# derive eps
|
||||
eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt)
|
||||
return eps
|
||||
|
||||
# function for compute x_0 and x_{t-1}
|
||||
def compute_x0(eps, t):
|
||||
# eps -> x0
|
||||
x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \
|
||||
_i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps
|
||||
|
||||
# deterministic sample
|
||||
alphas_prev = _i(self.alphas_cumprod, (t - stride).clamp(0), xt)
|
||||
direction = torch.sqrt(1 - alphas_prev) * eps
|
||||
mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1)))
|
||||
xt_1 = torch.sqrt(alphas_prev) * x0 + direction
|
||||
return xt_1, x0
|
||||
|
||||
# PLMS sample
|
||||
eps = compute_eps(xt, t)
|
||||
if len(eps_cache) == 0:
|
||||
# 2nd order pseudo improved Euler
|
||||
xt_1, x0 = compute_x0(eps, t)
|
||||
eps_next = compute_eps(xt_1, (t - stride).clamp(0))
|
||||
eps_prime = (eps + eps_next) / 2.0
|
||||
elif len(eps_cache) == 1:
|
||||
# 2nd order pseudo linear multistep (Adams-Bashforth)
|
||||
eps_prime = (3 * eps - eps_cache[-1]) / 2.0
|
||||
elif len(eps_cache) == 2:
|
||||
# 3nd order pseudo linear multistep (Adams-Bashforth)
|
||||
eps_prime = (23 * eps - 16 * eps_cache[-1] + 5 * eps_cache[-2]) / 12.0
|
||||
elif len(eps_cache) >= 3:
|
||||
# 4nd order pseudo linear multistep (Adams-Bashforth)
|
||||
eps_prime = (55 * eps - 59 * eps_cache[-1] + 37 * eps_cache[-2] - 9 * eps_cache[-3]) / 24.0
|
||||
xt_1, x0 = compute_x0(eps_prime, t)
|
||||
return xt_1, x0, eps
|
||||
|
||||
@torch.no_grad()
|
||||
def plms_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, plms_timesteps=20):
|
||||
# prepare input
|
||||
b = noise.size(0)
|
||||
xt = noise
|
||||
|
||||
# diffusion process
|
||||
steps = (1 + torch.arange(0, self.num_timesteps, self.num_timesteps // plms_timesteps)).clamp(0, self.num_timesteps - 1).flip(0)
|
||||
eps_cache = []
|
||||
for step in steps:
|
||||
# PLMS sampling step
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _, eps = self.plms_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale, plms_timesteps, eps_cache)
|
||||
|
||||
# update eps cache
|
||||
eps_cache.append(eps)
|
||||
if len(eps_cache) >= 4:
|
||||
eps_cache.pop(0)
|
||||
return xt
|
||||
|
||||
def _scale_timesteps(self, t):
|
||||
if self.rescale_timesteps:
|
||||
return t.float() * 1000.0 / self.num_timesteps
|
||||
return t
|
||||
#return t.float()
|
||||
|
||||
|
||||
|
||||
def ordered_halving(val):
|
||||
bin_str = f"{val:064b}"
|
||||
bin_flip = bin_str[::-1]
|
||||
as_int = int(bin_flip, 2)
|
||||
|
||||
return as_int / (1 << 64)
|
||||
|
||||
|
||||
def context_scheduler(
|
||||
step: int = ...,
|
||||
num_steps: Optional[int] = None,
|
||||
num_frames: int = ...,
|
||||
context_size: Optional[int] = None,
|
||||
context_stride: int = 3,
|
||||
context_overlap: int = 4,
|
||||
closed_loop: bool = False,
|
||||
):
|
||||
if num_frames <= context_size:
|
||||
yield list(range(num_frames))
|
||||
return
|
||||
|
||||
context_stride = min(
|
||||
context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1
|
||||
)
|
||||
|
||||
for context_step in 1 << np.arange(context_stride):
|
||||
pad = int(round(num_frames * ordered_halving(step)))
|
||||
for j in range(
|
||||
int(ordered_halving(step) * context_step) + pad,
|
||||
num_frames + pad + (0 if closed_loop else -context_overlap),
|
||||
(context_size * context_step - context_overlap),
|
||||
):
|
||||
|
||||
yield [
|
||||
e % num_frames
|
||||
for e in range(j, j + context_size * context_step, context_step)
|
||||
]
|
||||
|
||||
@@ -0,0 +1,498 @@
|
||||
"""
|
||||
GaussianDiffusion wraps operators for denoising diffusion models, including the
|
||||
diffusion and denoising processes, as well as the loss evaluation.
|
||||
"""
|
||||
import torch
|
||||
import torchsde
|
||||
import random
|
||||
from tqdm.auto import trange
|
||||
|
||||
|
||||
__all__ = ['GaussianDiffusion']
|
||||
|
||||
|
||||
def _i(tensor, t, x):
|
||||
"""
|
||||
Index tensor using t and format the output according to x.
|
||||
"""
|
||||
shape = (x.size(0), ) + (1, ) * (x.ndim - 1)
|
||||
return tensor[t.to(tensor.device)].view(shape).to(x.device)
|
||||
|
||||
|
||||
class BatchedBrownianTree:
|
||||
"""
|
||||
A wrapper around torchsde.BrownianTree that enables batches of entropy.
|
||||
"""
|
||||
def __init__(self, x, t0, t1, seed=None, **kwargs):
|
||||
t0, t1, self.sign = self.sort(t0, t1)
|
||||
w0 = kwargs.get('w0', torch.zeros_like(x))
|
||||
if seed is None:
|
||||
seed = torch.randint(0, 2 ** 63 - 1, []).item()
|
||||
self.batched = True
|
||||
try:
|
||||
assert len(seed) == x.shape[0]
|
||||
w0 = w0[0]
|
||||
except TypeError:
|
||||
seed = [seed]
|
||||
self.batched = False
|
||||
self.trees = [torchsde.BrownianTree(
|
||||
t0, w0, t1, entropy=s, **kwargs
|
||||
) for s in seed]
|
||||
|
||||
@staticmethod
|
||||
def sort(a, b):
|
||||
return (a, b, 1) if a < b else (b, a, -1)
|
||||
|
||||
def __call__(self, t0, t1):
|
||||
t0, t1, sign = self.sort(t0, t1)
|
||||
w = torch.stack([tree(t0, t1) for tree in self.trees]) * (self.sign * sign)
|
||||
return w if self.batched else w[0]
|
||||
|
||||
|
||||
class BrownianTreeNoiseSampler:
|
||||
"""
|
||||
A noise sampler backed by a torchsde.BrownianTree.
|
||||
|
||||
Args:
|
||||
x (Tensor): The tensor whose shape, device and dtype to use to generate
|
||||
random samples.
|
||||
sigma_min (float): The low end of the valid interval.
|
||||
sigma_max (float): The high end of the valid interval.
|
||||
seed (int or List[int]): The random seed. If a list of seeds is
|
||||
supplied instead of a single integer, then the noise sampler will
|
||||
use one BrownianTree per batch item, each with its own seed.
|
||||
transform (callable): A function that maps sigma to the sampler's
|
||||
internal timestep.
|
||||
"""
|
||||
def __init__(self, x, sigma_min, sigma_max, seed=None, transform=lambda x: x):
|
||||
self.transform = transform
|
||||
t0 = self.transform(torch.as_tensor(sigma_min))
|
||||
t1 = self.transform(torch.as_tensor(sigma_max))
|
||||
self.tree = BatchedBrownianTree(x, t0, t1, seed)
|
||||
|
||||
def __call__(self, sigma, sigma_next):
|
||||
t0 = self.transform(torch.as_tensor(sigma))
|
||||
t1 = self.transform(torch.as_tensor(sigma_next))
|
||||
return self.tree(t0, t1) / (t1 - t0).abs().sqrt()
|
||||
|
||||
|
||||
def get_scalings(sigma):
|
||||
c_out = -sigma
|
||||
c_in = 1 / (sigma ** 2 + 1. ** 2) ** 0.5
|
||||
return c_out, c_in
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_2m_sde(
|
||||
noise,
|
||||
model,
|
||||
sigmas,
|
||||
eta=1.,
|
||||
s_noise=1.,
|
||||
solver_type='midpoint',
|
||||
show_progress=True
|
||||
):
|
||||
"""
|
||||
DPM-Solver++ (2M) SDE.
|
||||
"""
|
||||
assert solver_type in {'heun', 'midpoint'}
|
||||
|
||||
x = noise * sigmas[0]
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[sigmas < float('inf')].max()
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max)
|
||||
old_denoised = None
|
||||
h_last = None
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
||||
if sigmas[i] == float('inf'):
|
||||
# Euler method
|
||||
denoised = model(noise, sigmas[i])
|
||||
x = denoised + sigmas[i + 1] * noise
|
||||
else:
|
||||
_, c_in = get_scalings(sigmas[i])
|
||||
denoised = model(x * c_in, sigmas[i])
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
# DPM-Solver++(2M) SDE
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
eta_h = eta * h
|
||||
|
||||
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + \
|
||||
(-h - eta_h).expm1().neg() * denoised
|
||||
|
||||
if old_denoised is not None:
|
||||
r = h_last / h
|
||||
if solver_type == 'heun':
|
||||
x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * \
|
||||
(1 / r) * (denoised - old_denoised)
|
||||
elif solver_type == 'midpoint':
|
||||
x = x + 0.5 * (-h - eta_h).expm1().neg() * \
|
||||
(1 / r) * (denoised - old_denoised)
|
||||
|
||||
x = x + noise_sampler(
|
||||
sigmas[i],
|
||||
sigmas[i + 1]
|
||||
) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
|
||||
|
||||
old_denoised = denoised
|
||||
h_last = h
|
||||
return x
|
||||
|
||||
|
||||
class GaussianDiffusion(object):
|
||||
|
||||
def __init__(self, sigmas, prediction_type='eps'):
|
||||
assert prediction_type in {'x0', 'eps', 'v'}
|
||||
self.sigmas = sigmas.float() # noise coefficients
|
||||
self.alphas = torch.sqrt(1 - sigmas ** 2).float() # signal coefficients
|
||||
self.num_timesteps = len(sigmas)
|
||||
self.prediction_type = prediction_type
|
||||
|
||||
def diffuse(self, x0, t, noise=None):
|
||||
"""
|
||||
Add Gaussian noise to signal x0 according to:
|
||||
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I).
|
||||
"""
|
||||
noise = torch.randn_like(x0) if noise is None else noise
|
||||
xt = _i(self.alphas, t, x0) * x0 + _i(self.sigmas, t, x0) * noise
|
||||
return xt
|
||||
|
||||
def denoise(
|
||||
self,
|
||||
xt,
|
||||
t,
|
||||
s,
|
||||
model,
|
||||
model_kwargs={},
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
percentile=None
|
||||
):
|
||||
"""
|
||||
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
|
||||
Since x0 is not available, estimate the denoising results using the learned
|
||||
distribution p(x_s | x_t, \hat{x}_0 == f(x_t)).
|
||||
"""
|
||||
s = t - 1 if s is None else s
|
||||
|
||||
# hyperparams
|
||||
sigmas = _i(self.sigmas, t, xt)
|
||||
alphas = _i(self.alphas, t, xt)
|
||||
alphas_s = _i(self.alphas, s.clamp(0), xt)
|
||||
alphas_s[s < 0] = 1.
|
||||
sigmas_s = torch.sqrt(1 - alphas_s ** 2)
|
||||
|
||||
# precompute variables
|
||||
betas = 1 - (alphas / alphas_s) ** 2
|
||||
coef1 = betas * alphas_s / sigmas ** 2
|
||||
coef2 = (alphas * sigmas_s ** 2) / (alphas_s * sigmas ** 2)
|
||||
var = betas * (sigmas_s / sigmas) ** 2
|
||||
log_var = torch.log(var).clamp_(-20, 20)
|
||||
|
||||
# prediction
|
||||
if guide_scale is None:
|
||||
assert isinstance(model_kwargs, dict)
|
||||
out = model(xt, t=t, **model_kwargs)
|
||||
else:
|
||||
# classifier-free guidance (arXiv:2207.12598)
|
||||
# model_kwargs[0]: conditional kwargs
|
||||
# model_kwargs[1]: non-conditional kwargs
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
|
||||
y_out = model(xt, t=t, **model_kwargs[0])
|
||||
if guide_scale == 1.:
|
||||
out = y_out
|
||||
else:
|
||||
u_out = model(xt, t=t, **model_kwargs[1])
|
||||
out = u_out + guide_scale * (y_out - u_out)
|
||||
|
||||
# rescale the output according to arXiv:2305.08891
|
||||
if guide_rescale is not None:
|
||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
||||
ratio = (y_out.flatten(1).std(dim=1) / (
|
||||
out.flatten(1).std(dim=1) + 1e-12
|
||||
)).view((-1, ) + (1, ) * (y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
|
||||
# compute x0
|
||||
if self.prediction_type == 'x0':
|
||||
x0 = out
|
||||
elif self.prediction_type == 'eps':
|
||||
x0 = (xt - sigmas * out) / alphas
|
||||
elif self.prediction_type == 'v':
|
||||
x0 = alphas * xt - sigmas * out
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'prediction_type {self.prediction_type} not implemented'
|
||||
)
|
||||
|
||||
# restrict the range of x0
|
||||
if percentile is not None:
|
||||
# NOTE: percentile should only be used when data is within range [-1, 1]
|
||||
assert percentile > 0 and percentile <= 1
|
||||
s = torch.quantile(x0.flatten(1).abs(), percentile, dim=1)
|
||||
s = s.clamp_(1.0).view((-1, ) + (1, ) * (xt.ndim - 1))
|
||||
x0 = torch.min(s, torch.max(-s, x0)) / s
|
||||
elif clamp is not None:
|
||||
x0 = x0.clamp(-clamp, clamp)
|
||||
|
||||
# recompute eps using the restricted x0
|
||||
eps = (xt - alphas * x0) / sigmas
|
||||
|
||||
# compute mu (mean of posterior distribution) using the restricted x0
|
||||
mu = coef1 * x0 + coef2 * xt
|
||||
return mu, var, log_var, x0, eps
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(
|
||||
self,
|
||||
noise,
|
||||
model,
|
||||
model_kwargs={},
|
||||
condition_fn=None,
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
solver='euler_a',
|
||||
steps=20,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discretization=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
show_progress=False,
|
||||
seed=-1,
|
||||
**kwargs
|
||||
):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
|
||||
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing')
|
||||
assert discard_penultimate_step in (None, True, False)
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
|
||||
# function of diffusion solver
|
||||
solver_fn = {
|
||||
# 'heun': sample_heun,
|
||||
'dpmpp_2m_sde': sample_dpmpp_2m_sde
|
||||
}[solver]
|
||||
|
||||
# options
|
||||
schedule = 'karras' if 'karras' in solver else None
|
||||
discretization = discretization or 'linspace'
|
||||
seed = seed if seed >= 0 else random.randint(0, 2 ** 31)
|
||||
if isinstance(steps, torch.LongTensor):
|
||||
discard_penultimate_step = False
|
||||
if discard_penultimate_step is None:
|
||||
discard_penultimate_step = True if solver in (
|
||||
'dpm2',
|
||||
'dpm2_ancestral',
|
||||
'dpmpp_2m_sde',
|
||||
'dpm2_karras',
|
||||
'dpm2_ancestral_karras',
|
||||
'dpmpp_2m_sde_karras'
|
||||
) else False
|
||||
|
||||
# function for denoising xt to get x0
|
||||
intermediates = []
|
||||
def model_fn(xt, sigma):
|
||||
# denoising
|
||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||
x0 = self.denoise(
|
||||
xt, t, None, model, model_kwargs, guide_scale, guide_rescale, clamp,
|
||||
percentile
|
||||
)[-2]
|
||||
|
||||
# collect intermediate outputs
|
||||
if return_intermediate == 'xt':
|
||||
intermediates.append(xt)
|
||||
elif return_intermediate == 'x0':
|
||||
intermediates.append(x0)
|
||||
return x0
|
||||
|
||||
# get timesteps
|
||||
if isinstance(steps, int):
|
||||
steps += 1 if discard_penultimate_step else 0
|
||||
t_max = self.num_timesteps - 1 if t_max is None else t_max
|
||||
t_min = 0 if t_min is None else t_min
|
||||
|
||||
# discretize timesteps
|
||||
if discretization == 'leading':
|
||||
steps = torch.arange(
|
||||
t_min, t_max + 1, (t_max - t_min + 1) / steps
|
||||
).flip(0)
|
||||
elif discretization == 'linspace':
|
||||
steps = torch.linspace(t_max, t_min, steps)
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1, -((t_max - t_min + 1) / steps))
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented'
|
||||
)
|
||||
steps = steps.clamp_(t_min, t_max)
|
||||
steps = torch.as_tensor(steps, dtype=torch.float32, device=noise.device)
|
||||
|
||||
# get sigmas
|
||||
sigmas = self._t_to_sigma(steps)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
if schedule == 'karras':
|
||||
if sigmas[0] == float('inf'):
|
||||
sigmas = karras_schedule(
|
||||
n=len(steps) - 1,
|
||||
sigma_min=sigmas[sigmas > 0].min().item(),
|
||||
sigma_max=sigmas[sigmas < float('inf')].max().item(),
|
||||
rho=7.
|
||||
).to(sigmas)
|
||||
sigmas = torch.cat([
|
||||
sigmas.new_tensor([float('inf')]), sigmas, sigmas.new_zeros([1])
|
||||
])
|
||||
else:
|
||||
sigmas = karras_schedule(
|
||||
n=len(steps),
|
||||
sigma_min=sigmas[sigmas > 0].min().item(),
|
||||
sigma_max=sigmas.max().item(),
|
||||
rho=7.
|
||||
).to(sigmas)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
if discard_penultimate_step:
|
||||
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
||||
|
||||
# sampling
|
||||
x0 = solver_fn(
|
||||
noise,
|
||||
model_fn,
|
||||
sigmas,
|
||||
show_progress=show_progress,
|
||||
**kwargs
|
||||
)
|
||||
return (x0, intermediates) if return_intermediate is not None else x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_reverse_sample(
|
||||
self,
|
||||
xt,
|
||||
t,
|
||||
model,
|
||||
model_kwargs={},
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
ddim_timesteps=20,
|
||||
reverse_steps=600
|
||||
):
|
||||
r"""Sample from p(x_{t+1} | x_t) using DDIM reverse ODE (deterministic).
|
||||
"""
|
||||
stride = reverse_steps // ddim_timesteps
|
||||
|
||||
# predict distribution of p(x_{t-1} | x_t)
|
||||
_, _, _, x0, eps = self.denoise(
|
||||
xt, t, None, model, model_kwargs, guide_scale, guide_rescale, clamp,
|
||||
percentile
|
||||
)
|
||||
# derive variables
|
||||
s = (t + stride).clamp(0, reverse_steps-1)
|
||||
# hyperparams
|
||||
sigmas = _i(self.sigmas, t, xt)
|
||||
alphas = _i(self.alphas, t, xt)
|
||||
alphas_s = _i(self.alphas, s.clamp(0), xt)
|
||||
alphas_s[s < 0] = 1.
|
||||
sigmas_s = torch.sqrt(1 - alphas_s ** 2)
|
||||
|
||||
# reverse sample
|
||||
mu = alphas_s * x0 + sigmas_s * eps
|
||||
return mu, x0
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_reverse_sample_loop(
|
||||
self,
|
||||
x0,
|
||||
model,
|
||||
model_kwargs={},
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
ddim_timesteps=20,
|
||||
reverse_steps=600
|
||||
):
|
||||
# prepare input
|
||||
b = x0.size(0)
|
||||
xt = x0
|
||||
|
||||
# reconstruction steps
|
||||
steps = torch.arange(0, reverse_steps, reverse_steps // ddim_timesteps)
|
||||
for step in steps:
|
||||
t = torch.full((b, ), step, dtype=torch.long, device=xt.device)
|
||||
xt, _ = self.ddim_reverse_sample(xt, t, model, model_kwargs, clamp, percentile, guide_scale, guide_rescale, ddim_timesteps, reverse_steps)
|
||||
return xt
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
if sigma == float('inf'):
|
||||
t = torch.full_like(sigma, len(self.sigmas) - 1)
|
||||
else:
|
||||
log_sigmas = torch.sqrt(
|
||||
self.sigmas ** 2 / (1 - self.sigmas ** 2)
|
||||
).log().to(sigma)
|
||||
log_sigma = sigma.log()
|
||||
dists = log_sigma - log_sigmas[:, None]
|
||||
low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp(
|
||||
max=log_sigmas.shape[0] - 2
|
||||
)
|
||||
high_idx = low_idx + 1
|
||||
low, high = log_sigmas[low_idx], log_sigmas[high_idx]
|
||||
w = (low - log_sigma) / (low - high)
|
||||
w = w.clamp(0, 1)
|
||||
t = (1 - w) * low_idx + w * high_idx
|
||||
t = t.view(sigma.shape)
|
||||
if t.ndim == 0:
|
||||
t = t.unsqueeze(0)
|
||||
return t
|
||||
|
||||
def _t_to_sigma(self, t):
|
||||
t = t.float()
|
||||
low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac()
|
||||
log_sigmas = torch.sqrt(self.sigmas ** 2 / (1 - self.sigmas ** 2)).log().to(t)
|
||||
log_sigma = (1 - w) * log_sigmas[low_idx] + w * log_sigmas[high_idx]
|
||||
log_sigma[torch.isnan(log_sigma) | torch.isinf(log_sigma)] = float('inf')
|
||||
return log_sigma.exp()
|
||||
|
||||
def prev_step(self, model_out, t, xt, inference_steps=50):
|
||||
prev_t = t - self.num_timesteps // inference_steps
|
||||
|
||||
sigmas = _i(self.sigmas, t, xt)
|
||||
alphas = _i(self.alphas, t, xt)
|
||||
alphas_prev = _i(self.alphas, prev_t.clamp(0), xt)
|
||||
alphas_prev[prev_t < 0] = 1.
|
||||
sigmas_prev = torch.sqrt(1 - alphas_prev ** 2)
|
||||
|
||||
x0 = alphas * xt - sigmas * model_out
|
||||
eps = (xt - alphas * x0) / sigmas
|
||||
prev_sample = alphas_prev * x0 + sigmas_prev * eps
|
||||
return prev_sample
|
||||
|
||||
def next_step(self, model_out, t, xt, inference_steps=50):
|
||||
t, next_t = min(t - self.num_timesteps // inference_steps, 999), t
|
||||
|
||||
sigmas = _i(self.sigmas, t, xt)
|
||||
alphas = _i(self.alphas, t, xt)
|
||||
alphas_next = _i(self.alphas, next_t.clamp(0), xt)
|
||||
alphas_next[next_t < 0] = 1.
|
||||
sigmas_next = torch.sqrt(1 - alphas_next ** 2)
|
||||
|
||||
x0 = alphas * xt - sigmas * model_out
|
||||
eps = (xt - alphas * x0) / sigmas
|
||||
next_sample = alphas_next * x0 + sigmas_next * eps
|
||||
return next_sample
|
||||
|
||||
def get_noise_pred_single(self, xt, t, model, model_kwargs):
|
||||
assert isinstance(model_kwargs, dict)
|
||||
out = model(xt, t=t, **model_kwargs)
|
||||
return out
|
||||
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
import math
|
||||
import torch
|
||||
|
||||
|
||||
def beta_schedule(schedule='cosine',
|
||||
num_timesteps=1000,
|
||||
zero_terminal_snr=False,
|
||||
**kwargs):
|
||||
# compute betas
|
||||
betas = {
|
||||
# 'logsnr_cosine_interp': logsnr_cosine_interp_schedule,
|
||||
'linear': linear_schedule,
|
||||
'linear_sd': linear_sd_schedule,
|
||||
'quadratic': quadratic_schedule,
|
||||
'cosine': cosine_schedule
|
||||
}[schedule](num_timesteps, **kwargs)
|
||||
|
||||
if zero_terminal_snr and abs(betas.max() - 1.0) > 0.0001:
|
||||
betas = rescale_zero_terminal_snr(betas)
|
||||
|
||||
return betas
|
||||
|
||||
|
||||
def sigma_schedule(schedule='cosine',
|
||||
num_timesteps=1000,
|
||||
zero_terminal_snr=False,
|
||||
**kwargs):
|
||||
# compute betas
|
||||
betas = {
|
||||
'logsnr_cosine_interp': logsnr_cosine_interp_schedule,
|
||||
'linear': linear_schedule,
|
||||
'linear_sd': linear_sd_schedule,
|
||||
'quadratic': quadratic_schedule,
|
||||
'cosine': cosine_schedule
|
||||
}[schedule](num_timesteps, **kwargs)
|
||||
if schedule == 'logsnr_cosine_interp':
|
||||
sigma = betas
|
||||
else:
|
||||
sigma = betas_to_sigmas(betas)
|
||||
if zero_terminal_snr and abs(sigma.max() - 1.0) > 0.0001:
|
||||
sigma = rescale_zero_terminal_snr(sigma)
|
||||
|
||||
return sigma
|
||||
|
||||
|
||||
def linear_schedule(num_timesteps, init_beta, last_beta, **kwargs):
|
||||
scale = 1000.0 / num_timesteps
|
||||
init_beta = init_beta or scale * 0.0001
|
||||
ast_beta = last_beta or scale * 0.02
|
||||
return torch.linspace(init_beta, last_beta, num_timesteps, dtype=torch.float64)
|
||||
|
||||
def logsnr_cosine_interp_schedule(
|
||||
num_timesteps,
|
||||
scale_min=2,
|
||||
scale_max=4,
|
||||
logsnr_min=-15,
|
||||
logsnr_max=15,
|
||||
**kwargs):
|
||||
return logsnrs_to_sigmas(
|
||||
_logsnr_cosine_interp(num_timesteps, logsnr_min, logsnr_max, scale_min, scale_max))
|
||||
|
||||
def linear_sd_schedule(num_timesteps, init_beta, last_beta, **kwargs):
|
||||
return torch.linspace(init_beta ** 0.5, last_beta ** 0.5, num_timesteps, dtype=torch.float64) ** 2
|
||||
|
||||
|
||||
def quadratic_schedule(num_timesteps, init_beta, last_beta, **kwargs):
|
||||
init_beta = init_beta or 0.0015
|
||||
last_beta = last_beta or 0.0195
|
||||
return torch.linspace(init_beta ** 0.5, last_beta ** 0.5, num_timesteps, dtype=torch.float64) ** 2
|
||||
|
||||
|
||||
def cosine_schedule(num_timesteps, cosine_s=0.008, **kwargs):
|
||||
betas = []
|
||||
for step in range(num_timesteps):
|
||||
t1 = step / num_timesteps
|
||||
t2 = (step + 1) / num_timesteps
|
||||
fn = lambda u: math.cos((u + cosine_s) / (1 + cosine_s) * math.pi / 2) ** 2
|
||||
betas.append(min(1.0 - fn(t2) / fn(t1), 0.999))
|
||||
return torch.tensor(betas, dtype=torch.float64)
|
||||
|
||||
|
||||
# def cosine_schedule(n, cosine_s=0.008, **kwargs):
|
||||
# ramp = torch.linspace(0, 1, n + 1)
|
||||
# square_alphas = torch.cos((ramp + cosine_s) / (1 + cosine_s) * torch.pi / 2) ** 2
|
||||
# betas = (1 - square_alphas[1:] / square_alphas[:-1]).clamp(max=0.999)
|
||||
# return betas_to_sigmas(betas)
|
||||
|
||||
|
||||
def betas_to_sigmas(betas):
|
||||
return torch.sqrt(1 - torch.cumprod(1 - betas, dim=0))
|
||||
|
||||
|
||||
def sigmas_to_betas(sigmas):
|
||||
square_alphas = 1 - sigmas**2
|
||||
betas = 1 - torch.cat(
|
||||
[square_alphas[:1], square_alphas[1:] / square_alphas[:-1]])
|
||||
return betas
|
||||
|
||||
|
||||
|
||||
def sigmas_to_logsnrs(sigmas):
|
||||
square_sigmas = sigmas**2
|
||||
return torch.log(square_sigmas / (1 - square_sigmas))
|
||||
|
||||
|
||||
def _logsnr_cosine(n, logsnr_min=-15, logsnr_max=15):
|
||||
t_min = math.atan(math.exp(-0.5 * logsnr_min))
|
||||
t_max = math.atan(math.exp(-0.5 * logsnr_max))
|
||||
t = torch.linspace(1, 0, n)
|
||||
logsnrs = -2 * torch.log(torch.tan(t_min + t * (t_max - t_min)))
|
||||
return logsnrs
|
||||
|
||||
|
||||
def _logsnr_cosine_shifted(n, logsnr_min=-15, logsnr_max=15, scale=2):
|
||||
logsnrs = _logsnr_cosine(n, logsnr_min, logsnr_max)
|
||||
logsnrs += 2 * math.log(1 / scale)
|
||||
return logsnrs
|
||||
|
||||
def karras_schedule(n, sigma_min=0.002, sigma_max=80.0, rho=7.0):
|
||||
ramp = torch.linspace(1, 0, n)
|
||||
min_inv_rho = sigma_min**(1 / rho)
|
||||
max_inv_rho = sigma_max**(1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
|
||||
sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2))
|
||||
return sigmas
|
||||
|
||||
def _logsnr_cosine_interp(n,
|
||||
logsnr_min=-15,
|
||||
logsnr_max=15,
|
||||
scale_min=2,
|
||||
scale_max=4):
|
||||
t = torch.linspace(1, 0, n)
|
||||
logsnrs_min = _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale_min)
|
||||
logsnrs_max = _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale_max)
|
||||
logsnrs = t * logsnrs_min + (1 - t) * logsnrs_max
|
||||
return logsnrs
|
||||
|
||||
|
||||
def logsnrs_to_sigmas(logsnrs):
|
||||
return torch.sqrt(torch.sigmoid(-logsnrs))
|
||||
|
||||
|
||||
def rescale_zero_terminal_snr(betas):
|
||||
"""
|
||||
Rescale Schedule to Zero Terminal SNR
|
||||
"""
|
||||
# Convert betas to alphas_bar_sqrt
|
||||
alphas = 1 - betas
|
||||
alphas_bar = alphas.cumprod(0)
|
||||
alphas_bar_sqrt = alphas_bar.sqrt()
|
||||
|
||||
# Store old values. 8 alphas_bar_sqrt_0 = a
|
||||
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
|
||||
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
|
||||
# Shift so last timestep is zero.
|
||||
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||
# Scale so first timestep is back to old value.
|
||||
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||
|
||||
# Convert alphas_bar_sqrt to betas
|
||||
alphas_bar = alphas_bar_sqrt ** 2
|
||||
alphas = alphas_bar[1:] / alphas_bar[:-1]
|
||||
alphas = torch.cat([alphas_bar[0:1], alphas])
|
||||
betas = 1 - alphas
|
||||
return betas
|
||||
|
||||
@@ -0,0 +1,643 @@
|
||||
import os
|
||||
import re
|
||||
import os.path as osp
|
||||
import sys
|
||||
sys.path.insert(0, '/'.join(osp.realpath(__file__).split('/')[:-4]))
|
||||
import json
|
||||
import math
|
||||
import torch
|
||||
import pynvml
|
||||
import logging
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from tqdm import tqdm
|
||||
import torch.cuda.amp as amp
|
||||
from importlib import reload
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
import random
|
||||
from einops import rearrange
|
||||
import torchvision.transforms as T
|
||||
import torchvision.transforms.functional as TF
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
|
||||
from .default_config import cfg
|
||||
from .model.autoencoder import get_first_stage_encoding
|
||||
|
||||
import utils.transforms as data
|
||||
from utils.seed import setup_seed
|
||||
from utils.multi_port import find_free_port
|
||||
from utils.distributed import generalized_all_gather
|
||||
from utils.video_op import save_video_multiple_conditions_not_gif_horizontal_3col, save_video_multiple_conditions_not_gif_horizontal_1col
|
||||
from utils.registry_class import INFER_ENGINE, MODEL, EMBEDDER, AUTO_ENCODER, DIFFUSION
|
||||
|
||||
from copy import copy
|
||||
import cv2, pickle
|
||||
|
||||
|
||||
@INFER_ENGINE.register_function()
|
||||
def inference_animate_x_entrance(cfg_update, **kwargs):
|
||||
for k, v in cfg_update.items():
|
||||
if isinstance(v, dict) and k in cfg:
|
||||
cfg[k].update(v)
|
||||
else:
|
||||
cfg[k] = v
|
||||
|
||||
if not 'MASTER_ADDR' in os.environ:
|
||||
os.environ['MASTER_ADDR']='localhost'
|
||||
os.environ['MASTER_PORT']= find_free_port()
|
||||
cfg.pmi_rank = int(os.getenv('RANK', 0))
|
||||
cfg.pmi_world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
|
||||
if cfg.debug:
|
||||
cfg.gpus_per_machine = 1
|
||||
cfg.world_size = 1
|
||||
else:
|
||||
cfg.gpus_per_machine = torch.cuda.device_count()
|
||||
cfg.world_size = cfg.pmi_world_size * cfg.gpus_per_machine
|
||||
|
||||
if cfg.world_size == 1:
|
||||
worker(0, cfg, cfg_update)
|
||||
else:
|
||||
mp.spawn(worker, nprocs=cfg.gpus_per_machine, args=(cfg, cfg_update))
|
||||
return cfg
|
||||
|
||||
def process_single_pose_embedding(dwpose_source_data):
|
||||
|
||||
|
||||
bodies = dwpose_source_data['bodies']['candidate'][:18]
|
||||
|
||||
results = np.swapaxes(bodies, 0, 1) # (32, 2, 128)
|
||||
return results
|
||||
|
||||
def make_masked_images(imgs, masks):
|
||||
masked_imgs = []
|
||||
for i, mask in enumerate(masks):
|
||||
# concatenation
|
||||
masked_imgs.append(torch.cat([imgs[i] * (1 - mask), (1 - mask)], dim=1))
|
||||
return torch.stack(masked_imgs, dim=0)
|
||||
|
||||
def process_single_pose_embedding_katong(dwpose_source_data, index):
|
||||
|
||||
|
||||
bodies = dwpose_source_data['bodies'][index][:18]
|
||||
|
||||
|
||||
|
||||
results = bodies
|
||||
results = np.swapaxes(results, 0, 1) # (32, 2, 128)
|
||||
return results
|
||||
|
||||
def load_video_frames(ref_image_path, pose_file_path, original_driven_video_path, pose_embedding_key, train_trans, vit_transforms, train_trans_pose, max_frames=32, frame_interval = 1, resolution=[512, 768], get_first_frame=True, vit_resolution=[224, 224]):
|
||||
|
||||
pose_embedding_dim = 18
|
||||
|
||||
|
||||
for _ in range(5):
|
||||
# try:
|
||||
dwpose_all = {}
|
||||
frames_all = {}
|
||||
|
||||
original_driven_video_all = {}
|
||||
original_driven_video_frame_all = {}
|
||||
pose_embedding_all = {}
|
||||
|
||||
|
||||
|
||||
|
||||
# 打开文件(以二进制读取模式)
|
||||
with open(pose_embedding_key, 'rb') as file:
|
||||
# 使用 pickle.load() 方法读取字典
|
||||
loaded_data = pickle.load(file)
|
||||
|
||||
try:
|
||||
ref_pose_embedding_key = pose_embedding_key.replace(".pkl", "_ref_pose.pkl")
|
||||
with open(ref_pose_embedding_key, 'rb') as file:
|
||||
# 使用 pickle.load() 方法读取字典
|
||||
ref_loaded_data = pickle.load(file)
|
||||
ref_pose_embedding = process_single_pose_embedding(ref_loaded_data)
|
||||
|
||||
except:
|
||||
ref_pose_embedding = process_single_pose_embedding_katong(loaded_data, 0)
|
||||
|
||||
|
||||
first_image = True
|
||||
for ii_index in sorted(os.listdir(pose_file_path)):
|
||||
# ii_index = ii_index.strip()
|
||||
|
||||
if ii_index != "ref_pose.jpg":
|
||||
|
||||
dwpose_all[ii_index] = Image.open(pose_file_path+"/"+ii_index)
|
||||
frames_all[ii_index] = Image.fromarray(cv2.cvtColor(cv2.imread(ref_image_path),cv2.COLOR_BGR2RGB))
|
||||
|
||||
try:
|
||||
i_index = int(ii_index.split('.')[0])
|
||||
except:
|
||||
i_index = int(ii_index.split('.')[0].split('_')[1])
|
||||
try:
|
||||
pose_embedding_all[ii_index] = process_single_pose_embedding( loaded_data[i_index]) # (2, 128)
|
||||
except:
|
||||
pose_embedding_all[ii_index] = process_single_pose_embedding_katong( loaded_data, i_index) # (2, 128)
|
||||
|
||||
for ii_index in sorted(os.listdir(original_driven_video_path)):
|
||||
|
||||
original_driven_video_all[ii_index] = Image.open(original_driven_video_path+"/"+ii_index)
|
||||
|
||||
|
||||
# frames_all[ii_index] = Image.open(ref_image_path)
|
||||
pose_ref_path = os.path.join(pose_file_path, "ref_pose.jpg")
|
||||
if os.path.exists(pose_ref_path) == False:
|
||||
pose_ref_path = os.path.join(pose_file_path, os.listdir(pose_file_path)[0])
|
||||
pose_ref = Image.open(pose_ref_path)
|
||||
first_eq_ref = False
|
||||
|
||||
# sample max_frames poses for video generation
|
||||
stride = frame_interval
|
||||
_total_frame_num = len(frames_all)
|
||||
cover_frame_num = (stride * (max_frames-1)+1)
|
||||
if _total_frame_num < cover_frame_num:
|
||||
print('_total_frame_num is smaller than cover_frame_num, the sampled frame interval is changed')
|
||||
start_frame = 0 # we set start_frame = 0 because the pose alignment is performed on the first frame
|
||||
end_frame = _total_frame_num
|
||||
stride = max((_total_frame_num-1//(max_frames-1)),1)
|
||||
end_frame = stride*max_frames
|
||||
else:
|
||||
start_frame = 0 # we set start_frame = 0 because the pose alignment is performed on the first frame
|
||||
end_frame = start_frame + cover_frame_num
|
||||
|
||||
frame_list = []
|
||||
dwpose_list = []
|
||||
original_driven_video_list = []
|
||||
pose_embedding_list = []
|
||||
random_ref_frame = frames_all[list(frames_all.keys())[0]]
|
||||
if random_ref_frame.mode != 'RGB':
|
||||
random_ref_frame = random_ref_frame.convert('RGB')
|
||||
random_ref_dwpose = pose_ref
|
||||
if random_ref_dwpose.mode != 'RGB':
|
||||
random_ref_dwpose = random_ref_dwpose.convert('RGB')
|
||||
for i_index in range(start_frame, end_frame, stride):
|
||||
# import pdb; pdb.set_trace()
|
||||
if i_index == start_frame and first_eq_ref:
|
||||
# print("i_index == start_frame and first_eq_ref:")
|
||||
i_key = list(frames_all.keys())[i_index]
|
||||
i_frame = frames_all[i_key]
|
||||
|
||||
if i_frame.mode != 'RGB':
|
||||
i_frame = i_frame.convert('RGB')
|
||||
i_dwpose = frames_pose_ref
|
||||
if i_dwpose.mode != 'RGB':
|
||||
i_dwpose = i_dwpose.convert('RGB')
|
||||
frame_list.append(i_frame)
|
||||
dwpose_list.append(i_dwpose)
|
||||
else:
|
||||
# added
|
||||
|
||||
if first_eq_ref:
|
||||
i_index = i_index - stride
|
||||
# print("key = list(frames_all.keys())[i_index]")
|
||||
i_key = list(frames_all.keys())[i_index]
|
||||
i_frame = frames_all[i_key]
|
||||
if i_frame.mode != 'RGB':
|
||||
i_frame = i_frame.convert('RGB')
|
||||
i_dwpose = dwpose_all[i_key]
|
||||
# ii_index = ii_index.strip()
|
||||
# print(original_driven_video_all.keys())
|
||||
i_original_driven_video = original_driven_video_all[i_key.strip()]
|
||||
i_pose_embedding = pose_embedding_all[i_key]
|
||||
if i_dwpose.mode != 'RGB':
|
||||
i_dwpose = i_dwpose.convert('RGB')
|
||||
|
||||
if i_original_driven_video.mode != 'RGB':
|
||||
i_original_driven_video = i_original_driven_video.convert('RGB')
|
||||
frame_list.append(i_frame)
|
||||
dwpose_list.append(i_dwpose)
|
||||
original_driven_video_list.append(i_original_driven_video)
|
||||
|
||||
|
||||
pose_embedding_list.append(i_pose_embedding)
|
||||
|
||||
have_frames = len(frame_list)>0
|
||||
middle_indix = 0
|
||||
if have_frames:
|
||||
ref_frame = frame_list[middle_indix]
|
||||
|
||||
vit_frame = vit_transforms(ref_frame)
|
||||
random_ref_frame_tmp = train_trans_pose(random_ref_frame)
|
||||
random_ref_dwpose_tmp = train_trans_pose(random_ref_dwpose)
|
||||
|
||||
original_driven_video_data_tmp = torch.stack([vit_transforms(ss) for ss in original_driven_video_list], dim=0)
|
||||
|
||||
|
||||
ref_pose_embedding_tmp = torch.from_numpy(ref_pose_embedding)
|
||||
|
||||
misc_data_tmp = torch.stack([train_trans_pose(ss) for ss in frame_list], dim=0)
|
||||
video_data_tmp = torch.stack([train_trans(ss) for ss in frame_list], dim=0)
|
||||
dwpose_data_tmp = torch.stack([train_trans_pose(ss) for ss in dwpose_list], dim=0)
|
||||
|
||||
pose_embedding_tmp = torch.stack([torch.from_numpy(ss) for ss in pose_embedding_list], dim=0)
|
||||
|
||||
|
||||
video_data = torch.zeros(max_frames, 3, resolution[1], resolution[0])
|
||||
dwpose_data = torch.zeros(max_frames, 3, resolution[1], resolution[0])
|
||||
|
||||
original_driven_video_data = torch.zeros(max_frames, 3, 224, 224)
|
||||
|
||||
pose_embedding = torch.zeros(max_frames, 2, pose_embedding_dim)
|
||||
ref_pose_embedding = torch.zeros(max_frames, 2, pose_embedding_dim)
|
||||
|
||||
misc_data = torch.zeros(max_frames, 3, resolution[1], resolution[0])
|
||||
random_ref_frame_data = torch.zeros(max_frames, 3, resolution[1], resolution[0]) # [32, 3, 512, 768]
|
||||
random_ref_dwpose_data = torch.zeros(max_frames, 3, resolution[1], resolution[0])
|
||||
if have_frames:
|
||||
video_data[:len(frame_list), ...] = video_data_tmp
|
||||
misc_data[:len(frame_list), ...] = misc_data_tmp
|
||||
dwpose_data[:len(frame_list), ...] = dwpose_data_tmp
|
||||
|
||||
original_driven_video_data[:len(frame_list), ...] = original_driven_video_data_tmp
|
||||
|
||||
pose_embedding[:len(frame_list), ...] = pose_embedding_tmp
|
||||
# print("random_ref_frame_tmp.shape", random_ref_frame_tmp.shape)
|
||||
random_ref_frame_data[:,...] = random_ref_frame_tmp
|
||||
# print("random_ref_frame_data.shape", random_ref_frame_data.shape)
|
||||
random_ref_dwpose_data[:,...] = random_ref_dwpose_tmp
|
||||
|
||||
|
||||
ref_pose_embedding[:,...] = ref_pose_embedding_tmp
|
||||
|
||||
break
|
||||
|
||||
# except Exception as e:
|
||||
# logging.info('{} read video frame failed with error: {}'.format(pose_file_path, e))
|
||||
# continue
|
||||
|
||||
return vit_frame, video_data, misc_data, dwpose_data, random_ref_frame_data, random_ref_dwpose_data, pose_embedding, ref_pose_embedding, original_driven_video_data
|
||||
|
||||
|
||||
|
||||
def worker(gpu, cfg, cfg_update):
|
||||
'''
|
||||
Inference worker for each gpu
|
||||
'''
|
||||
for k, v in cfg_update.items():
|
||||
if isinstance(v, dict) and k in cfg:
|
||||
cfg[k].update(v)
|
||||
else:
|
||||
cfg[k] = v
|
||||
|
||||
cfg.gpu = gpu
|
||||
cfg.seed = int(cfg.seed)
|
||||
cfg.rank = cfg.pmi_rank * cfg.gpus_per_machine + gpu
|
||||
setup_seed(cfg.seed + cfg.rank)
|
||||
|
||||
if not cfg.debug:
|
||||
torch.cuda.set_device(gpu)
|
||||
torch.backends.cudnn.benchmark = True
|
||||
if hasattr(cfg, "CPU_CLIP_VAE") and cfg.CPU_CLIP_VAE:
|
||||
torch.backends.cudnn.benchmark = False
|
||||
dist.init_process_group(backend='nccl', world_size=cfg.world_size, rank=cfg.rank)
|
||||
|
||||
# [Log] Save logging and make log dir
|
||||
log_dir = generalized_all_gather(cfg.log_dir)[0]
|
||||
inf_name = osp.basename(cfg.cfg_file).split('.')[0]
|
||||
test_model = osp.basename(cfg.test_model).split('.')[0].split('_')[-1]
|
||||
|
||||
cfg.log_dir = osp.join(cfg.log_dir, '%s' % (inf_name))
|
||||
os.makedirs(cfg.log_dir, exist_ok=True)
|
||||
log_file = osp.join(cfg.log_dir, 'log_%02d.txt' % (cfg.rank))
|
||||
cfg.log_file = log_file
|
||||
reload(logging)
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='[%(asctime)s] %(levelname)s: %(message)s',
|
||||
handlers=[
|
||||
logging.FileHandler(filename=log_file),
|
||||
logging.StreamHandler(stream=sys.stdout)])
|
||||
logging.info(cfg)
|
||||
logging.info(f"Running Animate-X inference on gpu {gpu}")
|
||||
|
||||
# [Diffusion]
|
||||
diffusion = DIFFUSION.build(cfg.Diffusion)
|
||||
|
||||
# [Data] Data Transform
|
||||
train_trans = data.Compose([
|
||||
data.Resize(cfg.resolution),
|
||||
data.ToTensor(),
|
||||
data.Normalize(mean=cfg.mean, std=cfg.std)
|
||||
])
|
||||
|
||||
train_trans_pose = data.Compose([
|
||||
data.Resize(cfg.resolution),
|
||||
data.ToTensor(),
|
||||
]
|
||||
)
|
||||
|
||||
vit_transforms = T.Compose([
|
||||
data.Resize(cfg.vit_resolution),
|
||||
T.ToTensor(),
|
||||
T.Normalize(mean=cfg.vit_mean, std=cfg.vit_std)])
|
||||
|
||||
# [Model] embedder
|
||||
clip_encoder = EMBEDDER.build(cfg.embedder)
|
||||
clip_encoder.model.to(gpu)
|
||||
with torch.no_grad():
|
||||
_, _, zero_y = clip_encoder(text="")
|
||||
|
||||
|
||||
# [Model] auotoencoder
|
||||
autoencoder = AUTO_ENCODER.build(cfg.auto_encoder)
|
||||
autoencoder.eval() # freeze
|
||||
for param in autoencoder.parameters():
|
||||
param.requires_grad = False
|
||||
autoencoder.cuda()
|
||||
|
||||
# [Model] UNet
|
||||
if "config" in cfg.UNet:
|
||||
cfg.UNet["config"] = cfg
|
||||
cfg.UNet["zero_y"] = zero_y
|
||||
model = MODEL.build(cfg.UNet)
|
||||
state_dict = torch.load(cfg.test_model, map_location='cpu')
|
||||
if 'state_dict' in state_dict:
|
||||
state_dict = state_dict['state_dict']
|
||||
if 'step' in state_dict:
|
||||
resume_step = state_dict['step']
|
||||
else:
|
||||
resume_step = 0
|
||||
try:
|
||||
status = model.load_state_dict(state_dict, strict=False)
|
||||
except:
|
||||
|
||||
for key in list(state_dict.keys()):
|
||||
if 'pose_embedding_before.pos_embed.pos_table' in key:
|
||||
del state_dict[key]
|
||||
status = model.load_state_dict(state_dict, strict=False)
|
||||
logging.info('Load model from {} with status {}'.format(cfg.test_model, status))
|
||||
model = model.to(gpu)
|
||||
model.eval()
|
||||
if hasattr(cfg, "CPU_CLIP_VAE") and cfg.CPU_CLIP_VAE:
|
||||
model.to(torch.float16)
|
||||
else:
|
||||
model = DistributedDataParallel(model, device_ids=[gpu]) if not cfg.debug else model
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
|
||||
test_list = cfg.test_list_path
|
||||
|
||||
# test_list.reverse()
|
||||
|
||||
num_videos = len(test_list)
|
||||
|
||||
|
||||
|
||||
logging.info(f'There are {num_videos} videos. with {cfg.round} times')
|
||||
# test_list = [item for item in test_list for _ in range(cfg.round)]
|
||||
test_list = [item for _ in range(cfg.round) for item in test_list]
|
||||
|
||||
for idx, file_path in enumerate(test_list):
|
||||
cfg.frame_interval, ref_image_key, pose_seq_key, original_driven_video_seq_key, pose_embedding_key = file_path[0], file_path[1], file_path[2], file_path[3], file_path[4]
|
||||
|
||||
try:
|
||||
current_seed = file_path[5]
|
||||
except:
|
||||
current_seed = int(cfg.seed)
|
||||
manual_seed = int(current_seed + cfg.rank + idx//num_videos)
|
||||
setup_seed(manual_seed)
|
||||
|
||||
logging.info(f"[{idx}]/[{len(test_list)}] Begin to sample {ref_image_key}, pose sequence from {pose_seq_key} init seed {manual_seed} ...")
|
||||
|
||||
|
||||
vit_frame, video_data, misc_data, dwpose_data, random_ref_frame_data, random_ref_dwpose_data, pose_embedding, ref_pose_embedding, original_driven_video_data = load_video_frames(ref_image_key, pose_seq_key, original_driven_video_seq_key, pose_embedding_key, train_trans, vit_transforms, train_trans_pose, max_frames=cfg.max_frames, frame_interval =cfg.frame_interval, resolution=cfg.resolution)
|
||||
|
||||
|
||||
|
||||
original_driven_video_data = torch.cat([vit_frame.unsqueeze(0), original_driven_video_data], 0)
|
||||
|
||||
misc_data = misc_data.unsqueeze(0).to(gpu)
|
||||
vit_frame = vit_frame.unsqueeze(0).to(gpu)
|
||||
dwpose_data = dwpose_data.unsqueeze(0).to(gpu)
|
||||
original_driven_video_data = original_driven_video_data.unsqueeze(0).to(gpu)
|
||||
random_ref_frame_data = random_ref_frame_data.unsqueeze(0).to(gpu)
|
||||
random_ref_dwpose_data = random_ref_dwpose_data.unsqueeze(0).to(gpu)
|
||||
|
||||
pose_embedding = pose_embedding.unsqueeze(0).to(gpu)
|
||||
ref_pose_embedding = ref_pose_embedding[0:1].unsqueeze(0).to(gpu)
|
||||
|
||||
pose_embedding = torch.cat([ref_pose_embedding, pose_embedding], dim = 1)
|
||||
|
||||
# print("pose_embedding.shape: ", pose_embedding.shape)
|
||||
|
||||
### save for visualization
|
||||
misc_backups = copy(misc_data)
|
||||
frames_num = misc_data.shape[1]
|
||||
misc_backups = rearrange(misc_backups, 'b f c h w -> b c f h w')
|
||||
mv_data_video = []
|
||||
|
||||
|
||||
### local image (first frame)
|
||||
image_local = []
|
||||
if 'local_image' in cfg.video_compositions:
|
||||
frames_num = misc_data.shape[1]
|
||||
bs_vd_local = misc_data.shape[0]
|
||||
image_local = misc_data[:,:1].clone().repeat(1,frames_num,1,1,1)
|
||||
image_local_clone = rearrange(image_local, 'b f c h w -> b c f h w', b = bs_vd_local)
|
||||
image_local = rearrange(image_local, 'b f c h w -> b c f h w', b = bs_vd_local)
|
||||
|
||||
# no
|
||||
if hasattr(cfg, "latent_local_image") and cfg.latent_local_image:
|
||||
with torch.no_grad():
|
||||
temporal_length = frames_num
|
||||
# print("video_data[:,0].shape", video_data[:,0].shape) #
|
||||
encoder_posterior = autoencoder.encode(video_data[:,0])
|
||||
local_image_data = get_first_stage_encoding(encoder_posterior).detach()
|
||||
image_local = local_image_data.unsqueeze(1).repeat(1,temporal_length,1,1,1) # [10, 16, 4, 64, 40]
|
||||
# print("image_local.shape", image_local.shape) #
|
||||
|
||||
|
||||
### encode the video_data
|
||||
bs_vd = misc_data.shape[0]
|
||||
misc_data = rearrange(misc_data, 'b f c h w -> (b f) c h w')
|
||||
misc_data_list = torch.chunk(misc_data, misc_data.shape[0]//cfg.chunk_size,dim=0)
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
|
||||
random_ref_frame = []
|
||||
if 'randomref' in cfg.video_compositions:
|
||||
random_ref_frame_clone = rearrange(random_ref_frame_data, 'b f c h w -> b c f h w')
|
||||
if hasattr(cfg, "latent_random_ref") and cfg.latent_random_ref:
|
||||
|
||||
temporal_length = random_ref_frame_data.shape[1]
|
||||
encoder_posterior = autoencoder.encode(random_ref_frame_data[:,0].sub(0.5).div_(0.5))
|
||||
random_ref_frame_data = get_first_stage_encoding(encoder_posterior).detach()
|
||||
random_ref_frame_data = random_ref_frame_data.unsqueeze(1).repeat(1,temporal_length,1,1,1) # [10, 16, 4, 64, 40]
|
||||
# print("random_ref_frame_data.shape", random_ref_frame_data.shape) #
|
||||
random_ref_frame = rearrange(random_ref_frame_data, 'b f c h w -> b c f h w')
|
||||
|
||||
|
||||
if 'dwpose' in cfg.video_compositions:
|
||||
bs_vd_local = dwpose_data.shape[0]
|
||||
dwpose_data_clone = rearrange(dwpose_data.clone(), 'b f c h w -> b c f h w', b = bs_vd_local)
|
||||
if 'randomref_pose' in cfg.video_compositions:
|
||||
dwpose_data = torch.cat([random_ref_dwpose_data[:,:1], dwpose_data], dim=1)
|
||||
dwpose_data = rearrange(dwpose_data, 'b f c h w -> b c f h w', b = bs_vd_local)
|
||||
# print("dwpose_data = rearrange(dwpose_dat.shape", dwpose_data.shape) #
|
||||
|
||||
y_visual = []
|
||||
if 'image' in cfg.video_compositions:
|
||||
with torch.no_grad():
|
||||
vit_frame = vit_frame.squeeze(1)
|
||||
y_visual = clip_encoder.encode_image(vit_frame).unsqueeze(1) # [60, 1024]
|
||||
y_visual0 = y_visual.clone()
|
||||
|
||||
batch_size, seq_len = original_driven_video_data.shape[0], original_driven_video_data.shape[1]
|
||||
|
||||
original_driven_video_data = original_driven_video_data.reshape(batch_size*seq_len,3,224,224)
|
||||
original_driven_video_data_embedding = clip_encoder.encode_image(original_driven_video_data).unsqueeze(1) # [60, 1024]
|
||||
|
||||
# print("original_driven_video_data_embedding.shape: ", original_driven_video_data_embedding.shape)
|
||||
original_driven_video_data_embedding = original_driven_video_data_embedding.clone()
|
||||
|
||||
with amp.autocast(enabled=True):
|
||||
pynvml.nvmlInit()
|
||||
handle=pynvml.nvmlDeviceGetHandleByIndex(0)
|
||||
meminfo=pynvml.nvmlDeviceGetMemoryInfo(handle)
|
||||
cur_seed = torch.initial_seed()
|
||||
logging.info(f"Current seed {cur_seed} ...")
|
||||
|
||||
noise = torch.randn([1, 4, cfg.max_frames, int(cfg.resolution[1]/cfg.scale), int(cfg.resolution[0]/cfg.scale)])
|
||||
noise = noise.to(gpu)
|
||||
|
||||
if hasattr(cfg.Diffusion, "noise_strength"):
|
||||
b, c, f, _, _= noise.shape
|
||||
offset_noise = torch.randn(b, c, f, 1, 1, device=noise.device)
|
||||
noise = noise + cfg.Diffusion.noise_strength * offset_noise
|
||||
|
||||
|
||||
full_model_kwargs=[{
|
||||
'y': None,
|
||||
'pose_embeddings': [pose_embedding, original_driven_video_data_embedding],
|
||||
"local_image": None if len(image_local) == 0 else image_local[:],
|
||||
'image': None if len(y_visual) == 0 else y_visual0[:],
|
||||
'dwpose': None if len(dwpose_data) == 0 else dwpose_data[:],
|
||||
'randomref': None if len(random_ref_frame) == 0 else random_ref_frame[:],
|
||||
},
|
||||
{
|
||||
'y': None,
|
||||
"local_image": None,
|
||||
'image': None,
|
||||
'randomref': None,
|
||||
'dwpose': None,
|
||||
"pose_embeddings": None,
|
||||
}]
|
||||
|
||||
# for visualization
|
||||
full_model_kwargs_vis =[{
|
||||
'y': None,
|
||||
"local_image": None if len(image_local) == 0 else image_local_clone[:],
|
||||
'image': None,
|
||||
'pose_embeddings': [pose_embedding, original_driven_video_data_embedding],
|
||||
'dwpose': None if len(dwpose_data_clone) == 0 else dwpose_data_clone[:],
|
||||
'randomref': None if len(random_ref_frame) == 0 else random_ref_frame_clone[:, :3],
|
||||
},
|
||||
{
|
||||
'y': None,
|
||||
"local_image": None,
|
||||
'image': None,
|
||||
'randomref': None,
|
||||
'dwpose': None,
|
||||
"pose_embeddings": None,
|
||||
}]
|
||||
|
||||
|
||||
partial_keys = [
|
||||
['image', 'randomref', "dwpose","pose_embeddings"],
|
||||
]
|
||||
if hasattr(cfg, "partial_keys") and cfg.partial_keys:
|
||||
partial_keys = cfg.partial_keys
|
||||
|
||||
for partial_keys_one in partial_keys:
|
||||
model_kwargs_one = prepare_model_kwargs(partial_keys = partial_keys_one,
|
||||
full_model_kwargs = full_model_kwargs,
|
||||
use_fps_condition = cfg.use_fps_condition)
|
||||
|
||||
|
||||
|
||||
|
||||
model_kwargs_one_vis = prepare_model_kwargs(partial_keys = partial_keys_one,
|
||||
full_model_kwargs = full_model_kwargs_vis,
|
||||
use_fps_condition = cfg.use_fps_condition)
|
||||
noise_one = noise
|
||||
|
||||
if hasattr(cfg, "CPU_CLIP_VAE") and cfg.CPU_CLIP_VAE:
|
||||
clip_encoder.cpu() # add this line
|
||||
autoencoder.cpu() # add this line
|
||||
torch.cuda.empty_cache() # add this line
|
||||
|
||||
video_data = diffusion.ddim_sample_loop(
|
||||
noise=noise_one,
|
||||
model=model.eval(),
|
||||
model_kwargs=model_kwargs_one,
|
||||
guide_scale=cfg.guide_scale,
|
||||
ddim_timesteps=cfg.ddim_timesteps,
|
||||
eta=0.0)
|
||||
|
||||
# print("video_data = diffusion.ddim_sample_", video_data.shape) #torch.Size([1, 4, 32, 96, 64])
|
||||
|
||||
if hasattr(cfg, "CPU_CLIP_VAE") and cfg.CPU_CLIP_VAE:
|
||||
# if run forward of autoencoder or clip_encoder second times, load them again
|
||||
clip_encoder.cuda()
|
||||
autoencoder.cuda()
|
||||
video_data = 1. / cfg.scale_factor * video_data
|
||||
video_data = rearrange(video_data, 'b c f h w -> (b f) c h w')
|
||||
chunk_size = min(cfg.decoder_bs, video_data.shape[0])
|
||||
video_data_list = torch.chunk(video_data, video_data.shape[0]//chunk_size, dim=0)
|
||||
decode_data = []
|
||||
for vd_data in video_data_list:
|
||||
gen_frames = autoencoder.decode(vd_data)
|
||||
decode_data.append(gen_frames)
|
||||
video_data = torch.cat(decode_data, dim=0)
|
||||
video_data = rearrange(video_data, '(b f) c h w -> b c f h w', b = cfg.batch_size).float()
|
||||
|
||||
text_size = cfg.resolution[-1]
|
||||
cap_name = re.sub(r'[^\w\s]', '', ref_image_key.split("/")[-1].split('.')[0]) # .replace(' ', '_')
|
||||
pose_name = re.sub(r'[^\w\s]', '', pose_seq_key.split("/")[-1].split('.')[0])
|
||||
name = f'seed_{cur_seed}'
|
||||
file_name = f'{cap_name}_{pose_name}_{name}_rank_{cfg.world_size:02d}_{cfg.rank:02d}_{idx:02d}_{cfg.resolution[1]}x{cfg.resolution[0]}.mp4'
|
||||
local_path = os.path.join(cfg.log_dir, f'{file_name}')
|
||||
local_path_1col = os.path.join(cfg.log_dir, f'{file_name[:-4]}_results_1col.mp4')
|
||||
os.makedirs(os.path.dirname(local_path), exist_ok=True)
|
||||
captions = "human"
|
||||
del model_kwargs_one_vis[0][list(model_kwargs_one_vis[0].keys())[0]]
|
||||
del model_kwargs_one_vis[1][list(model_kwargs_one_vis[1].keys())[0]]
|
||||
|
||||
del model_kwargs_one_vis[0]["pose_embeddings"]
|
||||
del model_kwargs_one_vis[1]["pose_embeddings"]
|
||||
|
||||
|
||||
|
||||
save_video_multiple_conditions_not_gif_horizontal_3col(local_path, video_data.cpu(), model_kwargs_one_vis, misc_backups,
|
||||
cfg.mean, cfg.std, nrow=1, save_fps=cfg.save_fps)
|
||||
|
||||
save_video_multiple_conditions_not_gif_horizontal_1col(local_path_1col, video_data.cpu(), model_kwargs_one_vis, misc_backups,
|
||||
cfg.mean, cfg.std, nrow=1, save_fps=cfg.save_fps)
|
||||
logging.info(f'video saved in {local_path}!')
|
||||
|
||||
|
||||
logging.info('Congratulations! The inference is completed!')
|
||||
# synchronize to finish some processes
|
||||
if not cfg.debug:
|
||||
torch.cuda.synchronize()
|
||||
dist.barrier()
|
||||
|
||||
def prepare_model_kwargs(partial_keys, full_model_kwargs, use_fps_condition=False):
|
||||
|
||||
if use_fps_condition is True:
|
||||
partial_keys.append('fps')
|
||||
|
||||
partial_model_kwargs = [{}, {}]
|
||||
for partial_key in partial_keys:
|
||||
partial_model_kwargs[0][partial_key] = full_model_kwargs[0][partial_key]
|
||||
partial_model_kwargs[1][partial_key] = full_model_kwargs[1][partial_key]
|
||||
|
||||
return partial_model_kwargs
|
||||
@@ -0,0 +1,225 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from .transformer import (
|
||||
TransformerEncoder,
|
||||
TransformerEncoderLayer,
|
||||
PositionalEncoding,
|
||||
TransformerDecoderLayer,
|
||||
TransformerDecoder,
|
||||
)
|
||||
|
||||
class ImageProjModel(nn.Module):
|
||||
"""Projection Model"""
|
||||
def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4):
|
||||
super().__init__()
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.clip_extra_context_tokens = clip_extra_context_tokens
|
||||
self.proj = nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
|
||||
self.norm = nn.LayerNorm(cross_attention_dim)
|
||||
|
||||
def forward(self, image_embeds):
|
||||
#embeds = image_embeds
|
||||
embeds = image_embeds.type(list(self.proj.parameters())[0].dtype)
|
||||
clip_extra_context_tokens = self.proj(embeds).reshape(-1, self.clip_extra_context_tokens, self.cross_attention_dim)
|
||||
clip_extra_context_tokens = self.norm(clip_extra_context_tokens)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
|
||||
# FFN
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
#(bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
||||
x = x.view(bs, length, heads, -1)
|
||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
||||
x = x.transpose(1, 2)
|
||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head**-0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
|
||||
class Resampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=1024,
|
||||
depth=8,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
num_queries=8,
|
||||
embedding_dim=1024,
|
||||
output_dim=1024,
|
||||
ff_mult=4,
|
||||
our_weight = 1,
|
||||
video_length=None, # using frame-wise version or not
|
||||
):
|
||||
super().__init__()
|
||||
## queries for a single frame / image
|
||||
self.num_queries = num_queries
|
||||
self.video_length = video_length
|
||||
self.our_weight = our_weight
|
||||
|
||||
## <num_queries> queries for each frame
|
||||
if video_length is not None:
|
||||
num_queries = num_queries * video_length
|
||||
|
||||
self.latents = nn.Parameter(torch.randn(1, 4, dim) / dim**0.5)
|
||||
|
||||
|
||||
# print('self.latents.shape:', self.latents.shape)
|
||||
|
||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
||||
self.proj_out = nn.Linear(dim, output_dim)
|
||||
self.norm_out = nn.LayerNorm(output_dim)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x, latents):
|
||||
# import pdb; pdb.set_trace()
|
||||
latents = self.latents.repeat(x.size(0), 1, 1) * self.our_weight+ latents ## B (T L) C
|
||||
x = self.proj_in(x)
|
||||
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
|
||||
latents = self.proj_out(latents)
|
||||
latents = self.norm_out(latents) # B L C or B (T L) C
|
||||
|
||||
return latents
|
||||
|
||||
def _reset_parameters(model):
|
||||
for p in model.parameters():
|
||||
if p.dim() > 1:
|
||||
nn.init.xavier_uniform_(p)
|
||||
|
||||
class DwposeEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model=1024,
|
||||
nhead=8,
|
||||
num_encoder_layers=3,
|
||||
dim_feedforward=1024,
|
||||
dropout=0.1,
|
||||
activation="relu",
|
||||
normalize_before=False,
|
||||
pos_embed_len=33,
|
||||
input_dim=36,
|
||||
num_tokens=4,
|
||||
|
||||
):
|
||||
super().__init__()
|
||||
encoder_layer = TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, activation, normalize_before)
|
||||
encoder_norm = nn.LayerNorm(d_model) if normalize_before else None
|
||||
self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers, encoder_norm)
|
||||
_reset_parameters(self.encoder)
|
||||
|
||||
self.pos_embed = PositionalEncoding(d_model, pos_embed_len)
|
||||
|
||||
increase_embed_dim = [nn.Linear(input_dim, d_model//2), nn.Linear(d_model//2, d_model)]
|
||||
self.increase_embed_dim = nn.Sequential(*increase_embed_dim)
|
||||
|
||||
self.pre_image_condition_p = nn.Sequential(
|
||||
nn.Linear(d_model, d_model),
|
||||
nn.SiLU(),
|
||||
nn.Linear(d_model, d_model*num_tokens))
|
||||
|
||||
nn.init.zeros_(self.pre_image_condition_p[-1].weight)
|
||||
nn.init.zeros_(self.pre_image_condition_p[-1].bias)
|
||||
|
||||
self.d_model = d_model
|
||||
def forward(self, x, pad_mask=None):
|
||||
"""
|
||||
|
||||
Args:
|
||||
x (_type_): (B, num_frames(L), C_exp)
|
||||
pad_mask: (B, num_frames)
|
||||
|
||||
Returns:
|
||||
style_code: (B, C_model)
|
||||
"""
|
||||
batch_size, seq_len = x.shape[0], x.shape[1]
|
||||
x = self.increase_embed_dim(x)
|
||||
# (B, L, C)
|
||||
x = x.permute(1, 0, 2)
|
||||
# (L, B, C)
|
||||
|
||||
pos = self.pos_embed(x.shape[0])
|
||||
pos = pos.permute(1, 0, 2)
|
||||
# (L, 1, C)
|
||||
|
||||
style = self.encoder(x, pos=pos, src_key_padding_mask=pad_mask)
|
||||
# (L, B, C)
|
||||
style = style.reshape(batch_size* seq_len,1,self.d_model)
|
||||
style = self.pre_image_condition_p(style)
|
||||
# style = style.permute(1, 0, 2)
|
||||
style = style.reshape(batch_size* seq_len,-1,self.d_model)
|
||||
|
||||
return style
|
||||
@@ -0,0 +1 @@
|
||||
from .unet_animate_x import *
|
||||
@@ -0,0 +1,690 @@
|
||||
import torch
|
||||
import logging
|
||||
import collections
|
||||
import numpy as np
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from utils.registry_class import AUTO_ENCODER,DISTRIBUTION
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
# swish
|
||||
return x*torch.sigmoid(x)
|
||||
|
||||
def Normalize(in_channels, num_groups=32):
|
||||
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def get_first_stage_encoding(encoder_posterior, scale_factor=0.18215):
|
||||
if isinstance(encoder_posterior, DiagonalGaussianDistribution):
|
||||
z = encoder_posterior.sample()
|
||||
elif isinstance(encoder_posterior, torch.Tensor):
|
||||
z = encoder_posterior
|
||||
else:
|
||||
raise NotImplementedError(f"encoder_posterior of type '{type(encoder_posterior)}' not yet implemented")
|
||||
return scale_factor * z
|
||||
|
||||
|
||||
@AUTO_ENCODER.register_class()
|
||||
class AutoencoderKL(nn.Module):
|
||||
def __init__(self,
|
||||
ddconfig,
|
||||
embed_dim,
|
||||
pretrained=None,
|
||||
ignore_keys=[],
|
||||
image_key="image",
|
||||
colorize_nlabels=None,
|
||||
monitor=None,
|
||||
ema_decay=None,
|
||||
learn_logvar=False,
|
||||
use_vid_decoder=False,
|
||||
**kwargs):
|
||||
super().__init__()
|
||||
self.learn_logvar = learn_logvar
|
||||
self.image_key = image_key
|
||||
self.encoder = Encoder(**ddconfig)
|
||||
self.decoder = Decoder(**ddconfig)
|
||||
assert ddconfig["double_z"]
|
||||
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
|
||||
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
||||
self.embed_dim = embed_dim
|
||||
if colorize_nlabels is not None:
|
||||
assert type(colorize_nlabels)==int
|
||||
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
|
||||
if monitor is not None:
|
||||
self.monitor = monitor
|
||||
|
||||
self.use_ema = ema_decay is not None
|
||||
|
||||
if pretrained is not None:
|
||||
self.init_from_ckpt(pretrained, ignore_keys=ignore_keys)
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
||||
sd = torch.load(path, map_location="cpu")["state_dict"]
|
||||
keys = list(sd.keys())
|
||||
sd_new = collections.OrderedDict()
|
||||
for k in keys:
|
||||
if k.find('first_stage_model') >= 0:
|
||||
k_new = k.split('first_stage_model.')[-1]
|
||||
sd_new[k_new] = sd[k]
|
||||
self.load_state_dict(sd_new, strict=True)
|
||||
logging.info(f"Restored from {path}")
|
||||
|
||||
def on_train_batch_end(self, *args, **kwargs):
|
||||
if self.use_ema:
|
||||
self.model_ema(self)
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def encode_firsr_stage(self, x, scale_factor=1.0):
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
z = get_first_stage_encoding(posterior, scale_factor)
|
||||
return z
|
||||
|
||||
def encode_ms(self, x):
|
||||
hs = self.encoder(x, True)
|
||||
h = hs[-1]
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
hs[-1] = h
|
||||
return hs
|
||||
|
||||
def decode(self, z, **kwargs):
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z, **kwargs)
|
||||
return dec
|
||||
|
||||
|
||||
def forward(self, input, sample_posterior=True):
|
||||
posterior = self.encode(input)
|
||||
if sample_posterior:
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z)
|
||||
return dec, posterior
|
||||
|
||||
def get_input(self, batch, k):
|
||||
x = batch[k]
|
||||
if len(x.shape) == 3:
|
||||
x = x[..., None]
|
||||
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
|
||||
return x
|
||||
|
||||
def get_last_layer(self):
|
||||
return self.decoder.conv_out.weight
|
||||
|
||||
@torch.no_grad()
|
||||
def log_images(self, batch, only_inputs=False, log_ema=False, **kwargs):
|
||||
log = dict()
|
||||
x = self.get_input(batch, self.image_key)
|
||||
x = x.to(self.device)
|
||||
if not only_inputs:
|
||||
xrec, posterior = self(x)
|
||||
if x.shape[1] > 3:
|
||||
# colorize with random projection
|
||||
assert xrec.shape[1] > 3
|
||||
x = self.to_rgb(x)
|
||||
xrec = self.to_rgb(xrec)
|
||||
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
|
||||
log["reconstructions"] = xrec
|
||||
if log_ema or self.use_ema:
|
||||
with self.ema_scope():
|
||||
xrec_ema, posterior_ema = self(x)
|
||||
if x.shape[1] > 3:
|
||||
# colorize with random projection
|
||||
assert xrec_ema.shape[1] > 3
|
||||
xrec_ema = self.to_rgb(xrec_ema)
|
||||
log["samples_ema"] = self.decode(torch.randn_like(posterior_ema.sample()))
|
||||
log["reconstructions_ema"] = xrec_ema
|
||||
log["inputs"] = x
|
||||
return log
|
||||
|
||||
def to_rgb(self, x):
|
||||
assert self.image_key == "segmentation"
|
||||
if not hasattr(self, "colorize"):
|
||||
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
|
||||
x = F.conv2d(x, weight=self.colorize)
|
||||
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
|
||||
return x
|
||||
|
||||
|
||||
@AUTO_ENCODER.register_class()
|
||||
class AutoencoderVideo(AutoencoderKL):
|
||||
def __init__(self,
|
||||
ddconfig,
|
||||
embed_dim,
|
||||
pretrained=None,
|
||||
ignore_keys=[],
|
||||
image_key="image",
|
||||
colorize_nlabels=None,
|
||||
monitor=None,
|
||||
ema_decay=None,
|
||||
use_vid_decoder=True,
|
||||
learn_logvar=False,
|
||||
**kwargs):
|
||||
use_vid_decoder = True
|
||||
super().__init__(ddconfig, embed_dim, pretrained, ignore_keys, image_key, colorize_nlabels, monitor, ema_decay, learn_logvar, use_vid_decoder, **kwargs)
|
||||
|
||||
def decode(self, z, **kwargs):
|
||||
# z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z, **kwargs)
|
||||
return dec
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
# moments = self.quant_conv(h)
|
||||
moments = h
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
|
||||
class IdentityFirstStage(torch.nn.Module):
|
||||
def __init__(self, *args, vq_interface=False, **kwargs):
|
||||
self.vq_interface = vq_interface
|
||||
super().__init__()
|
||||
|
||||
def encode(self, x, *args, **kwargs):
|
||||
return x
|
||||
|
||||
def decode(self, x, *args, **kwargs):
|
||||
return x
|
||||
|
||||
def quantize(self, x, *args, **kwargs):
|
||||
if self.vq_interface:
|
||||
return x, None, [None, None, None]
|
||||
return x
|
||||
|
||||
def forward(self, x, *args, **kwargs):
|
||||
return x
|
||||
|
||||
|
||||
|
||||
@DISTRIBUTION.register_class()
|
||||
class DiagonalGaussianDistribution(object):
|
||||
def __init__(self, parameters, deterministic=False):
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if self.deterministic:
|
||||
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
|
||||
|
||||
def sample(self):
|
||||
x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
|
||||
return x
|
||||
|
||||
def kl(self, other=None):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.])
|
||||
else:
|
||||
if other is None:
|
||||
return 0.5 * torch.sum(torch.pow(self.mean, 2)
|
||||
+ self.var - 1.0 - self.logvar,
|
||||
dim=[1, 2, 3])
|
||||
else:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean - other.mean, 2) / other.var
|
||||
+ self.var / other.var - 1.0 - self.logvar + other.logvar,
|
||||
dim=[1, 2, 3])
|
||||
|
||||
def nll(self, sample, dims=[1,2,3]):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
return 0.5 * torch.sum(
|
||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims)
|
||||
|
||||
def mode(self):
|
||||
return self.mean
|
||||
|
||||
|
||||
# -------------------------------modules--------------------------------
|
||||
|
||||
class Downsample(nn.Module):
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
# no asymmetric padding in torch conv, must do it ourselves
|
||||
self.conv = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=0)
|
||||
|
||||
def forward(self, x):
|
||||
if self.with_conv:
|
||||
pad = (0,1,0,1)
|
||||
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
|
||||
x = self.conv(x)
|
||||
else:
|
||||
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
|
||||
return x
|
||||
|
||||
class ResnetBlock(nn.Module):
|
||||
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
|
||||
dropout, temb_channels=512):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
|
||||
self.norm1 = Normalize(in_channels)
|
||||
self.conv1 = torch.nn.Conv2d(in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
if temb_channels > 0:
|
||||
self.temb_proj = torch.nn.Linear(temb_channels,
|
||||
out_channels)
|
||||
self.norm2 = Normalize(out_channels)
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = torch.nn.Conv2d(out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
if self.in_channels != self.out_channels:
|
||||
if self.use_conv_shortcut:
|
||||
self.conv_shortcut = torch.nn.Conv2d(in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
else:
|
||||
self.nin_shortcut = torch.nn.Conv2d(in_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
|
||||
def forward(self, x, temb):
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
|
||||
if temb is not None:
|
||||
h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
|
||||
|
||||
h = self.norm2(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.dropout(h)
|
||||
h = self.conv2(h)
|
||||
|
||||
if self.in_channels != self.out_channels:
|
||||
if self.use_conv_shortcut:
|
||||
x = self.conv_shortcut(x)
|
||||
else:
|
||||
x = self.nin_shortcut(x)
|
||||
|
||||
return x+h
|
||||
|
||||
|
||||
class AttnBlock(nn.Module):
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.k = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.v = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.proj_out = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
|
||||
def forward(self, x):
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b,c,h,w = q.shape
|
||||
q = q.reshape(b,c,h*w)
|
||||
q = q.permute(0,2,1) # b,hw,c
|
||||
k = k.reshape(b,c,h*w) # b,c,hw
|
||||
w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
|
||||
w_ = w_ * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = v.reshape(b,c,h*w)
|
||||
w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q)
|
||||
h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
|
||||
h_ = h_.reshape(b,c,h,w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x+h_
|
||||
|
||||
class AttnBlock(nn.Module):
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.k = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.v = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.proj_out = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
|
||||
def forward(self, x):
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b,c,h,w = q.shape
|
||||
q = q.reshape(b,c,h*w)
|
||||
q = q.permute(0,2,1) # b,hw,c
|
||||
k = k.reshape(b,c,h*w) # b,c,hw
|
||||
w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
|
||||
w_ = w_ * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = v.reshape(b,c,h*w)
|
||||
w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q)
|
||||
h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
|
||||
h_ = h_.reshape(b,c,h,w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x+h_
|
||||
|
||||
class Upsample(nn.Module):
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
self.conv = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
def forward(self, x):
|
||||
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
|
||||
if self.with_conv:
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class Downsample(nn.Module):
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
# no asymmetric padding in torch conv, must do it ourselves
|
||||
self.conv = torch.nn.Conv2d(in_channels,
|
||||
in_channels,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=0)
|
||||
|
||||
def forward(self, x):
|
||||
if self.with_conv:
|
||||
pad = (0,1,0,1)
|
||||
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
|
||||
x = self.conv(x)
|
||||
else:
|
||||
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
|
||||
return x
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
|
||||
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
|
||||
resolution, z_channels, double_z=True, use_linear_attn=False, attn_type="vanilla",
|
||||
**ignore_kwargs):
|
||||
super().__init__()
|
||||
if use_linear_attn: attn_type = "linear"
|
||||
self.ch = ch
|
||||
self.temb_ch = 0
|
||||
self.num_resolutions = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.resolution = resolution
|
||||
self.in_channels = in_channels
|
||||
|
||||
# downsampling
|
||||
self.conv_in = torch.nn.Conv2d(in_channels,
|
||||
self.ch,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
curr_res = resolution
|
||||
in_ch_mult = (1,)+tuple(ch_mult)
|
||||
self.in_ch_mult = in_ch_mult
|
||||
self.down = nn.ModuleList()
|
||||
for i_level in range(self.num_resolutions):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = ch*in_ch_mult[i_level]
|
||||
block_out = ch*ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks):
|
||||
block.append(ResnetBlock(in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout))
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if i_level != self.num_resolutions-1:
|
||||
down.downsample = Downsample(block_in, resamp_with_conv)
|
||||
curr_res = curr_res // 2
|
||||
self.down.append(down)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout)
|
||||
self.mid.attn_1 = AttnBlock(block_in)
|
||||
self.mid.block_2 = ResnetBlock(in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(block_in,
|
||||
2*z_channels if double_z else z_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
def forward(self, x, return_feat=False):
|
||||
# timestep embedding
|
||||
temb = None
|
||||
|
||||
# downsampling
|
||||
hs = [self.conv_in(x)]
|
||||
for i_level in range(self.num_resolutions):
|
||||
for i_block in range(self.num_res_blocks):
|
||||
h = self.down[i_level].block[i_block](hs[-1], temb)
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
hs.append(h)
|
||||
if i_level != self.num_resolutions-1:
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
|
||||
# middle
|
||||
h = hs[-1]
|
||||
h = self.mid.block_1(h, temb)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h, temb)
|
||||
|
||||
# end
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
if return_feat:
|
||||
hs[-1] = h
|
||||
return hs
|
||||
else:
|
||||
return h
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
|
||||
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
|
||||
resolution, z_channels, give_pre_end=False, tanh_out=False, use_linear_attn=False,
|
||||
attn_type="vanilla", **ignorekwargs):
|
||||
super().__init__()
|
||||
if use_linear_attn: attn_type = "linear"
|
||||
self.ch = ch
|
||||
self.temb_ch = 0
|
||||
self.num_resolutions = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.resolution = resolution
|
||||
self.in_channels = in_channels
|
||||
self.give_pre_end = give_pre_end
|
||||
self.tanh_out = tanh_out
|
||||
|
||||
# compute in_ch_mult, block_in and curr_res at lowest res
|
||||
in_ch_mult = (1,)+tuple(ch_mult)
|
||||
block_in = ch*ch_mult[self.num_resolutions-1]
|
||||
curr_res = resolution // 2**(self.num_resolutions-1)
|
||||
self.z_shape = (1,z_channels, curr_res, curr_res)
|
||||
# logging.info("Working with z of shape {} = {} dimensions.".format(self.z_shape, np.prod(self.z_shape)))
|
||||
|
||||
# z to block_in
|
||||
self.conv_in = torch.nn.Conv2d(z_channels,
|
||||
block_in,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout)
|
||||
self.mid.attn_1 = AttnBlock(block_in)
|
||||
self.mid.block_2 = ResnetBlock(in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout)
|
||||
|
||||
# upsampling
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = ch*ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks+1):
|
||||
block.append(ResnetBlock(in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout))
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if i_level != 0:
|
||||
up.upsample = Upsample(block_in, resamp_with_conv)
|
||||
curr_res = curr_res * 2
|
||||
self.up.insert(0, up) # prepend to get consistent order
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(block_in,
|
||||
out_ch,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
def forward(self, z, **kwargs):
|
||||
#assert z.shape[1:] == self.z_shape[1:]
|
||||
self.last_z_shape = z.shape
|
||||
|
||||
# timestep embedding
|
||||
temb = None
|
||||
|
||||
# z to block_in
|
||||
h = self.conv_in(z)
|
||||
|
||||
# middle
|
||||
h = self.mid.block_1(h, temb)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h, temb)
|
||||
|
||||
# upsampling
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
for i_block in range(self.num_res_blocks+1):
|
||||
h = self.up[i_level].block[i_block](h, temb)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h)
|
||||
if i_level != 0:
|
||||
h = self.up[i_level].upsample(h)
|
||||
|
||||
# end
|
||||
if self.give_pre_end:
|
||||
return h
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
if self.tanh_out:
|
||||
h = torch.tanh(h)
|
||||
return h
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
import os
|
||||
import torch
|
||||
import logging
|
||||
import open_clip
|
||||
import numpy as np
|
||||
import torch.nn as nn
|
||||
import torchvision.transforms as T
|
||||
|
||||
from utils.registry_class import EMBEDDER
|
||||
import kornia
|
||||
|
||||
|
||||
@EMBEDDER.register_class()
|
||||
class FrozenOpenCLIPEmbedder(nn.Module):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
LAYERS = [
|
||||
#"pooled",
|
||||
"last",
|
||||
"penultimate"
|
||||
]
|
||||
def __init__(self, pretrained, arch="ViT-H-14", device="cuda", max_length=77,
|
||||
freeze=True, layer="last"):
|
||||
super().__init__()
|
||||
assert layer in self.LAYERS
|
||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=pretrained)
|
||||
del model.visual
|
||||
self.model = model
|
||||
|
||||
self.device = device
|
||||
self.max_length = max_length
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
if self.layer == "last":
|
||||
self.layer_idx = 0
|
||||
elif self.layer == "penultimate":
|
||||
self.layer_idx = 1
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, text):
|
||||
tokens = open_clip.tokenize(text)
|
||||
z = self.encode_with_transformer(tokens.to(self.device))
|
||||
return z
|
||||
|
||||
def encode_with_transformer(self, text):
|
||||
x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model]
|
||||
x = x + self.model.positional_embedding
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.model.ln_final(x)
|
||||
return x
|
||||
|
||||
def text_transformer_forward(self, x: torch.Tensor, attn_mask = None):
|
||||
for i, r in enumerate(self.model.transformer.resblocks):
|
||||
if i == len(self.model.transformer.resblocks) - self.layer_idx:
|
||||
break
|
||||
if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting():
|
||||
x = checkpoint(r, x, attn_mask)
|
||||
else:
|
||||
x = r(x, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
|
||||
@EMBEDDER.register_class()
|
||||
class FrozenOpenCLIPVisualEmbedder(nn.Module):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
LAYERS = [
|
||||
#"pooled",
|
||||
"last",
|
||||
"penultimate"
|
||||
]
|
||||
def __init__(self, pretrained, vit_resolution=(224, 224), arch="ViT-H-14", device="cuda", max_length=77,
|
||||
freeze=True, layer="last"):
|
||||
super().__init__()
|
||||
assert layer in self.LAYERS
|
||||
model, _, preprocess = open_clip.create_model_and_transforms(
|
||||
arch, device=torch.device('cpu'), pretrained=pretrained)
|
||||
|
||||
del model.transformer
|
||||
self.model = model
|
||||
data_white = np.ones((vit_resolution[0], vit_resolution[1], 3), dtype=np.uint8)*255
|
||||
self.white_image = preprocess(T.ToPILImage()(data_white)).unsqueeze(0)
|
||||
|
||||
self.device = device
|
||||
self.max_length = max_length # 77
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer # 'penultimate'
|
||||
if self.layer == "last":
|
||||
self.layer_idx = 0
|
||||
elif self.layer == "penultimate":
|
||||
self.layer_idx = 1
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, image):
|
||||
# tokens = open_clip.tokenize(text)
|
||||
z = self.model.encode_image(image.to(self.device))
|
||||
return z
|
||||
|
||||
def encode_with_transformer(self, text):
|
||||
x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model]
|
||||
x = x + self.model.positional_embedding
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.model.ln_final(x)
|
||||
|
||||
return x
|
||||
|
||||
def text_transformer_forward(self, x: torch.Tensor, attn_mask = None):
|
||||
for i, r in enumerate(self.model.transformer.resblocks):
|
||||
if i == len(self.model.transformer.resblocks) - self.layer_idx:
|
||||
break
|
||||
if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting():
|
||||
x = checkpoint(r, x, attn_mask)
|
||||
else:
|
||||
x = r(x, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
|
||||
|
||||
@EMBEDDER.register_class()
|
||||
class FrozenOpenCLIPTextVisualEmbedder(nn.Module):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
LAYERS = [
|
||||
#"pooled",
|
||||
"last",
|
||||
"penultimate"
|
||||
]
|
||||
def __init__(self, pretrained, arch="ViT-H-14", device="cuda", max_length=77,
|
||||
freeze=True, layer="last", **kwargs):
|
||||
super().__init__()
|
||||
assert layer in self.LAYERS
|
||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=pretrained)
|
||||
self.model = model
|
||||
|
||||
self.device = device
|
||||
self.max_length = max_length
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
if self.layer == "last":
|
||||
self.layer_idx = 0
|
||||
elif self.layer == "penultimate":
|
||||
self.layer_idx = 1
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
|
||||
def forward(self, image=None, text=None):
|
||||
|
||||
xi = self.model.encode_image(image.to(self.device)) if image is not None else None
|
||||
tokens = open_clip.tokenize(text)
|
||||
xt, x = self.encode_with_transformer(tokens.to(self.device))
|
||||
return xi, xt, x
|
||||
|
||||
def encode_with_transformer(self, text):
|
||||
x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model]
|
||||
x = x + self.model.positional_embedding
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.model.ln_final(x)
|
||||
xt = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.model.text_projection
|
||||
return xt, x
|
||||
|
||||
|
||||
def encode_image(self, image):
|
||||
return self.model.visual(image)
|
||||
|
||||
def text_transformer_forward(self, x: torch.Tensor, attn_mask = None):
|
||||
for i, r in enumerate(self.model.transformer.resblocks):
|
||||
if i == len(self.model.transformer.resblocks) - self.layer_idx:
|
||||
break
|
||||
if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting():
|
||||
x = checkpoint(r, x, attn_mask)
|
||||
else:
|
||||
x = r(x, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
def encode(self, text):
|
||||
|
||||
return self(text)
|
||||
|
||||
|
||||
|
||||
|
||||
class AbstractEncoder(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def encode(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
|
||||
"""
|
||||
Uses the OpenCLIP vision transformer encoder for images
|
||||
"""
|
||||
|
||||
def __init__(self, pretrained="", arch="ViT-H-14", device="cuda",
|
||||
freeze=True, layer="pooled", antialias=True):
|
||||
super().__init__()
|
||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
|
||||
pretrained=pretrained, )
|
||||
del model.transformer
|
||||
self.model = model
|
||||
self.device = device
|
||||
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
if self.layer == "penultimate":
|
||||
raise NotImplementedError()
|
||||
self.layer_idx = 1
|
||||
|
||||
self.antialias = antialias
|
||||
|
||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
||||
|
||||
|
||||
def preprocess(self, x):
|
||||
# normalize to [0,1]
|
||||
x = kornia.geometry.resize(x, (224, 224),
|
||||
interpolation='bicubic', align_corners=True,
|
||||
antialias=self.antialias)
|
||||
x = (x + 1.) / 2.
|
||||
# renormalize according to clip
|
||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
||||
return x
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.model.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, image, no_dropout=False):
|
||||
## image: b c h w
|
||||
z = self.encode_with_vision_transformer(image)
|
||||
return z
|
||||
|
||||
def encode_with_vision_transformer(self, x):
|
||||
x = self.preprocess(x)
|
||||
|
||||
# to patches - whether to use dual patchnorm
|
||||
if self.model.visual.input_patchnorm:
|
||||
# einops - rearrange(x, 'b c (h p1) (w p2) -> b (h w) (c p1 p2)')
|
||||
x = x.reshape(x.shape[0], x.shape[1], self.model.visual.grid_size[0], self.model.visual.patch_size[0], self.model.visual.grid_size[1], self.model.visual.patch_size[1])
|
||||
x = x.permute(0, 2, 4, 1, 3, 5)
|
||||
x = x.reshape(x.shape[0], self.model.visual.grid_size[0] * self.model.visual.grid_size[1], -1)
|
||||
x = self.model.visual.patchnorm_pre_ln(x)
|
||||
x = self.model.visual.conv1(x)
|
||||
else:
|
||||
x = self.model.visual.conv1(x) # shape = [*, width, grid, grid]
|
||||
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
|
||||
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
|
||||
|
||||
# class embeddings and positional embeddings
|
||||
x = torch.cat(
|
||||
[self.model.visual.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device),
|
||||
x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
||||
x = x + self.model.visual.positional_embedding.to(x.dtype)
|
||||
|
||||
# a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in
|
||||
x = self.model.visual.patch_dropout(x)
|
||||
x = self.model.visual.ln_pre(x)
|
||||
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.model.visual.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
|
||||
return x
|
||||
|
||||
@@ -0,0 +1,300 @@
|
||||
import torch.nn as nn
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
import copy
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
def __init__(self, d_hid, n_position=200):
|
||||
super(PositionalEncoding, self).__init__()
|
||||
|
||||
# Not a parameter
|
||||
self.register_buffer("pos_table", self._get_sinusoid_encoding_table(n_position, d_hid))
|
||||
|
||||
def _get_sinusoid_encoding_table(self, n_position, d_hid):
|
||||
"""Sinusoid position encoding table"""
|
||||
# TODO: make it with torch instead of numpy
|
||||
|
||||
def get_position_angle_vec(position):
|
||||
return [position / np.power(10000, 2 * (hid_j // 2) / d_hid) for hid_j in range(d_hid)]
|
||||
|
||||
sinusoid_table = np.array([get_position_angle_vec(pos_i) for pos_i in range(n_position)])
|
||||
sinusoid_table[:, 0::2] = np.sin(sinusoid_table[:, 0::2]) # dim 2i
|
||||
sinusoid_table[:, 1::2] = np.cos(sinusoid_table[:, 1::2]) # dim 2i+1
|
||||
|
||||
return torch.FloatTensor(sinusoid_table).unsqueeze(0)
|
||||
|
||||
def forward(self, winsize):
|
||||
return self.pos_table[:, :winsize].clone().detach()
|
||||
|
||||
|
||||
def _get_activation_fn(activation):
|
||||
"""Return an activation function given a string"""
|
||||
if activation == "relu":
|
||||
return F.relu
|
||||
if activation == "gelu":
|
||||
return F.gelu
|
||||
if activation == "glu":
|
||||
return F.glu
|
||||
raise RuntimeError(f"activation should be relu/gelu, not {activation}.")
|
||||
|
||||
|
||||
def _get_clones(module, N):
|
||||
return nn.ModuleList([copy.deepcopy(module) for i in range(N)])
|
||||
|
||||
|
||||
class Transformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model=512,
|
||||
nhead=8,
|
||||
num_encoder_layers=6,
|
||||
num_decoder_layers=6,
|
||||
dim_feedforward=2048,
|
||||
dropout=0.1,
|
||||
activation="relu",
|
||||
normalize_before=False,
|
||||
return_intermediate_dec=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
encoder_layer = TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, activation, normalize_before)
|
||||
encoder_norm = nn.LayerNorm(d_model) if normalize_before else None
|
||||
self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers, encoder_norm)
|
||||
|
||||
decoder_layer = TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout, activation, normalize_before)
|
||||
decoder_norm = nn.LayerNorm(d_model)
|
||||
self.decoder = TransformerDecoder(
|
||||
decoder_layer, num_decoder_layers, decoder_norm, return_intermediate=return_intermediate_dec
|
||||
)
|
||||
|
||||
self._reset_parameters()
|
||||
|
||||
self.d_model = d_model
|
||||
self.nhead = nhead
|
||||
|
||||
def _reset_parameters(self):
|
||||
for p in self.parameters():
|
||||
if p.dim() > 1:
|
||||
nn.init.xavier_uniform_(p)
|
||||
|
||||
def forward(self, opt, src, query_embed, pos_embed):
|
||||
# flatten NxCxHxW to HWxNxC
|
||||
|
||||
src = src.permute(1, 0, 2)
|
||||
pos_embed = pos_embed.permute(1, 0, 2)
|
||||
query_embed = query_embed.permute(1, 0, 2)
|
||||
|
||||
tgt = torch.zeros_like(query_embed)
|
||||
memory = self.encoder(src, pos=pos_embed)
|
||||
|
||||
hs = self.decoder(tgt, memory, pos=pos_embed, query_pos=query_embed)
|
||||
return hs
|
||||
|
||||
|
||||
class TransformerEncoder(nn.Module):
|
||||
def __init__(self, encoder_layer, num_layers, norm=None):
|
||||
super().__init__()
|
||||
self.layers = _get_clones(encoder_layer, num_layers)
|
||||
self.num_layers = num_layers
|
||||
self.norm = norm
|
||||
|
||||
def forward(self, src, mask=None, src_key_padding_mask=None, pos=None):
|
||||
output = src + pos
|
||||
|
||||
for layer in self.layers:
|
||||
output = layer(output, src_mask=mask, src_key_padding_mask=src_key_padding_mask, pos=pos)
|
||||
|
||||
if self.norm is not None:
|
||||
output = self.norm(output)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class TransformerDecoder(nn.Module):
|
||||
def __init__(self, decoder_layer, num_layers, norm=None, return_intermediate=False):
|
||||
super().__init__()
|
||||
self.layers = _get_clones(decoder_layer, num_layers)
|
||||
self.num_layers = num_layers
|
||||
self.norm = norm
|
||||
self.return_intermediate = return_intermediate
|
||||
|
||||
def forward(
|
||||
self,
|
||||
tgt,
|
||||
memory,
|
||||
tgt_mask=None,
|
||||
memory_mask=None,
|
||||
tgt_key_padding_mask=None,
|
||||
memory_key_padding_mask=None,
|
||||
pos=None,
|
||||
query_pos=None,
|
||||
):
|
||||
output = tgt + pos + query_pos
|
||||
|
||||
intermediate = []
|
||||
|
||||
for layer in self.layers:
|
||||
output = layer(
|
||||
output,
|
||||
memory,
|
||||
tgt_mask=tgt_mask,
|
||||
memory_mask=memory_mask,
|
||||
tgt_key_padding_mask=tgt_key_padding_mask,
|
||||
memory_key_padding_mask=memory_key_padding_mask,
|
||||
pos=pos,
|
||||
query_pos=query_pos,
|
||||
)
|
||||
if self.return_intermediate:
|
||||
intermediate.append(self.norm(output))
|
||||
|
||||
if self.norm is not None:
|
||||
output = self.norm(output)
|
||||
if self.return_intermediate:
|
||||
intermediate.pop()
|
||||
intermediate.append(output)
|
||||
|
||||
if self.return_intermediate:
|
||||
return torch.stack(intermediate)
|
||||
|
||||
return output.unsqueeze(0)
|
||||
|
||||
|
||||
class TransformerEncoderLayer(nn.Module):
|
||||
def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1, activation="relu", normalize_before=False):
|
||||
super().__init__()
|
||||
self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
||||
# Implementation of Feedforward model
|
||||
self.linear1 = nn.Linear(d_model, dim_feedforward)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.linear2 = nn.Linear(dim_feedforward, d_model)
|
||||
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
self.dropout1 = nn.Dropout(dropout)
|
||||
self.dropout2 = nn.Dropout(dropout)
|
||||
|
||||
self.activation = _get_activation_fn(activation)
|
||||
self.normalize_before = normalize_before
|
||||
|
||||
def with_pos_embed(self, tensor, pos):
|
||||
return tensor if pos is None else tensor + pos
|
||||
|
||||
def forward_post(self, src, src_mask=None, src_key_padding_mask=None, pos=None):
|
||||
# q = k = self.with_pos_embed(src, pos)
|
||||
src2 = self.self_attn(src, src, value=src, attn_mask=src_mask, key_padding_mask=src_key_padding_mask)[0]
|
||||
src = src + self.dropout1(src2)
|
||||
src = self.norm1(src)
|
||||
src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))
|
||||
src = src + self.dropout2(src2)
|
||||
src = self.norm2(src)
|
||||
return src
|
||||
|
||||
def forward_pre(self, src, src_mask=None, src_key_padding_mask=None, pos=None):
|
||||
src2 = self.norm1(src)
|
||||
# q = k = self.with_pos_embed(src2, pos)
|
||||
src2 = self.self_attn(src2, src2, value=src2, attn_mask=src_mask, key_padding_mask=src_key_padding_mask)[0]
|
||||
src = src + self.dropout1(src2)
|
||||
src2 = self.norm2(src)
|
||||
src2 = self.linear2(self.dropout(self.activation(self.linear1(src2))))
|
||||
src = src + self.dropout2(src2)
|
||||
return src
|
||||
|
||||
def forward(self, src, src_mask=None, src_key_padding_mask=None, pos=None):
|
||||
if self.normalize_before:
|
||||
return self.forward_pre(src, src_mask, src_key_padding_mask, pos)
|
||||
return self.forward_post(src, src_mask, src_key_padding_mask, pos)
|
||||
|
||||
|
||||
class TransformerDecoderLayer(nn.Module):
|
||||
def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1, activation="relu", normalize_before=False):
|
||||
super().__init__()
|
||||
self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
||||
self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
|
||||
# Implementation of Feedforward model
|
||||
self.linear1 = nn.Linear(d_model, dim_feedforward)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.linear2 = nn.Linear(dim_feedforward, d_model)
|
||||
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
self.norm3 = nn.LayerNorm(d_model)
|
||||
self.dropout1 = nn.Dropout(dropout)
|
||||
self.dropout2 = nn.Dropout(dropout)
|
||||
self.dropout3 = nn.Dropout(dropout)
|
||||
|
||||
self.activation = _get_activation_fn(activation)
|
||||
self.normalize_before = normalize_before
|
||||
|
||||
def with_pos_embed(self, tensor, pos):
|
||||
return tensor if pos is None else tensor + pos
|
||||
|
||||
def forward_post(
|
||||
self,
|
||||
tgt,
|
||||
memory,
|
||||
tgt_mask=None,
|
||||
memory_mask=None,
|
||||
tgt_key_padding_mask=None,
|
||||
memory_key_padding_mask=None,
|
||||
pos=None,
|
||||
query_pos=None,
|
||||
):
|
||||
# q = k = self.with_pos_embed(tgt, query_pos)
|
||||
tgt2 = self.self_attn(tgt, tgt, value=tgt, attn_mask=tgt_mask, key_padding_mask=tgt_key_padding_mask)[0]
|
||||
tgt = tgt + self.dropout1(tgt2)
|
||||
tgt = self.norm1(tgt)
|
||||
tgt2 = self.multihead_attn(
|
||||
query=tgt, key=memory, value=memory, attn_mask=memory_mask, key_padding_mask=memory_key_padding_mask
|
||||
)[0]
|
||||
tgt = tgt + self.dropout2(tgt2)
|
||||
tgt = self.norm2(tgt)
|
||||
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
|
||||
tgt = tgt + self.dropout3(tgt2)
|
||||
tgt = self.norm3(tgt)
|
||||
return tgt
|
||||
|
||||
def forward_pre(
|
||||
self,
|
||||
tgt,
|
||||
memory,
|
||||
tgt_mask=None,
|
||||
memory_mask=None,
|
||||
tgt_key_padding_mask=None,
|
||||
memory_key_padding_mask=None,
|
||||
pos=None,
|
||||
query_pos=None,
|
||||
):
|
||||
tgt2 = self.norm1(tgt)
|
||||
# q = k = self.with_pos_embed(tgt2, query_pos)
|
||||
tgt2 = self.self_attn(tgt2, tgt2, value=tgt2, attn_mask=tgt_mask, key_padding_mask=tgt_key_padding_mask)[0]
|
||||
tgt = tgt + self.dropout1(tgt2)
|
||||
tgt2 = self.norm2(tgt)
|
||||
tgt2 = self.multihead_attn(
|
||||
query=tgt2, key=memory, value=memory, attn_mask=memory_mask, key_padding_mask=memory_key_padding_mask
|
||||
)[0]
|
||||
tgt = tgt + self.dropout2(tgt2)
|
||||
tgt2 = self.norm3(tgt)
|
||||
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))
|
||||
tgt = tgt + self.dropout3(tgt2)
|
||||
return tgt
|
||||
|
||||
def forward(
|
||||
self,
|
||||
tgt,
|
||||
memory,
|
||||
tgt_mask=None,
|
||||
memory_mask=None,
|
||||
tgt_key_padding_mask=None,
|
||||
memory_key_padding_mask=None,
|
||||
pos=None,
|
||||
query_pos=None,
|
||||
):
|
||||
if self.normalize_before:
|
||||
return self.forward_pre(
|
||||
tgt, memory, tgt_mask, memory_mask, tgt_key_padding_mask, memory_key_padding_mask, pos, query_pos
|
||||
)
|
||||
return self.forward_post(
|
||||
tgt, memory, tgt_mask, memory_mask, tgt_key_padding_mask, memory_key_padding_mask, pos, query_pos
|
||||
)
|
||||
@@ -0,0 +1,714 @@
|
||||
import math
|
||||
import torch
|
||||
import xformers
|
||||
import xformers.ops
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
import torch.nn.functional as F
|
||||
from rotary_embedding_torch import RotaryEmbedding
|
||||
from fairscale.nn.checkpoint import checkpoint_wrapper
|
||||
|
||||
from .util import *
|
||||
from .IPI_module import DwposeEncoder, Resampler
|
||||
from utils.registry_class import MODEL
|
||||
|
||||
|
||||
USE_TEMPORAL_TRANSFORMER = True
|
||||
|
||||
|
||||
|
||||
class PreNormattention(nn.Module):
|
||||
def __init__(self, dim, fn):
|
||||
super().__init__()
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
self.fn = fn
|
||||
def forward(self, x, **kwargs):
|
||||
return self.fn(self.norm(x), **kwargs) + x
|
||||
|
||||
class PreNormattention_qkv(nn.Module):
|
||||
def __init__(self, dim, fn):
|
||||
super().__init__()
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
self.fn = fn
|
||||
def forward(self, q, k, v, **kwargs):
|
||||
return self.fn(self.norm(q), self.norm(k), self.norm(v), **kwargs) + q
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
project_out = not (heads == 1 and dim_head == dim)
|
||||
|
||||
self.heads = heads
|
||||
self.scale = dim_head ** -0.5
|
||||
|
||||
self.attend = nn.Softmax(dim = -1)
|
||||
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(inner_dim, dim),
|
||||
nn.Dropout(dropout)
|
||||
) if project_out else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
b, n, _, h = *x.shape, self.heads
|
||||
qkv = self.to_qkv(x).chunk(3, dim = -1)
|
||||
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), qkv)
|
||||
|
||||
dots = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale
|
||||
|
||||
attn = self.attend(dots)
|
||||
|
||||
out = einsum('b h i j, b h j d -> b h i d', attn, v)
|
||||
out = rearrange(out, 'b h n d -> b n (h d)')
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Attention_qkv(nn.Module):
|
||||
def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
project_out = not (heads == 1 and dim_head == dim)
|
||||
|
||||
self.heads = heads
|
||||
self.scale = dim_head ** -0.5
|
||||
|
||||
self.attend = nn.Softmax(dim = -1)
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias = False)
|
||||
self.to_k = nn.Linear(dim, inner_dim, bias = False)
|
||||
self.to_v = nn.Linear(dim, inner_dim, bias = False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(inner_dim, dim),
|
||||
nn.Dropout(dropout)
|
||||
) if project_out else nn.Identity()
|
||||
|
||||
def forward(self, q, k, v):
|
||||
b, n, _, h = *q.shape, self.heads
|
||||
bk = k.shape[0]
|
||||
|
||||
q = self.to_q(q)
|
||||
k = self.to_k(k)
|
||||
v = self.to_v(v)
|
||||
q = rearrange(q, 'b n (h d) -> b h n d', h = h)
|
||||
k = rearrange(k, 'b n (h d) -> b h n d', b=bk, h = h)
|
||||
v = rearrange(v, 'b n (h d) -> b h n d', b=bk, h = h)
|
||||
|
||||
dots = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale
|
||||
|
||||
attn = self.attend(dots)
|
||||
|
||||
out = einsum('b h i j, b h j d -> b h i d', attn, v)
|
||||
out = rearrange(out, 'b h n d -> b n (h d)')
|
||||
return self.to_out(out)
|
||||
|
||||
class PostNormattention(nn.Module):
|
||||
def __init__(self, dim, fn):
|
||||
super().__init__()
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
self.fn = fn
|
||||
def forward(self, x, **kwargs):
|
||||
return self.norm(self.fn(x, **kwargs) + x)
|
||||
|
||||
|
||||
|
||||
|
||||
class Transformer_v2(nn.Module):
|
||||
def __init__(self, heads=8, dim=2048, dim_head_k=256, dim_head_v=256, dropout_atte = 0.05, mlp_dim=2048, dropout_ffn = 0.05, depth=1):
|
||||
super().__init__()
|
||||
self.layers = nn.ModuleList([])
|
||||
self.depth = depth
|
||||
for _ in range(depth):
|
||||
self.layers.append(nn.ModuleList([
|
||||
PreNormattention(dim, Attention(dim, heads = heads, dim_head = dim_head_k, dropout = dropout_atte)),
|
||||
FeedForward(dim, mlp_dim, dropout = dropout_ffn),
|
||||
]))
|
||||
def forward(self, x):
|
||||
for attn, ff in self.layers[:1]:
|
||||
x = attn(x)
|
||||
x = ff(x) + x
|
||||
if self.depth > 1:
|
||||
for attn, ff in self.layers[1:]:
|
||||
x = attn(x)
|
||||
x = ff(x) + x
|
||||
return x
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
r"""DropPath but without rescaling and supports optional all-zero and/or all-keep.
|
||||
"""
|
||||
def __init__(self, p):
|
||||
super(DropPath, self).__init__()
|
||||
self.p = p
|
||||
|
||||
def forward(self, *args, zero=None, keep=None):
|
||||
if not self.training:
|
||||
return args[0] if len(args) == 1 else args
|
||||
|
||||
# params
|
||||
x = args[0]
|
||||
b = x.size(0)
|
||||
n = (torch.rand(b) < self.p).sum()
|
||||
|
||||
# non-zero and non-keep mask
|
||||
mask = x.new_ones(b, dtype=torch.bool)
|
||||
if keep is not None:
|
||||
mask[keep] = False
|
||||
if zero is not None:
|
||||
mask[zero] = False
|
||||
|
||||
# drop-path index
|
||||
index = torch.where(mask)[0]
|
||||
index = index[torch.randperm(len(index))[:n]]
|
||||
if zero is not None:
|
||||
index = torch.cat([index, torch.where(zero)[0]], dim=0)
|
||||
|
||||
# drop-path multiplier
|
||||
multiplier = x.new_ones(b)
|
||||
multiplier[index] = 0.0
|
||||
output = tuple(u * self.broadcast(multiplier, u) for u in args)
|
||||
return output[0] if len(args) == 1 else output
|
||||
|
||||
def broadcast(self, src, dst):
|
||||
assert src.size(0) == dst.size(0)
|
||||
shape = (dst.size(0), ) + (1, ) * (dst.ndim - 1)
|
||||
return src.view(shape)
|
||||
|
||||
|
||||
|
||||
@MODEL.register_class()
|
||||
class UNetSD_Animate_X(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config=None,
|
||||
in_dim=4,
|
||||
dim=512,
|
||||
y_dim=512,
|
||||
context_dim=1024,
|
||||
hist_dim = 156,
|
||||
concat_dim = 8,
|
||||
out_dim=6,
|
||||
dim_mult=[1, 2, 3, 4],
|
||||
num_heads=None,
|
||||
head_dim=64,
|
||||
num_res_blocks=3,
|
||||
attn_scales=[1 / 2, 1 / 4, 1 / 8],
|
||||
use_scale_shift_norm=True,
|
||||
dropout=0.1,
|
||||
temporal_attn_times=1,
|
||||
pose_attention = True,
|
||||
temporal_attention = True,
|
||||
use_checkpoint=False,
|
||||
use_image_dataset=False,
|
||||
use_fps_condition= False,
|
||||
use_sim_mask = False,
|
||||
misc_dropout = 0.5,
|
||||
training=True,
|
||||
inpainting=True,
|
||||
p_all_zero=0.1,
|
||||
p_all_keep=0.1,
|
||||
zero_y = None,
|
||||
black_image_feature = None,
|
||||
adapter_transformer_layers = 1,
|
||||
num_tokens=4,
|
||||
no_hand = False,
|
||||
use_pose_transformer = True,
|
||||
num = 0,
|
||||
our_weight = 1,
|
||||
seq_len = None,
|
||||
**kwargs
|
||||
):
|
||||
embed_dim = dim * 4
|
||||
num_heads=num_heads if num_heads else dim//32
|
||||
super(UNetSD_Animate_X, self).__init__()
|
||||
self.zero_y = zero_y
|
||||
self.black_image_feature = black_image_feature
|
||||
self.our_weight = our_weight
|
||||
|
||||
if no_hand:
|
||||
self.pose_embedding_dim = 18
|
||||
else:
|
||||
self.pose_embedding_dim = 128
|
||||
|
||||
if num!=0:
|
||||
self.pose_embedding_dim = num*17+18
|
||||
|
||||
self.use_pose_transformer = use_pose_transformer
|
||||
|
||||
|
||||
self.video_compositions = ['image', 'local_image', 'dwpose', 'randomref', 'randomref_pose', 'pose_embedding']
|
||||
self.resolution = [512, 768]
|
||||
|
||||
|
||||
self.in_dim = in_dim
|
||||
self.dim = dim
|
||||
self.y_dim = y_dim
|
||||
self.context_dim = context_dim
|
||||
self.num_tokens = num_tokens
|
||||
self.hist_dim = hist_dim
|
||||
self.concat_dim = concat_dim
|
||||
self.embed_dim = embed_dim
|
||||
self.out_dim = out_dim
|
||||
self.dim_mult = dim_mult
|
||||
|
||||
self.num_heads = num_heads
|
||||
|
||||
self.head_dim = head_dim
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.use_scale_shift_norm = use_scale_shift_norm
|
||||
self.temporal_attn_times = temporal_attn_times
|
||||
self.temporal_attention = temporal_attention
|
||||
|
||||
self.pose_attention = pose_attention
|
||||
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.use_image_dataset = use_image_dataset
|
||||
self.use_fps_condition = use_fps_condition
|
||||
self.use_sim_mask = use_sim_mask
|
||||
self.training=training
|
||||
self.inpainting = inpainting
|
||||
|
||||
self.misc_dropout = misc_dropout
|
||||
self.p_all_zero = p_all_zero
|
||||
self.p_all_keep = p_all_keep
|
||||
|
||||
use_linear_in_temporal = False
|
||||
transformer_depth = 1
|
||||
disabled_sa = False
|
||||
# params
|
||||
enc_dims = [dim * u for u in [1] + dim_mult]
|
||||
dec_dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
||||
shortcut_dims = []
|
||||
scale = 1.0
|
||||
|
||||
|
||||
|
||||
# embeddings
|
||||
self.time_embed = nn.Sequential(
|
||||
nn.Linear(dim, embed_dim),
|
||||
nn.SiLU(),
|
||||
nn.Linear(embed_dim, embed_dim))
|
||||
if 'image' in self.video_compositions:
|
||||
self.pre_image_condition = nn.Sequential(
|
||||
nn.Linear(self.context_dim, self.context_dim),
|
||||
nn.SiLU(),
|
||||
nn.Linear(self.context_dim, self.context_dim*self.num_tokens))
|
||||
|
||||
if 'pose_embedding' in self.video_compositions:
|
||||
if seq_len!=None:
|
||||
self.pose_embedding_before = DwposeEncoder(pos_embed_len = seq_len)
|
||||
else:
|
||||
self.pose_embedding_before = DwposeEncoder()
|
||||
self.pose_embedding_after = Resampler()
|
||||
|
||||
if 'local_image' in self.video_compositions:
|
||||
self.local_image_embedding = nn.Sequential(
|
||||
nn.Conv2d(3, concat_dim * 4, 3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.AdaptiveAvgPool2d((self.resolution[1]//2, self.resolution[0]//2)),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim, 3, stride=2, padding=1))
|
||||
self.local_image_embedding_after = Transformer_v2(heads=2, dim=concat_dim, dim_head_k=concat_dim, dim_head_v=concat_dim, dropout_atte = 0.05, mlp_dim=concat_dim, dropout_ffn = 0.05, depth=adapter_transformer_layers)
|
||||
|
||||
if 'dwpose' in self.video_compositions:
|
||||
self.dwpose_embedding = nn.Sequential(
|
||||
nn.Conv2d(3, concat_dim * 4, 3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.AdaptiveAvgPool2d((self.resolution[1]//2, self.resolution[0]//2)),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim, 3, stride=2, padding=1))
|
||||
# print(f'dwpose transformer input dim: dim={concat_dim}, dim_head_k={concat_dim}, dim_head_v={concat_dim}, dropout_atte = 0.05, mlp_dim={concat_dim}, dropout_ffn = 0.05, depth={adapter_transformer_layers}')
|
||||
self.dwpose_embedding_after = Transformer_v2(heads=2, dim=concat_dim, dim_head_k=concat_dim, dim_head_v=concat_dim, dropout_atte = 0.05, mlp_dim=concat_dim, dropout_ffn = 0.05, depth=adapter_transformer_layers)
|
||||
|
||||
if 'randomref_pose' in self.video_compositions:
|
||||
randomref_dim = 4
|
||||
self.randomref_pose2_embedding = nn.Sequential(
|
||||
nn.Conv2d(3, concat_dim * 4, 3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.AdaptiveAvgPool2d((self.resolution[1]//2, self.resolution[0]//2)),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim+randomref_dim, 3, stride=2, padding=1))
|
||||
self.randomref_pose2_embedding_after = Transformer_v2(heads=2, dim=concat_dim+randomref_dim, dim_head_k=concat_dim+randomref_dim, dim_head_v=concat_dim+randomref_dim, dropout_atte = 0.05, mlp_dim=concat_dim+randomref_dim, dropout_ffn = 0.05, depth=adapter_transformer_layers)
|
||||
|
||||
if 'randomref' in self.video_compositions:
|
||||
randomref_dim = 4
|
||||
self.randomref_embedding2 = nn.Sequential(
|
||||
nn.Conv2d(randomref_dim, concat_dim * 4, 3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim+randomref_dim, 3, stride=1, padding=1))
|
||||
self.randomref_embedding_after2 = Transformer_v2(heads=2, dim=concat_dim+randomref_dim, dim_head_k=concat_dim+randomref_dim, dim_head_v=concat_dim+randomref_dim, dropout_atte = 0.05, mlp_dim=concat_dim+randomref_dim, dropout_ffn = 0.05, depth=adapter_transformer_layers)
|
||||
|
||||
### Condition Dropout
|
||||
self.misc_dropout = DropPath(misc_dropout)
|
||||
|
||||
|
||||
if temporal_attention and not USE_TEMPORAL_TRANSFORMER:
|
||||
self.rotary_emb = RotaryEmbedding(min(32, head_dim))
|
||||
self.time_rel_pos_bias = RelativePositionBias(heads = num_heads, max_distance = 32) # realistically will not be able to generate that many frames of video... yet
|
||||
|
||||
if self.use_fps_condition:
|
||||
self.fps_embedding = nn.Sequential(
|
||||
nn.Linear(dim, embed_dim),
|
||||
nn.SiLU(),
|
||||
nn.Linear(embed_dim, embed_dim))
|
||||
nn.init.zeros_(self.fps_embedding[-1].weight)
|
||||
nn.init.zeros_(self.fps_embedding[-1].bias)
|
||||
|
||||
# encoder
|
||||
self.input_blocks = nn.ModuleList()
|
||||
self.pre_image = nn.Sequential()
|
||||
init_block = nn.ModuleList([nn.Conv2d(self.in_dim + concat_dim, dim, 3, padding=1)])
|
||||
|
||||
#### need an initial temporal attention?
|
||||
if temporal_attention:
|
||||
|
||||
if USE_TEMPORAL_TRANSFORMER:
|
||||
init_block.append(TemporalTransformer(dim, num_heads, head_dim, depth=transformer_depth, context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, multiply_zero=use_image_dataset))
|
||||
else:
|
||||
init_block.append(TemporalAttentionMultiBlock(dim, num_heads, head_dim, rotary_emb=self.rotary_emb, temporal_attn_times=temporal_attn_times, use_image_dataset=use_image_dataset))
|
||||
|
||||
self.input_blocks.append(init_block)
|
||||
shortcut_dims.append(dim)
|
||||
for i, (in_dim, out_dim) in enumerate(zip(enc_dims[:-1], enc_dims[1:])):
|
||||
for j in range(num_res_blocks):
|
||||
|
||||
block = nn.ModuleList([ResBlock(in_dim, embed_dim, dropout, out_channels=out_dim, use_scale_shift_norm=False, use_image_dataset=use_image_dataset,)])
|
||||
|
||||
if scale in attn_scales:
|
||||
block.append(
|
||||
SpatialTransformer(
|
||||
out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim,
|
||||
disable_self_attn=False, use_linear=True, pose_attention = self.pose_attention, our_weight = self.our_weight
|
||||
)
|
||||
)
|
||||
|
||||
if self.temporal_attention:
|
||||
if USE_TEMPORAL_TRANSFORMER:
|
||||
block.append(TemporalTransformer(out_dim, out_dim // head_dim, head_dim, depth=transformer_depth, context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, multiply_zero=use_image_dataset))
|
||||
else:
|
||||
block.append(TemporalAttentionMultiBlock(out_dim, num_heads, head_dim, rotary_emb = self.rotary_emb, use_image_dataset=use_image_dataset, use_sim_mask=use_sim_mask, temporal_attn_times=temporal_attn_times))
|
||||
in_dim = out_dim
|
||||
self.input_blocks.append(block)
|
||||
shortcut_dims.append(out_dim)
|
||||
|
||||
# downsample
|
||||
if i != len(dim_mult) - 1 and j == num_res_blocks - 1:
|
||||
downsample = Downsample(
|
||||
out_dim, True, dims=2, out_channels=out_dim
|
||||
)
|
||||
shortcut_dims.append(out_dim)
|
||||
scale /= 2.0
|
||||
self.input_blocks.append(downsample)
|
||||
|
||||
# middle
|
||||
self.middle_block = nn.ModuleList([
|
||||
ResBlock(out_dim, embed_dim, dropout, use_scale_shift_norm=False, use_image_dataset=use_image_dataset,),
|
||||
SpatialTransformer(
|
||||
out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim,
|
||||
disable_self_attn=False, use_linear=True, pose_attention = self.pose_attention , our_weight = self.our_weight
|
||||
)])
|
||||
# print(f"middle_block, out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim,", out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim)
|
||||
# middle_block, out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim, 1280 20 64 1024
|
||||
|
||||
|
||||
if self.temporal_attention:
|
||||
if USE_TEMPORAL_TRANSFORMER:
|
||||
self.middle_block.append(
|
||||
TemporalTransformer(
|
||||
out_dim, out_dim // head_dim, head_dim, depth=transformer_depth, context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal,
|
||||
multiply_zero=use_image_dataset,
|
||||
)
|
||||
)
|
||||
else:
|
||||
self.middle_block.append(TemporalAttentionMultiBlock(out_dim, num_heads, head_dim, rotary_emb = self.rotary_emb, use_image_dataset=use_image_dataset, use_sim_mask=use_sim_mask, temporal_attn_times=temporal_attn_times))
|
||||
|
||||
self.middle_block.append(ResBlock(out_dim, embed_dim, dropout, use_scale_shift_norm=False))
|
||||
|
||||
|
||||
# decoder
|
||||
self.output_blocks = nn.ModuleList()
|
||||
for i, (in_dim, out_dim) in enumerate(zip(dec_dims[:-1], dec_dims[1:])):
|
||||
for j in range(num_res_blocks + 1):
|
||||
|
||||
block = nn.ModuleList([ResBlock(in_dim + shortcut_dims.pop(), embed_dim, dropout, out_dim, use_scale_shift_norm=False, use_image_dataset=use_image_dataset, )])
|
||||
if scale in attn_scales:
|
||||
block.append(
|
||||
SpatialTransformer(
|
||||
out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=1024,
|
||||
disable_self_attn=False, use_linear=True, pose_attention = self.pose_attention, our_weight = self.our_weight
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if self.temporal_attention:
|
||||
if USE_TEMPORAL_TRANSFORMER:
|
||||
block.append(
|
||||
TemporalTransformer(
|
||||
out_dim, out_dim // head_dim, head_dim, depth=transformer_depth, context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, multiply_zero=use_image_dataset
|
||||
)
|
||||
)
|
||||
else:
|
||||
block.append(TemporalAttentionMultiBlock(out_dim, num_heads, head_dim, rotary_emb =self.rotary_emb, use_image_dataset=use_image_dataset, use_sim_mask=use_sim_mask, temporal_attn_times=temporal_attn_times))
|
||||
in_dim = out_dim
|
||||
|
||||
# upsample
|
||||
if i != len(dim_mult) - 1 and j == num_res_blocks:
|
||||
upsample = Upsample(out_dim, True, dims=2.0, out_channels=out_dim)
|
||||
scale *= 2.0
|
||||
block.append(upsample)
|
||||
self.output_blocks.append(block)
|
||||
|
||||
# head
|
||||
self.out = nn.Sequential(
|
||||
nn.GroupNorm(32, out_dim),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(out_dim, self.out_dim, 3, padding=1))
|
||||
|
||||
# zero out the last layer params
|
||||
nn.init.zeros_(self.out[-1].weight)
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
t,
|
||||
y = None,
|
||||
pose_embeddings = None,
|
||||
depth = None,
|
||||
image = None,
|
||||
motion = None,
|
||||
local_image = None,
|
||||
single_sketch = None,
|
||||
masked = None,
|
||||
canny = None,
|
||||
sketch = None,
|
||||
dwpose = None,
|
||||
randomref = None,
|
||||
histogram = None,
|
||||
fps = None,
|
||||
video_mask = None,
|
||||
focus_present_mask = None,
|
||||
prob_focus_present = 0., # probability at which a given batch sample will focus on the present (0. is all off, 1. is completely arrested attention across time)
|
||||
mask_last_frame_num = 0 # mask last frame num
|
||||
):
|
||||
|
||||
|
||||
assert self.inpainting or masked is None, 'inpainting is not supported'
|
||||
|
||||
batch, c, f, h, w= x.shape
|
||||
frames = f
|
||||
device = x.device
|
||||
self.batch = batch
|
||||
|
||||
#### image and video joint training, if mask_last_frame_num is set, prob_focus_present will be ignored
|
||||
if mask_last_frame_num > 0:
|
||||
focus_present_mask = None
|
||||
video_mask[-mask_last_frame_num:] = False
|
||||
else:
|
||||
focus_present_mask = default(focus_present_mask, lambda: prob_mask_like((batch,), prob_focus_present, device = device))
|
||||
|
||||
if self.temporal_attention and not USE_TEMPORAL_TRANSFORMER:
|
||||
time_rel_pos_bias = self.time_rel_pos_bias(x.shape[2], device = x.device)
|
||||
else:
|
||||
time_rel_pos_bias = None
|
||||
|
||||
|
||||
# all-zero and all-keep masks
|
||||
zero = torch.zeros(batch, dtype=torch.bool).to(x.device)
|
||||
keep = torch.zeros(batch, dtype=torch.bool).to(x.device)
|
||||
if self.training:
|
||||
nzero = (torch.rand(batch) < self.p_all_zero).sum()
|
||||
nkeep = (torch.rand(batch) < self.p_all_keep).sum()
|
||||
index = torch.randperm(batch)
|
||||
zero[index[0:nzero]] = True
|
||||
keep[index[nzero:nzero + nkeep]] = True
|
||||
assert not (zero & keep).any()
|
||||
misc_dropout = partial(self.misc_dropout, zero = zero, keep = keep)
|
||||
|
||||
|
||||
concat = x.new_zeros(batch, self.concat_dim, f, h, w)
|
||||
|
||||
|
||||
# local_image_embedding (first frame)
|
||||
if local_image is not None:
|
||||
local_image = rearrange(local_image, 'b c f h w -> (b f) c h w')
|
||||
local_image = self.local_image_embedding(local_image)
|
||||
|
||||
h = local_image.shape[2]
|
||||
local_image = self.local_image_embedding_after(rearrange(local_image, '(b f) c h w -> (b h w) f c', b = batch))
|
||||
local_image = rearrange(local_image, '(b h w) f c -> b c f h w', b = batch, h = h)
|
||||
|
||||
concat = concat + misc_dropout(local_image)
|
||||
|
||||
|
||||
|
||||
|
||||
if dwpose is not None:
|
||||
if 'randomref_pose' in self.video_compositions:
|
||||
dwpose_random_ref = dwpose[:,:,:1].clone()
|
||||
dwpose = dwpose[:,:,1:]
|
||||
dwpose = rearrange(dwpose, 'b c f h w -> (b f) c h w') # 32 * 3 * 768, 512
|
||||
#
|
||||
dwpose = self.dwpose_embedding(dwpose)
|
||||
h = dwpose.shape[2]
|
||||
dwpose = self.dwpose_embedding_after(rearrange(dwpose, '(b f) c h w -> (b h w) f c', b = batch))
|
||||
dwpose = rearrange(dwpose, '(b h w) f c -> b c f h w', b = batch, h = h)
|
||||
|
||||
concat = concat + misc_dropout(dwpose)
|
||||
randomref_b = x.new_zeros(batch, self.concat_dim+4, 1, h, w)
|
||||
if randomref is not None:
|
||||
randomref = rearrange(randomref[:,:,:1,], 'b c f h w -> (b f) c h w')
|
||||
randomref = self.randomref_embedding2(randomref)
|
||||
|
||||
h = randomref.shape[2]
|
||||
randomref = self.randomref_embedding_after2(rearrange(randomref, '(b f) c h w -> (b h w) f c', b = batch))
|
||||
if 'randomref_pose' in self.video_compositions:
|
||||
dwpose_random_ref = rearrange(dwpose_random_ref, 'b c f h w -> (b f) c h w')
|
||||
dwpose_random_ref = self.randomref_pose2_embedding(dwpose_random_ref)
|
||||
dwpose_random_ref = self.randomref_pose2_embedding_after(rearrange(dwpose_random_ref, '(b f) c h w -> (b h w) f c', b = batch))
|
||||
randomref = randomref + dwpose_random_ref
|
||||
|
||||
randomref_a = rearrange(randomref, '(b h w) f c -> b c f h w', b = batch, h = h)
|
||||
randomref_b = randomref_b + randomref_a
|
||||
|
||||
|
||||
x = torch.cat([randomref_b, torch.cat([x, concat], dim=1)], dim=2)
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
x = self.pre_image(x)
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = batch)
|
||||
|
||||
# embeddings
|
||||
if self.use_fps_condition and fps is not None:
|
||||
e = self.time_embed(sinusoidal_embedding(t, self.dim)) + self.fps_embedding(sinusoidal_embedding(fps, self.dim))
|
||||
else:
|
||||
e = self.time_embed(sinusoidal_embedding(t, self.dim))
|
||||
|
||||
context = x.new_zeros(batch, 0, self.context_dim)
|
||||
|
||||
|
||||
if image is not None:
|
||||
y_context = self.zero_y.repeat(batch, 1, 1)
|
||||
context = torch.cat([context, y_context], dim=1)
|
||||
|
||||
image_context = misc_dropout(self.pre_image_condition(image).view(-1, self.num_tokens, self.context_dim)) # torch.cat([y[:,:-1,:], self.pre_image_condition(y[:,-1:,:]) ], dim=1)
|
||||
context = torch.cat([context, image_context], dim=1)
|
||||
else:
|
||||
y_context = self.zero_y.repeat(batch, 1, 1)
|
||||
context = torch.cat([context, y_context], dim=1)
|
||||
image_context = torch.zeros_like(self.zero_y.repeat(batch, 1, 1))[:,:self.num_tokens]
|
||||
context = torch.cat([context, image_context], dim=1)
|
||||
|
||||
# repeat f times for spatial e and context
|
||||
e = e.repeat_interleave(repeats=f+1, dim=0)
|
||||
context = context.repeat_interleave(repeats=f+1, dim=0)
|
||||
|
||||
|
||||
|
||||
# if pose_embedding is not None:
|
||||
if pose_embeddings is None:
|
||||
pose_embedding = torch.zeros([batch, f+1, 2, self.pose_embedding_dim]).half().cuda() # 2 33 2 18
|
||||
driving_image_feature = torch.zeros([batch*(f+1), 1, 1024]).half().cuda() # self.black_image_feature.repeat(batch,1,1)
|
||||
else:
|
||||
pose_embedding, driving_image_feature = pose_embeddings
|
||||
pose_embedding = rearrange(pose_embedding, 'b f dim c -> b f (dim c)') # dim =2 c =18
|
||||
pose_embedding = self.pose_embedding_before(pose_embedding) # 2 33 1024 18
|
||||
|
||||
|
||||
pose_embedding = self.pose_embedding_after(driving_image_feature, pose_embedding)
|
||||
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
# encoder
|
||||
xs = []
|
||||
|
||||
for block in self.input_blocks:
|
||||
x = self._forward_single(block, x, e, context, pose_embedding, time_rel_pos_bias, focus_present_mask, video_mask)
|
||||
xs.append(x)
|
||||
|
||||
# middle
|
||||
for block in self.middle_block:
|
||||
x = self._forward_single(block, x, e, context, pose_embedding, time_rel_pos_bias,focus_present_mask, video_mask)
|
||||
|
||||
# decoder
|
||||
for block in self.output_blocks:
|
||||
x = torch.cat([x, xs.pop()], dim=1)
|
||||
x = self._forward_single(block, x, e, context, pose_embedding, time_rel_pos_bias,focus_present_mask, video_mask, reference=xs[-1] if len(xs) > 0 else None)
|
||||
|
||||
# head
|
||||
x = self.out(x)
|
||||
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = batch)
|
||||
return x[:,:,1:]
|
||||
|
||||
def _forward_single(self, module, x, e, context, pose_embedding, time_rel_pos_bias, focus_present_mask, video_mask, reference=None):
|
||||
|
||||
if isinstance(module, ResidualBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = x.contiguous()
|
||||
x = module(x, e, reference)
|
||||
elif isinstance(module, ResBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = x.contiguous()
|
||||
x = module(x, e, self.batch)
|
||||
elif isinstance(module, SpatialTransformer):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = module(x, context, pose_embedding)
|
||||
|
||||
elif isinstance(module, TemporalTransformer):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch)
|
||||
x = module(x, context)
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
elif isinstance(module, CrossAttention):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = module(x, context)
|
||||
elif isinstance(module, MemoryEfficientCrossAttention):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = module(x, context)
|
||||
elif isinstance(module, BasicTransformerBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = module(x, context)
|
||||
elif isinstance(module, FeedForward):
|
||||
x = module(x, context)
|
||||
elif isinstance(module, Upsample):
|
||||
x = module(x)
|
||||
elif isinstance(module, Downsample):
|
||||
x = module(x)
|
||||
elif isinstance(module, Resample):
|
||||
x = module(x, reference)
|
||||
elif isinstance(module, TemporalAttentionBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch)
|
||||
x = module(x, time_rel_pos_bias, focus_present_mask, video_mask)
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
elif isinstance(module, TemporalAttentionMultiBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch)
|
||||
x = module(x, time_rel_pos_bias, focus_present_mask, video_mask)
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
elif isinstance(module, InitTemporalConvBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch)
|
||||
x = module(x)
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
elif isinstance(module, TemporalConvBlock):
|
||||
module = checkpoint_wrapper(module) if self.use_checkpoint else module
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch)
|
||||
x = module(x)
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
elif isinstance(module, nn.ModuleList):
|
||||
for block in module:
|
||||
x = self._forward_single(block, x, e, context, pose_embedding, time_rel_pos_bias, focus_present_mask, video_mask, reference)
|
||||
else:
|
||||
x = module(x)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
|
||||
|
||||
# manual setting
|
||||
max_frames: 32
|
||||
resolution: [512, 768]
|
||||
round: 5
|
||||
ddim_timesteps: 30
|
||||
seed: 13
|
||||
test_list_path: [
|
||||
[2, "data/images/1.jpg", "data/saved_pose/dance_1","data/saved_frames/dance_1","data/saved_pkl/dance_1.pkl", 14], # seed 14
|
||||
[2, "data/images/4.png", "data/saved_pose/dance_1","data/saved_frames/dance_1","data/saved_pkl/dance_1.pkl", 14], # seed 14
|
||||
[2, "data/images/3.jpeg", "data/saved_pose/dance_2","data/saved_frames/dance_2","data/saved_pkl/dance_2.pkl", 13], # seed 13
|
||||
[2, "data/images/10.jpeg", "data/saved_pose/dance_2","data/saved_frames/dance_2","data/saved_pkl/dance_2.pkl", 14], #
|
||||
[2, "data/images/zhu.png", "data/saved_pose/dance_3","data/saved_frames/dance_3","data/saved_pkl/dance_3.pkl", 14],
|
||||
[2, "data/images/5.jpeg", "data/saved_pose/dance_3","data/saved_frames/dance_3","data/saved_pkl/dance_3.pkl", 14],
|
||||
[1, "data/images/2.jpeg", "data/saved_pose/dance_3_right","data/saved_frames/dance_3_right","data/saved_pkl/dance_3_right.pkl", 13], # seed 13
|
||||
[4, "data/images/7.jpeg", "data/saved_pose/ubc","data/saved_frames/ubc","data/saved_pkl/ubc.pkl", 14],
|
||||
]
|
||||
|
||||
|
||||
log_dir: 'results' # save dir
|
||||
|
||||
test_model: checkpoints/animate-x_ckpt.pth
|
||||
partial_keys: [
|
||||
['image','local_image', 'dwpose','pose_embeddings'],
|
||||
]
|
||||
|
||||
|
||||
|
||||
|
||||
# default settings
|
||||
TASK_TYPE: inference_animate_x_entrance
|
||||
use_fp16: True
|
||||
guide_scale: 2.5
|
||||
vit_resolution: [224, 224]
|
||||
use_fp16: True
|
||||
batch_size: 1
|
||||
latent_random_ref: True
|
||||
chunk_size: 2
|
||||
decoder_bs: 2
|
||||
scale: 8
|
||||
use_fps_condition: False
|
||||
|
||||
embedder: {
|
||||
'type': 'FrozenOpenCLIPTextVisualEmbedder',
|
||||
'layer': 'penultimate',
|
||||
'pretrained': 'checkpoints/open_clip_pytorch_model.bin'
|
||||
}
|
||||
|
||||
auto_encoder: {
|
||||
'type': 'AutoencoderKL',
|
||||
'ddconfig': {
|
||||
'double_z': True,
|
||||
'z_channels': 4,
|
||||
'resolution': 256,
|
||||
'in_channels': 3,
|
||||
'out_ch': 3,
|
||||
'ch': 128,
|
||||
'ch_mult': [1, 2, 4, 4],
|
||||
'num_res_blocks': 2,
|
||||
'attn_resolutions': [],
|
||||
'dropout': 0.0,
|
||||
'video_kernel_size': [3, 1, 1]
|
||||
},
|
||||
'embed_dim': 4,
|
||||
'pretrained': 'checkpoints/v2-1_512-ema-pruned.ckpt'
|
||||
}
|
||||
|
||||
|
||||
UNet: {
|
||||
'type': 'UNetSD_Animate_X',
|
||||
'config': None,
|
||||
'in_dim': 4,
|
||||
'num': 0,
|
||||
'no_hand': True,
|
||||
'dim': 320,
|
||||
'y_dim': 1024,
|
||||
'context_dim': 1024,
|
||||
'out_dim': 4,
|
||||
'dim_mult': [1, 2, 4, 4],
|
||||
'num_heads': 8,
|
||||
'head_dim': 64,
|
||||
'num_res_blocks': 2,
|
||||
'dropout': 0.1,
|
||||
'temporal_attention': True,
|
||||
'num_tokens': 4,
|
||||
'temporal_attn_times': 1,
|
||||
'use_checkpoint': True,
|
||||
'use_fps_condition': False,
|
||||
'use_sim_mask': False
|
||||
}
|
||||
video_compositions: ['image', 'local_image', 'dwpose', 'randomref', 'randomref_pose', 'pose_embedding']
|
||||
Diffusion: {
|
||||
'type': 'DiffusionDDIM',
|
||||
'schedule': 'linear_sd',
|
||||
'schedule_param': {
|
||||
'num_timesteps': 1000,
|
||||
"init_beta": 0.00085,
|
||||
"last_beta": 0.0120,
|
||||
'zero_terminal_snr': True,
|
||||
},
|
||||
'mean_type': 'v',
|
||||
'loss_type': 'mse',
|
||||
'var_type': 'fixed_small',
|
||||
'rescale_timesteps': False,
|
||||
'noise_strength': 0.1
|
||||
}
|
||||
use_DiffusionDPM: False
|
||||
CPU_CLIP_VAE: True
|
||||
|
||||
|
After Width: | Height: | Size: 3.6 MiB |
|
After Width: | Height: | Size: 3.7 MiB |
|
After Width: | Height: | Size: 88 KiB |
|
After Width: | Height: | Size: 1.0 MiB |
|
After Width: | Height: | Size: 139 KiB |
|
After Width: | Height: | Size: 951 KiB |
|
After Width: | Height: | Size: 575 KiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 405 KiB |
|
After Width: | Height: | Size: 1.6 MiB |
|
After Width: | Height: | Size: 191 KiB |
|
After Width: | Height: | Size: 1.7 MiB |
|
After Width: | Height: | Size: 294 KiB |
|
After Width: | Height: | Size: 1.3 MiB |
@@ -0,0 +1,127 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import onnxruntime
|
||||
|
||||
def nms(boxes, scores, nms_thr):
|
||||
"""Single class NMS implemented in Numpy."""
|
||||
x1 = boxes[:, 0]
|
||||
y1 = boxes[:, 1]
|
||||
x2 = boxes[:, 2]
|
||||
y2 = boxes[:, 3]
|
||||
|
||||
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
|
||||
order = scores.argsort()[::-1]
|
||||
|
||||
keep = []
|
||||
while order.size > 0:
|
||||
i = order[0]
|
||||
keep.append(i)
|
||||
xx1 = np.maximum(x1[i], x1[order[1:]])
|
||||
yy1 = np.maximum(y1[i], y1[order[1:]])
|
||||
xx2 = np.minimum(x2[i], x2[order[1:]])
|
||||
yy2 = np.minimum(y2[i], y2[order[1:]])
|
||||
|
||||
w = np.maximum(0.0, xx2 - xx1 + 1)
|
||||
h = np.maximum(0.0, yy2 - yy1 + 1)
|
||||
inter = w * h
|
||||
ovr = inter / (areas[i] + areas[order[1:]] - inter)
|
||||
|
||||
inds = np.where(ovr <= nms_thr)[0]
|
||||
order = order[inds + 1]
|
||||
|
||||
return keep
|
||||
|
||||
def multiclass_nms(boxes, scores, nms_thr, score_thr):
|
||||
"""Multiclass NMS implemented in Numpy. Class-aware version."""
|
||||
final_dets = []
|
||||
num_classes = scores.shape[1]
|
||||
for cls_ind in range(num_classes):
|
||||
cls_scores = scores[:, cls_ind]
|
||||
valid_score_mask = cls_scores > score_thr
|
||||
if valid_score_mask.sum() == 0:
|
||||
continue
|
||||
else:
|
||||
valid_scores = cls_scores[valid_score_mask]
|
||||
valid_boxes = boxes[valid_score_mask]
|
||||
keep = nms(valid_boxes, valid_scores, nms_thr)
|
||||
if len(keep) > 0:
|
||||
cls_inds = np.ones((len(keep), 1)) * cls_ind
|
||||
dets = np.concatenate(
|
||||
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
|
||||
)
|
||||
final_dets.append(dets)
|
||||
if len(final_dets) == 0:
|
||||
return None
|
||||
return np.concatenate(final_dets, 0)
|
||||
|
||||
def demo_postprocess(outputs, img_size, p6=False):
|
||||
grids = []
|
||||
expanded_strides = []
|
||||
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
|
||||
|
||||
hsizes = [img_size[0] // stride for stride in strides]
|
||||
wsizes = [img_size[1] // stride for stride in strides]
|
||||
|
||||
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
|
||||
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
|
||||
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
|
||||
grids.append(grid)
|
||||
shape = grid.shape[:2]
|
||||
expanded_strides.append(np.full((*shape, 1), stride))
|
||||
|
||||
grids = np.concatenate(grids, 1)
|
||||
expanded_strides = np.concatenate(expanded_strides, 1)
|
||||
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
|
||||
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
|
||||
|
||||
return outputs
|
||||
|
||||
def preprocess(img, input_size, swap=(2, 0, 1)):
|
||||
if len(img.shape) == 3:
|
||||
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
|
||||
else:
|
||||
padded_img = np.ones(input_size, dtype=np.uint8) * 114
|
||||
|
||||
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
|
||||
resized_img = cv2.resize(
|
||||
img,
|
||||
(int(img.shape[1] * r), int(img.shape[0] * r)),
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
).astype(np.uint8)
|
||||
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
|
||||
|
||||
padded_img = padded_img.transpose(swap)
|
||||
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
|
||||
return padded_img, r
|
||||
|
||||
def inference_detector(session, oriImg):
|
||||
input_shape = (640,640)
|
||||
img, ratio = preprocess(oriImg, input_shape)
|
||||
|
||||
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
|
||||
|
||||
output = session.run(None, ort_inputs)
|
||||
|
||||
predictions = demo_postprocess(output[0], input_shape)[0]
|
||||
|
||||
boxes = predictions[:, :4]
|
||||
scores = predictions[:, 4:5] * predictions[:, 5:]
|
||||
|
||||
boxes_xyxy = np.ones_like(boxes)
|
||||
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
|
||||
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
|
||||
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
|
||||
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
|
||||
boxes_xyxy /= ratio
|
||||
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
|
||||
if dets is not None:
|
||||
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
|
||||
isscore = final_scores>0.3
|
||||
iscat = final_cls_inds == 0
|
||||
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
|
||||
final_boxes = final_boxes[isbbox]
|
||||
else:
|
||||
final_boxes = np.array([])
|
||||
|
||||
return final_boxes
|
||||
@@ -0,0 +1,360 @@
|
||||
from typing import List, Tuple
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
def preprocess(
|
||||
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
|
||||
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||
"""Do preprocessing for RTMPose model inference.
|
||||
|
||||
Args:
|
||||
img (np.ndarray): Input image in shape.
|
||||
input_size (tuple): Input image size in shape (w, h).
|
||||
|
||||
Returns:
|
||||
tuple:
|
||||
- resized_img (np.ndarray): Preprocessed image.
|
||||
- center (np.ndarray): Center of image.
|
||||
- scale (np.ndarray): Scale of image.
|
||||
"""
|
||||
# get shape of image
|
||||
img_shape = img.shape[:2]
|
||||
out_img, out_center, out_scale = [], [], []
|
||||
if len(out_bbox) == 0:
|
||||
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
|
||||
for i in range(len(out_bbox)):
|
||||
x0 = out_bbox[i][0]
|
||||
y0 = out_bbox[i][1]
|
||||
x1 = out_bbox[i][2]
|
||||
y1 = out_bbox[i][3]
|
||||
bbox = np.array([x0, y0, x1, y1])
|
||||
|
||||
# get center and scale
|
||||
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
|
||||
|
||||
# do affine transformation
|
||||
resized_img, scale = top_down_affine(input_size, scale, center, img)
|
||||
|
||||
# normalize image
|
||||
mean = np.array([123.675, 116.28, 103.53])
|
||||
std = np.array([58.395, 57.12, 57.375])
|
||||
resized_img = (resized_img - mean) / std
|
||||
|
||||
out_img.append(resized_img)
|
||||
out_center.append(center)
|
||||
out_scale.append(scale)
|
||||
|
||||
return out_img, out_center, out_scale
|
||||
|
||||
|
||||
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
|
||||
"""Inference RTMPose model.
|
||||
|
||||
Args:
|
||||
sess (ort.InferenceSession): ONNXRuntime session.
|
||||
img (np.ndarray): Input image in shape.
|
||||
|
||||
Returns:
|
||||
outputs (np.ndarray): Output of RTMPose model.
|
||||
"""
|
||||
all_out = []
|
||||
# build input
|
||||
for i in range(len(img)):
|
||||
input = [img[i].transpose(2, 0, 1)]
|
||||
|
||||
# build output
|
||||
sess_input = {sess.get_inputs()[0].name: input}
|
||||
sess_output = []
|
||||
for out in sess.get_outputs():
|
||||
sess_output.append(out.name)
|
||||
|
||||
# run model
|
||||
outputs = sess.run(sess_output, sess_input)
|
||||
all_out.append(outputs)
|
||||
|
||||
return all_out
|
||||
|
||||
|
||||
def postprocess(outputs: List[np.ndarray],
|
||||
model_input_size: Tuple[int, int],
|
||||
center: Tuple[int, int],
|
||||
scale: Tuple[int, int],
|
||||
simcc_split_ratio: float = 2.0
|
||||
) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Postprocess for RTMPose model output.
|
||||
|
||||
Args:
|
||||
outputs (np.ndarray): Output of RTMPose model.
|
||||
model_input_size (tuple): RTMPose model Input image size.
|
||||
center (tuple): Center of bbox in shape (x, y).
|
||||
scale (tuple): Scale of bbox in shape (w, h).
|
||||
simcc_split_ratio (float): Split ratio of simcc.
|
||||
|
||||
Returns:
|
||||
tuple:
|
||||
- keypoints (np.ndarray): Rescaled keypoints.
|
||||
- scores (np.ndarray): Model predict scores.
|
||||
"""
|
||||
all_key = []
|
||||
all_score = []
|
||||
for i in range(len(outputs)):
|
||||
# use simcc to decode
|
||||
simcc_x, simcc_y = outputs[i]
|
||||
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
|
||||
|
||||
# rescale keypoints
|
||||
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
|
||||
all_key.append(keypoints[0])
|
||||
all_score.append(scores[0])
|
||||
|
||||
return np.array(all_key), np.array(all_score)
|
||||
|
||||
|
||||
def bbox_xyxy2cs(bbox: np.ndarray,
|
||||
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Transform the bbox format from (x,y,w,h) into (center, scale)
|
||||
|
||||
Args:
|
||||
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
|
||||
as (left, top, right, bottom)
|
||||
padding (float): BBox padding factor that will be multilied to scale.
|
||||
Default: 1.0
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
|
||||
(n, 2)
|
||||
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
|
||||
(n, 2)
|
||||
"""
|
||||
# convert single bbox from (4, ) to (1, 4)
|
||||
dim = bbox.ndim
|
||||
if dim == 1:
|
||||
bbox = bbox[None, :]
|
||||
|
||||
# get bbox center and scale
|
||||
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
|
||||
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
|
||||
scale = np.hstack([x2 - x1, y2 - y1]) * padding
|
||||
|
||||
if dim == 1:
|
||||
center = center[0]
|
||||
scale = scale[0]
|
||||
|
||||
return center, scale
|
||||
|
||||
|
||||
def _fix_aspect_ratio(bbox_scale: np.ndarray,
|
||||
aspect_ratio: float) -> np.ndarray:
|
||||
"""Extend the scale to match the given aspect ratio.
|
||||
|
||||
Args:
|
||||
scale (np.ndarray): The image scale (w, h) in shape (2, )
|
||||
aspect_ratio (float): The ratio of ``w/h``
|
||||
|
||||
Returns:
|
||||
np.ndarray: The reshaped image scale in (2, )
|
||||
"""
|
||||
w, h = np.hsplit(bbox_scale, [1])
|
||||
bbox_scale = np.where(w > h * aspect_ratio,
|
||||
np.hstack([w, w / aspect_ratio]),
|
||||
np.hstack([h * aspect_ratio, h]))
|
||||
return bbox_scale
|
||||
|
||||
|
||||
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
|
||||
"""Rotate a point by an angle.
|
||||
|
||||
Args:
|
||||
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
|
||||
angle_rad (float): rotation angle in radian
|
||||
|
||||
Returns:
|
||||
np.ndarray: Rotated point in shape (2, )
|
||||
"""
|
||||
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
|
||||
rot_mat = np.array([[cs, -sn], [sn, cs]])
|
||||
return rot_mat @ pt
|
||||
|
||||
|
||||
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
|
||||
"""To calculate the affine matrix, three pairs of points are required. This
|
||||
function is used to get the 3rd point, given 2D points a & b.
|
||||
|
||||
The 3rd point is defined by rotating vector `a - b` by 90 degrees
|
||||
anticlockwise, using b as the rotation center.
|
||||
|
||||
Args:
|
||||
a (np.ndarray): The 1st point (x,y) in shape (2, )
|
||||
b (np.ndarray): The 2nd point (x,y) in shape (2, )
|
||||
|
||||
Returns:
|
||||
np.ndarray: The 3rd point.
|
||||
"""
|
||||
direction = a - b
|
||||
c = b + np.r_[-direction[1], direction[0]]
|
||||
return c
|
||||
|
||||
|
||||
def get_warp_matrix(center: np.ndarray,
|
||||
scale: np.ndarray,
|
||||
rot: float,
|
||||
output_size: Tuple[int, int],
|
||||
shift: Tuple[float, float] = (0., 0.),
|
||||
inv: bool = False) -> np.ndarray:
|
||||
"""Calculate the affine transformation matrix that can warp the bbox area
|
||||
in the input image to the output size.
|
||||
|
||||
Args:
|
||||
center (np.ndarray[2, ]): Center of the bounding box (x, y).
|
||||
scale (np.ndarray[2, ]): Scale of the bounding box
|
||||
wrt [width, height].
|
||||
rot (float): Rotation angle (degree).
|
||||
output_size (np.ndarray[2, ] | list(2,)): Size of the
|
||||
destination heatmaps.
|
||||
shift (0-100%): Shift translation ratio wrt the width/height.
|
||||
Default (0., 0.).
|
||||
inv (bool): Option to inverse the affine transform direction.
|
||||
(inv=False: src->dst or inv=True: dst->src)
|
||||
|
||||
Returns:
|
||||
np.ndarray: A 2x3 transformation matrix
|
||||
"""
|
||||
shift = np.array(shift)
|
||||
src_w = scale[0]
|
||||
dst_w = output_size[0]
|
||||
dst_h = output_size[1]
|
||||
|
||||
# compute transformation matrix
|
||||
rot_rad = np.deg2rad(rot)
|
||||
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
|
||||
dst_dir = np.array([0., dst_w * -0.5])
|
||||
|
||||
# get four corners of the src rectangle in the original image
|
||||
src = np.zeros((3, 2), dtype=np.float32)
|
||||
src[0, :] = center + scale * shift
|
||||
src[1, :] = center + src_dir + scale * shift
|
||||
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
|
||||
|
||||
# get four corners of the dst rectangle in the input image
|
||||
dst = np.zeros((3, 2), dtype=np.float32)
|
||||
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
|
||||
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
|
||||
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
|
||||
|
||||
if inv:
|
||||
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
|
||||
else:
|
||||
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
|
||||
|
||||
return warp_mat
|
||||
|
||||
|
||||
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
|
||||
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Get the bbox image as the model input by affine transform.
|
||||
|
||||
Args:
|
||||
input_size (dict): The input size of the model.
|
||||
bbox_scale (dict): The bbox scale of the img.
|
||||
bbox_center (dict): The bbox center of the img.
|
||||
img (np.ndarray): The original image.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: img after affine transform.
|
||||
- np.ndarray[float32]: bbox scale after affine transform.
|
||||
"""
|
||||
w, h = input_size
|
||||
warp_size = (int(w), int(h))
|
||||
|
||||
# reshape bbox to fixed aspect ratio
|
||||
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
|
||||
|
||||
# get the affine matrix
|
||||
center = bbox_center
|
||||
scale = bbox_scale
|
||||
rot = 0
|
||||
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
|
||||
|
||||
# do affine transform
|
||||
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
|
||||
|
||||
return img, bbox_scale
|
||||
|
||||
|
||||
def get_simcc_maximum(simcc_x: np.ndarray,
|
||||
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Get maximum response location and value from simcc representations.
|
||||
|
||||
Note:
|
||||
instance number: N
|
||||
num_keypoints: K
|
||||
heatmap height: H
|
||||
heatmap width: W
|
||||
|
||||
Args:
|
||||
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
|
||||
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
|
||||
|
||||
Returns:
|
||||
tuple:
|
||||
- locs (np.ndarray): locations of maximum heatmap responses in shape
|
||||
(K, 2) or (N, K, 2)
|
||||
- vals (np.ndarray): values of maximum heatmap responses in shape
|
||||
(K,) or (N, K)
|
||||
"""
|
||||
N, K, Wx = simcc_x.shape
|
||||
simcc_x = simcc_x.reshape(N * K, -1)
|
||||
simcc_y = simcc_y.reshape(N * K, -1)
|
||||
|
||||
# get maximum value locations
|
||||
x_locs = np.argmax(simcc_x, axis=1)
|
||||
y_locs = np.argmax(simcc_y, axis=1)
|
||||
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
|
||||
max_val_x = np.amax(simcc_x, axis=1)
|
||||
max_val_y = np.amax(simcc_y, axis=1)
|
||||
|
||||
# get maximum value across x and y axis
|
||||
mask = max_val_x > max_val_y
|
||||
max_val_x[mask] = max_val_y[mask]
|
||||
vals = max_val_x
|
||||
locs[vals <= 0.] = -1
|
||||
|
||||
# reshape
|
||||
locs = locs.reshape(N, K, 2)
|
||||
vals = vals.reshape(N, K)
|
||||
|
||||
return locs, vals
|
||||
|
||||
|
||||
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
|
||||
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Modulate simcc distribution with Gaussian.
|
||||
|
||||
Args:
|
||||
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
|
||||
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
|
||||
simcc_split_ratio (int): The split ratio of simcc.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
|
||||
- np.ndarray[float32]: scores in shape (K,) or (n, K)
|
||||
"""
|
||||
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
|
||||
keypoints /= simcc_split_ratio
|
||||
|
||||
return keypoints, scores
|
||||
|
||||
|
||||
def inference_pose(session, out_bbox, oriImg):
|
||||
h, w = session.get_inputs()[0].shape[2:]
|
||||
model_input_size = (w, h)
|
||||
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
|
||||
outputs = inference(session, resized_img)
|
||||
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
|
||||
|
||||
return keypoints, scores
|
||||
@@ -0,0 +1,335 @@
|
||||
import math
|
||||
import numpy as np
|
||||
import matplotlib
|
||||
import cv2
|
||||
|
||||
|
||||
eps = 0.01
|
||||
|
||||
|
||||
def smart_resize(x, s):
|
||||
Ht, Wt = s
|
||||
if x.ndim == 2:
|
||||
Ho, Wo = x.shape
|
||||
Co = 1
|
||||
else:
|
||||
Ho, Wo, Co = x.shape
|
||||
if Co == 3 or Co == 1:
|
||||
k = float(Ht + Wt) / float(Ho + Wo)
|
||||
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
|
||||
else:
|
||||
return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2)
|
||||
|
||||
|
||||
def smart_resize_k(x, fx, fy):
|
||||
if x.ndim == 2:
|
||||
Ho, Wo = x.shape
|
||||
Co = 1
|
||||
else:
|
||||
Ho, Wo, Co = x.shape
|
||||
Ht, Wt = Ho * fy, Wo * fx
|
||||
if Co == 3 or Co == 1:
|
||||
k = float(Ht + Wt) / float(Ho + Wo)
|
||||
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
|
||||
else:
|
||||
return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2)
|
||||
|
||||
|
||||
def padRightDownCorner(img, stride, padValue):
|
||||
h = img.shape[0]
|
||||
w = img.shape[1]
|
||||
|
||||
pad = 4 * [None]
|
||||
pad[0] = 0 # up
|
||||
pad[1] = 0 # left
|
||||
pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down
|
||||
pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right
|
||||
|
||||
img_padded = img
|
||||
pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1))
|
||||
img_padded = np.concatenate((pad_up, img_padded), axis=0)
|
||||
pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1))
|
||||
img_padded = np.concatenate((pad_left, img_padded), axis=1)
|
||||
pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1))
|
||||
img_padded = np.concatenate((img_padded, pad_down), axis=0)
|
||||
pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1))
|
||||
img_padded = np.concatenate((img_padded, pad_right), axis=1)
|
||||
|
||||
return img_padded, pad
|
||||
|
||||
|
||||
def transfer(model, model_weights):
|
||||
transfered_model_weights = {}
|
||||
for weights_name in model.state_dict().keys():
|
||||
transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])]
|
||||
return transfered_model_weights
|
||||
|
||||
|
||||
def draw_bodypose(canvas, candidate, subset):
|
||||
H, W, C = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
|
||||
stickwidth = 4
|
||||
|
||||
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
|
||||
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
|
||||
[1, 16], [16, 18], [3, 17], [6, 18]]
|
||||
|
||||
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
|
||||
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
|
||||
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
|
||||
|
||||
for i in range(17):
|
||||
for n in range(len(subset)):
|
||||
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||
if -1 in index:
|
||||
continue
|
||||
Y = candidate[index.astype(int), 0] * float(W)
|
||||
X = candidate[index.astype(int), 1] * float(H)
|
||||
mX = np.mean(X)
|
||||
mY = np.mean(Y)
|
||||
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||
cv2.fillConvexPoly(canvas, polygon, colors[i])
|
||||
|
||||
canvas = (canvas * 0.6).astype(np.uint8)
|
||||
|
||||
for i in range(18):
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i])
|
||||
if index == -1:
|
||||
continue
|
||||
x, y = candidate[index][0:2]
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_body_and_foot(canvas, candidate, subset):
|
||||
H, W, C = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
|
||||
stickwidth = 4
|
||||
|
||||
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
|
||||
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
|
||||
[1, 16], [16, 18], [14,19], [11, 20]]
|
||||
|
||||
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
|
||||
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
|
||||
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85], [170, 255, 255], [255, 255, 0]]
|
||||
|
||||
for i in range(19):
|
||||
for n in range(len(subset)):
|
||||
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||
if -1 in index:
|
||||
continue
|
||||
Y = candidate[index.astype(int), 0] * float(W)
|
||||
X = candidate[index.astype(int), 1] * float(H)
|
||||
mX = np.mean(X)
|
||||
mY = np.mean(Y)
|
||||
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||
cv2.fillConvexPoly(canvas, polygon, colors[i])
|
||||
|
||||
canvas = (canvas * 0.6).astype(np.uint8)
|
||||
|
||||
for i in range(20):
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i])
|
||||
if index == -1:
|
||||
continue
|
||||
x, y = candidate[index][0:2]
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_handpose(canvas, all_hand_peaks):
|
||||
H, W, C = canvas.shape
|
||||
|
||||
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
|
||||
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
|
||||
|
||||
for peaks in all_hand_peaks:
|
||||
peaks = np.array(peaks)
|
||||
|
||||
for ie, e in enumerate(edges):
|
||||
x1, y1 = peaks[e[0]]
|
||||
x2, y2 = peaks[e[1]]
|
||||
x1 = int(x1 * W)
|
||||
y1 = int(y1 * H)
|
||||
x2 = int(x2 * W)
|
||||
y2 = int(y2 * H)
|
||||
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||
cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2)
|
||||
|
||||
for i, keyponit in enumerate(peaks):
|
||||
x, y = keyponit
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_facepose(canvas, all_lmks):
|
||||
H, W, C = canvas.shape
|
||||
for lmks in all_lmks:
|
||||
lmks = np.array(lmks)
|
||||
for lmk in lmks:
|
||||
x, y = lmk
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
# detect hand according to body pose keypoints
|
||||
def handDetect(candidate, subset, oriImg):
|
||||
# right hand: wrist 4, elbow 3, shoulder 2
|
||||
# left hand: wrist 7, elbow 6, shoulder 5
|
||||
ratioWristElbow = 0.33
|
||||
detect_result = []
|
||||
image_height, image_width = oriImg.shape[0:2]
|
||||
for person in subset.astype(int):
|
||||
# if any of three not detected
|
||||
has_left = np.sum(person[[5, 6, 7]] == -1) == 0
|
||||
has_right = np.sum(person[[2, 3, 4]] == -1) == 0
|
||||
if not (has_left or has_right):
|
||||
continue
|
||||
hands = []
|
||||
#left hand
|
||||
if has_left:
|
||||
left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]]
|
||||
x1, y1 = candidate[left_shoulder_index][:2]
|
||||
x2, y2 = candidate[left_elbow_index][:2]
|
||||
x3, y3 = candidate[left_wrist_index][:2]
|
||||
hands.append([x1, y1, x2, y2, x3, y3, True])
|
||||
# right hand
|
||||
if has_right:
|
||||
right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]]
|
||||
x1, y1 = candidate[right_shoulder_index][:2]
|
||||
x2, y2 = candidate[right_elbow_index][:2]
|
||||
x3, y3 = candidate[right_wrist_index][:2]
|
||||
hands.append([x1, y1, x2, y2, x3, y3, False])
|
||||
|
||||
for x1, y1, x2, y2, x3, y3, is_left in hands:
|
||||
|
||||
x = x3 + ratioWristElbow * (x3 - x2)
|
||||
y = y3 + ratioWristElbow * (y3 - y2)
|
||||
distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2)
|
||||
distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)
|
||||
width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder)
|
||||
# x-y refers to the center --> offset to topLeft point
|
||||
# handRectangle.x -= handRectangle.width / 2.f;
|
||||
# handRectangle.y -= handRectangle.height / 2.f;
|
||||
x -= width / 2
|
||||
y -= width / 2 # width = height
|
||||
# overflow the image
|
||||
if x < 0: x = 0
|
||||
if y < 0: y = 0
|
||||
width1 = width
|
||||
width2 = width
|
||||
if x + width > image_width: width1 = image_width - x
|
||||
if y + width > image_height: width2 = image_height - y
|
||||
width = min(width1, width2)
|
||||
# the max hand box value is 20 pixels
|
||||
if width >= 20:
|
||||
detect_result.append([int(x), int(y), int(width), is_left])
|
||||
|
||||
'''
|
||||
return value: [[x, y, w, True if left hand else False]].
|
||||
width=height since the network require squared input.
|
||||
x, y is the coordinate of top left
|
||||
'''
|
||||
return detect_result
|
||||
|
||||
|
||||
# Written by Lvmin
|
||||
def faceDetect(candidate, subset, oriImg):
|
||||
# left right eye ear 14 15 16 17
|
||||
detect_result = []
|
||||
image_height, image_width = oriImg.shape[0:2]
|
||||
for person in subset.astype(int):
|
||||
has_head = person[0] > -1
|
||||
if not has_head:
|
||||
continue
|
||||
|
||||
has_left_eye = person[14] > -1
|
||||
has_right_eye = person[15] > -1
|
||||
has_left_ear = person[16] > -1
|
||||
has_right_ear = person[17] > -1
|
||||
|
||||
if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear):
|
||||
continue
|
||||
|
||||
head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]]
|
||||
|
||||
width = 0.0
|
||||
x0, y0 = candidate[head][:2]
|
||||
|
||||
if has_left_eye:
|
||||
x1, y1 = candidate[left_eye][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 3.0)
|
||||
|
||||
if has_right_eye:
|
||||
x1, y1 = candidate[right_eye][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 3.0)
|
||||
|
||||
if has_left_ear:
|
||||
x1, y1 = candidate[left_ear][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 1.5)
|
||||
|
||||
if has_right_ear:
|
||||
x1, y1 = candidate[right_ear][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 1.5)
|
||||
|
||||
x, y = x0, y0
|
||||
|
||||
x -= width
|
||||
y -= width
|
||||
|
||||
if x < 0:
|
||||
x = 0
|
||||
|
||||
if y < 0:
|
||||
y = 0
|
||||
|
||||
width1 = width * 2
|
||||
width2 = width * 2
|
||||
|
||||
if x + width > image_width:
|
||||
width1 = image_width - x
|
||||
|
||||
if y + width > image_height:
|
||||
width2 = image_height - y
|
||||
|
||||
width = min(width1, width2)
|
||||
|
||||
if width >= 20:
|
||||
detect_result.append([int(x), int(y), int(width)])
|
||||
|
||||
return detect_result
|
||||
|
||||
|
||||
# get max index of 2d array
|
||||
def npmax(array):
|
||||
arrayindex = array.argmax(1)
|
||||
arrayvalue = array.max(1)
|
||||
i = arrayvalue.argmax()
|
||||
j = arrayindex[i]
|
||||
return i, j
|
||||
@@ -0,0 +1,48 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import onnxruntime as ort
|
||||
from dwpose.onnxdet import inference_detector
|
||||
from dwpose.onnxpose import inference_pose
|
||||
|
||||
class Wholebody:
|
||||
def __init__(self):
|
||||
device = 'cuda' # 'cpu' #
|
||||
providers = ['CPUExecutionProvider'
|
||||
] if device == 'cpu' else ['CUDAExecutionProvider']
|
||||
onnx_det = 'checkpoints/yolox_l.onnx'
|
||||
onnx_pose = 'checkpoints/dw-ll_ucoco_384.onnx'
|
||||
|
||||
self.session_det = ort.InferenceSession(path_or_bytes=onnx_det, providers=providers)
|
||||
self.session_pose = ort.InferenceSession(path_or_bytes=onnx_pose, providers=providers)
|
||||
|
||||
def __call__(self, oriImg):
|
||||
det_result = inference_detector(self.session_det, oriImg)
|
||||
keypoints, scores = inference_pose(self.session_pose, det_result, oriImg)
|
||||
|
||||
keypoints_info = np.concatenate(
|
||||
(keypoints, scores[..., None]), axis=-1)
|
||||
# compute neck joint
|
||||
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
|
||||
# neck score when visualizing pred
|
||||
neck[:, 2:4] = np.logical_and(
|
||||
keypoints_info[:, 5, 2:4] > 0.3,
|
||||
keypoints_info[:, 6, 2:4] > 0.3).astype(int)
|
||||
new_keypoints_info = np.insert(
|
||||
keypoints_info, 17, neck, axis=1)
|
||||
mmpose_idx = [
|
||||
17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3
|
||||
]
|
||||
openpose_idx = [
|
||||
1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17
|
||||
]
|
||||
new_keypoints_info[:, openpose_idx] = \
|
||||
new_keypoints_info[:, mmpose_idx]
|
||||
keypoints_info = new_keypoints_info
|
||||
|
||||
keypoints, scores = keypoints_info[
|
||||
..., :2], keypoints_info[..., 2]
|
||||
|
||||
return keypoints, scores
|
||||
|
||||
|
||||
@@ -0,0 +1,236 @@
|
||||
name: /mnt/user/miniconda3/envs/dtrans
|
||||
channels:
|
||||
- http://mirrors.aliyun.com/anaconda/pkgs/main
|
||||
- defaults
|
||||
dependencies:
|
||||
- _libgcc_mutex=0.1=main
|
||||
- _openmp_mutex=5.1=1_gnu
|
||||
- ca-certificates=2023.12.12=h06a4308_0
|
||||
- ld_impl_linux-64=2.38=h1181459_1
|
||||
- libffi=3.4.4=h6a678d5_0
|
||||
- libgcc-ng=11.2.0=h1234567_1
|
||||
- libgomp=11.2.0=h1234567_1
|
||||
- libstdcxx-ng=11.2.0=h1234567_1
|
||||
- ncurses=6.4=h6a678d5_0
|
||||
- openssl=3.0.12=h7f8727e_0
|
||||
- pip=23.3.1=py39h06a4308_0
|
||||
- python=3.9.18=h955ad1f_0
|
||||
- readline=8.2=h5eee18b_0
|
||||
- setuptools=68.2.2=py39h06a4308_0
|
||||
- sqlite=3.41.2=h5eee18b_0
|
||||
- tk=8.6.12=h1ccaba5_0
|
||||
- wheel=0.41.2=py39h06a4308_0
|
||||
- xz=5.4.5=h5eee18b_0
|
||||
- zlib=1.2.13=h5eee18b_0
|
||||
- pip:
|
||||
- aiofiles==23.2.1
|
||||
- aiohttp==3.9.1
|
||||
- aiosignal==1.3.1
|
||||
- aliyun-python-sdk-core==2.14.0
|
||||
- aliyun-python-sdk-kms==2.16.2
|
||||
- altair==5.2.0
|
||||
- annotated-types==0.6.0
|
||||
- antlr4-python3-runtime==4.9.3
|
||||
- anyio==4.2.0
|
||||
- argparse==1.4.0
|
||||
- asttokens==2.4.1
|
||||
- async-timeout==4.0.3
|
||||
- attrs==23.2.0
|
||||
- automat==22.10.0
|
||||
- beartype==0.16.4
|
||||
- blessed==1.20.0
|
||||
- buildtools==1.0.6
|
||||
- causal-conv1d==1.1.3.post1
|
||||
- certifi==2023.11.17
|
||||
- cffi==1.16.0
|
||||
- chardet==5.2.0
|
||||
- charset-normalizer==3.3.2
|
||||
- clean-fid==0.1.35
|
||||
- click==8.1.7
|
||||
- clip==1.0
|
||||
- cmake==3.28.1
|
||||
- colorama==0.4.6
|
||||
- coloredlogs==15.0.1
|
||||
- constantly==23.10.4
|
||||
- contourpy==1.2.0
|
||||
- crcmod==1.7
|
||||
- cryptography==41.0.7
|
||||
- cycler==0.12.1
|
||||
- decorator==5.1.1
|
||||
- decord==0.6.0
|
||||
- diffusers==0.26.3
|
||||
- docopt==0.6.2
|
||||
- easydict==1.11
|
||||
- einops==0.7.0
|
||||
- exceptiongroup==1.2.0
|
||||
- executing==2.0.1
|
||||
- fairscale==0.4.13
|
||||
- fastapi==0.109.0
|
||||
- ffmpeg==1.4
|
||||
- ffmpy==0.3.1
|
||||
- filelock==3.13.1
|
||||
- flatbuffers==24.3.25
|
||||
- fonttools==4.47.2
|
||||
- frozenlist==1.4.1
|
||||
- fsspec==2023.12.2
|
||||
- ftfy==6.1.3
|
||||
- furl==2.1.3
|
||||
- gpustat==1.1.1
|
||||
- gradio==4.14.0
|
||||
- gradio-client==0.8.0
|
||||
- greenlet==3.0.3
|
||||
- h11==0.14.0
|
||||
- httpcore==1.0.2
|
||||
- httpx==0.26.0
|
||||
- huggingface-hub==0.20.2
|
||||
- humanfriendly==10.0
|
||||
- hyperlink==21.0.0
|
||||
- idna==3.6
|
||||
- imageio==2.33.1
|
||||
- imageio-ffmpeg==0.4.9
|
||||
- importlib-metadata==7.0.1
|
||||
- importlib-resources==6.1.1
|
||||
- incremental==22.10.0
|
||||
- ipdb==0.13.13
|
||||
- ipython==8.18.1
|
||||
- jedi==0.19.1
|
||||
- jinja2==3.1.3
|
||||
- jmespath==0.10.0
|
||||
- joblib==1.3.2
|
||||
- jsonschema==4.21.0
|
||||
- jsonschema-specifications==2023.12.1
|
||||
- kiwisolver==1.4.5
|
||||
- kornia==0.7.1
|
||||
- lazy-loader==0.3
|
||||
- lightning-utilities==0.10.0
|
||||
- lit==17.0.6
|
||||
- lpips==0.1.4
|
||||
- mamba-ssm==1.1.4
|
||||
- markdown-it-py==3.0.0
|
||||
- markupsafe==2.1.3
|
||||
- matplotlib==3.8.2
|
||||
- matplotlib-inline==0.1.6
|
||||
- mdurl==0.1.2
|
||||
- motion-vector-extractor==1.0.6
|
||||
- mpmath==1.3.0
|
||||
- multidict==6.0.4
|
||||
- mypy-extensions==1.0.0
|
||||
- networkx==3.2.1
|
||||
- ninja==1.11.1.1
|
||||
- numpy==1.26.3
|
||||
- nvidia-cublas-cu11==11.10.3.66
|
||||
- nvidia-cublas-cu12==12.1.3.1
|
||||
- nvidia-cuda-cupti-cu11==11.7.101
|
||||
- nvidia-cuda-cupti-cu12==12.1.105
|
||||
- nvidia-cuda-nvrtc-cu11==11.7.99
|
||||
- nvidia-cuda-nvrtc-cu12==12.1.105
|
||||
- nvidia-cuda-runtime-cu11==11.7.99
|
||||
- nvidia-cuda-runtime-cu12==12.1.105
|
||||
- nvidia-cudnn-cu11==8.5.0.96
|
||||
- nvidia-cudnn-cu12==8.9.2.26
|
||||
- nvidia-cufft-cu11==10.9.0.58
|
||||
- nvidia-cufft-cu12==11.0.2.54
|
||||
- nvidia-curand-cu11==10.2.10.91
|
||||
- nvidia-curand-cu12==10.3.2.106
|
||||
- nvidia-cusolver-cu11==11.4.0.1
|
||||
- nvidia-cusolver-cu12==11.4.5.107
|
||||
- nvidia-cusparse-cu11==11.7.4.91
|
||||
- nvidia-cusparse-cu12==12.1.0.106
|
||||
- nvidia-ml-py==12.535.133
|
||||
- nvidia-nccl-cu11==2.14.3
|
||||
- nvidia-nccl-cu12==2.19.3
|
||||
- nvidia-nvjitlink-cu12==12.3.101
|
||||
- nvidia-nvtx-cu11==11.7.91
|
||||
- nvidia-nvtx-cu12==12.1.105
|
||||
- omegaconf==2.3.0
|
||||
- onnxruntime==1.18.0
|
||||
- open-clip-torch==2.24.0
|
||||
- opencv-python==4.5.3.56
|
||||
- opencv-python-headless==4.9.0.80
|
||||
- orderedmultidict==1.0.1
|
||||
- orjson==3.9.10
|
||||
- oss2==2.18.4
|
||||
- packaging==23.2
|
||||
- pandas==2.1.4
|
||||
- parso==0.8.3
|
||||
- pexpect==4.9.0
|
||||
- pillow==10.2.0
|
||||
- piq==0.8.0
|
||||
- pkgconfig==1.5.5
|
||||
- prompt-toolkit==3.0.43
|
||||
- protobuf==4.25.2
|
||||
- psutil==5.9.8
|
||||
- ptflops==0.7.2.2
|
||||
- ptyprocess==0.7.0
|
||||
- pure-eval==0.2.2
|
||||
- pycparser==2.21
|
||||
- pycryptodome==3.20.0
|
||||
- pydantic==2.5.3
|
||||
- pydantic-core==2.14.6
|
||||
- pydub==0.25.1
|
||||
- pygments==2.17.2
|
||||
- pynvml==11.5.0
|
||||
- pyparsing==3.1.1
|
||||
- pyre-extensions==0.0.29
|
||||
- python-dateutil==2.8.2
|
||||
- python-multipart==0.0.6
|
||||
- pytorch-lightning==2.1.3
|
||||
- pytz==2023.3.post1
|
||||
- pyyaml==6.0.1
|
||||
- redo==2.0.4
|
||||
- referencing==0.32.1
|
||||
- regex==2023.12.25
|
||||
- requests==2.31.0
|
||||
- rich==13.7.0
|
||||
- rotary-embedding-torch==0.5.3
|
||||
- rpds-py==0.17.1
|
||||
- ruff==0.2.0
|
||||
- safetensors==0.4.1
|
||||
- scikit-image==0.22.0
|
||||
- scikit-learn==1.4.0
|
||||
- scipy==1.11.4
|
||||
- semantic-version==2.10.0
|
||||
- sentencepiece==0.1.99
|
||||
- shellingham==1.5.4
|
||||
- simplejson==3.19.2
|
||||
- six==1.16.0
|
||||
- sk-video==1.1.10
|
||||
- sniffio==1.3.0
|
||||
- sqlalchemy==2.0.27
|
||||
- stack-data==0.6.3
|
||||
- starlette==0.35.1
|
||||
- sympy==1.12
|
||||
- thop==0.1.1-2209072238
|
||||
- threadpoolctl==3.2.0
|
||||
- tifffile==2023.12.9
|
||||
- timm==0.9.12
|
||||
- tokenizers==0.15.0
|
||||
- tomli==2.0.1
|
||||
- tomlkit==0.12.0
|
||||
- toolz==0.12.0
|
||||
- torch==2.0.1+cu118
|
||||
- torchaudio==2.0.2+cu118
|
||||
- torchdiffeq==0.2.3
|
||||
- torchmetrics==1.3.0.post0
|
||||
- torchsde==0.2.6
|
||||
- torchvision==0.15.2+cu118
|
||||
- tqdm==4.66.1
|
||||
- traitlets==5.14.1
|
||||
- trampoline==0.1.2
|
||||
- transformers==4.36.2
|
||||
- triton==2.0.0
|
||||
- twisted==23.10.0
|
||||
- typer==0.9.0
|
||||
- typing-extensions==4.9.0
|
||||
- typing-inspect==0.9.0
|
||||
- tzdata==2023.4
|
||||
- urllib3==2.1.0
|
||||
- uvicorn==0.26.0
|
||||
- wcwidth==0.2.13
|
||||
- websockets==11.0.3
|
||||
- xformers==0.0.20
|
||||
- yarl==1.9.4
|
||||
- zipp==3.17.0
|
||||
- zope-interface==6.2
|
||||
- onnxruntime-gpu==1.13.1
|
||||
prefix: /mnt/user/miniconda3/envs/dtrans
|
||||
@@ -0,0 +1,16 @@
|
||||
import os
|
||||
import sys
|
||||
import copy
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import logging
|
||||
import itertools
|
||||
import numpy as np
|
||||
|
||||
from utils.config import Config
|
||||
from utils.registry_class import INFER_ENGINE
|
||||
|
||||
if __name__ == '__main__':
|
||||
cfg_update = Config(load=True)
|
||||
INFER_ENGINE.build(dict(type=cfg_update.TASK_TYPE), cfg_update=cfg_update.cfg_dict)
|
||||
@@ -0,0 +1,324 @@
|
||||
import os
|
||||
os.environ["KMP_DUPLICATE_LIB_OK"]="TRUE"
|
||||
import cv2
|
||||
import torch
|
||||
import numpy as np
|
||||
import json
|
||||
import copy
|
||||
import torch
|
||||
import random
|
||||
import argparse
|
||||
import shutil
|
||||
import tempfile
|
||||
import subprocess
|
||||
import numpy as np
|
||||
import math
|
||||
|
||||
import torch.multiprocessing as mp
|
||||
import torch.distributed as dist
|
||||
import pickle
|
||||
import logging
|
||||
from io import BytesIO
|
||||
import oss2 as oss
|
||||
import os.path as osp
|
||||
|
||||
import sys
|
||||
import dwpose.util as util
|
||||
from dwpose.wholebody import Wholebody
|
||||
|
||||
import pickle
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def smoothing_factor(t_e, cutoff):
|
||||
r = 2 * math.pi * cutoff * t_e
|
||||
return r / (r + 1)
|
||||
|
||||
|
||||
def exponential_smoothing(a, x, x_prev):
|
||||
return a * x + (1 - a) * x_prev
|
||||
|
||||
|
||||
class OneEuroFilter:
|
||||
def __init__(self, t0, x0, dx0=0.0, min_cutoff=1.0, beta=0.0,
|
||||
d_cutoff=1.0):
|
||||
"""Initialize the one euro filter."""
|
||||
# The parameters.
|
||||
self.min_cutoff = float(min_cutoff)
|
||||
self.beta = float(beta)
|
||||
self.d_cutoff = float(d_cutoff)
|
||||
# Previous values.
|
||||
self.x_prev = x0
|
||||
self.dx_prev = float(dx0)
|
||||
self.t_prev = float(t0)
|
||||
|
||||
def __call__(self, t, x):
|
||||
"""Compute the filtered signal."""
|
||||
t_e = t - self.t_prev
|
||||
|
||||
# The filtered derivative of the signal.
|
||||
a_d = smoothing_factor(t_e, self.d_cutoff)
|
||||
dx = (x - self.x_prev) / t_e
|
||||
dx_hat = exponential_smoothing(a_d, dx, self.dx_prev)
|
||||
|
||||
# The filtered signal.
|
||||
cutoff = self.min_cutoff + self.beta * abs(dx_hat)
|
||||
a = smoothing_factor(t_e, cutoff)
|
||||
x_hat = exponential_smoothing(a, x, self.x_prev)
|
||||
|
||||
# Memorize the previous values.
|
||||
self.x_prev = x_hat
|
||||
self.dx_prev = dx_hat
|
||||
self.t_prev = t
|
||||
|
||||
return x_hat
|
||||
|
||||
|
||||
|
||||
def get_logger(name="essmc2"):
|
||||
logger = logging.getLogger(name)
|
||||
logger.propagate = False
|
||||
if len(logger.handlers) == 0:
|
||||
std_handler = logging.StreamHandler(sys.stdout)
|
||||
formatter = logging.Formatter(
|
||||
'%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
std_handler.setFormatter(formatter)
|
||||
std_handler.setLevel(logging.INFO)
|
||||
logger.setLevel(logging.INFO)
|
||||
logger.addHandler(std_handler)
|
||||
return logger
|
||||
|
||||
class DWposeDetector:
|
||||
def __init__(self):
|
||||
|
||||
self.pose_estimation = Wholebody()
|
||||
|
||||
def __call__(self, oriImg):
|
||||
oriImg = oriImg.copy()
|
||||
H, W, C = oriImg.shape
|
||||
with torch.no_grad():
|
||||
candidate, subset = self.pose_estimation(oriImg)
|
||||
candidate = candidate[0][np.newaxis, :, :]
|
||||
subset = subset[0][np.newaxis, :]
|
||||
nums, keys, locs = candidate.shape
|
||||
candidate[..., 0] /= float(W)
|
||||
candidate[..., 1] /= float(H)
|
||||
body = candidate[:,:18].copy()
|
||||
body = body.reshape(nums*18, locs)
|
||||
score = subset[:,:18].copy()
|
||||
|
||||
for i in range(len(score)):
|
||||
for j in range(len(score[i])):
|
||||
if score[i][j] > 0.3:
|
||||
score[i][j] = int(18*i+j)
|
||||
else:
|
||||
score[i][j] = -1
|
||||
|
||||
un_visible = subset<0.3
|
||||
candidate[un_visible] = -1
|
||||
|
||||
bodyfoot_score = subset[:,:24].copy()
|
||||
for i in range(len(bodyfoot_score)):
|
||||
for j in range(len(bodyfoot_score[i])):
|
||||
if bodyfoot_score[i][j] > 0.3:
|
||||
bodyfoot_score[i][j] = int(18*i+j)
|
||||
else:
|
||||
bodyfoot_score[i][j] = -1
|
||||
if -1 not in bodyfoot_score[:,18] and -1 not in bodyfoot_score[:,19]:
|
||||
bodyfoot_score[:,18] = np.array([18.])
|
||||
else:
|
||||
bodyfoot_score[:,18] = np.array([-1.])
|
||||
if -1 not in bodyfoot_score[:,21] and -1 not in bodyfoot_score[:,22]:
|
||||
bodyfoot_score[:,19] = np.array([19.])
|
||||
else:
|
||||
bodyfoot_score[:,19] = np.array([-1.])
|
||||
bodyfoot_score = bodyfoot_score[:, :20]
|
||||
|
||||
bodyfoot = candidate[:,:24].copy()
|
||||
|
||||
for i in range(nums):
|
||||
if -1 not in bodyfoot[i][18] and -1 not in bodyfoot[i][19]:
|
||||
bodyfoot[i][18] = (bodyfoot[i][18]+bodyfoot[i][19])/2
|
||||
else:
|
||||
bodyfoot[i][18] = np.array([-1., -1.])
|
||||
if -1 not in bodyfoot[i][21] and -1 not in bodyfoot[i][22]:
|
||||
bodyfoot[i][19] = (bodyfoot[i][21]+bodyfoot[i][22])/2
|
||||
else:
|
||||
bodyfoot[i][19] = np.array([-1., -1.])
|
||||
|
||||
bodyfoot = bodyfoot[:,:20,:]
|
||||
bodyfoot = bodyfoot.reshape(nums*20, locs)
|
||||
|
||||
foot = candidate[:,18:24]
|
||||
|
||||
faces = candidate[:,24:92]
|
||||
|
||||
hands = candidate[:,92:113]
|
||||
hands = np.vstack([hands, candidate[:,113:]])
|
||||
|
||||
# bodies = dict(candidate=body, subset=score)
|
||||
|
||||
# print(body.shape)
|
||||
# print(bodyfoot.shape)
|
||||
|
||||
# print(body == bodyfoot[:18])
|
||||
|
||||
bodies = dict(candidate=bodyfoot, subset=bodyfoot_score)
|
||||
pose = dict(bodies=bodies, hands=hands, faces=faces)
|
||||
|
||||
# return draw_pose(pose, H, W)
|
||||
return pose
|
||||
|
||||
def draw_pose(pose, H, W):
|
||||
bodies = pose['bodies']
|
||||
faces = pose['faces']
|
||||
hands = pose['hands']
|
||||
candidate = bodies['candidate']
|
||||
subset = bodies['subset']
|
||||
canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8)
|
||||
|
||||
canvas = util.draw_body_and_foot(canvas, candidate, subset)
|
||||
|
||||
canvas = util.draw_handpose(canvas, hands)
|
||||
|
||||
canvas_without_face = copy.deepcopy(canvas)
|
||||
|
||||
canvas = util.draw_facepose(canvas, faces)
|
||||
|
||||
return canvas_without_face, canvas
|
||||
|
||||
def dw_func(_id, frame, dwpose_model, dwpose_woface_folder='tmp_dwpose_wo_face', dwpose_withface_folder='tmp_dwpose_with_face'):
|
||||
|
||||
# frame = cv2.imread(frame_name, cv2.IMREAD_COLOR)
|
||||
pose = dwpose_model(frame)
|
||||
|
||||
return pose
|
||||
|
||||
|
||||
def video2img(video_path, img_dir):
|
||||
# pdb.set_trace()
|
||||
video_capture = cv2.VideoCapture(video_path)
|
||||
|
||||
os.makedirs(img_dir, exist_ok=True)
|
||||
|
||||
# Extract frames from the video
|
||||
success, image = video_capture.read()
|
||||
count = 0
|
||||
while success:
|
||||
# Save frame as JPEG file
|
||||
# cv2.imwrite(os.path.join(img_dir, f'{count:03d}.jpg'), image)
|
||||
if os.path.exists(os.path.join(img_dir, f'frame_{count:04d}.jpg')) == False:
|
||||
cv2.imwrite(os.path.join(img_dir, f'frame_{count:04d}.jpg'), image)
|
||||
|
||||
success, image = video_capture.read()
|
||||
count += 1
|
||||
print("frame: ", count)
|
||||
|
||||
|
||||
def mp_main(args):
|
||||
os.makedirs(args.saved_pose_dir, exist_ok = True)
|
||||
if args.source_video_paths.endswith('mp4'):
|
||||
video_paths = [args.source_video_paths]
|
||||
else:
|
||||
# video list
|
||||
video_paths = [os.path.join(args.source_video_paths, frame_name) for frame_name in os.listdir(args.source_video_paths)]
|
||||
|
||||
logger.info("There are {} videos for extracting poses".format(len(video_paths)))
|
||||
|
||||
logger.info('LOAD: DW Pose Model')
|
||||
dwpose_model = DWposeDetector()
|
||||
|
||||
results_vis = []
|
||||
for i, file_path in enumerate(video_paths):
|
||||
try:
|
||||
|
||||
|
||||
|
||||
logger.info(f"{i}/{len(video_paths)}, {file_path}")
|
||||
|
||||
save_frame_dir = os.path.join(args.saved_frame_dir, os.path.basename(file_path)[:-4])
|
||||
os.makedirs(save_frame_dir, exist_ok = True)
|
||||
video2img(file_path, save_frame_dir)
|
||||
|
||||
videoCapture = cv2.VideoCapture(file_path)
|
||||
cur_output_dir = os.path.join(args.saved_pose, os.path.basename(file_path)[:-4])
|
||||
os.makedirs(cur_output_dir, exist_ok = True)
|
||||
fps = int(videoCapture.get(cv2.CAP_PROP_FPS))
|
||||
bodies = []
|
||||
body_indices = []
|
||||
hands = []
|
||||
faces = []
|
||||
|
||||
idx = 0
|
||||
while videoCapture.isOpened():
|
||||
# get a frame
|
||||
ret, frame = videoCapture.read()
|
||||
|
||||
# print(frame.shape)
|
||||
# import pdb; pdb.set_trace()
|
||||
if ret:
|
||||
size = frame.shape # (1216, 832, 3)
|
||||
pose = dw_func(i, frame, dwpose_model)
|
||||
bodies.append(pose['bodies']['candidate'][:18])
|
||||
body_indices.append(pose['bodies']['subset'][0][:18])
|
||||
faces.append(pose['faces'][0])
|
||||
hands.append(pose['hands'])
|
||||
# results_vis.append(pose)
|
||||
(H,W,_) = size
|
||||
dwpose_woface, dwpose_wface = draw_pose(
|
||||
pose,
|
||||
H,
|
||||
W
|
||||
# draw_face=False,
|
||||
)
|
||||
# output_transformed = cv2.cvtColor(output_transformed, cv2.COLOR_BGR2RGB)
|
||||
# output_transformed = cv2.resize(output_transformed, (W, H))
|
||||
|
||||
# img = Image.fromarray(output_transformed)
|
||||
cv2.imwrite(os.path.join(cur_output_dir, f"frame_{idx:04d}.jpg"), dwpose_woface)
|
||||
# img.save(os.path.join(cur_output_dir, f"frame_{idx:04d}.jpg"))
|
||||
idx += 1
|
||||
# import pdb; pdb.set_trace()
|
||||
else:
|
||||
break
|
||||
logger.info(f'all frames in {file_path} have been read.')
|
||||
videoCapture.release()
|
||||
|
||||
|
||||
new_dict = {}
|
||||
new_dict['bodies'] = np.array(bodies)
|
||||
new_dict['body_indices'] = np.array(body_indices)
|
||||
new_dict['faces'] = np.array(faces)
|
||||
new_dict['hands'] = np.array(hands)
|
||||
new_dict['size'] = size
|
||||
new_dict['fps'] = fps
|
||||
|
||||
|
||||
save_pkl_path = os.path.join(args.saved_pose_dir, os.path.basename(file_path)[:-4]+'.pkl')
|
||||
print(save_pkl_path)
|
||||
with open(save_pkl_path, 'wb') as file:
|
||||
# 使用 pickle.dump() 方法将字典写入文件
|
||||
pickle.dump(new_dict, file)
|
||||
|
||||
# import pdb; pdb.set_trace()
|
||||
|
||||
except:
|
||||
print(file_path," wrong")
|
||||
logger = get_logger('dw pose extraction')
|
||||
|
||||
|
||||
# python
|
||||
|
||||
if __name__=='__main__':
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Simple example of a training script.")
|
||||
parser.add_argument("--source_video_paths", type=str, default="data/videos",)
|
||||
parser.add_argument("--saved_pose_dir", type=str, default="data/saved_pkl",)
|
||||
parser.add_argument("--saved_pose", type=str, default="data/saved_pose",)
|
||||
parser.add_argument("--saved_frame_dir", type=str, default="data/saved_frames",)
|
||||
args = parser.parse_args()
|
||||
|
||||
return args
|
||||
|
||||
args = parse_args()
|
||||
mp_main(args)
|
||||
@@ -0,0 +1,208 @@
|
||||
aiofiles==23.2.1
|
||||
aiohttp==3.9.1
|
||||
aiosignal==1.3.1
|
||||
aliyun-python-sdk-core==2.14.0
|
||||
aliyun-python-sdk-kms==2.16.2
|
||||
altair==5.2.0
|
||||
annotated-types==0.6.0
|
||||
antlr4-python3-runtime==4.9.3
|
||||
anyio==4.2.0
|
||||
asttokens==2.4.1
|
||||
async-timeout==4.0.3
|
||||
attrs==23.2.0
|
||||
Automat==22.10.0
|
||||
beartype==0.16.4
|
||||
blessed==1.20.0
|
||||
buildtools==1.0.6
|
||||
# causal-conv1d==1.1.3.post1
|
||||
certifi==2023.11.17
|
||||
cffi==1.16.0
|
||||
chardet==5.2.0
|
||||
charset-normalizer==3.3.2
|
||||
clean-fid==0.1.35
|
||||
click==8.1.7
|
||||
# clip==1.0
|
||||
cmake==3.28.1
|
||||
colorama==0.4.6
|
||||
coloredlogs==15.0.1
|
||||
constantly==23.10.4
|
||||
contourpy==1.2.0
|
||||
crcmod==1.7
|
||||
cryptography==41.0.7
|
||||
cycler==0.12.1
|
||||
decorator==5.1.1
|
||||
decord==0.6.0
|
||||
diffusers==0.26.3
|
||||
docopt==0.6.2
|
||||
easydict==1.11
|
||||
einops==0.7.0
|
||||
exceptiongroup==1.2.0
|
||||
executing==2.0.1
|
||||
fairscale==0.4.13
|
||||
fastapi==0.109.0
|
||||
ffmpeg==1.4
|
||||
ffmpy==0.3.1
|
||||
filelock==3.13.1
|
||||
flatbuffers==24.3.25
|
||||
fonttools==4.47.2
|
||||
frozenlist==1.4.1
|
||||
fsspec==2023.12.2
|
||||
ftfy==6.1.3
|
||||
furl==2.1.3
|
||||
gpustat==1.1.1
|
||||
gradio==4.14.0
|
||||
gradio_client==0.8.0
|
||||
greenlet==3.0.3
|
||||
h11==0.14.0
|
||||
httpcore==1.0.2
|
||||
httpx==0.26.0
|
||||
huggingface-hub==0.20.2
|
||||
humanfriendly==10.0
|
||||
hyperlink==21.0.0
|
||||
idna==3.6
|
||||
imageio==2.33.1
|
||||
imageio-ffmpeg==0.4.9
|
||||
importlib-metadata==7.0.1
|
||||
importlib-resources==6.1.1
|
||||
incremental==22.10.0
|
||||
ipdb==0.13.13
|
||||
ipython==8.18.1
|
||||
jedi==0.19.1
|
||||
Jinja2==3.1.3
|
||||
jmespath==0.10.0
|
||||
joblib==1.3.2
|
||||
jsonschema==4.21.0
|
||||
jsonschema-specifications==2023.12.1
|
||||
kiwisolver==1.4.5
|
||||
kornia==0.7.1
|
||||
lazy_loader==0.3
|
||||
lightning-utilities==0.10.0
|
||||
lit==17.0.6
|
||||
lpips==0.1.4
|
||||
markdown-it-py==3.0.0
|
||||
MarkupSafe==2.1.3
|
||||
matplotlib==3.8.2
|
||||
matplotlib-inline==0.1.6
|
||||
mdurl==0.1.2
|
||||
# motion-vector-extractor==1.0.6
|
||||
mpmath==1.3.0
|
||||
multidict==6.0.4
|
||||
mypy-extensions==1.0.0
|
||||
networkx==3.2.1
|
||||
ninja==1.11.1.1
|
||||
numpy==1.26.3
|
||||
nvidia-cublas-cu11==11.10.3.66
|
||||
nvidia-cublas-cu12==12.1.3.1
|
||||
nvidia-cuda-cupti-cu11==11.7.101
|
||||
nvidia-cuda-cupti-cu12==12.1.105
|
||||
nvidia-cuda-nvrtc-cu11==11.7.99
|
||||
nvidia-cuda-nvrtc-cu12==12.1.105
|
||||
nvidia-cuda-runtime-cu11==11.7.99
|
||||
nvidia-cuda-runtime-cu12==12.1.105
|
||||
nvidia-cudnn-cu11==8.5.0.96
|
||||
nvidia-cudnn-cu12==8.9.2.26
|
||||
nvidia-cufft-cu11==10.9.0.58
|
||||
nvidia-cufft-cu12==11.0.2.54
|
||||
nvidia-curand-cu11==10.2.10.91
|
||||
nvidia-curand-cu12==10.3.2.106
|
||||
nvidia-cusolver-cu11==11.4.0.1
|
||||
nvidia-cusolver-cu12==11.4.5.107
|
||||
nvidia-cusparse-cu11==11.7.4.91
|
||||
nvidia-cusparse-cu12==12.1.0.106
|
||||
nvidia-ml-py==12.535.133
|
||||
nvidia-nccl-cu11==2.14.3
|
||||
nvidia-nccl-cu12==2.19.3
|
||||
nvidia-nvjitlink-cu12==12.3.101
|
||||
nvidia-nvtx-cu11==11.7.91
|
||||
nvidia-nvtx-cu12==12.1.105
|
||||
omegaconf==2.3.0
|
||||
onnxruntime==1.18.0
|
||||
open-clip-torch==2.24.0
|
||||
opencv-python==4.5.3.56
|
||||
opencv-python-headless==4.9.0.80
|
||||
orderedmultidict==1.0.1
|
||||
orjson==3.9.10
|
||||
oss2==2.18.4
|
||||
# packaging==23.2
|
||||
pandas==2.1.4
|
||||
parso==0.8.3
|
||||
pexpect==4.9.0
|
||||
pillow==10.2.0
|
||||
piq==0.8.0
|
||||
pkgconfig==1.5.5
|
||||
prompt-toolkit==3.0.43
|
||||
protobuf==4.25.2
|
||||
psutil==5.9.8
|
||||
ptflops==0.7.2.2
|
||||
ptyprocess==0.7.0
|
||||
pure-eval==0.2.2
|
||||
pycparser==2.21
|
||||
pycryptodome==3.20.0
|
||||
pydantic==2.5.3
|
||||
pydantic_core==2.14.6
|
||||
pydub==0.25.1
|
||||
Pygments==2.17.2
|
||||
pynvml==11.5.0
|
||||
pyparsing==3.1.1
|
||||
pyre-extensions==0.0.29
|
||||
python-dateutil==2.8.2
|
||||
python-multipart==0.0.6
|
||||
pytorch-lightning==2.1.3
|
||||
pytz==2023.3.post1
|
||||
PyYAML==6.0.1
|
||||
redo==2.0.4
|
||||
referencing==0.32.1
|
||||
regex==2023.12.25
|
||||
requests==2.31.0
|
||||
rich==13.7.0
|
||||
rotary-embedding-torch==0.5.3
|
||||
rpds-py==0.17.1
|
||||
ruff==0.2.0
|
||||
safetensors==0.4.1
|
||||
scikit-image==0.22.0
|
||||
scikit-learn==1.4.0
|
||||
scipy==1.11.4
|
||||
semantic-version==2.10.0
|
||||
sentencepiece==0.1.99
|
||||
shellingham==1.5.4
|
||||
simplejson==3.19.2
|
||||
six==1.16.0
|
||||
sk-video==1.1.10
|
||||
sniffio==1.3.0
|
||||
SQLAlchemy==2.0.27
|
||||
stack-data==0.6.3
|
||||
starlette==0.35.1
|
||||
sympy==1.12
|
||||
thop==0.1.1.post2209072238
|
||||
threadpoolctl==3.2.0
|
||||
tifffile==2023.12.9
|
||||
timm==0.9.12
|
||||
tokenizers==0.15.0
|
||||
tomli==2.0.1
|
||||
tomlkit==0.12.0
|
||||
toolz==0.12.0
|
||||
# torch==2.0.1+cu118
|
||||
# torchaudio==2.0.2+cu118
|
||||
torchdiffeq==0.2.3
|
||||
torchmetrics==1.3.0.post0
|
||||
torchsde==0.2.6
|
||||
# torchvision==0.15.2+cu118
|
||||
tqdm==4.66.1
|
||||
traitlets==5.14.1
|
||||
trampoline==0.1.2
|
||||
transformers==4.36.2
|
||||
triton==2.0.0
|
||||
Twisted==23.10.0
|
||||
typer==0.9.0
|
||||
typing-inspect==0.9.0
|
||||
typing_extensions==4.9.0
|
||||
tzdata==2023.4
|
||||
urllib3==2.1.0
|
||||
uvicorn==0.26.0
|
||||
wcwidth==0.2.13
|
||||
websockets==11.0.3
|
||||
xformers==0.0.20
|
||||
yarl==1.9.4
|
||||
zipp==3.17.0
|
||||
zope.interface==6.2
|
||||
onnxruntime-gpu==1.13.1
|
||||
@@ -0,0 +1,230 @@
|
||||
import os
|
||||
import yaml
|
||||
import json
|
||||
import copy
|
||||
import argparse
|
||||
|
||||
import utils.logging as logging
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
class Config(object):
|
||||
def __init__(self, load=True, cfg_dict=None, cfg_level=None):
|
||||
self._level = "cfg" + ("." + cfg_level if cfg_level is not None else "")
|
||||
if load:
|
||||
self.args = self._parse_args()
|
||||
logger.info("Loading config from {}.".format(self.args.cfg_file))
|
||||
self.need_initialization = True
|
||||
cfg_base = self._load_yaml(self.args) # self._initialize_cfg()
|
||||
cfg_dict = self._load_yaml(self.args)
|
||||
cfg_dict = self._merge_cfg_from_base(cfg_base, cfg_dict)
|
||||
cfg_dict = self._update_from_args(cfg_dict)
|
||||
self.cfg_dict = cfg_dict
|
||||
self._update_dict(cfg_dict)
|
||||
|
||||
def _parse_args(self):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Argparser for configuring [code base name to think of] codebase"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cfg",
|
||||
dest="cfg_file",
|
||||
help="Path to the configuration file",
|
||||
default='configs/Animate_X_infer.yaml'
|
||||
)
|
||||
parser.add_argument(
|
||||
"--init_method",
|
||||
help="Initialization method, includes TCP or shared file-system",
|
||||
default="tcp://localhost:9999",
|
||||
type=str,
|
||||
)
|
||||
parser.add_argument(
|
||||
'--debug',
|
||||
action='store_true',
|
||||
default=False,
|
||||
help='Into debug information'
|
||||
)
|
||||
parser.add_argument(
|
||||
"opts",
|
||||
help="other configurations",
|
||||
default=None,
|
||||
nargs=argparse.REMAINDER)
|
||||
return parser.parse_args()
|
||||
|
||||
def _path_join(self, path_list):
|
||||
path = ""
|
||||
for p in path_list:
|
||||
path+= p + '/'
|
||||
return path[:-1]
|
||||
|
||||
def _update_from_args(self, cfg_dict):
|
||||
args = self.args
|
||||
for var in vars(args):
|
||||
cfg_dict[var] = getattr(args, var)
|
||||
return cfg_dict
|
||||
|
||||
def _initialize_cfg(self):
|
||||
if self.need_initialization:
|
||||
self.need_initialization = False
|
||||
if os.path.exists('./configs/base.yaml'):
|
||||
with open("./configs/base.yaml", 'r') as f:
|
||||
cfg = yaml.load(f.read(), Loader=yaml.SafeLoader)
|
||||
else:
|
||||
with open(os.path.realpath(__file__).split('/')[-3] + "/configs/base.yaml", 'r') as f:
|
||||
cfg = yaml.load(f.read(), Loader=yaml.SafeLoader)
|
||||
return cfg
|
||||
|
||||
def _load_yaml(self, args, file_name=""):
|
||||
assert args.cfg_file is not None
|
||||
if not file_name == "": # reading from base file
|
||||
with open(file_name, 'r') as f:
|
||||
cfg = yaml.load(f.read(), Loader=yaml.SafeLoader)
|
||||
else:
|
||||
if os.getcwd().split("/")[-1] == args.cfg_file.split("/")[0]:
|
||||
args.cfg_file = args.cfg_file.replace(os.getcwd().split("/")[-1], "./")
|
||||
with open(args.cfg_file, 'r') as f:
|
||||
cfg = yaml.load(f.read(), Loader=yaml.SafeLoader)
|
||||
file_name = args.cfg_file
|
||||
|
||||
if "_BASE_RUN" not in cfg.keys() and "_BASE_MODEL" not in cfg.keys() and "_BASE" not in cfg.keys():
|
||||
# return cfg if the base file is being accessed
|
||||
cfg = self._merge_cfg_from_command_update(args, cfg)
|
||||
return cfg
|
||||
|
||||
if "_BASE" in cfg.keys():
|
||||
if cfg["_BASE"][1] == '.':
|
||||
prev_count = cfg["_BASE"].count('..')
|
||||
cfg_base_file = self._path_join(file_name.split('/')[:(-1-cfg["_BASE"].count('..'))] + cfg["_BASE"].split('/')[prev_count:])
|
||||
else:
|
||||
cfg_base_file = cfg["_BASE"].replace(
|
||||
"./",
|
||||
args.cfg_file.replace(args.cfg_file.split('/')[-1], "")
|
||||
)
|
||||
cfg_base = self._load_yaml(args, cfg_base_file)
|
||||
cfg = self._merge_cfg_from_base(cfg_base, cfg)
|
||||
else:
|
||||
if "_BASE_RUN" in cfg.keys():
|
||||
if cfg["_BASE_RUN"][1] == '.':
|
||||
prev_count = cfg["_BASE_RUN"].count('..')
|
||||
cfg_base_file = self._path_join(file_name.split('/')[:(-1-prev_count)] + cfg["_BASE_RUN"].split('/')[prev_count:])
|
||||
else:
|
||||
cfg_base_file = cfg["_BASE_RUN"].replace(
|
||||
"./",
|
||||
args.cfg_file.replace(args.cfg_file.split('/')[-1], "")
|
||||
)
|
||||
cfg_base = self._load_yaml(args, cfg_base_file)
|
||||
cfg = self._merge_cfg_from_base(cfg_base, cfg, preserve_base=True)
|
||||
if "_BASE_MODEL" in cfg.keys():
|
||||
if cfg["_BASE_MODEL"][1] == '.':
|
||||
prev_count = cfg["_BASE_MODEL"].count('..')
|
||||
cfg_base_file = self._path_join(file_name.split('/')[:(-1-cfg["_BASE_MODEL"].count('..'))] + cfg["_BASE_MODEL"].split('/')[prev_count:])
|
||||
else:
|
||||
cfg_base_file = cfg["_BASE_MODEL"].replace(
|
||||
"./",
|
||||
args.cfg_file.replace(args.cfg_file.split('/')[-1], "")
|
||||
)
|
||||
cfg_base = self._load_yaml(args, cfg_base_file)
|
||||
cfg = self._merge_cfg_from_base(cfg_base, cfg)
|
||||
cfg = self._merge_cfg_from_command(args, cfg)
|
||||
return cfg
|
||||
|
||||
def _merge_cfg_from_base(self, cfg_base, cfg_new, preserve_base=False):
|
||||
for k,v in cfg_new.items():
|
||||
if k in cfg_base.keys():
|
||||
if isinstance(v, dict):
|
||||
self._merge_cfg_from_base(cfg_base[k], v)
|
||||
else:
|
||||
cfg_base[k] = v
|
||||
else:
|
||||
if "BASE" not in k or preserve_base:
|
||||
cfg_base[k] = v
|
||||
return cfg_base
|
||||
|
||||
def _merge_cfg_from_command_update(self, args, cfg):
|
||||
if len(args.opts) == 0:
|
||||
return cfg
|
||||
|
||||
assert len(args.opts) % 2 == 0, 'Override list {} has odd length: {}.'.format(
|
||||
args.opts, len(args.opts)
|
||||
)
|
||||
keys = args.opts[0::2]
|
||||
vals = args.opts[1::2]
|
||||
|
||||
for key, val in zip(keys, vals):
|
||||
cfg[key] = val
|
||||
|
||||
return cfg
|
||||
|
||||
def _merge_cfg_from_command(self, args, cfg):
|
||||
assert len(args.opts) % 2 == 0, 'Override list {} has odd length: {}.'.format(
|
||||
args.opts, len(args.opts)
|
||||
)
|
||||
keys = args.opts[0::2]
|
||||
vals = args.opts[1::2]
|
||||
|
||||
# maximum supported depth 3
|
||||
for idx, key in enumerate(keys):
|
||||
key_split = key.split('.')
|
||||
assert len(key_split) <= 4, 'Key depth error. \nMaximum depth: 3\n Get depth: {}'.format(
|
||||
len(key_split)
|
||||
)
|
||||
assert key_split[0] in cfg.keys(), 'Non-existant key: {}.'.format(
|
||||
key_split[0]
|
||||
)
|
||||
if len(key_split) == 2:
|
||||
assert key_split[1] in cfg[key_split[0]].keys(), 'Non-existant key: {}.'.format(
|
||||
key
|
||||
)
|
||||
elif len(key_split) == 3:
|
||||
assert key_split[1] in cfg[key_split[0]].keys(), 'Non-existant key: {}.'.format(
|
||||
key
|
||||
)
|
||||
assert key_split[2] in cfg[key_split[0]][key_split[1]].keys(), 'Non-existant key: {}.'.format(
|
||||
key
|
||||
)
|
||||
elif len(key_split) == 4:
|
||||
assert key_split[1] in cfg[key_split[0]].keys(), 'Non-existant key: {}.'.format(
|
||||
key
|
||||
)
|
||||
assert key_split[2] in cfg[key_split[0]][key_split[1]].keys(), 'Non-existant key: {}.'.format(
|
||||
key
|
||||
)
|
||||
assert key_split[3] in cfg[key_split[0]][key_split[1]][key_split[2]].keys(), 'Non-existant key: {}.'.format(
|
||||
key
|
||||
)
|
||||
if len(key_split) == 1:
|
||||
cfg[key_split[0]] = vals[idx]
|
||||
elif len(key_split) == 2:
|
||||
cfg[key_split[0]][key_split[1]] = vals[idx]
|
||||
elif len(key_split) == 3:
|
||||
cfg[key_split[0]][key_split[1]][key_split[2]] = vals[idx]
|
||||
elif len(key_split) == 4:
|
||||
cfg[key_split[0]][key_split[1]][key_split[2]][key_split[3]] = vals[idx]
|
||||
return cfg
|
||||
|
||||
def _update_dict(self, cfg_dict):
|
||||
def recur(key, elem):
|
||||
if type(elem) is dict:
|
||||
return key, Config(load=False, cfg_dict=elem, cfg_level=key)
|
||||
else:
|
||||
if type(elem) is str and elem[1:3]=="e-":
|
||||
elem = float(elem)
|
||||
return key, elem
|
||||
dic = dict(recur(k, v) for k, v in cfg_dict.items())
|
||||
self.__dict__.update(dic)
|
||||
|
||||
def get_args(self):
|
||||
return self.args
|
||||
|
||||
def __repr__(self):
|
||||
return "{}\n".format(self.dump())
|
||||
|
||||
def dump(self):
|
||||
return json.dumps(self.cfg_dict, indent=2)
|
||||
|
||||
def deep_copy(self):
|
||||
return copy.deepcopy(self)
|
||||
|
||||
if __name__ == '__main__':
|
||||
# debug
|
||||
cfg = Config(load=True)
|
||||
print(cfg.DATA)
|
||||
@@ -0,0 +1,430 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torch.distributed as dist
|
||||
import functools
|
||||
import pickle
|
||||
import numpy as np
|
||||
from collections import OrderedDict
|
||||
from torch.autograd import Function
|
||||
|
||||
__all__ = ['is_dist_initialized',
|
||||
'get_world_size',
|
||||
'get_rank',
|
||||
'new_group',
|
||||
'destroy_process_group',
|
||||
'barrier',
|
||||
'broadcast',
|
||||
'all_reduce',
|
||||
'reduce',
|
||||
'gather',
|
||||
'all_gather',
|
||||
'reduce_dict',
|
||||
'get_global_gloo_group',
|
||||
'generalized_all_gather',
|
||||
'generalized_gather',
|
||||
'scatter',
|
||||
'reduce_scatter',
|
||||
'send',
|
||||
'recv',
|
||||
'isend',
|
||||
'irecv',
|
||||
'shared_random_seed',
|
||||
'diff_all_gather',
|
||||
'diff_all_reduce',
|
||||
'diff_scatter',
|
||||
'diff_copy',
|
||||
'spherical_kmeans',
|
||||
'sinkhorn']
|
||||
|
||||
#-------------------------------- Distributed operations --------------------------------#
|
||||
|
||||
def is_dist_initialized():
|
||||
return dist.is_available() and dist.is_initialized()
|
||||
|
||||
def get_world_size(group=None):
|
||||
return dist.get_world_size(group) if is_dist_initialized() else 1
|
||||
|
||||
def get_rank(group=None):
|
||||
return dist.get_rank(group) if is_dist_initialized() else 0
|
||||
|
||||
def new_group(ranks=None, **kwargs):
|
||||
if is_dist_initialized():
|
||||
return dist.new_group(ranks, **kwargs)
|
||||
return None
|
||||
|
||||
def destroy_process_group():
|
||||
if is_dist_initialized():
|
||||
dist.destroy_process_group()
|
||||
|
||||
def barrier(group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
dist.barrier(group, **kwargs)
|
||||
|
||||
def broadcast(tensor, src, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
return dist.broadcast(tensor, src, group, **kwargs)
|
||||
|
||||
def all_reduce(tensor, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
return dist.all_reduce(tensor, op, group, **kwargs)
|
||||
|
||||
def reduce(tensor, dst, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
return dist.reduce(tensor, dst, op, group, **kwargs)
|
||||
|
||||
def gather(tensor, dst=0, group=None, **kwargs):
|
||||
rank = get_rank() # global rank
|
||||
world_size = get_world_size(group)
|
||||
if world_size == 1:
|
||||
return [tensor]
|
||||
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)] if rank == dst else None
|
||||
dist.gather(tensor, tensor_list, dst, group, **kwargs)
|
||||
return tensor_list
|
||||
|
||||
def all_gather(tensor, uniform_size=True, group=None, **kwargs):
|
||||
world_size = get_world_size(group)
|
||||
if world_size == 1:
|
||||
return [tensor]
|
||||
assert tensor.is_contiguous(), 'ops.all_gather requires the tensor to be contiguous()'
|
||||
|
||||
if uniform_size:
|
||||
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||
dist.all_gather(tensor_list, tensor, group, **kwargs)
|
||||
return tensor_list
|
||||
else:
|
||||
# collect tensor shapes across GPUs
|
||||
shape = tuple(tensor.shape)
|
||||
shape_list = generalized_all_gather(shape, group)
|
||||
|
||||
# flatten the tensor
|
||||
tensor = tensor.reshape(-1)
|
||||
size = int(np.prod(shape))
|
||||
size_list = [int(np.prod(u)) for u in shape_list]
|
||||
max_size = max(size_list)
|
||||
|
||||
# pad to maximum size
|
||||
if size != max_size:
|
||||
padding = tensor.new_zeros(max_size - size)
|
||||
tensor = torch.cat([tensor, padding], dim=0)
|
||||
|
||||
# all_gather
|
||||
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||
dist.all_gather(tensor_list, tensor, group, **kwargs)
|
||||
|
||||
# reshape tensors
|
||||
tensor_list = [t[:n].view(s) for t, n, s in zip(
|
||||
tensor_list, size_list, shape_list)]
|
||||
return tensor_list
|
||||
|
||||
@torch.no_grad()
|
||||
def reduce_dict(input_dict, group=None, reduction='mean', **kwargs):
|
||||
assert reduction in ['mean', 'sum']
|
||||
world_size = get_world_size(group)
|
||||
if world_size == 1:
|
||||
return input_dict
|
||||
|
||||
# ensure that the orders of keys are consistent across processes
|
||||
if isinstance(input_dict, OrderedDict):
|
||||
keys = list(input_dict.keys)
|
||||
else:
|
||||
keys = sorted(input_dict.keys())
|
||||
vals = [input_dict[key] for key in keys]
|
||||
vals = torch.stack(vals, dim=0)
|
||||
dist.reduce(vals, dst=0, group=group, **kwargs)
|
||||
if dist.get_rank(group) == 0 and reduction == 'mean':
|
||||
vals /= world_size
|
||||
dist.broadcast(vals, src=0, group=group, **kwargs)
|
||||
reduced_dict = type(input_dict)([
|
||||
(key, val) for key, val in zip(keys, vals)])
|
||||
return reduced_dict
|
||||
|
||||
@functools.lru_cache()
|
||||
def get_global_gloo_group():
|
||||
backend = dist.get_backend()
|
||||
assert backend in ['gloo', 'nccl']
|
||||
if backend == 'nccl':
|
||||
return dist.new_group(backend='gloo')
|
||||
else:
|
||||
return dist.group.WORLD
|
||||
|
||||
def _serialize_to_tensor(data, group):
|
||||
backend = dist.get_backend(group)
|
||||
assert backend in ['gloo', 'nccl']
|
||||
device = torch.device('cpu' if backend == 'gloo' else 'cuda')
|
||||
|
||||
buffer = pickle.dumps(data)
|
||||
if len(buffer) > 1024 ** 3:
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.warning(
|
||||
'Rank {} trying to all-gather {:.2f} GB of data on device'
|
||||
'{}'.format(get_rank(), len(buffer) / (1024 ** 3), device))
|
||||
storage = torch.ByteStorage.from_buffer(buffer)
|
||||
tensor = torch.ByteTensor(storage).to(device=device)
|
||||
return tensor
|
||||
|
||||
def _pad_to_largest_tensor(tensor, group):
|
||||
world_size = dist.get_world_size(group=group)
|
||||
assert world_size >= 1, \
|
||||
'gather/all_gather must be called from ranks within' \
|
||||
'the give group!'
|
||||
local_size = torch.tensor(
|
||||
[tensor.numel()], dtype=torch.int64, device=tensor.device)
|
||||
size_list = [torch.zeros(
|
||||
[1], dtype=torch.int64, device=tensor.device)
|
||||
for _ in range(world_size)]
|
||||
|
||||
# gather tensors and compute the maximum size
|
||||
dist.all_gather(size_list, local_size, group=group)
|
||||
size_list = [int(size.item()) for size in size_list]
|
||||
max_size = max(size_list)
|
||||
|
||||
# pad tensors to the same size
|
||||
if local_size != max_size:
|
||||
padding = torch.zeros(
|
||||
(max_size - local_size, ),
|
||||
dtype=torch.uint8, device=tensor.device)
|
||||
tensor = torch.cat((tensor, padding), dim=0)
|
||||
return size_list, tensor
|
||||
|
||||
def generalized_all_gather(data, group=None):
|
||||
if get_world_size(group) == 1:
|
||||
return [data]
|
||||
if group is None:
|
||||
group = get_global_gloo_group()
|
||||
|
||||
tensor = _serialize_to_tensor(data, group)
|
||||
size_list, tensor = _pad_to_largest_tensor(tensor, group)
|
||||
max_size = max(size_list)
|
||||
|
||||
# receiving tensors from all ranks
|
||||
tensor_list = [torch.empty(
|
||||
(max_size, ), dtype=torch.uint8, device=tensor.device)
|
||||
for _ in size_list]
|
||||
dist.all_gather(tensor_list, tensor, group=group)
|
||||
|
||||
data_list = []
|
||||
for size, tensor in zip(size_list, tensor_list):
|
||||
buffer = tensor.cpu().numpy().tobytes()[:size]
|
||||
data_list.append(pickle.loads(buffer))
|
||||
return data_list
|
||||
|
||||
def generalized_gather(data, dst=0, group=None):
|
||||
world_size = get_world_size(group)
|
||||
if world_size == 1:
|
||||
return [data]
|
||||
if group is None:
|
||||
group = get_global_gloo_group()
|
||||
rank = dist.get_rank() # global rank
|
||||
|
||||
tensor = _serialize_to_tensor(data, group)
|
||||
size_list, tensor = _pad_to_largest_tensor(tensor, group)
|
||||
|
||||
# receiving tensors from all ranks to dst
|
||||
if rank == dst:
|
||||
max_size = max(size_list)
|
||||
tensor_list = [torch.empty(
|
||||
(max_size, ), dtype=torch.uint8, device=tensor.device)
|
||||
for _ in size_list]
|
||||
dist.gather(tensor, tensor_list, dst=dst, group=group)
|
||||
|
||||
data_list = []
|
||||
for size, tensor in zip(size_list, tensor_list):
|
||||
buffer = tensor.cpu().numpy().tobytes()[:size]
|
||||
data_list.append(pickle.loads(buffer))
|
||||
return data_list
|
||||
else:
|
||||
dist.gather(tensor, [], dst=dst, group=group)
|
||||
return []
|
||||
|
||||
def scatter(data, scatter_list=None, src=0, group=None, **kwargs):
|
||||
r"""NOTE: only supports CPU tensor communication.
|
||||
"""
|
||||
if get_world_size(group) > 1:
|
||||
return dist.scatter(data, scatter_list, src, group, **kwargs)
|
||||
|
||||
def reduce_scatter(output, input_list, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
return dist.reduce_scatter(output, input_list, op, group, **kwargs)
|
||||
|
||||
def send(tensor, dst, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
assert tensor.is_contiguous(), 'ops.send requires the tensor to be contiguous()'
|
||||
return dist.send(tensor, dst, group, **kwargs)
|
||||
|
||||
def recv(tensor, src=None, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
assert tensor.is_contiguous(), 'ops.recv requires the tensor to be contiguous()'
|
||||
return dist.recv(tensor, src, group, **kwargs)
|
||||
|
||||
def isend(tensor, dst, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
assert tensor.is_contiguous(), 'ops.isend requires the tensor to be contiguous()'
|
||||
return dist.isend(tensor, dst, group, **kwargs)
|
||||
|
||||
def irecv(tensor, src=None, group=None, **kwargs):
|
||||
if get_world_size(group) > 1:
|
||||
assert tensor.is_contiguous(), 'ops.irecv requires the tensor to be contiguous()'
|
||||
return dist.irecv(tensor, src, group, **kwargs)
|
||||
|
||||
def shared_random_seed(group=None):
|
||||
seed = np.random.randint(2 ** 31)
|
||||
all_seeds = generalized_all_gather(seed, group)
|
||||
return all_seeds[0]
|
||||
|
||||
#-------------------------------- Differentiable operations --------------------------------#
|
||||
|
||||
def _all_gather(x):
|
||||
if not (dist.is_available() and dist.is_initialized()) or dist.get_world_size() == 1:
|
||||
return x
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
tensors = [torch.empty_like(x) for _ in range(world_size)]
|
||||
tensors[rank] = x
|
||||
dist.all_gather(tensors, x)
|
||||
return torch.cat(tensors, dim=0).contiguous()
|
||||
|
||||
def _all_reduce(x):
|
||||
if not (dist.is_available() and dist.is_initialized()) or dist.get_world_size() == 1:
|
||||
return x
|
||||
dist.all_reduce(x)
|
||||
return x
|
||||
|
||||
def _split(x):
|
||||
if not (dist.is_available() and dist.is_initialized()) or dist.get_world_size() == 1:
|
||||
return x
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
return x.chunk(world_size, dim=0)[rank].contiguous()
|
||||
|
||||
class DiffAllGather(Function):
|
||||
r"""Differentiable all-gather.
|
||||
"""
|
||||
@staticmethod
|
||||
def symbolic(graph, input):
|
||||
return _all_gather(input)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
return _all_gather(input)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return _split(grad_output)
|
||||
|
||||
class DiffAllReduce(Function):
|
||||
r"""Differentiable all-reducd.
|
||||
"""
|
||||
@staticmethod
|
||||
def symbolic(graph, input):
|
||||
return _all_reduce(input)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
return _all_reduce(input)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return grad_output
|
||||
|
||||
class DiffScatter(Function):
|
||||
r"""Differentiable scatter.
|
||||
"""
|
||||
@staticmethod
|
||||
def symbolic(graph, input):
|
||||
return _split(input)
|
||||
|
||||
@staticmethod
|
||||
def symbolic(ctx, input):
|
||||
return _split(input)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return _all_gather(grad_output)
|
||||
|
||||
class DiffCopy(Function):
|
||||
r"""Differentiable copy that reduces all gradients during backward.
|
||||
"""
|
||||
@staticmethod
|
||||
def symbolic(graph, input):
|
||||
return input
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
return input
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return _all_reduce(grad_output)
|
||||
|
||||
diff_all_gather = DiffAllGather.apply
|
||||
diff_all_reduce = DiffAllReduce.apply
|
||||
diff_scatter = DiffScatter.apply
|
||||
diff_copy = DiffCopy.apply
|
||||
|
||||
#-------------------------------- Distributed algorithms --------------------------------#
|
||||
|
||||
@torch.no_grad()
|
||||
def spherical_kmeans(feats, num_clusters, num_iters=10):
|
||||
k, n, c = num_clusters, *feats.size()
|
||||
ones = feats.new_ones(n, dtype=torch.long)
|
||||
|
||||
# distributed settings
|
||||
rank = get_rank()
|
||||
world_size = get_world_size()
|
||||
|
||||
# init clusters
|
||||
rand_inds = torch.randperm(n)[:int(np.ceil(k / world_size))]
|
||||
clusters = torch.cat(all_gather(feats[rand_inds]), dim=0)[:k]
|
||||
|
||||
# variables
|
||||
new_clusters = feats.new_zeros(k, c)
|
||||
counts = feats.new_zeros(k, dtype=torch.long)
|
||||
|
||||
# iterative Expectation-Maximization
|
||||
for step in range(num_iters + 1):
|
||||
# Expectation step
|
||||
simmat = torch.mm(feats, clusters.t())
|
||||
scores, assigns = simmat.max(dim=1)
|
||||
if step == num_iters:
|
||||
break
|
||||
|
||||
# Maximization step
|
||||
new_clusters.zero_().scatter_add_(0, assigns.unsqueeze(1).repeat(1, c), feats)
|
||||
all_reduce(new_clusters)
|
||||
|
||||
counts.zero_()
|
||||
counts.index_add_(0, assigns, ones)
|
||||
all_reduce(counts)
|
||||
|
||||
mask = (counts > 0)
|
||||
clusters[mask] = new_clusters[mask] / counts[mask].view(-1, 1)
|
||||
clusters = F.normalize(clusters, p=2, dim=1)
|
||||
return clusters, assigns, scores
|
||||
|
||||
@torch.no_grad()
|
||||
def sinkhorn(Q, eps=0.5, num_iters=3):
|
||||
# normalize Q
|
||||
Q = torch.exp(Q / eps).t()
|
||||
sum_Q = Q.sum()
|
||||
all_reduce(sum_Q)
|
||||
Q /= sum_Q
|
||||
|
||||
# variables
|
||||
n, m = Q.size()
|
||||
u = Q.new_zeros(n)
|
||||
r = Q.new_ones(n) / n
|
||||
c = Q.new_ones(m) / (m * get_world_size())
|
||||
|
||||
# iterative update
|
||||
cur_sum = Q.sum(dim=1)
|
||||
all_reduce(cur_sum)
|
||||
for i in range(num_iters):
|
||||
u = cur_sum
|
||||
Q *= (r / u).unsqueeze(1)
|
||||
Q *= (c / Q.sum(dim=0)).unsqueeze(0)
|
||||
cur_sum = Q.sum(dim=1)
|
||||
all_reduce(cur_sum)
|
||||
return (Q / Q.sum(dim=0, keepdim=True)).t().float()
|
||||
@@ -0,0 +1,90 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
|
||||
|
||||
"""Logging."""
|
||||
|
||||
import builtins
|
||||
import decimal
|
||||
import functools
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import simplejson
|
||||
# from fvcore.common.file_io import PathManager
|
||||
|
||||
import utils.distributed as du
|
||||
|
||||
|
||||
def _suppress_print():
|
||||
"""
|
||||
Suppresses printing from the current process.
|
||||
"""
|
||||
|
||||
def print_pass(*objects, sep=" ", end="\n", file=sys.stdout, flush=False):
|
||||
pass
|
||||
|
||||
builtins.print = print_pass
|
||||
|
||||
|
||||
# @functools.lru_cache(maxsize=None)
|
||||
# def _cached_log_stream(filename):
|
||||
# return PathManager.open(filename, "a")
|
||||
|
||||
|
||||
def setup_logging(cfg, log_file):
|
||||
"""
|
||||
Sets up the logging for multiple processes. Only enable the logging for the
|
||||
master process, and suppress logging for the non-master processes.
|
||||
"""
|
||||
if du.is_master_proc():
|
||||
# Enable logging for the master process.
|
||||
logging.root.handlers = []
|
||||
else:
|
||||
# Suppress logging for non-master processes.
|
||||
_suppress_print()
|
||||
|
||||
logger = logging.getLogger()
|
||||
logger.setLevel(logging.INFO)
|
||||
logger.propagate = False
|
||||
plain_formatter = logging.Formatter(
|
||||
"[%(asctime)s][%(levelname)s] %(name)s: %(lineno)4d: %(message)s",
|
||||
datefmt="%m/%d %H:%M:%S",
|
||||
)
|
||||
|
||||
if du.is_master_proc():
|
||||
ch = logging.StreamHandler(stream=sys.stdout)
|
||||
ch.setLevel(logging.DEBUG)
|
||||
ch.setFormatter(plain_formatter)
|
||||
logger.addHandler(ch)
|
||||
|
||||
if log_file is not None and du.is_master_proc(du.get_world_size()):
|
||||
filename = os.path.join(cfg.OUTPUT_DIR, log_file)
|
||||
fh = logging.FileHandler(filename)
|
||||
fh.setLevel(logging.DEBUG)
|
||||
fh.setFormatter(plain_formatter)
|
||||
logger.addHandler(fh)
|
||||
|
||||
|
||||
def get_logger(name):
|
||||
"""
|
||||
Retrieve the logger with the specified name or, if name is None, return a
|
||||
logger which is the root logger of the hierarchy.
|
||||
Args:
|
||||
name (string): name of the logger.
|
||||
"""
|
||||
return logging.getLogger(name)
|
||||
|
||||
|
||||
def log_json_stats(stats):
|
||||
"""
|
||||
Logs json stats.
|
||||
Args:
|
||||
stats (dict): a dictionary of statistical information to log.
|
||||
"""
|
||||
stats = {
|
||||
k: decimal.Decimal("{:.6f}".format(v)) if isinstance(v, float) else v
|
||||
for k, v in stats.items()
|
||||
}
|
||||
json_stats = simplejson.dumps(stats, sort_keys=True, use_decimal=True)
|
||||
logger = get_logger(__name__)
|
||||
logger.info("{:s}".format(json_stats))
|
||||
@@ -0,0 +1,8 @@
|
||||
import socket
|
||||
from contextlib import closing
|
||||
|
||||
def find_free_port():
|
||||
with closing(socket.socket(socket.AF_INET, socket.SOCK_STREAM)) as s:
|
||||
s.bind(('', 0))
|
||||
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
return str(s.getsockname()[1])
|
||||
@@ -0,0 +1,152 @@
|
||||
import copy
|
||||
import inspect
|
||||
import warnings
|
||||
|
||||
|
||||
def build_from_config(cfg, registry, **kwargs):
|
||||
""" Default builder function.
|
||||
|
||||
Args:
|
||||
cfg (dict): A dict which contains parameters passes to target class or function.
|
||||
Must contains key 'type', indicates the target class or function name.
|
||||
registry (Registry): An registry to search target class or function.
|
||||
kwargs (dict, optional): Other params not in config dict.
|
||||
|
||||
Returns:
|
||||
Target class object or object returned by invoking function.
|
||||
|
||||
Raises:
|
||||
TypeError:
|
||||
KeyError:
|
||||
Exception:
|
||||
"""
|
||||
|
||||
|
||||
if not isinstance(cfg, dict):
|
||||
raise TypeError(f"config must be type dict, got {type(cfg)}")
|
||||
if "type" not in cfg:
|
||||
raise KeyError(f"config must contain key type, got {cfg}")
|
||||
if not isinstance(registry, Registry):
|
||||
raise TypeError(f"registry must be type Registry, got {type(registry)}")
|
||||
|
||||
cfg = copy.deepcopy(cfg)
|
||||
|
||||
req_type = cfg.pop("type")
|
||||
req_type_entry = req_type
|
||||
if isinstance(req_type, str):
|
||||
req_type_entry = registry.get(req_type)
|
||||
if req_type_entry is None:
|
||||
try:
|
||||
print(f"For Windows users, we explicitly import registry function {req_type} !!!")
|
||||
from animatex.diffusion.diffusion_ddim import DiffusionDDIM
|
||||
from animatex.diffusion.diffusion_ddim import DiffusionDDIMLong
|
||||
|
||||
from animatex.model.autoencoder import AutoencoderKL
|
||||
from animatex.model.clip_embedder import FrozenOpenCLIPTextVisualEmbedder
|
||||
from animatex.model.unet_animate_x import UNetSD_Animate_X
|
||||
|
||||
from animatex.inference_animate_x_entrance import inference_animate_x_entrance
|
||||
|
||||
req_type_entry = eval(req_type)
|
||||
except:
|
||||
raise KeyError(f"{req_type} not found in {registry.name} registry")
|
||||
|
||||
if kwargs is not None:
|
||||
cfg.update(kwargs)
|
||||
|
||||
if inspect.isclass(req_type_entry):
|
||||
try:
|
||||
return req_type_entry(**cfg)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to init class {req_type_entry}, with {e}")
|
||||
elif inspect.isfunction(req_type_entry):
|
||||
try:
|
||||
return req_type_entry(**cfg)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to invoke function {req_type_entry}, with {e}")
|
||||
else:
|
||||
raise TypeError(f"type must be str or class, got {type(req_type_entry)}")
|
||||
|
||||
|
||||
class Registry(object):
|
||||
""" A registry maps key to classes or functions.
|
||||
|
||||
Example:
|
||||
>>> MODELS = Registry('MODELS')
|
||||
>>> @MODELS.register_class()
|
||||
>>> class ResNet(object):
|
||||
>>> pass
|
||||
>>> resnet = MODELS.build(dict(type="ResNet"))
|
||||
>>>
|
||||
>>> import torchvision
|
||||
>>> @MODELS.register_function("InceptionV3")
|
||||
>>> def get_inception_v3(pretrained=False, progress=True):
|
||||
>>> return torchvision.models.inception_v3(pretrained=pretrained, progress=progress)
|
||||
>>> inception_v3 = MODELS.build(dict(type='InceptionV3', pretrained=True))
|
||||
|
||||
Args:
|
||||
name (str): Registry name.
|
||||
build_func (func, None): Instance construct function. Default is build_from_config.
|
||||
allow_types (tuple): Indicates how to construct the instance, by constructing class or invoking function.
|
||||
"""
|
||||
|
||||
def __init__(self, name, build_func=None, allow_types=("class", "function")):
|
||||
self.name = name
|
||||
self.allow_types = allow_types
|
||||
self.class_map = {}
|
||||
self.func_map = {}
|
||||
self.build_func = build_func or build_from_config
|
||||
|
||||
def get(self, req_type):
|
||||
return self.class_map.get(req_type) or self.func_map.get(req_type)
|
||||
|
||||
def build(self, *args, **kwargs):
|
||||
return self.build_func(*args, **kwargs, registry=self)
|
||||
|
||||
def register_class(self, name=None):
|
||||
def _register(cls):
|
||||
if not inspect.isclass(cls):
|
||||
raise TypeError(f"Module must be type class, got {type(cls)}")
|
||||
if "class" not in self.allow_types:
|
||||
raise TypeError(f"Register {self.name} only allows type {self.allow_types}, got class")
|
||||
module_name = name or cls.__name__
|
||||
if module_name in self.class_map:
|
||||
warnings.warn(f"Class {module_name} already registered by {self.class_map[module_name]}, "
|
||||
f"will be replaced by {cls}")
|
||||
self.class_map[module_name] = cls
|
||||
return cls
|
||||
|
||||
return _register
|
||||
|
||||
def register_function(self, name=None):
|
||||
def _register(func):
|
||||
if not inspect.isfunction(func):
|
||||
raise TypeError(f"Registry must be type function, got {type(func)}")
|
||||
if "function" not in self.allow_types:
|
||||
raise TypeError(f"Registry {self.name} only allows type {self.allow_types}, got function")
|
||||
func_name = name or func.__name__
|
||||
if func_name in self.class_map:
|
||||
warnings.warn(f"Function {func_name} already registered by {self.func_map[func_name]}, "
|
||||
f"will be replaced by {func}")
|
||||
self.func_map[func_name] = func
|
||||
return func
|
||||
|
||||
return _register
|
||||
|
||||
def _list(self):
|
||||
keys = sorted(list(self.class_map.keys()) + list(self.func_map.keys()))
|
||||
descriptions = []
|
||||
for key in keys:
|
||||
if key in self.class_map:
|
||||
descriptions.append(f"{key}: {self.class_map[key]}")
|
||||
else:
|
||||
descriptions.append(
|
||||
f"{key}: <function '{self.func_map[key].__module__}.{self.func_map[key].__name__}'>")
|
||||
return "\n".join(descriptions)
|
||||
|
||||
def __repr__(self):
|
||||
description = self._list()
|
||||
description = '\n'.join(['\t' + s for s in description.split('\n')])
|
||||
return f"{self.__class__.__name__} [{self.name}], \n" + description
|
||||
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
from .registry import Registry, build_from_config
|
||||
|
||||
def build_func(cfg, registry, **kwargs):
|
||||
"""
|
||||
Except for config, if passing a list of dataset config, then return the concat type of it
|
||||
"""
|
||||
|
||||
|
||||
return build_from_config(cfg, registry, **kwargs)
|
||||
|
||||
AUTO_ENCODER = Registry("AUTO_ENCODER", build_func=build_func)
|
||||
DATASETS = Registry("DATASETS", build_func=build_func)
|
||||
DIFFUSION = Registry("DIFFUSION", build_func=build_func)
|
||||
DISTRIBUTION = Registry("DISTRIBUTION", build_func=build_func)
|
||||
EMBEDDER = Registry("EMBEDDER", build_func=build_func)
|
||||
ENGINE = Registry("ENGINE", build_func=build_func)
|
||||
INFER_ENGINE = Registry("INFER_ENGINE", build_func=build_func)
|
||||
MODEL = Registry("MODEL", build_func=build_func)
|
||||
PRETRAIN = Registry("PRETRAIN", build_func=build_func)
|
||||
VISUAL = Registry("VISUAL", build_func=build_func)
|
||||
@@ -0,0 +1,11 @@
|
||||
import torch
|
||||
import random
|
||||
import numpy as np
|
||||
|
||||
|
||||
def setup_seed(seed):
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
np.random.seed(seed)
|
||||
random.seed(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
@@ -0,0 +1,353 @@
|
||||
import torch
|
||||
import torchvision.transforms.functional as F
|
||||
import random
|
||||
import math
|
||||
import numpy as np
|
||||
from PIL import Image, ImageFilter
|
||||
|
||||
__all__ = ['Compose', 'Resize', 'Rescale', 'CenterCrop', 'CenterCropV2', 'CenterCropWide', 'RandomCrop', 'RandomCropV2', 'RandomHFlip',\
|
||||
'GaussianBlur', 'ColorJitter', 'RandomGray', 'ToTensor', 'Normalize', "ResizeRandomCrop", "ExtractResizeRandomCrop", "ExtractResizeAssignCrop"]
|
||||
|
||||
|
||||
class Compose(object):
|
||||
|
||||
def __init__(self, transforms):
|
||||
self.transforms = transforms
|
||||
|
||||
def __getitem__(self, index):
|
||||
if isinstance(index, slice):
|
||||
return Compose(self.transforms[index])
|
||||
else:
|
||||
return self.transforms[index]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.transforms)
|
||||
|
||||
def __call__(self, rgb):
|
||||
for t in self.transforms:
|
||||
rgb = t(rgb)
|
||||
return rgb
|
||||
|
||||
class Resize(object):
|
||||
|
||||
def __init__(self, size=256):
|
||||
if isinstance(size, int):
|
||||
size = (size, size)
|
||||
self.size = size
|
||||
|
||||
def __call__(self, rgb):
|
||||
if isinstance(rgb, list):
|
||||
rgb = [u.resize(self.size, Image.BILINEAR) for u in rgb]
|
||||
else:
|
||||
rgb = rgb.resize(self.size, Image.BILINEAR)
|
||||
return rgb
|
||||
|
||||
class Rescale(object):
|
||||
|
||||
def __init__(self, size=256, interpolation=Image.BILINEAR):
|
||||
self.size = size
|
||||
self.interpolation = interpolation
|
||||
|
||||
def __call__(self, rgb):
|
||||
w, h = rgb[0].size
|
||||
scale = self.size / min(w, h)
|
||||
out_w, out_h = int(round(w * scale)), int(round(h * scale))
|
||||
rgb = [u.resize((out_w, out_h), self.interpolation) for u in rgb]
|
||||
return rgb
|
||||
|
||||
class CenterCrop(object):
|
||||
|
||||
def __init__(self, size=224):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, rgb):
|
||||
w, h = rgb[0].size
|
||||
assert min(w, h) >= self.size
|
||||
x1 = (w - self.size) // 2
|
||||
y1 = (h - self.size) // 2
|
||||
rgb = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in rgb]
|
||||
return rgb
|
||||
|
||||
class ResizeRandomCrop(object):
|
||||
|
||||
def __init__(self, size=256, size_short=292):
|
||||
self.size = size
|
||||
# self.min_area = min_area
|
||||
self.size_short = size_short
|
||||
|
||||
def __call__(self, rgb):
|
||||
|
||||
# consistent crop between rgb and m
|
||||
while min(rgb[0].size) >= 2 * self.size_short:
|
||||
rgb = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in rgb]
|
||||
scale = self.size_short / min(rgb[0].size)
|
||||
rgb = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in rgb]
|
||||
out_w = self.size
|
||||
out_h = self.size
|
||||
w, h = rgb[0].size # (518, 292)
|
||||
x1 = random.randint(0, w - out_w)
|
||||
y1 = random.randint(0, h - out_h)
|
||||
|
||||
rgb = [u.crop((x1, y1, x1 + out_w, y1 + out_h)) for u in rgb]
|
||||
# rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb]
|
||||
# # center crop
|
||||
# x1 = (img[0].width - self.size) // 2
|
||||
# y1 = (img[0].height - self.size) // 2
|
||||
# img = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in img]
|
||||
return rgb
|
||||
|
||||
|
||||
|
||||
class ExtractResizeRandomCrop(object):
|
||||
|
||||
def __init__(self, size=256, size_short=292):
|
||||
self.size = size
|
||||
# self.min_area = min_area
|
||||
self.size_short = size_short
|
||||
|
||||
def __call__(self, rgb):
|
||||
|
||||
# consistent crop between rgb and m
|
||||
while min(rgb[0].size) >= 2 * self.size_short:
|
||||
rgb = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in rgb]
|
||||
scale = self.size_short / min(rgb[0].size)
|
||||
rgb = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in rgb]
|
||||
out_w = self.size
|
||||
out_h = self.size
|
||||
w, h = rgb[0].size # (518, 292)
|
||||
x1 = random.randint(0, w - out_w)
|
||||
y1 = random.randint(0, h - out_h)
|
||||
|
||||
rgb = [u.crop((x1, y1, x1 + out_w, y1 + out_h)) for u in rgb]
|
||||
wh = [x1, y1, x1 + out_w, y1 + out_h]
|
||||
return rgb, wh
|
||||
|
||||
|
||||
class ExtractResizeAssignCrop(object):
|
||||
|
||||
def __init__(self, size=256, size_short=292):
|
||||
self.size = size
|
||||
# self.min_area = min_area
|
||||
self.size_short = size_short
|
||||
|
||||
def __call__(self, rgb, wh):
|
||||
|
||||
# consistent crop between rgb and m
|
||||
while min(rgb[0].size) >= 2 * self.size_short:
|
||||
rgb = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in rgb]
|
||||
scale = self.size_short / min(rgb[0].size)
|
||||
rgb = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in rgb]
|
||||
|
||||
rgb = [u.crop(wh) for u in rgb]
|
||||
rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb]
|
||||
return rgb
|
||||
|
||||
class CenterCropV2(object):
|
||||
def __init__(self, size):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, img):
|
||||
# fast resize
|
||||
while min(img[0].size) >= 2 * self.size:
|
||||
img = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in img]
|
||||
scale = self.size / min(img[0].size)
|
||||
img = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in img]
|
||||
|
||||
# center crop
|
||||
x1 = (img[0].width - self.size) // 2
|
||||
y1 = (img[0].height - self.size) // 2
|
||||
img = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in img]
|
||||
return img
|
||||
|
||||
|
||||
class CenterCropWide(object):
|
||||
def __init__(self, size, interpolation=Image.BOX):
|
||||
self.size = size
|
||||
self.interpolation = interpolation
|
||||
|
||||
def __call__(self, img):
|
||||
if isinstance(img, list):
|
||||
scale = min(img[0].size[0]/self.size[0], img[0].size[1]/self.size[1])
|
||||
img = [u.resize((round(u.width // scale), round(u.height // scale)), resample=self.interpolation) for u in img]
|
||||
|
||||
# center crop
|
||||
x1 = (img[0].width - self.size[0]) // 2
|
||||
y1 = (img[0].height - self.size[1]) // 2
|
||||
img = [u.crop((x1, y1, x1 + self.size[0], y1 + self.size[1])) for u in img]
|
||||
return img
|
||||
else:
|
||||
scale = min(img.size[0]/self.size[0], img.size[1]/self.size[1])
|
||||
img = img.resize((round(img.width // scale), round(img.height // scale)), resample=self.interpolation)
|
||||
x1 = (img.width - self.size[0]) // 2
|
||||
y1 = (img.height - self.size[1]) // 2
|
||||
img = img.crop((x1, y1, x1 + self.size[0], y1 + self.size[1]))
|
||||
return img
|
||||
|
||||
|
||||
|
||||
class RandomCrop(object):
|
||||
|
||||
def __init__(self, size=224, min_area=0.4):
|
||||
self.size = size
|
||||
self.min_area = min_area
|
||||
|
||||
def __call__(self, rgb):
|
||||
|
||||
# consistent crop between rgb and m
|
||||
w, h = rgb[0].size
|
||||
area = w * h
|
||||
out_w, out_h = float('inf'), float('inf')
|
||||
while out_w > w or out_h > h:
|
||||
target_area = random.uniform(self.min_area, 1.0) * area
|
||||
aspect_ratio = random.uniform(3. / 4., 4. / 3.)
|
||||
out_w = int(round(math.sqrt(target_area * aspect_ratio)))
|
||||
out_h = int(round(math.sqrt(target_area / aspect_ratio)))
|
||||
x1 = random.randint(0, w - out_w)
|
||||
y1 = random.randint(0, h - out_h)
|
||||
|
||||
rgb = [u.crop((x1, y1, x1 + out_w, y1 + out_h)) for u in rgb]
|
||||
rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb]
|
||||
|
||||
return rgb
|
||||
|
||||
class RandomCropV2(object):
|
||||
|
||||
def __init__(self, size=224, min_area=0.4, ratio=(3. / 4., 4. / 3.)):
|
||||
if isinstance(size, (tuple, list)):
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
self.min_area = min_area
|
||||
self.ratio = ratio
|
||||
|
||||
def _get_params(self, img):
|
||||
width, height = img.size
|
||||
area = height * width
|
||||
|
||||
for _ in range(10):
|
||||
target_area = random.uniform(self.min_area, 1.0) * area
|
||||
log_ratio = (math.log(self.ratio[0]), math.log(self.ratio[1]))
|
||||
aspect_ratio = math.exp(random.uniform(*log_ratio))
|
||||
|
||||
w = int(round(math.sqrt(target_area * aspect_ratio)))
|
||||
h = int(round(math.sqrt(target_area / aspect_ratio)))
|
||||
|
||||
if 0 < w <= width and 0 < h <= height:
|
||||
i = random.randint(0, height - h)
|
||||
j = random.randint(0, width - w)
|
||||
return i, j, h, w
|
||||
|
||||
# Fallback to central crop
|
||||
in_ratio = float(width) / float(height)
|
||||
if (in_ratio < min(self.ratio)):
|
||||
w = width
|
||||
h = int(round(w / min(self.ratio)))
|
||||
elif (in_ratio > max(self.ratio)):
|
||||
h = height
|
||||
w = int(round(h * max(self.ratio)))
|
||||
else: # whole image
|
||||
w = width
|
||||
h = height
|
||||
i = (height - h) // 2
|
||||
j = (width - w) // 2
|
||||
return i, j, h, w
|
||||
|
||||
def __call__(self, rgb):
|
||||
i, j, h, w = self._get_params(rgb[0])
|
||||
rgb = [F.resized_crop(u, i, j, h, w, self.size) for u in rgb]
|
||||
return rgb
|
||||
|
||||
class RandomHFlip(object):
|
||||
|
||||
def __init__(self, p=0.5):
|
||||
self.p = p
|
||||
|
||||
def __call__(self, rgb):
|
||||
if random.random() < self.p:
|
||||
rgb = [u.transpose(Image.FLIP_LEFT_RIGHT) for u in rgb]
|
||||
return rgb
|
||||
|
||||
class GaussianBlur(object):
|
||||
|
||||
def __init__(self, sigmas=[0.1, 2.0], p=0.5):
|
||||
self.sigmas = sigmas
|
||||
self.p = p
|
||||
|
||||
def __call__(self, rgb):
|
||||
if random.random() < self.p:
|
||||
sigma = random.uniform(*self.sigmas)
|
||||
rgb = [u.filter(ImageFilter.GaussianBlur(radius=sigma)) for u in rgb]
|
||||
return rgb
|
||||
|
||||
class ColorJitter(object):
|
||||
|
||||
def __init__(self, brightness=0.4, contrast=0.4, saturation=0.4, hue=0.1, p=0.5):
|
||||
self.brightness = brightness
|
||||
self.contrast = contrast
|
||||
self.saturation = saturation
|
||||
self.hue = hue
|
||||
self.p = p
|
||||
|
||||
def __call__(self, rgb):
|
||||
if random.random() < self.p:
|
||||
brightness, contrast, saturation, hue = self._random_params()
|
||||
transforms = [
|
||||
lambda f: F.adjust_brightness(f, brightness),
|
||||
lambda f: F.adjust_contrast(f, contrast),
|
||||
lambda f: F.adjust_saturation(f, saturation),
|
||||
lambda f: F.adjust_hue(f, hue)]
|
||||
random.shuffle(transforms)
|
||||
for t in transforms:
|
||||
rgb = [t(u) for u in rgb]
|
||||
|
||||
return rgb
|
||||
|
||||
def _random_params(self):
|
||||
brightness = random.uniform(
|
||||
max(0, 1 - self.brightness), 1 + self.brightness)
|
||||
contrast = random.uniform(
|
||||
max(0, 1 - self.contrast), 1 + self.contrast)
|
||||
saturation = random.uniform(
|
||||
max(0, 1 - self.saturation), 1 + self.saturation)
|
||||
hue = random.uniform(-self.hue, self.hue)
|
||||
return brightness, contrast, saturation, hue
|
||||
|
||||
class RandomGray(object):
|
||||
|
||||
def __init__(self, p=0.2):
|
||||
self.p = p
|
||||
|
||||
def __call__(self, rgb):
|
||||
if random.random() < self.p:
|
||||
rgb = [u.convert('L').convert('RGB') for u in rgb]
|
||||
return rgb
|
||||
|
||||
class ToTensor(object):
|
||||
|
||||
def __call__(self, rgb):
|
||||
if isinstance(rgb, list):
|
||||
rgb = torch.stack([F.to_tensor(u) for u in rgb], dim=0)
|
||||
else:
|
||||
rgb = F.to_tensor(rgb)
|
||||
|
||||
return rgb
|
||||
|
||||
class Normalize(object):
|
||||
|
||||
def __init__(self, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
def __call__(self, rgb):
|
||||
rgb = rgb.clone()
|
||||
rgb.clamp_(0, 1)
|
||||
if not isinstance(self.mean, torch.Tensor):
|
||||
self.mean = rgb.new_tensor(self.mean).view(-1)
|
||||
if not isinstance(self.std, torch.Tensor):
|
||||
self.std = rgb.new_tensor(self.std).view(-1)
|
||||
if rgb.dim() == 4:
|
||||
rgb.sub_(self.mean.view(1, -1, 1, 1)).div_(self.std.view(1, -1, 1, 1))
|
||||
elif rgb.dim() == 3:
|
||||
rgb.sub_(self.mean.view(-1, 1, 1)).div_(self.std.view(-1, 1, 1))
|
||||
return rgb
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
import os
|
||||
import os.path as osp
|
||||
import sys
|
||||
import cv2
|
||||
import glob
|
||||
import math
|
||||
import torch
|
||||
import gzip
|
||||
import copy
|
||||
import time
|
||||
import json
|
||||
import pickle
|
||||
import base64
|
||||
import imageio
|
||||
import hashlib
|
||||
import requests
|
||||
import binascii
|
||||
import zipfile
|
||||
# import skvideo.io
|
||||
import numpy as np
|
||||
from io import BytesIO
|
||||
import urllib.request
|
||||
import torch.nn.functional as F
|
||||
import torchvision.utils as tvutils
|
||||
from multiprocessing.pool import ThreadPool as Pool
|
||||
from einops import rearrange
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
|
||||
@torch.no_grad()
|
||||
def save_video_multiple_conditions_not_gif_horizontal_1col(local_path, video_tensor, model_kwargs, source_imgs,
|
||||
mean=[0.5,0.5,0.5], std=[0.5,0.5,0.5], nrow=8, retry=5, save_fps=8):
|
||||
mean=torch.tensor(mean,device=video_tensor.device).view(1,-1,1,1,1)#ncfhw
|
||||
std=torch.tensor(std,device=video_tensor.device).view(1,-1,1,1,1)#ncfhw
|
||||
video_tensor = video_tensor.mul_(std).add_(mean) #### unnormalize back to [0,1]
|
||||
video_tensor.clamp_(0, 1)
|
||||
|
||||
b, c, n, h, w = video_tensor.shape
|
||||
source_imgs = F.adaptive_avg_pool3d(source_imgs, (n, h, w))
|
||||
source_imgs = source_imgs.cpu()
|
||||
|
||||
model_kwargs_channel3 = {}
|
||||
for key, conditions in model_kwargs[0].items():
|
||||
|
||||
|
||||
if conditions.size(1) == 1:
|
||||
conditions = torch.cat([conditions, conditions, conditions], dim=1)
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
if conditions.size(1) == 2:
|
||||
conditions = torch.cat([conditions, conditions[:,:1,]], dim=1)
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
elif conditions.size(1) == 3:
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
elif conditions.size(1) == 4: # means it is a mask.
|
||||
color = ((conditions[:, 0:3] + 1.)/2.) # .astype(np.float32)
|
||||
alpha = conditions[:, 3:4] # .astype(np.float32)
|
||||
conditions = color * alpha + 1.0 * (1.0 - alpha)
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
model_kwargs_channel3[key] = conditions.cpu() if conditions.is_cuda else conditions
|
||||
|
||||
# filename = rand_name(suffix='.gif')
|
||||
for _ in [None] * retry:
|
||||
try:
|
||||
vid_gif = rearrange(video_tensor, '(i j) c f h w -> c f (i h) (j w)', i = nrow)
|
||||
|
||||
# cons_list = [rearrange(con, '(i j) c f h w -> c f (i h) (j w)', i = nrow) for _, con in model_kwargs_channel3.items()]
|
||||
# vid_gif = torch.cat(cons_list + [vid_gif,], dim=3)
|
||||
|
||||
vid_gif = vid_gif.permute(1,2,3,0)
|
||||
|
||||
images = vid_gif * 255.0
|
||||
images = [(img.numpy()).astype('uint8') for img in images]
|
||||
if len(images) == 1:
|
||||
|
||||
local_path = local_path.replace('.mp4', '.png')
|
||||
cv2.imwrite(local_path, images[0][:,:,::-1], [int(cv2.IMWRITE_JPEG_QUALITY), 100])
|
||||
# local_path
|
||||
# bucket.put_object_from_file(oss_key, local_path)
|
||||
else:
|
||||
|
||||
outputs = []
|
||||
for image_name in images:
|
||||
x = Image.fromarray(image_name)
|
||||
outputs.append(x)
|
||||
from pathlib import Path
|
||||
save_fmt = Path(local_path).suffix
|
||||
|
||||
if save_fmt == ".mp4":
|
||||
with imageio.get_writer(local_path, fps=save_fps) as writer:
|
||||
for img in outputs:
|
||||
img_array = np.array(img) # Convert PIL Image to numpy array
|
||||
writer.append_data(img_array)
|
||||
|
||||
elif save_fmt == ".gif":
|
||||
outputs[0].save(
|
||||
fp=local_path,
|
||||
format="GIF",
|
||||
append_images=outputs[1:],
|
||||
save_all=True,
|
||||
duration=(1 / save_fps * 1000),
|
||||
loop=0,
|
||||
)
|
||||
else:
|
||||
raise ValueError("Unsupported file type. Use .mp4 or .gif.")
|
||||
|
||||
# fourcc = cv2.VideoWriter_fourcc(*'mp4v')
|
||||
# fps = save_fps
|
||||
# image = images[0]
|
||||
# media_writer = cv2.VideoWriter(local_path, fourcc, fps, (image.shape[1],image.shape[0]))
|
||||
# for image_name in images:
|
||||
# im = image_name[:,:,::-1]
|
||||
# media_writer.write(im)
|
||||
# media_writer.release()
|
||||
|
||||
|
||||
exception = None
|
||||
break
|
||||
except Exception as e:
|
||||
exception = e
|
||||
continue
|
||||
if exception is not None:
|
||||
print('save video to {} failed, error: {}'.format(local_path, exception), flush=True)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def save_video_multiple_conditions_not_gif_horizontal_3col(local_path, video_tensor, model_kwargs, source_imgs,
|
||||
mean=[0.5,0.5,0.5], std=[0.5,0.5,0.5], nrow=8, retry=5, save_fps=8):
|
||||
mean=torch.tensor(mean,device=video_tensor.device).view(1,-1,1,1,1)#ncfhw
|
||||
std=torch.tensor(std,device=video_tensor.device).view(1,-1,1,1,1)#ncfhw
|
||||
video_tensor = video_tensor.mul_(std).add_(mean) #### unnormalize back to [0,1]
|
||||
video_tensor.clamp_(0, 1)
|
||||
|
||||
b, c, n, h, w = video_tensor.shape
|
||||
source_imgs = F.adaptive_avg_pool3d(source_imgs, (n, h, w))
|
||||
source_imgs = source_imgs.cpu()
|
||||
|
||||
model_kwargs_channel3 = {}
|
||||
for key, conditions in model_kwargs[0].items():
|
||||
|
||||
|
||||
if conditions.size(1) == 1:
|
||||
conditions = torch.cat([conditions, conditions, conditions], dim=1)
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
if conditions.size(1) == 2:
|
||||
conditions = torch.cat([conditions, conditions[:,:1,]], dim=1)
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
elif conditions.size(1) == 3:
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
elif conditions.size(1) == 4: # means it is a mask.
|
||||
color = ((conditions[:, 0:3] + 1.)/2.) # .astype(np.float32)
|
||||
alpha = conditions[:, 3:4] # .astype(np.float32)
|
||||
conditions = color * alpha + 1.0 * (1.0 - alpha)
|
||||
conditions = F.adaptive_avg_pool3d(conditions, (n, h, w))
|
||||
model_kwargs_channel3[key] = conditions.cpu() if conditions.is_cuda else conditions
|
||||
|
||||
# filename = rand_name(suffix='.gif')
|
||||
for _ in [None] * retry:
|
||||
try:
|
||||
vid_gif = rearrange(video_tensor, '(i j) c f h w -> c f (i h) (j w)', i = nrow)
|
||||
|
||||
cons_list = [rearrange(con, '(i j) c f h w -> c f (i h) (j w)', i = nrow) for _, con in model_kwargs_channel3.items()]
|
||||
vid_gif = torch.cat(cons_list + [vid_gif,], dim=3)
|
||||
|
||||
vid_gif = vid_gif.permute(1,2,3,0)
|
||||
|
||||
images = vid_gif * 255.0
|
||||
images = [(img.numpy()).astype('uint8') for img in images]
|
||||
if len(images) == 1:
|
||||
|
||||
local_path = local_path.replace('.mp4', '.png')
|
||||
cv2.imwrite(local_path, images[0][:,:,::-1], [int(cv2.IMWRITE_JPEG_QUALITY), 100])
|
||||
# local_path
|
||||
# bucket.put_object_from_file(oss_key, local_path)
|
||||
else:
|
||||
|
||||
outputs = []
|
||||
for image_name in images:
|
||||
x = Image.fromarray(image_name)
|
||||
outputs.append(x)
|
||||
from pathlib import Path
|
||||
save_fmt = Path(local_path).suffix
|
||||
|
||||
if save_fmt == ".mp4":
|
||||
with imageio.get_writer(local_path, fps=save_fps) as writer:
|
||||
for img in outputs:
|
||||
img_array = np.array(img) # Convert PIL Image to numpy array
|
||||
writer.append_data(img_array)
|
||||
|
||||
elif save_fmt == ".gif":
|
||||
outputs[0].save(
|
||||
fp=local_path,
|
||||
format="GIF",
|
||||
append_images=outputs[1:],
|
||||
save_all=True,
|
||||
duration=(1 / save_fps * 1000),
|
||||
loop=0,
|
||||
)
|
||||
else:
|
||||
raise ValueError("Unsupported file type. Use .mp4 or .gif.")
|
||||
|
||||
# fourcc = cv2.VideoWriter_fourcc(*'mp4v')
|
||||
# fps = save_fps
|
||||
# image = images[0]
|
||||
# media_writer = cv2.VideoWriter(local_path, fourcc, fps, (image.shape[1],image.shape[0]))
|
||||
# for image_name in images:
|
||||
# im = image_name[:,:,::-1]
|
||||
# media_writer.write(im)
|
||||
# media_writer.release()
|
||||
|
||||
|
||||
exception = None
|
||||
break
|
||||
except Exception as e:
|
||||
exception = e
|
||||
continue
|
||||
if exception is not None:
|
||||
print('save video to {} failed, error: {}'.format(local_path, exception), flush=True)
|
||||