Add Animate-X (see details)

https://lucaria-academy.github.io/Animate-X/

https://github.com/antgroup/animate-x

@ fdc80909f911d8e487cb4e6847f2e4c6501b62af
This commit is contained in:
Brandon Thomas
2025-02-03 18:22:05 -05:00
parent 5ab5b1c036
commit 046ab789c0
59 changed files with 9720 additions and 0 deletions
+18
View File
@@ -0,0 +1,18 @@
*.pkl
*.pt
*.mov
*.pth
*.mov
*.npz
*.npy
*.boj
*.onnx
*.tar
*.bin
cache*
.DS_Store
*DS_Store
outputs/
**/__pycache__
***/__pycache__
*/__pycache__
+201
View File
@@ -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.
+181
View File
@@ -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 &nbsp; | &nbsp; </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>
## &#x1F4CC; 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> -->
## &#x1F304; 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>
## &#x1F680; Installation
Install with `conda`:
```bash
conda env create -f environment.yaml
conda activate animate-x
```
## &#x1F680; 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
```
## &#x1F4A1; 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`.
**&#10004; 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`)
## &#x1F4E7; 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.
## &#x2696; License
This repository is released under the Apache-2.0 license as found in the [LICENSE](LICENSE) file.
## &#x1F4DA; 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}
}
```
+1
View File
@@ -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
File diff suppressed because it is too large Load Diff
@@ -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
Binary file not shown.

After

Width:  |  Height:  |  Size: 3.6 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 3.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 88 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.0 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 139 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 951 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 575 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 405 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.6 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 191 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 294 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+127
View File
@@ -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
+360
View File
@@ -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
+335
View File
@@ -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
+48
View File
@@ -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
+236
View File
@@ -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
+16
View File
@@ -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)
+324
View File
@@ -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)
+208
View File
@@ -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
+230
View File
@@ -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)
+430
View File
@@ -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()
+90
View File
@@ -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))
+8
View File
@@ -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])
+152
View File
@@ -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)
+11
View File
@@ -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
+353
View File
@@ -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
+216
View File
@@ -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)