init
@@ -0,0 +1,29 @@
|
||||
wandb/
|
||||
debugs/
|
||||
outputs/
|
||||
samples/
|
||||
__pycache__/
|
||||
ossutil_output/
|
||||
.ossutil_checkpoint/
|
||||
.idea/
|
||||
.DS_Store/
|
||||
|
||||
scripts/
|
||||
!scripts/animate.py
|
||||
|
||||
*.ipynb
|
||||
*.safetensors
|
||||
*.ckpt
|
||||
.idea
|
||||
*.csv
|
||||
|
||||
outputs/*
|
||||
!models/StableDiffusion/
|
||||
models/StableDiffusion/*
|
||||
!models/StableDiffusion/*.txt
|
||||
!models/Motion_Module/
|
||||
!models/Motion_Module/*.txt
|
||||
!models/DreamBooth_LoRA/
|
||||
!models/DreamBooth_LoRA/*.txt
|
||||
!models/MotionLoRA/
|
||||
!models/MotionLoRA/*.txt
|
||||
@@ -0,0 +1,201 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
@@ -0,0 +1,239 @@
|
||||
# CameraCtrl
|
||||
|
||||
This repository is the official implementation of [CameraCtrl](http://arxiv.org/abs/2404.02101).
|
||||
|
||||
> **CameraCtrl: Enabling Camera Control for Text-to-Video Generation** <br>
|
||||
> [Hao He](https://hehao13.github.io), [Yinghao Xu](https://justimyhxu.github.io), [Yuwei Guo](https://guoyww.github.io), [Gordon Wetzstein](https://web.stanford.edu/~gordonwz/), [Bo Dai](http://daibo.info), [Hongsheng Li](https://www.ee.cuhk.edu.hk/~hsli/), [Ceyuan Yang](https://ceyuan.me)<br>
|
||||
|
||||
## [[Paper](http://arxiv.org/abs/2404.02101)] [[Project Page]( https://hehao13.github.io/projects-CameraCtrl/)] [[Weights](https://huggingface.co/hehao13/CameraCtrl/tree/main)]
|
||||
|
||||
## Todo List
|
||||
- [x] Release inference code.
|
||||
- [x] Release pretrained models on [AnimateDiffV3](https://github.com/guoyww/AnimateDiff).
|
||||
- [x] Release training code.
|
||||
- [ ] Release Gradio Demo.
|
||||
- [ ] Release pretrained models on SVD.
|
||||
|
||||
## Configurations
|
||||
### Environment
|
||||
* 64-bit Python 3.10 and PyTorch 1.13.0 or higher.
|
||||
* CUDA 11.7
|
||||
* Users can use the following commands to install the packages
|
||||
```bash
|
||||
conda env create -f environment.yaml
|
||||
conda activate cameractrl
|
||||
```
|
||||
|
||||
### Dataset
|
||||
- Download the camera trajectories and videos from [RealEstate10K](https://google.github.io/realestate10k/download.html).
|
||||
- Run `tools/gather_realestate.py` to get all the clips for each video.
|
||||
- Run `tools/get_realestate_clips.py` to get the video clips from the original videos.
|
||||
- Using [LAVIS](https://github.com/salesforce/LAVIS) or other methods to generate a caption for each video clip. We provide our extracted captions in [Google Drive](https://drive.google.com/file/d/1nytBYjTa0bJ-8AMJWVCtKT2XwkJR3Jra/view?usp=share_link) and [Google Drive](https://drive.google.com/file/d/1AGEJYbfip0jcp-ymgU9uCjUHzqETivYP/view?usp=share_link).
|
||||
- Run `tools/generate_realestate_json.py` to generate the json files for training and test, you can construct the validation json file by randomly sampling some item from the training json file.
|
||||
- After the above steps, you can get the dataset folder like this
|
||||
```angular2html
|
||||
- RealEstate10k
|
||||
- annotations
|
||||
- test.json
|
||||
- train.json
|
||||
- validation.json
|
||||
- pose_files
|
||||
- 0000cc6d8b108390.txt
|
||||
- 00028da87cc5a4c4.txt
|
||||
- 0002b126b0a8a685.txt
|
||||
- 0003a9bce989e532.txt
|
||||
- 000465ebe46a98d2.txt
|
||||
- ...
|
||||
- video_clips
|
||||
- 00ccbtp2aSQ
|
||||
- 00rMZpGSeOI
|
||||
- 01bTY_glskw
|
||||
- 01PJ3skCZPo
|
||||
- 01uaDoluhzo
|
||||
- ...
|
||||
```
|
||||
|
||||
## Inferences
|
||||
|
||||
### Prepare Models
|
||||
- Download Stable Diffusion V1.5 (SD1.5) from [HuggingFace](https://huggingface.co/runwayml/stable-diffusion-v1-5/tree/main).
|
||||
- Download the checkpoints of AnimateDiffV3 (ADV3) adaptor and motion module from [AnimateDiff](https://github.com/guoyww/AnimateDiff).
|
||||
- Download the pretrained camera control model from [HuggingFace](https://huggingface.co/hehao13/CameraCtrl/blob/main/CameraCtrl.ckpt).
|
||||
- Run `tools/merge_lora2unet.py` to merge the ADV3 adaptor weights into SD1.5 unet and save results to new subfolder (like, `unet_webvidlora_v3`) under the SD1.5 folder.
|
||||
- (Optional) Download the pretrained image LoRA model on RealEstate10K dataset from [HuggingFace](https://huggingface.co/hehao13/CameraCtrl/blob/main/RealEstate10K_LoRA.ckpt) to sample videos on indoor and outdoor estates.
|
||||
- (Optional) Download the personalized base model, like [Realistic Vision](https://civitai.com/models/4201?modelVersionId=130072) from [CivitAI](https://civitai.com).
|
||||
|
||||
### Prepare camera trajectory & prompts
|
||||
- Adopt `tools/select_realestate_clips.py` to prepare trajectory txt file, some example trajectories and corresponding reference videos are in `assets/pose_files` and `assets/reference_videos`, respectively. The generated trajectories can be visualized with `tools/visualize_trajectory.py`.
|
||||
- Prepare the prompts (negative prompts, specific seeds), one example is `assets/cameractrl_prompts.json`.
|
||||
|
||||
### Inference
|
||||
- Run `inference.py` to sample videos
|
||||
```shell
|
||||
python -m torch.distributed.launch --nproc_per_node=8 --master_port=25000 inference.py \
|
||||
--out_root ${OUTPUT_PATH} \
|
||||
--ori_model_path ${SD1.5_PATH} \
|
||||
--unet_subfolder ${SUBFOUDER_NAME} \
|
||||
--motion_module_ckpt ${ADV3_MM_CKPT} \
|
||||
--pose_adaptor_ckpt ${CAMERACTRL_CKPT} \
|
||||
--model_config configs/train_cameractrl/adv3_256_384_cameractrl_relora.yaml \
|
||||
--visualization_captions assets/cameractrl_prompts.json \
|
||||
--use_specific_seeds \
|
||||
--trajectory_file assets/pose_files/0f47577ab3441480.txt \
|
||||
--n_procs 8
|
||||
```
|
||||
|
||||
where
|
||||
|
||||
- `OUTPUT_PATH` refers to the path to save resules.
|
||||
- `SD1.5_PATH` refers to the root path of the downloaded SD1.5 model.
|
||||
- `SUBFOUDER_NAME` refers to the subfolder name of unet in the `SD1.5_PATH`, default is `unet`. Here we adopt the name specified by `tools/merge_lora2unet.py`.
|
||||
- `ADV3_MM_CKPT` refers to the path of the downloaded AnimateDiffV3 motion module checkpoint.
|
||||
- `CAMERACTRL_CKPT` refers to the
|
||||
|
||||
The above inference example is used to generate videos in the original T2V model domain. The `inference.py` script supports
|
||||
generate videos in other domains with image LoRAs (`args.image_lora_rank` and `args.image_lora_ckpt`), like the [RealEstate10K](https://huggingface.co/hehao13/CameraCtrl/blob/main/RealEstate10K_LoRA.ckpt) LoRA or some personalized base models (`args.personalized_base_model`), like the [Realistic Vision](https://civitai.com/models/4201?modelVersionId=130072). please refer to the code for detail.
|
||||
|
||||
### Results
|
||||
- Same text prompt with different camera trajectories
|
||||
<table>
|
||||
<tr>
|
||||
<th width=13.3% style="text-align:center">Camera Trajectory</th>
|
||||
<th width=20% style="text-align:center">Video</th>
|
||||
<th width=13.3% style="text-align:center">Camera Trajectory</th>
|
||||
<th width=20% style="text-align:center">Video</th>
|
||||
<th width=13.3% style="text-align:center">Camera Trajectory</th>
|
||||
<th width=20% style="text-align:center">Video</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=13.3% ><img src="assets/images/horse_1.png" alt="horse1_traj" width="90%"></td>
|
||||
<td width=20%><img src="assets/gifs/horse_1.gif" alt="horse1_vid" width="90%" ></td>
|
||||
<td width=13.3%><img src="assets/images/horse_2.png" alt="horse2_traj" width="90%"></td>
|
||||
<td width=20%><img src="assets/gifs/horse_2.gif" alt="horse2_vid" width="90%" ></td>
|
||||
<td width=13.3%><img src="assets/images/horse_3.png" alt="horse3_traj" width="90%"></td>
|
||||
<td width=20%><img src="assets/gifs/horse_3.gif" alt="horse3_vid" width="90%"></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=13.3%><img src="assets/images/horse_4.png" alt="horse4_traj" width="90%"></td>
|
||||
<td width=20%><img src="assets/gifs/horse_4.gif" alt="horse4_vid" width="90%" ></td>
|
||||
<td width=13.3%><img src="assets/images/horse_5.png" alt="horse5_traj" width="90%"></td>
|
||||
<td width=20%><img src="assets/gifs/horse_5.gif" alt="horse5_vid" width="90%"></td>
|
||||
<td width=13.3%><img src="assets/images/horse_6.png" alt="horse6_traj" width="90%"></td>
|
||||
<td width=20%><img src="assets/gifs/horse_6.gif" alt="horse6_vid" width="90%"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
- Camera control on different domains' videos
|
||||
<table>
|
||||
<tr>
|
||||
<th width=11.7% style="text-align:center">Generator</th>
|
||||
<th width=11.7% style="text-align:center">Camera Trajectory</th>
|
||||
<th width=17.6% style="text-align:center">Video</th>
|
||||
<th width=11.7% style="text-align:center">Camera Trajectory</th>
|
||||
<th width=17.6% style="text-align:center">Video</th>
|
||||
<th width=11.7% style="text-align:center">Camera Trajectory</th>
|
||||
<th width=17.6% style="text-align:center">Video</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=11.7% style="text-align:center" width="90%">SD1.5</td>
|
||||
<td width=11.7%><img src="assets/images/dd1.png" alt="dd1_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/0aa284f8166e19e4_A fish is swimming in the aquarium tank.gif" alt="dd1_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd2.png" alt="dd2_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/sunflowers_3b9420585a1e66fc.gif" alt="dd2_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd3.png" alt="dd3_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/massive, multi-tiered elven palace adorned with flowing waterfalls, its cascades forming staircases between ethereal realms_2f25826f0d0ef09a.gif" alt="dd3_vid" width="90%"></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=11.7% style="text-align:center" width="90%">SD1.5 + RealEstate LoRA </td>
|
||||
<td width=11.7%><img src="assets/images/dd4.png" alt="dd4_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/a_kitchen_with_wooden_cabinets_and_a_black_stove_0bf152ef84195293.gif" alt="dd4_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd5.png" alt="dd5_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/a_living_room_with_leather_couches_and_a_fireplace_2cc5f95fbe24ffe5.gif" alt="dd5_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd6.png" alt="dd6_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/an_empty_room_with_a_desk_and_chair_in_it_0f47577ab3441480.gif" alt="dd6_vid" width="90%"></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=11.7% style="text-align:center" width="90%">Realistic Vision</td>
|
||||
<td width=11.7%><img src="assets/images/dd7.png" alt="dd7_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/photo of coastline, rocks, storm weather, wind, waves, lightning, soft lighting_ 9d022c4ec370112a.gif" alt="dd7_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd8.png" alt="dd8_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/close up photo of a rabbit, forest, haze, halation, bloom, dramatic atmosphere, centred_3f79dc32d575bcdc.gif" alt="dd8_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd9.png" alt="dd9_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/realestate_horizontal_uniform.gif" alt="dd9_vid" width="90%"></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=11.7% style="text-align:center" width="90%">ToonYou</td>
|
||||
<td width=11.7%><img src="assets/images/dd10.png" alt="dd10_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/toonyou_ 62feb0ed164ebcbe.gif" alt="dd10_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd11.png" alt="dd11_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/0f47577ab3441480_mkvd, 1girl, turtleneck sweater, sweater yellow, happy, looking at viewer_.gif" alt="dd8_vid" width="90%"></td>
|
||||
<td width=11.7%><img src="assets/images/dd12.png" alt="dd12_traj" width="90%"></td>
|
||||
<td width=17.6%><img src="assets/gifs/closeup face photo of man in black clothes, night city street, bokeh, fireworks in background_4e012c05fdf8f9b3.gif" alt="dd12_vid" width="90%"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
Note that, each image paired with the video represents the camera trajectory. Each small tetrahedron on the image represents the position and orientation of the camera for one video frame. Its vertex stands for the camera location, while the base represents the imaging plane of the camera. The red arrows indicate the movement of camera **position**. The camera rotation can be observed through the orientation of the tetrahedrons.
|
||||
|
||||
|
||||
## Training
|
||||
|
||||
### Step1 (RealEstate10K image LoRA)
|
||||
Update the below paths to data and pretrained model of the config `configs/train_image_lora/realestate_lora.yaml`
|
||||
```shell
|
||||
pretrained_model_path: "[replace with SD1.5 root path]"
|
||||
train_data:
|
||||
root_path: "[replace RealEstate10K root path]"
|
||||
```
|
||||
Other training parameters (lr, epochs, validation settings, etc.) are also included in the config files.
|
||||
|
||||
Then, launch the image LoRA training using slurm
|
||||
```shell
|
||||
./slurm_run.sh ${PARTITION} image_lora 8 configs/train_image_lora/realestate_lora.yaml train_image_lora.py
|
||||
```
|
||||
or PyTorch
|
||||
```shell
|
||||
./dist_run.sh configs/train_image_lora/realestate_lora.yaml 8 train_image_lora.py
|
||||
```
|
||||
We provide our pretrained checkpoint of the RealEstate10K LoRA model in [HuggingFace](https://huggingface.co/hehao13/CameraCtrl/blob/main/RealEstate10K_LoRA.ckpt).
|
||||
|
||||
### Step2 (Camera control model)
|
||||
Update the below paths to data and pretrained model of the config `configs/train_cameractrl/adv3_256_384_cameractrl_relora.yaml`
|
||||
```shell
|
||||
pretrained_model_path: "[replace with SD1.5 root path]"
|
||||
train_data:
|
||||
root_path: "[replace RealEstate10K root path]"
|
||||
validation_data:
|
||||
root_path: "[replace RealEstate10K root path]"
|
||||
lora_ckpt: "[Replace with RealEstate10k image LoRA ckpt]"
|
||||
motion_module_ckpt: "[Replace with ADV3 motion module]"
|
||||
```
|
||||
Other training parameters (lr, epochs, validation settings, etc.) are also included in the config files.
|
||||
|
||||
|
||||
Then, launch the camera control model training using slurm
|
||||
```shell
|
||||
./slurm_run.sh ${PARTITION} cameractrl 8 configs/train_cameractrl/adv3_256_384_cameractrl_relora.yaml train_camera_control.py
|
||||
```
|
||||
or PyTorch
|
||||
```shell
|
||||
./dist_run.sh configs/train_cameractrl/adv3_256_384_cameractrl_relora.yaml 8 train_camera_control.py
|
||||
```
|
||||
|
||||
## Disclaimer
|
||||
This project is released for academic use. We disclaim responsibility for user-generated content. Users are solely liable for their actions. The project contributors are not legally affiliated with, nor accountable for, users' behaviors. Use the generative model responsibly, adhering to ethical and legal standards.
|
||||
|
||||
|
||||
## Acknowledgement
|
||||
We thank [AnimateDiff](https://github.com/guoyww/AnimateDiff) for their amazing codes and models.
|
||||
|
||||
|
||||
## BibTeX
|
||||
|
||||
```bibtex
|
||||
@article{he2024cameractrl,
|
||||
title={CameraCtrl: Enabling Camera Control for Text-to-Video Generation},
|
||||
author={Hao He and Yinghao Xu and Yuwei Guo and Gordon Wetzstein and Bo Dai and Hongsheng Li and Ceyuan Yang},
|
||||
journal={arXiv preprint arXiv:2404.02101},
|
||||
year={2024}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS']
|
||||
@@ -0,0 +1,22 @@
|
||||
{"prompts":["A serene mountain lake at sunrise, with mist hovering over the water.",
|
||||
"A still life of vintage objects on a wooden table.",
|
||||
"Ripe apples on a wooden table.",
|
||||
"Turtle swimming in ocean.",
|
||||
"A horse is eating grass on the grassland.",
|
||||
"Natural hot spring the steam floats up.",
|
||||
"A fireworks display illuminating the night sky.",
|
||||
"Rocky coastline with crashing waves.",
|
||||
"A fish is swimming in the aquarium tank.",
|
||||
"massive, multi-tiered elven palace adorned with flowing waterfalls, its cascades forming staircases between ethereal realms",
|
||||
"The sunflower swaying in the wind."],
|
||||
"seeds":[6426918851609095805,
|
||||
6426918851609095805,
|
||||
10051121533489271199,
|
||||
2639703735194618768,
|
||||
6426918851609095805,
|
||||
7236506129240698163,
|
||||
9543750473687926035,
|
||||
11689686181370386440,
|
||||
2795087382974079784,
|
||||
6479874378964516444,
|
||||
10800640933498641993]}
|
||||
@@ -0,0 +1,10 @@
|
||||
{"prompts":["a kitchen with wooden cabinets and a black stove",
|
||||
"a bathroom with a shower curtain and a toilet",
|
||||
"a kitchen with a sliding glass door leading to a deck",
|
||||
"an empty room with a desk and chair in it",
|
||||
"a kitchen with wooden cabinets and stainless steel appliances",
|
||||
"a view of a living room through a doorway",
|
||||
"a bedroom with a bed, desk and television",
|
||||
"a bedroom with a large bed and three windows",
|
||||
"a bedroom with a bed and dresser in it",
|
||||
"a kitchen with a stove top oven and a sink"]}
|
||||
|
After Width: | Height: | Size: 771 KiB |
|
After Width: | Height: | Size: 588 KiB |
|
After Width: | Height: | Size: 819 KiB |
|
After Width: | Height: | Size: 903 KiB |
|
After Width: | Height: | Size: 750 KiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
After Width: | Height: | Size: 762 KiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 1.3 MiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 1.3 MiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 1.4 MiB |
|
After Width: | Height: | Size: 847 KiB |
|
After Width: | Height: | Size: 1.0 MiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 737 KiB |
|
After Width: | Height: | Size: 70 KiB |
|
After Width: | Height: | Size: 53 KiB |
|
After Width: | Height: | Size: 63 KiB |
|
After Width: | Height: | Size: 67 KiB |
|
After Width: | Height: | Size: 62 KiB |
|
After Width: | Height: | Size: 51 KiB |
|
After Width: | Height: | Size: 74 KiB |
|
After Width: | Height: | Size: 68 KiB |
|
After Width: | Height: | Size: 62 KiB |
|
After Width: | Height: | Size: 60 KiB |
|
After Width: | Height: | Size: 46 KiB |
|
After Width: | Height: | Size: 55 KiB |
|
After Width: | Height: | Size: 73 KiB |
|
After Width: | Height: | Size: 72 KiB |
|
After Width: | Height: | Size: 26 KiB |
|
After Width: | Height: | Size: 75 KiB |
|
After Width: | Height: | Size: 68 KiB |
|
After Width: | Height: | Size: 48 KiB |
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=QShWPZxTDoE
|
||||
158692025 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.780003667 0.059620168 -0.622928321 0.726968666 -0.062449891 0.997897983 0.017311305 0.217967188 0.622651041 0.025398925 0.782087326 -1.002211444
|
||||
158958959 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.743836701 0.064830206 -0.665209770 0.951841944 -0.068305343 0.997446954 0.020830527 0.206496789 0.664861917 0.029942872 0.746365905 -1.084913992
|
||||
159225893 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.697046876 0.070604131 -0.713540971 1.208789672 -0.074218854 0.996899366 0.026138915 0.196421447 0.713174045 0.034738146 0.700125754 -1.130142078
|
||||
159526193 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.635762572 0.077846259 -0.767949164 1.465161122 -0.080595709 0.996158004 0.034256749 0.157107229 0.767665446 0.040114246 0.639594078 -1.136893070
|
||||
159793126 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.593250692 0.083153486 -0.800711632 1.635091834 -0.085384794 0.995539784 0.040124334 0.135863998 0.800476789 0.044564810 0.597704709 -1.166997229
|
||||
160093427 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.555486798 0.087166689 -0.826943994 1.803789619 -0.089439675 0.994984210 0.044799786 0.145490422 0.826701283 0.049075913 0.560496747 -1.243827350
|
||||
160360360 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.523399472 0.090266660 -0.847292721 1.945815368 -0.093254104 0.994468153 0.048340045 0.174777447 0.846969128 0.053712368 0.528921843 -1.336914479
|
||||
160660661 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.491546303 0.092127070 -0.865964711 2.093852892 -0.095617607 0.994085968 0.051482171 0.196702533 0.865586221 0.057495601 0.497448236 -1.439709380
|
||||
160927594 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.475284129 0.093297184 -0.874871790 2.200792438 -0.096743606 0.993874133 0.053430639 0.209217395 0.874497354 0.059243519 0.481398523 -1.547068315
|
||||
161227895 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.464444131 0.093880348 -0.880612373 2.324141986 -0.097857766 0.993716478 0.054326952 0.220651207 0.880179226 0.060942926 0.470712721 -1.712512928
|
||||
161494828 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.458157241 0.093640216 -0.883925021 2.443100890 -0.098046601 0.993691206 0.054448847 0.257385043 0.883447111 0.061719712 0.464447916 -1.885672329
|
||||
161795128 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.457354397 0.093508720 -0.884354591 2.543246338 -0.097820736 0.993711591 0.054482624 0.281562244 0.883888066 0.061590351 0.463625461 -2.094829165
|
||||
162062062 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.465170115 0.093944497 -0.880222261 2.606377358 -0.097235762 0.993758380 0.054675922 0.277376127 0.879864752 0.060155477 0.471401453 -2.299280675
|
||||
162362362 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.511845231 0.090872414 -0.854257941 2.576774100 -0.093636356 0.994366586 0.049672548 0.270516319 0.853959382 0.054564942 0.517470777 -2.624374352
|
||||
162629296 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.590568483 0.083218277 -0.802685261 2.398318316 -0.085610889 0.995516419 0.040222570 0.282138215 0.802433550 0.044964414 0.595045030 -3.012309268
|
||||
162929596 0.474812461 0.844111024 0.500000000 0.500000000 0.000000000 0.000000000 0.684302032 0.072693504 -0.725566208 2.086323553 -0.074529484 0.996780157 0.029575195 0.310959312 0.725379944 0.033837710 0.687516510 -3.456740526
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=a-Unpcomk5k
|
||||
89889800 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.959632158 -0.051068146 0.276583046 0.339363991 0.046715312 0.998659134 0.022308502 0.111317310 -0.277351439 -0.008487292 0.960731030 -0.353512177
|
||||
90156733 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.939171016 -0.057914909 0.338531673 0.380727498 0.052699961 0.998307705 0.024584483 0.134404073 -0.339382589 -0.005248427 0.940633774 -0.477942109
|
||||
90423667 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.913449824 -0.063028678 0.402040780 0.393354042 0.056629892 0.998008251 0.027794635 0.151535333 -0.402991891 -0.002621480 0.915199816 -0.622810637
|
||||
90723967 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.879072070 -0.069992281 0.471522361 0.381271678 0.062575974 0.997545719 0.031412520 0.175549569 -0.472563744 0.001892101 0.881294429 -0.821022008
|
||||
90990900 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.846152365 -0.078372896 0.527146876 0.360267421 0.071291871 0.996883452 0.033775900 0.212440374 -0.528151155 0.009001731 0.849102676 -1.013792538
|
||||
91291200 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.806246638 -0.086898506 0.585162342 0.297888150 0.078344196 0.996124208 0.039983708 0.243578507 -0.586368918 0.013607344 0.809929788 -1.248063630
|
||||
91558133 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.771091938 -0.093814306 0.629774630 0.223948432 0.087357447 0.995320201 0.041307874 0.293608807 -0.630702674 0.023163332 0.775678813 -1.459775674
|
||||
91858433 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.737968326 -0.099363215 0.667480111 0.145501271 0.093257703 0.994626462 0.044957232 0.329381977 -0.668360531 0.029070651 0.743269205 -1.688460978
|
||||
92125367 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.716826320 -0.101809755 0.689778805 0.086545731 0.098867603 0.994127929 0.043986596 0.379651732 -0.690206647 0.036666028 0.722682774 -1.885393814
|
||||
92425667 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.703360021 -0.101482928 0.703552365 0.039205180 0.098760851 0.994108558 0.044659954 0.417778776 -0.703939617 0.038071405 0.709238708 -2.106152155
|
||||
92692600 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.699525177 -0.101035394 0.707429409 0.029387371 0.096523918 0.994241416 0.046552572 0.439027166 -0.708059072 0.035719164 0.705249250 -2.314481674
|
||||
92992900 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.698582709 -0.101620331 0.708276451 0.018437890 0.096638583 0.994193733 0.047326516 0.478349552 -0.708973348 0.035385344 0.704347014 -2.540820022
|
||||
93259833 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.704948425 -0.098988213 0.702316940 0.047566428 0.095107265 0.994462848 0.044701166 0.517456396 -0.702853024 0.035283424 0.710459530 -2.724204596
|
||||
93560133 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.714113414 -0.104848787 0.692133486 0.107161588 0.100486010 0.993833601 0.046875130 0.568063228 -0.692780316 0.036075566 0.720245779 -2.948379150
|
||||
93827067 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.717699587 -0.112323314 0.687234104 0.118765931 0.105546549 0.993049562 0.052081093 0.593900230 -0.688307464 0.035156611 0.724566638 -3.140363331
|
||||
94127367 0.485388169 0.862912326 0.500000000 0.500000000 0.000000000 0.000000000 0.715531290 -0.122954883 0.687675118 0.089455249 0.115526602 0.991661787 0.057100743 0.643643035 -0.688961923 0.038587399 0.723769605 -3.310401931
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=_ca03xP_KUU
|
||||
211244000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.984322786 0.006958477 -0.176239252 0.004217217 -0.005594095 0.999950409 0.008237306 -0.107944544 0.176287830 -0.007122268 0.984312892 -0.571743822
|
||||
211511000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.981951714 0.008860772 -0.188924149 0.000856103 -0.007234470 0.999930620 0.009296093 -0.149397579 0.188993424 -0.007761548 0.981947660 -0.776566486
|
||||
211778000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.981318414 0.010160952 -0.192122281 -0.005546933 -0.008323869 0.999911606 0.010366773 -0.170816348 0.192210630 -0.008573905 0.981316268 -0.981924227
|
||||
212078000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.981108844 0.010863926 -0.193151161 0.019480142 -0.008781361 0.999893725 0.011634931 -0.185801323 0.193257034 -0.009719004 0.981100023 -1.207220396
|
||||
212345000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.981263518 0.010073495 -0.192407012 0.069708411 -0.008015377 0.999902070 0.011472094 -0.203594876 0.192503735 -0.009714933 0.981248140 -1.408936391
|
||||
212646000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.980964184 0.009669405 -0.193947718 0.166020848 -0.007467276 0.999899149 0.012082115 -0.219176122 0.194044977 -0.010403861 0.980937481 -1.602649833
|
||||
212913000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.980841637 0.009196524 -0.194589555 0.262465567 -0.006609587 0.999880970 0.013939449 -0.224018296 0.194694594 -0.012386235 0.980785728 -1.740759996
|
||||
213212000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.980620921 0.008805701 -0.195716679 0.389752858 -0.006055873 0.999874413 0.014644019 -0.230312701 0.195821062 -0.013174997 0.980551124 -1.890949759
|
||||
213479000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.980327129 0.009317402 -0.197159693 0.505632551 -0.006113928 0.999839306 0.016850581 -0.230702867 0.197285011 -0.015313662 0.980226576 -2.016199670
|
||||
213779000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.980493963 0.009960363 -0.196296573 0.623893674 -0.006936011 0.999846518 0.016088497 -0.223079036 0.196426690 -0.014413159 0.980412602 -2.137999468
|
||||
214046000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.980032921 0.010318150 -0.198567480 0.754726451 -0.007264129 0.999843955 0.016102606 -0.222246314 0.198702648 -0.014338664 0.979954958 -2.230292399
|
||||
214347000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.976653159 0.010179597 -0.214580998 0.946523963 -0.006709154 0.999834776 0.016895246 -0.210005171 0.214717537 -0.015061138 0.976560056 -2.305666573
|
||||
214614000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.971478105 0.011535713 -0.236848563 1.096604956 -0.007706031 0.999824286 0.017088750 -0.192895049 0.237004071 -0.014776184 0.971396267 -2.365701917
|
||||
214914000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.965282261 0.014877280 -0.260785013 1.237534109 -0.014124592 0.999888897 0.004760279 -0.136261458 0.260826856 -0.000911531 0.965385139 -2.458136272
|
||||
215181000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.961933076 0.016891202 -0.272762626 1.331672110 -0.022902885 0.999559581 -0.018870916 -0.076291319 0.272323757 0.024399608 0.961896241 -2.579417067
|
||||
215481000 0.479272232 0.852039479 0.500000000 0.500000000 0.000000000 0.000000000 0.959357142 0.017509742 -0.281651050 1.417338469 -0.039949402 0.996448219 -0.074127860 0.083949011 0.279352754 0.082366876 0.956649244 -2.712094466
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=in69BD2eZqg
|
||||
195562033 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.999749303 -0.004518872 0.021929268 0.038810557 0.004613766 0.999980211 -0.004278630 0.328177052 -0.021909500 0.004378735 0.999750376 -0.278403591
|
||||
195828967 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.999336481 -0.006239665 0.035883281 0.034735125 0.006456365 0.999961615 -0.005926326 0.417233500 -0.035844926 0.006154070 0.999338388 -0.270773664
|
||||
196095900 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.998902142 -0.007417044 0.046254709 0.033849936 0.007582225 0.999965489 -0.003396692 0.504852301 -0.046227921 0.003743677 0.998923898 -0.256677740
|
||||
196396200 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.998096347 -0.008753631 0.061049398 0.026475959 0.009088391 0.999945164 -0.005207890 0.583593760 -0.061000463 0.005752816 0.998121142 -0.236166024
|
||||
196663133 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.997463286 -0.009214416 0.070583619 0.014842158 0.009590282 0.999941587 -0.004988078 0.634675512 -0.070533529 0.005652342 0.997493386 -0.198663134
|
||||
196963433 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.996699810 -0.009558053 0.080611102 0.003250557 0.009986609 0.999938071 -0.004914839 0.670145924 -0.080559134 0.005703651 0.996733487 -0.141256339
|
||||
197230367 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.996102691 -0.010129508 0.087617576 -0.013035317 0.010638822 0.999929130 -0.005347892 0.673139255 -0.087557197 0.006259197 0.996139824 -0.073934910
|
||||
197530667 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.995880842 -0.009925503 0.090126604 -0.036202423 0.010367444 0.999936402 -0.004436717 0.655632681 -0.090076834 0.005352824 0.995920420 0.017267095
|
||||
197797600 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.995802402 -0.010077500 0.090972595 -0.060858524 0.010445373 0.999939084 -0.003568561 0.618604505 -0.090931088 0.004503824 0.995846987 0.133592270
|
||||
198097900 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.995846093 -0.010148350 0.090484887 -0.077962281 0.010412642 0.999942780 -0.002449236 0.561822755 -0.090454854 0.003381249 0.995894790 0.274195378
|
||||
198364833 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.995989919 -0.009936163 0.088912196 -0.082315587 0.010200773 0.999944806 -0.002522171 0.520613290 -0.088882230 0.003419030 0.996036291 0.395169547
|
||||
198665133 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.997159958 -0.009351323 0.074730076 -0.068472873 0.009822783 0.999934077 -0.005943770 0.466061412 -0.074669570 0.006660947 0.997186065 0.549834051
|
||||
198932067 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.998626053 -0.008290987 0.051742285 -0.037270541 0.008407482 0.999962568 -0.002034174 0.410440195 -0.051723484 0.002466401 0.998658419 0.690111645
|
||||
199232367 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.999980092 -0.004756952 0.004140501 -0.005957613 0.004773445 0.999980688 -0.003982662 0.354437092 -0.004121476 0.004002347 0.999983490 0.842797271
|
||||
199499300 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.998872638 0.001147069 -0.047456335 0.002603018 -0.001435435 0.999980688 -0.006042828 0.295339877 0.047448486 0.006104136 0.998855054 0.988644188
|
||||
199799600 0.507650910 0.902490531 0.500000000 0.500000000 0.000000000 0.000000000 0.992951691 0.008710741 -0.118199304 -0.030798243 -0.009495872 0.999936402 -0.006080875 0.208803899 0.118138820 0.007160421 0.992971301 1.161643267
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=-aldZQifF2U
|
||||
103736967 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.804089785 -0.073792785 0.589910388 -2.686968354 0.081914566 0.996554494 0.013005137 0.128970374 -0.588837504 0.037864953 0.807363987 -1.789505608
|
||||
104003900 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.772824645 -0.077280566 0.629896700 -2.856354365 0.084460691 0.996253133 0.018602582 0.115028772 -0.628974140 0.038824979 0.776456118 -1.799931844
|
||||
104270833 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.740043461 -0.078656308 0.667943776 -3.017167990 0.086847030 0.995998919 0.021066183 0.116867188 -0.666928232 0.042419042 0.743913531 -1.815074499
|
||||
104571133 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.696879685 -0.073477358 0.713414192 -3.221640235 0.086792909 0.996067226 0.017807571 0.133618379 -0.711916924 0.049509555 0.700516284 -1.784051774
|
||||
104838067 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.654997289 -0.066671766 0.752684176 -3.418233112 0.086666502 0.996154904 0.012819566 0.161623584 -0.750644684 0.056835718 0.658256948 -1.733288907
|
||||
105138367 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.603833497 -0.059696361 0.794871926 -3.619566170 0.087576874 0.996123314 0.008281946 0.184519895 -0.792284906 0.064611480 0.606720686 -1.643568460
|
||||
105405300 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.555575073 -0.055402864 0.829618514 -3.768244320 0.089813948 0.995938241 0.006363695 0.197587954 -0.826601386 0.070975810 0.558294415 -1.559717271
|
||||
105705600 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.501615226 -0.052979972 0.863467038 -3.914896511 0.093892507 0.995560884 0.006539768 0.201989601 -0.859980464 0.077792637 0.504362881 -1.476983336
|
||||
105972533 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.454045177 -0.052372806 0.889438093 -4.034987790 0.099656843 0.994991958 0.007714771 0.211683202 -0.885387778 0.085135736 0.456990600 -1.405070279
|
||||
106272833 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.397668689 -0.051785514 0.916066527 -4.178181130 0.105599925 0.994354606 0.010369749 0.208751884 -0.911431968 0.092612833 0.400892258 -1.295093582
|
||||
106539767 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.345666498 -0.052948993 0.936862350 -4.285116664 0.110631727 0.993743002 0.015344846 0.195070069 -0.931812882 0.098342501 0.349361509 -1.182773054
|
||||
106840067 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.284817457 -0.055293880 0.956985712 -4.392320606 0.115495987 0.993041575 0.023003323 0.168523273 -0.951598525 0.103976257 0.289221793 -1.053514096
|
||||
107107000 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.228878200 -0.056077410 0.971838534 -4.485196000 0.120451130 0.992298782 0.028890507 0.159180748 -0.965974271 0.110446639 0.233870149 -0.923927626
|
||||
107407300 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.162932962 -0.053445265 0.985188544 -4.601126217 0.124115810 0.991709769 0.033272449 0.152041098 -0.978799343 0.116856292 0.168215603 -0.758111250
|
||||
107674233 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.102818660 -0.051196381 0.993381739 -4.691710857 0.127722457 0.991087139 0.037858382 0.141352300 -0.986466050 0.122984610 0.108441174 -0.599244073
|
||||
107974533 0.474175212 0.842978122 0.500000000 0.500000000 0.000000000 0.000000000 0.034108389 -0.050325166 0.998150289 -4.758242879 0.132215530 0.990180492 0.045405328 0.118994547 -0.990633965 0.130422264 0.040427230 -0.433560831
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=sLIFyXD2ujI
|
||||
77444033 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.980310440 0.030424286 -0.195104495 -0.195846403 -0.034550700 0.999244750 -0.017780757 0.034309913 0.194416180 0.024171660 0.980621278 -0.178639121
|
||||
77610867 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.973806441 0.034138829 -0.224801064 -0.221452338 -0.039088678 0.999080658 -0.017603843 0.038706263 0.223993421 0.025929911 0.974245667 -0.219951444
|
||||
77777700 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.965910375 0.037603889 -0.256131083 -0.242696017 -0.043735024 0.998875856 -0.018281631 0.046505467 0.255155712 0.028860316 0.966469169 -0.265310453
|
||||
77944533 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.956261098 0.040829532 -0.289650917 -0.252766079 -0.048421524 0.998644531 -0.019089982 0.054620904 0.288478881 0.032280345 0.956941962 -0.321621308
|
||||
78144733 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.941536248 0.043692805 -0.334066480 -0.250198162 -0.053955212 0.998311937 -0.021497937 0.069548726 0.332563221 0.038265716 0.942304313 -0.401964240
|
||||
78311567 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.926658213 0.046350952 -0.373036444 -0.239336491 -0.058738846 0.998033047 -0.021904159 0.077439241 0.371287435 0.042209402 0.927558064 -0.474019461
|
||||
78478400 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.909880757 0.048629351 -0.412009954 -0.218247042 -0.063676558 0.997708619 -0.022863906 0.088967126 0.409954011 0.047038805 0.910892427 -0.543114491
|
||||
78645233 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.891359746 0.050869841 -0.450433195 -0.185763327 -0.067926541 0.997452736 -0.021771761 0.093745158 0.448178291 0.050002839 0.892544627 -0.611223637
|
||||
78845433 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.870108902 0.054302387 -0.489858925 -0.153515269 -0.074510135 0.996981323 -0.021829695 0.107765162 0.487194777 0.055493668 0.871528387 -0.691303250
|
||||
79012267 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.852772951 0.056338910 -0.519234240 -0.128052677 -0.078825951 0.996660411 -0.021319628 0.116291007 0.516299069 0.059109934 0.854366004 -0.760654136
|
||||
79179100 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.835254073 0.059146367 -0.546673834 -0.101344556 -0.084243484 0.996225357 -0.020929486 0.126763936 0.543372452 0.063535146 0.837083995 -0.832841061
|
||||
79345933 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.818755865 0.062443536 -0.570736051 -0.077325807 -0.089739971 0.995768547 -0.019791666 0.136091605 0.567085147 0.067422375 0.820895016 -0.908256727
|
||||
79546133 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.798365474 0.066207208 -0.598522484 -0.043774887 -0.096616283 0.995144248 -0.018795265 0.150808225 0.594371796 0.072832510 0.800885499 -0.994657638
|
||||
79712967 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.781648815 0.069040783 -0.619885862 -0.013285614 -0.101820730 0.994646847 -0.017611075 0.161173621 0.615351617 0.076882906 0.784494340 -1.070102980
|
||||
79879800 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.765168309 0.072694756 -0.639713168 0.019850080 -0.108554602 0.993946910 -0.016894773 0.177612448 0.634612799 0.082371153 0.768428028 -1.147576811
|
||||
80080000 0.483930168 0.860320329 0.500000000 0.500000000 0.000000000 0.000000000 0.745406330 0.077463314 -0.662094295 0.062107046 -0.117075674 0.993000031 -0.015629012 0.200140798 0.656248927 0.089165099 0.749257565 -1.238600776
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=t-mlAKnESzQ
|
||||
167200000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.991872013 -0.011311784 0.126735851 0.400533760 0.012037775 0.999915242 -0.004963919 -0.047488550 -0.126668960 0.006449190 0.991924107 -0.414499612
|
||||
167467000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.991945148 -0.011409644 0.126153216 0.506974565 0.012122569 0.999914587 -0.004884966 -0.069421149 -0.126086697 0.006374919 0.991998732 -0.517325825
|
||||
167734000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.992271781 -0.010751382 0.123616949 0.590358341 0.011312660 0.999928653 -0.003839425 -0.085158661 -0.123566844 0.005208189 0.992322564 -0.599035085
|
||||
168034000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.993287027 -0.009973313 0.115245141 0.673577580 0.010455138 0.999938965 -0.003577147 -0.104263255 -0.115202427 0.004758038 0.993330657 -0.691557669
|
||||
168301000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.993988216 -0.009955904 0.109033749 0.753843765 0.010435819 0.999938190 -0.003831771 -0.106670354 -0.108988866 0.004946592 0.994030654 -0.805538867
|
||||
168602000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.994774222 -0.010583352 0.101549298 0.846176230 0.011122120 0.999926925 -0.004740742 -0.089426372 -0.101491705 0.005845411 0.994819224 -0.933449460
|
||||
168869000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.995415390 -0.010595482 0.095057681 0.913119395 0.011053002 0.999929726 -0.004287821 -0.072756893 -0.095005572 0.005318835 0.995462537 -1.037255409
|
||||
169169000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.996029556 -0.009414902 0.088523701 0.977259045 0.009879347 0.999939620 -0.004809874 -0.042104006 -0.088473074 0.005665333 0.996062458 -1.127427189
|
||||
169436000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.996695757 -0.008890423 0.080737323 1.025351476 0.009221899 0.999950528 -0.003733651 -0.007486727 -0.080700137 0.004465866 0.996728420 -1.188659636
|
||||
169736000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.997404695 -0.008404067 0.071506783 1.073562767 0.008649707 0.999957681 -0.003126226 0.054879890 -0.071477488 0.003736625 0.997435212 -1.216979926
|
||||
170003000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.997887254 -0.008228444 0.064446673 1.110116903 0.008409287 0.999961436 -0.002535321 0.124372514 -0.064423330 0.003071915 0.997917950 -1.231904045
|
||||
170303000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.998332024 -0.007790270 0.057205603 1.136173895 0.007975516 0.999963641 -0.003010646 0.212542522 -0.057180069 0.003461868 0.998357892 -1.242942079
|
||||
170570000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.998471320 -0.007715963 0.054730706 1.159189486 0.007868989 0.999965727 -0.002581036 0.310163907 -0.054708913 0.003007766 0.998497844 -1.245661417
|
||||
170871000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.998552144 -0.007742116 0.053231847 1.173763753 0.007991423 0.999958038 -0.004472161 0.412779543 -0.053194992 0.004891084 0.998572171 -1.229165757
|
||||
171137000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.998553872 -0.007909958 0.053175092 1.179029258 0.008138723 0.999958515 -0.004086919 0.509089997 -0.053140558 0.004513786 0.998576820 -1.196146494
|
||||
171438000 0.470983989 0.837304886 0.500000000 0.500000000 0.000000000 0.000000000 0.998469293 -0.008281939 0.054685175 1.181414517 0.008542870 0.999953210 -0.004539483 0.618089736 -0.054645021 0.004999703 0.998493314 -1.159911786
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=bJyPo9pESu0
|
||||
189622767 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.966956913 0.041186374 -0.251590967 0.235831829 -0.037132759 0.999092996 0.020840336 0.069818943 0.252221137 -0.010809440 0.967609227 -0.850289525
|
||||
189789600 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.967445135 0.041703269 -0.249621317 0.217678822 -0.037349533 0.999056637 0.022154763 0.078295447 0.250309765 -0.012110277 0.968090057 -0.818677483
|
||||
189956433 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.967587769 0.043305319 -0.248794302 0.196350216 -0.038503598 0.998966932 0.024136283 0.085749990 0.249582499 -0.013774496 0.968255579 -0.778043636
|
||||
190123267 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.967742383 0.044170257 -0.248039767 0.170234078 -0.039154600 0.998917341 0.025120445 0.090556068 0.248880804 -0.014598221 0.968424082 -0.733500964
|
||||
190323467 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.973553717 0.043272153 -0.224322766 0.120337922 -0.038196862 0.998907626 0.026917407 0.091227451 0.225242496 -0.017637115 0.974143088 -0.680520640
|
||||
190490300 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.984184802 0.039637893 -0.172653258 0.065019106 -0.035194401 0.998967648 0.028723357 0.090969669 0.173613548 -0.022192664 0.984563768 -0.638603728
|
||||
190657133 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.993078411 0.035358477 -0.112004772 0.011571313 -0.032207530 0.999036312 0.029818388 0.092482656 0.112951167 -0.026004599 0.993260205 -0.588118143
|
||||
190823967 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.997807920 0.031166473 -0.058378015 -0.027908508 -0.029339414 0.999060452 0.031897116 0.092538838 0.059317287 -0.030114418 0.997784853 -0.529325066
|
||||
191024167 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.999654651 0.026247790 -0.001263706 -0.064570799 -0.026190240 0.999087334 0.033742432 0.091922841 0.002148218 -0.033697683 0.999429762 -0.448626929
|
||||
191191000 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.998773992 0.022079065 0.044305529 -0.084478169 -0.023622099 0.999121666 0.034611158 0.087434649 -0.043502431 -0.035615314 0.998418272 -0.371306296
|
||||
191357833 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.995725632 0.017435640 0.090699598 -0.094868572 -0.020876031 0.999092638 0.037122324 0.082208324 -0.089970052 -0.038857099 0.995186150 -0.290596011
|
||||
191524667 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.989503622 0.013347236 0.143890470 -0.096537122 -0.019140780 0.999057651 0.038954727 0.079283141 -0.143234938 -0.041300017 0.988826632 -0.207477308
|
||||
191724867 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.975660741 0.006981443 0.219174415 -0.085240259 -0.016479453 0.999001026 0.041537181 0.072219148 -0.218665481 -0.044138070 0.974801123 -0.112100139
|
||||
191891700 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.955792487 -0.000511726 0.294041574 -0.064476318 -0.012924311 0.998958945 0.043749433 0.061688334 -0.293757826 -0.045615666 0.954790831 -0.034724173
|
||||
192058533 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.925362229 -0.009678445 0.378960580 -0.029417786 -0.008889219 0.998845160 0.047216032 0.058476640 -0.378979892 -0.047060598 0.924207509 0.042010383
|
||||
192258733 0.474122545 0.842884498 0.500000000 0.500000000 0.000000000 0.000000000 0.872846186 -0.021581186 0.487517983 0.038433307 -0.004890797 0.998584569 0.052961230 0.057516307 -0.487970918 -0.048611358 0.871505201 0.124675285
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=1qVpRlWxam4
|
||||
86319567 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999183893 0.038032386 -0.013605987 -0.249154748 -0.038085770 0.999267697 -0.003686040 0.047875167 0.013455833 0.004201226 0.999900639 -0.566803149
|
||||
86586500 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999392629 0.034676589 -0.003445767 -0.282371175 -0.034685481 0.999395013 -0.002555777 0.057086778 0.003355056 0.002673743 0.999990821 -0.624021456
|
||||
86853433 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999498725 0.028301919 0.014187563 -0.320995587 -0.028301118 0.999599397 -0.000257314 0.061367205 -0.014189162 -0.000144339 0.999899328 -0.706664680
|
||||
87153733 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999064565 0.022049030 0.037200645 -0.371910835 -0.022201553 0.999746680 0.003691827 0.063911726 -0.037109818 -0.004514286 0.999301016 -0.799748814
|
||||
87420667 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.998171926 0.018552339 0.057520505 -0.440220060 -0.018887693 0.999807596 0.005291941 0.070160264 -0.057411261 -0.006368696 0.998330295 -0.853433007
|
||||
87720967 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.997675776 0.016262729 0.066170901 -0.486385324 -0.016915560 0.999813497 0.009317505 0.069230577 -0.066007033 -0.010415167 0.997764826 -0.912234761
|
||||
87987900 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.998019218 0.015867118 0.060876362 -0.497549423 -0.016505934 0.999813735 0.010005167 0.076295227 -0.060706269 -0.010990170 0.998095155 -0.980435972
|
||||
88288200 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999152124 0.018131699 0.036962789 -0.468507446 -0.018461898 0.999792457 0.008611582 0.087696066 -0.036798976 -0.009286684 0.999279559 -1.074633197
|
||||
88555133 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999717414 0.022977378 0.006097841 -0.420528982 -0.023013741 0.999717355 0.005961678 0.101216630 -0.005959134 -0.006100327 0.999963641 -1.169004730
|
||||
88855433 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.999106526 0.030726369 -0.029017152 -0.374249594 -0.030677194 0.999527037 0.002138488 0.120936030 0.029069137 -0.001246413 0.999576628 -1.251082317
|
||||
89122367 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.997359693 0.039784521 -0.060752310 -0.335843098 -0.039773725 0.999207735 0.001387495 0.132824955 0.060759377 0.001032514 0.998151898 -1.312258423
|
||||
89422667 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.992973983 0.050480653 -0.107025139 -0.253623964 -0.050627887 0.998716712 0.001342622 0.144421611 0.106955573 0.004085269 0.994255424 -1.394020432
|
||||
89689600 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.986886561 0.059628733 -0.149997801 -0.173418608 -0.059660275 0.998209476 0.004293700 0.142984494 0.149985254 0.004711515 0.988677025 -1.462588413
|
||||
89989900 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.978200734 0.067550205 -0.196367815 -0.089199207 -0.067402542 0.997698128 0.007442682 0.141665403 0.196418539 0.005955252 0.980502069 -1.524381413
|
||||
90256833 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.967793405 0.073765829 -0.240695804 0.013635864 -0.073441446 0.997246027 0.010330606 0.134276795 0.240794986 0.007679154 0.970545650 -1.588498428
|
||||
90557133 0.487278048 0.866272132 0.500000000 0.500000000 0.000000000 0.000000000 0.953711152 0.081056722 -0.289594263 0.148156165 -0.081249826 0.996628821 0.011376631 0.129987979 0.289540142 0.012679463 0.957081914 -1.633951355
|
||||
@@ -0,0 +1,17 @@
|
||||
https://www.youtube.com/watch?v=mGFQkgadzRQ
|
||||
123665000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.996869564 0.002875770 -0.079011612 -0.427841466 -0.002861131 0.999995887 0.000298484 -0.005788880 0.079012141 -0.000071487 0.996873677 0.132732609
|
||||
123999000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.993462563 0.003229393 -0.114112593 -0.472377562 -0.003208589 0.999994814 0.000365978 -0.005932507 0.114113182 0.000002555 0.993467748 0.123959606
|
||||
124332000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.988605380 0.003602870 -0.150487319 -0.517270184 -0.003599323 0.999993503 0.000295953 -0.005751638 0.150487408 0.000249071 0.988611877 0.113156366
|
||||
124708000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.981692851 0.004048047 -0.190427750 -0.566330350 -0.004096349 0.999991596 0.000139980 -0.007622665 0.190426722 0.000642641 0.981701195 0.098572887
|
||||
125041000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.974759340 0.004326052 -0.223216295 -0.606091424 -0.004403458 0.999990284 0.000150970 -0.009427620 0.223214790 0.000835764 0.974768937 0.084984909
|
||||
125417000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.965238512 0.004419941 -0.261333257 -0.651601078 -0.004571608 0.999989569 0.000027561 -0.007437027 0.261330664 0.001168111 0.965248644 0.068577736
|
||||
125750000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.953956902 0.004390486 -0.299911648 -0.697081969 -0.004806366 0.999988258 -0.000648964 -0.003676960 0.299905270 0.002060569 0.953966737 0.050264043
|
||||
126126000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.940579295 0.004839818 -0.339539677 -0.744385684 -0.005527717 0.999984145 -0.001058831 -0.001820489 0.339529186 0.002872794 0.940591156 0.028560147
|
||||
126459000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.928297341 0.004980532 -0.371805429 -0.781716025 -0.005848793 0.999982178 -0.001207554 -0.001832299 0.371792793 0.003295582 0.928309917 0.009470658
|
||||
126835000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.913324535 0.005156573 -0.407199889 -0.824074795 -0.006227055 0.999979734 -0.001303667 -0.001894351 0.407184929 0.003726327 0.913338125 -0.013179829
|
||||
127168000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.898822486 0.005400294 -0.438279599 -0.860775204 -0.006702366 0.999976516 -0.001423908 -0.001209170 0.438261628 0.004217350 0.898837566 -0.034594674
|
||||
127544000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.880397439 0.005455900 -0.474205226 -0.903308447 -0.007032821 0.999974072 -0.001551900 -0.000798134 0.474184483 0.004701289 0.880412936 -0.061250069
|
||||
127877000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.862660766 0.005402398 -0.505754173 -0.939888304 -0.007276668 0.999972045 -0.001730187 -0.000489221 0.505730629 0.005172769 0.862675905 -0.086411685
|
||||
128253000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.841714203 0.005442667 -0.539895892 -0.978630821 -0.007698633 0.999968529 -0.001921765 0.000975953 0.539868414 0.005774037 0.841729641 -0.115983579
|
||||
128587000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.823229551 0.005282366 -0.567684054 -1.010071242 -0.007977336 0.999965608 -0.002263572 0.002284809 0.567652583 0.006392045 0.823243380 -0.141444392
|
||||
128962000 0.591609280 1.051749871 0.500000000 0.500000000 0.000000000 0.000000000 0.802855015 0.005112482 -0.596152425 -1.042319682 -0.008217614 0.999963105 -0.002491409 0.003637235 0.596117735 0.006899191 0.802867413 -0.169369454
|
||||
|
After Width: | Height: | Size: 25 KiB |
|
After Width: | Height: | Size: 40 KiB |
|
After Width: | Height: | Size: 36 KiB |
|
After Width: | Height: | Size: 26 KiB |
|
After Width: | Height: | Size: 34 KiB |
|
After Width: | Height: | Size: 36 KiB |
|
After Width: | Height: | Size: 22 KiB |
|
After Width: | Height: | Size: 40 KiB |
|
After Width: | Height: | Size: 18 KiB |
|
After Width: | Height: | Size: 45 KiB |
|
After Width: | Height: | Size: 30 KiB |
@@ -0,0 +1,348 @@
|
||||
import os
|
||||
import random
|
||||
import json
|
||||
import torch
|
||||
|
||||
import torch.nn as nn
|
||||
import torchvision.transforms as transforms
|
||||
import torchvision.transforms.functional as F
|
||||
import numpy as np
|
||||
|
||||
from decord import VideoReader
|
||||
from torch.utils.data.dataset import Dataset
|
||||
from packaging import version as pver
|
||||
|
||||
|
||||
class RandomHorizontalFlipWithPose(nn.Module):
|
||||
def __init__(self, p=0.5):
|
||||
super(RandomHorizontalFlipWithPose, self).__init__()
|
||||
self.p = p
|
||||
|
||||
def get_flip_flag(self, n_image):
|
||||
return torch.rand(n_image) < self.p
|
||||
|
||||
def forward(self, image, flip_flag=None):
|
||||
n_image = image.shape[0]
|
||||
if flip_flag is not None:
|
||||
assert n_image == flip_flag.shape[0]
|
||||
else:
|
||||
flip_flag = self.get_flip_flag(n_image)
|
||||
|
||||
ret_images = []
|
||||
for fflag, img in zip(flip_flag, image):
|
||||
if fflag:
|
||||
ret_images.append(F.hflip(img))
|
||||
else:
|
||||
ret_images.append(img)
|
||||
return torch.stack(ret_images, dim=0)
|
||||
|
||||
|
||||
class Camera(object):
|
||||
def __init__(self, entry):
|
||||
fx, fy, cx, cy = entry[1:5]
|
||||
self.fx = fx
|
||||
self.fy = fy
|
||||
self.cx = cx
|
||||
self.cy = cy
|
||||
w2c_mat = np.array(entry[7:]).reshape(3, 4)
|
||||
w2c_mat_4x4 = np.eye(4)
|
||||
w2c_mat_4x4[:3, :] = w2c_mat
|
||||
self.w2c_mat = w2c_mat_4x4
|
||||
self.c2w_mat = np.linalg.inv(w2c_mat_4x4)
|
||||
|
||||
|
||||
def custom_meshgrid(*args):
|
||||
# ref: https://pytorch.org/docs/stable/generated/torch.meshgrid.html?highlight=meshgrid#torch.meshgrid
|
||||
if pver.parse(torch.__version__) < pver.parse('1.10'):
|
||||
return torch.meshgrid(*args)
|
||||
else:
|
||||
return torch.meshgrid(*args, indexing='ij')
|
||||
|
||||
|
||||
def ray_condition(K, c2w, H, W, device, flip_flag=None):
|
||||
# c2w: B, V, 4, 4
|
||||
# K: B, V, 4
|
||||
|
||||
B, V = K.shape[:2]
|
||||
|
||||
j, i = custom_meshgrid(
|
||||
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
|
||||
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
|
||||
)
|
||||
i = i.reshape([1, 1, H * W]).expand([B, V, H * W]) + 0.5 # [B, V, HxW]
|
||||
j = j.reshape([1, 1, H * W]).expand([B, V, H * W]) + 0.5 # [B, V, HxW]
|
||||
|
||||
n_flip = torch.sum(flip_flag).item() if flip_flag is not None else 0
|
||||
if n_flip > 0:
|
||||
j_flip, i_flip = custom_meshgrid(
|
||||
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
|
||||
torch.linspace(W - 1, 0, W, device=device, dtype=c2w.dtype)
|
||||
)
|
||||
i_flip = i_flip.reshape([1, 1, H * W]).expand(B, 1, H * W) + 0.5
|
||||
j_flip = j_flip.reshape([1, 1, H * W]).expand(B, 1, H * W) + 0.5
|
||||
i[:, flip_flag, ...] = i_flip
|
||||
j[:, flip_flag, ...] = j_flip
|
||||
|
||||
fx, fy, cx, cy = K.chunk(4, dim=-1) # B,V, 1
|
||||
|
||||
zs = torch.ones_like(i) # [B, V, HxW]
|
||||
xs = (i - cx) / fx * zs
|
||||
ys = (j - cy) / fy * zs
|
||||
zs = zs.expand_as(ys)
|
||||
|
||||
directions = torch.stack((xs, ys, zs), dim=-1) # B, V, HW, 3
|
||||
directions = directions / directions.norm(dim=-1, keepdim=True) # B, V, HW, 3
|
||||
|
||||
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2) # B, V, HW, 3
|
||||
rays_o = c2w[..., :3, 3] # B, V, 3
|
||||
rays_o = rays_o[:, :, None].expand_as(rays_d) # B, V, HW, 3
|
||||
# c2w @ dirctions
|
||||
rays_dxo = torch.cross(rays_o, rays_d) # B, V, HW, 3
|
||||
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
|
||||
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6) # B, V, H, W, 6
|
||||
# plucker = plucker.permute(0, 1, 4, 2, 3)
|
||||
return plucker
|
||||
|
||||
|
||||
class RealEstate10K(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
root_path,
|
||||
annotation_json,
|
||||
sample_stride=4,
|
||||
sample_n_frames=16,
|
||||
sample_size=[256, 384],
|
||||
is_image=False,
|
||||
):
|
||||
self.root_path = root_path
|
||||
self.sample_stride = sample_stride
|
||||
self.sample_n_frames = sample_n_frames
|
||||
self.is_image = is_image
|
||||
|
||||
self.dataset = json.load(open(os.path.join(root_path, annotation_json), 'r'))
|
||||
self.length = len(self.dataset)
|
||||
|
||||
sample_size = tuple(sample_size) if not isinstance(sample_size, int) else (sample_size, sample_size)
|
||||
pixel_transforms = [transforms.Resize(sample_size),
|
||||
transforms.RandomHorizontalFlip(),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)]
|
||||
|
||||
self.pixel_transforms = transforms.Compose(pixel_transforms)
|
||||
|
||||
def load_video_reader(self, idx):
|
||||
video_dict = self.dataset[idx]
|
||||
|
||||
video_path = os.path.join(self.root_path, video_dict['clip_path'])
|
||||
video_reader = VideoReader(video_path)
|
||||
return video_reader, video_dict['caption']
|
||||
|
||||
def get_batch(self, idx):
|
||||
video_reader, video_caption = self.load_video_reader(idx)
|
||||
total_frames = len(video_reader)
|
||||
|
||||
if self.is_image:
|
||||
frame_indice = [random.randint(0, total_frames - 1)]
|
||||
else:
|
||||
if isinstance(self.sample_stride, int):
|
||||
current_sample_stride = self.sample_stride
|
||||
else:
|
||||
assert len(self.sample_stride) == 2
|
||||
assert (self.sample_stride[0] >= 1) and (self.sample_stride[1] >= self.sample_stride[0])
|
||||
current_sample_stride = random.randint(self.sample_stride[0], self.sample_stride[1])
|
||||
|
||||
cropped_length = self.sample_n_frames * current_sample_stride
|
||||
start_frame_ind = random.randint(0, max(0, total_frames - cropped_length - 1))
|
||||
end_frame_ind = min(start_frame_ind + cropped_length, total_frames)
|
||||
|
||||
assert end_frame_ind - start_frame_ind >= self.sample_n_frames
|
||||
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, self.sample_n_frames, dtype=int)
|
||||
|
||||
pixel_values = torch.from_numpy(video_reader.get_batch(frame_indice).asnumpy()).permute(0, 3, 1, 2).contiguous()
|
||||
pixel_values = pixel_values / 255.
|
||||
|
||||
if self.is_image:
|
||||
pixel_values = pixel_values[0]
|
||||
|
||||
return pixel_values, video_caption
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
while True:
|
||||
try:
|
||||
video, video_caption = self.get_batch(idx)
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
idx = random.randint(0, self.length - 1)
|
||||
|
||||
video = self.pixel_transforms(video)
|
||||
sample = dict(pixel_values=video, caption=video_caption)
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
class RealEstate10KPose(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
root_path,
|
||||
annotation_json,
|
||||
sample_stride=4,
|
||||
minimum_sample_stride=1,
|
||||
sample_n_frames=16,
|
||||
relative_pose=False,
|
||||
zero_t_first_frame=False,
|
||||
sample_size=[256, 384],
|
||||
rescale_fxy=False,
|
||||
shuffle_frames=False,
|
||||
use_flip=False,
|
||||
return_clip_name=False,
|
||||
):
|
||||
self.root_path = root_path
|
||||
self.relative_pose = relative_pose
|
||||
self.zero_t_first_frame = zero_t_first_frame
|
||||
self.sample_stride = sample_stride
|
||||
self.minimum_sample_stride = minimum_sample_stride
|
||||
self.sample_n_frames = sample_n_frames
|
||||
self.return_clip_name = return_clip_name
|
||||
|
||||
self.dataset = json.load(open(os.path.join(root_path, annotation_json), 'r'))
|
||||
self.length = len(self.dataset)
|
||||
|
||||
sample_size = tuple(sample_size) if not isinstance(sample_size, int) else (sample_size, sample_size)
|
||||
self.sample_size = sample_size
|
||||
if use_flip:
|
||||
pixel_transforms = [transforms.Resize(sample_size),
|
||||
RandomHorizontalFlipWithPose(),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)]
|
||||
else:
|
||||
pixel_transforms = [transforms.Resize(sample_size),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)]
|
||||
self.rescale_fxy = rescale_fxy
|
||||
self.sample_wh_ratio = sample_size[1] / sample_size[0]
|
||||
|
||||
self.pixel_transforms = pixel_transforms
|
||||
self.shuffle_frames = shuffle_frames
|
||||
self.use_flip = use_flip
|
||||
|
||||
def get_relative_pose(self, cam_params):
|
||||
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
|
||||
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
|
||||
source_cam_c2w = abs_c2ws[0]
|
||||
if self.zero_t_first_frame:
|
||||
cam_to_origin = 0
|
||||
else:
|
||||
cam_to_origin = np.linalg.norm(source_cam_c2w[:3, 3])
|
||||
target_cam_c2w = np.array([
|
||||
[1, 0, 0, 0],
|
||||
[0, 1, 0, -cam_to_origin],
|
||||
[0, 0, 1, 0],
|
||||
[0, 0, 0, 1]
|
||||
])
|
||||
abs2rel = target_cam_c2w @ abs_w2cs[0]
|
||||
ret_poses = [target_cam_c2w, ] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
|
||||
ret_poses = np.array(ret_poses, dtype=np.float32)
|
||||
return ret_poses
|
||||
|
||||
def load_video_reader(self, idx):
|
||||
video_dict = self.dataset[idx]
|
||||
|
||||
video_path = os.path.join(self.root_path, video_dict['clip_path'])
|
||||
video_reader = VideoReader(video_path)
|
||||
return video_dict['clip_name'], video_reader, video_dict['caption']
|
||||
|
||||
def load_cameras(self, idx):
|
||||
video_dict = self.dataset[idx]
|
||||
pose_file = os.path.join(self.root_path, video_dict['pose_file'])
|
||||
with open(pose_file, 'r') as f:
|
||||
poses = f.readlines()
|
||||
poses = [pose.strip().split(' ') for pose in poses[1:]]
|
||||
cam_params = [[float(x) for x in pose] for pose in poses]
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
return cam_params
|
||||
|
||||
def get_batch(self, idx):
|
||||
clip_name, video_reader, video_caption = self.load_video_reader(idx)
|
||||
cam_params = self.load_cameras(idx)
|
||||
assert len(cam_params) >= self.sample_n_frames
|
||||
total_frames = len(cam_params)
|
||||
|
||||
current_sample_stride = self.sample_stride
|
||||
|
||||
if total_frames < self.sample_n_frames * current_sample_stride:
|
||||
maximum_sample_stride = int(total_frames // self.sample_n_frames)
|
||||
current_sample_stride = random.randint(self.minimum_sample_stride, maximum_sample_stride)
|
||||
|
||||
cropped_length = self.sample_n_frames * current_sample_stride
|
||||
start_frame_ind = random.randint(0, max(0, total_frames - cropped_length - 1))
|
||||
end_frame_ind = min(start_frame_ind + cropped_length, total_frames)
|
||||
|
||||
assert end_frame_ind - start_frame_ind >= self.sample_n_frames
|
||||
frame_indices = np.linspace(start_frame_ind, end_frame_ind - 1, self.sample_n_frames, dtype=int)
|
||||
|
||||
if self.shuffle_frames:
|
||||
perm = np.random.permutation(self.sample_n_frames)
|
||||
frame_indices = frame_indices[perm]
|
||||
|
||||
pixel_values = torch.from_numpy(video_reader.get_batch(frame_indices).asnumpy()).permute(0, 3, 1, 2).contiguous()
|
||||
pixel_values = pixel_values / 255.
|
||||
|
||||
cam_params = [cam_params[indice] for indice in frame_indices]
|
||||
if self.rescale_fxy:
|
||||
ori_h, ori_w = pixel_values.shape[-2:]
|
||||
ori_wh_ratio = ori_w / ori_h
|
||||
if ori_wh_ratio > self.sample_wh_ratio: # rescale fx
|
||||
resized_ori_w = self.sample_size[0] * ori_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / self.sample_size[1]
|
||||
else: # rescale fy
|
||||
resized_ori_h = self.sample_size[1] / ori_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / self.sample_size[0]
|
||||
intrinsics = np.asarray([[cam_param.fx * self.sample_size[1],
|
||||
cam_param.fy * self.sample_size[0],
|
||||
cam_param.cx * self.sample_size[1],
|
||||
cam_param.cy * self.sample_size[0]]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
intrinsics = torch.as_tensor(intrinsics)[None] # [1, n_frame, 4]
|
||||
if self.relative_pose:
|
||||
c2w_poses = self.get_relative_pose(cam_params)
|
||||
else:
|
||||
c2w_poses = np.array([cam_param.c2w_mat for cam_param in cam_params], dtype=np.float32)
|
||||
c2w = torch.as_tensor(c2w_poses)[None] # [1, n_frame, 4, 4]
|
||||
if self.use_flip:
|
||||
flip_flag = self.pixel_transforms[1].get_flip_flag(self.sample_n_frames)
|
||||
else:
|
||||
flip_flag = torch.zeros(self.sample_n_frames, dtype=torch.bool, device=c2w.device)
|
||||
plucker_embedding = ray_condition(intrinsics, c2w, self.sample_size[0], self.sample_size[1], device='cpu',
|
||||
flip_flag=flip_flag)[0].permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
return pixel_values, video_caption, plucker_embedding, flip_flag, clip_name
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
while True:
|
||||
try:
|
||||
video, video_caption, plucker_embedding, flip_flag, clip_name = self.get_batch(idx)
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
idx = random.randint(0, self.length - 1)
|
||||
|
||||
if self.use_flip:
|
||||
video = self.pixel_transforms[0](video)
|
||||
video = self.pixel_transforms[1](video, flip_flag)
|
||||
video = self.pixel_transforms[2](video)
|
||||
else:
|
||||
for transform in self.pixel_transforms:
|
||||
video = transform(video)
|
||||
if self.return_clip_name:
|
||||
sample = dict(pixel_values=video, text=video_caption, plucker_embedding=plucker_embedding, clip_name=clip_name)
|
||||
else:
|
||||
sample = dict(pixel_values=video, text=video_caption, plucker_embedding=plucker_embedding)
|
||||
|
||||
return sample
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention.py
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.models.attention import BasicTransformerBlock
|
||||
from einops import rearrange, repeat
|
||||
|
||||
|
||||
@dataclass
|
||||
class Transformer3DModelOutput(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
class Transformer3DModel(ModelMixin, ConfigMixin):
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
dropout: float = 0.0,
|
||||
norm_num_groups: int = 32,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
activation_fn: str = "geglu",
|
||||
num_embeds_ada_norm: Optional[int] = None,
|
||||
use_linear_projection: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
norm_type: str = "layer_norm",
|
||||
norm_elementwise_affine: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self.use_linear_projection = use_linear_projection
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
# Define input layers
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
if use_linear_projection:
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
else:
|
||||
self.proj_in = nn.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
# Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout=dropout,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||
attention_bias=attention_bias,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
norm_type=norm_type,
|
||||
norm_elementwise_affine=norm_elementwise_affine,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Define output layers
|
||||
if use_linear_projection:
|
||||
self.proj_out = nn.Linear(in_channels, inner_dim)
|
||||
else:
|
||||
self.proj_out = nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, return_dict: bool = True):
|
||||
# Input
|
||||
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||
batch_size, _, video_length = hidden_states.shape[:3]
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
|
||||
if encoder_hidden_states.shape[0] == batch_size:
|
||||
encoder_hidden_states = repeat(encoder_hidden_states, 'b n c -> (b f) n c', f=video_length)
|
||||
|
||||
elif encoder_hidden_states.shape[0] == batch_size * video_length:
|
||||
pass
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
batch, channel, height, weight = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
if not self.use_linear_projection:
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||
else:
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
|
||||
# Output
|
||||
if not self.use_linear_projection:
|
||||
hidden_states = (
|
||||
hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
else:
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = (
|
||||
hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
)
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return Transformer3DModelOutput(sample=output)
|
||||
@@ -0,0 +1,412 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.init as init
|
||||
import logging
|
||||
from diffusers.models.lora import LoRALinearLayer
|
||||
from diffusers.models.attention import Attention
|
||||
from diffusers.utils import USE_PEFT_BACKEND
|
||||
from typing import Optional
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AttnProcessor:
|
||||
r"""
|
||||
Default processor for performing attention-related computations.
|
||||
"""
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: Optional[torch.FloatTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
temb: Optional[torch.FloatTensor] = None,
|
||||
scale: float = 1.0,
|
||||
pose_feature=None
|
||||
) -> torch.Tensor:
|
||||
residual = hidden_states
|
||||
|
||||
args = () if USE_PEFT_BACKEND else (scale,)
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states, *args)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states, *args)
|
||||
value = attn.to_v(encoder_hidden_states, *args)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states, *args)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class LoRAAttnProcessor(nn.Module):
|
||||
r"""
|
||||
Default processor for performing attention-related computations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size=None,
|
||||
cross_attention_dim=None,
|
||||
rank=4,
|
||||
network_alpha=None,
|
||||
lora_scale=1.0,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.rank = rank
|
||||
self.lora_scale = lora_scale
|
||||
|
||||
self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
self.to_k_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
self.to_v_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
self.to_out_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
pose_feature=None,
|
||||
scale=None
|
||||
):
|
||||
lora_scale = self.lora_scale if scale is None else scale
|
||||
residual = hidden_states
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states) + lora_scale * self.to_q_lora(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states) + lora_scale * self.to_k_lora(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states) + lora_scale * self.to_v_lora(encoder_hidden_states)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states) + lora_scale * self.to_out_lora(hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class PoseAdaptorAttnProcessor(nn.Module):
|
||||
def __init__(self,
|
||||
hidden_size, # dimension of hidden state
|
||||
pose_feature_dim=None, # dimension of the pose feature
|
||||
cross_attention_dim=None, # dimension of the text embedding
|
||||
query_condition=False,
|
||||
key_value_condition=False,
|
||||
scale=1.0):
|
||||
super().__init__()
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.pose_feature_dim = pose_feature_dim
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
self.query_condition = query_condition
|
||||
self.key_value_condition = key_value_condition
|
||||
assert hidden_size == pose_feature_dim
|
||||
if self.query_condition and self.key_value_condition:
|
||||
self.qkv_merge = nn.Linear(hidden_size, hidden_size)
|
||||
init.zeros_(self.qkv_merge.weight)
|
||||
init.zeros_(self.qkv_merge.bias)
|
||||
elif self.query_condition:
|
||||
self.q_merge = nn.Linear(hidden_size, hidden_size)
|
||||
init.zeros_(self.q_merge.weight)
|
||||
init.zeros_(self.q_merge.bias)
|
||||
else:
|
||||
self.kv_merge = nn.Linear(hidden_size, hidden_size)
|
||||
init.zeros_(self.kv_merge.weight)
|
||||
init.zeros_(self.kv_merge.bias)
|
||||
|
||||
def forward(self,
|
||||
attn,
|
||||
hidden_states,
|
||||
pose_feature,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
scale=None,):
|
||||
assert pose_feature is not None
|
||||
pose_embedding_scale = (scale or self.scale)
|
||||
|
||||
residual = hidden_states
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
if hidden_states.dim == 5:
|
||||
hidden_states = rearrange(hidden_states, 'b c f h w -> (b f) (h w) c')
|
||||
elif hidden_states.ndim == 4:
|
||||
hidden_states = rearrange(hidden_states, 'b c h w -> b (h w) c')
|
||||
else:
|
||||
assert hidden_states.ndim == 3
|
||||
|
||||
if self.query_condition and self.key_value_condition:
|
||||
assert encoder_hidden_states is None
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
|
||||
if encoder_hidden_states.ndim == 5:
|
||||
encoder_hidden_states = rearrange(encoder_hidden_states, 'b c f h w -> (b f) (h w) c')
|
||||
elif encoder_hidden_states.ndim == 4:
|
||||
encoder_hidden_states = rearrange(encoder_hidden_states, 'b c h w -> b (h w) c')
|
||||
else:
|
||||
assert encoder_hidden_states.ndim == 3
|
||||
if pose_feature.ndim == 5:
|
||||
pose_feature = rearrange(pose_feature, "b c f h w -> (b f) (h w) c")
|
||||
elif pose_feature.ndim == 4:
|
||||
pose_feature = rearrange(pose_feature, "b c h w -> b (h w) c")
|
||||
else:
|
||||
assert pose_feature.ndim == 3
|
||||
|
||||
batch_size, ehs_sequence_length, _ = encoder_hidden_states.shape
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, ehs_sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
if attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
if self.query_condition and self.key_value_condition: # only self attention
|
||||
query_hidden_state = self.qkv_merge(hidden_states + pose_feature) * pose_embedding_scale + hidden_states
|
||||
key_value_hidden_state = query_hidden_state
|
||||
elif self.query_condition:
|
||||
query_hidden_state = self.q_merge(hidden_states + pose_feature) * pose_embedding_scale + hidden_states
|
||||
key_value_hidden_state = encoder_hidden_states
|
||||
else:
|
||||
key_value_hidden_state = self.kv_merge(encoder_hidden_states + pose_feature) * pose_embedding_scale + encoder_hidden_states
|
||||
query_hidden_state = hidden_states
|
||||
|
||||
# original attention
|
||||
query = attn.to_q(query_hidden_state)
|
||||
key = attn.to_k(key_value_hidden_state)
|
||||
value = attn.to_v(key_value_hidden_state)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class LORAPoseAdaptorAttnProcessor(nn.Module):
|
||||
def __init__(self,
|
||||
hidden_size, # dimension of hidden state
|
||||
pose_feature_dim=None, # dimension of the pose feature
|
||||
cross_attention_dim=None, # dimension of the text embedding
|
||||
query_condition=False,
|
||||
key_value_condition=False,
|
||||
scale=1.0,
|
||||
# lora keywords
|
||||
rank=4,
|
||||
network_alpha=None,
|
||||
lora_scale=1.0):
|
||||
super().__init__()
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.pose_feature_dim = pose_feature_dim
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
self.query_condition = query_condition
|
||||
self.key_value_condition = key_value_condition
|
||||
assert hidden_size == pose_feature_dim
|
||||
if self.query_condition and self.key_value_condition:
|
||||
self.qkv_merge = nn.Linear(hidden_size, hidden_size)
|
||||
init.zeros_(self.qkv_merge.weight)
|
||||
init.zeros_(self.qkv_merge.bias)
|
||||
elif self.query_condition:
|
||||
self.q_merge = nn.Linear(hidden_size, hidden_size)
|
||||
init.zeros_(self.q_merge.weight)
|
||||
init.zeros_(self.q_merge.bias)
|
||||
else:
|
||||
self.kv_merge = nn.Linear(hidden_size, hidden_size)
|
||||
init.zeros_(self.kv_merge.weight)
|
||||
init.zeros_(self.kv_merge.bias)
|
||||
# lora
|
||||
self.rank = rank
|
||||
self.lora_scale = lora_scale
|
||||
self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
self.to_k_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
self.to_v_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
self.to_out_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
|
||||
def __call__(self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
scale=1.0,
|
||||
pose_feature=None,
|
||||
):
|
||||
assert pose_feature is not None
|
||||
lora_scale = self.lora_scale if scale is None else scale
|
||||
residual = hidden_states
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
if hidden_states.dim == 5:
|
||||
hidden_states = rearrange(hidden_states, 'b c f h w -> (b f) (h w) c')
|
||||
elif hidden_states.ndim == 4:
|
||||
hidden_states = rearrange(hidden_states, 'b c h w -> b (h w) c')
|
||||
else:
|
||||
assert hidden_states.ndim == 3
|
||||
|
||||
if self.query_condition and self.key_value_condition:
|
||||
assert encoder_hidden_states is None
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
|
||||
if encoder_hidden_states.ndim == 5:
|
||||
encoder_hidden_states = rearrange(encoder_hidden_states, 'b c f h w -> (b f) (h w) c')
|
||||
elif encoder_hidden_states.ndim == 4:
|
||||
encoder_hidden_states = rearrange(encoder_hidden_states, 'b c h w -> b (h w) c')
|
||||
else:
|
||||
assert encoder_hidden_states.ndim == 3
|
||||
if pose_feature.ndim == 5:
|
||||
pose_feature = rearrange(pose_feature, "b c f h w -> (b f) (h w) c")
|
||||
elif pose_feature.ndim == 4:
|
||||
pose_feature = rearrange(pose_feature, "b c h w -> b (h w) c")
|
||||
else:
|
||||
assert pose_feature.ndim == 3
|
||||
|
||||
batch_size, ehs_sequence_length, _ = encoder_hidden_states.shape
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, ehs_sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
if attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
if self.query_condition and self.key_value_condition: # only self attention
|
||||
query_hidden_state = self.qkv_merge(hidden_states + pose_feature) * self.scale + hidden_states
|
||||
key_value_hidden_state = query_hidden_state
|
||||
elif self.query_condition:
|
||||
query_hidden_state = self.q_merge(hidden_states + pose_feature) * self.scale + hidden_states
|
||||
key_value_hidden_state = encoder_hidden_states
|
||||
else:
|
||||
key_value_hidden_state = self.kv_merge(encoder_hidden_states + pose_feature) * self.scale + encoder_hidden_states
|
||||
query_hidden_state = hidden_states
|
||||
|
||||
# original attention
|
||||
query = attn.to_q(query_hidden_state) + lora_scale * self.to_q_lora(query_hidden_state)
|
||||
key = attn.to_k(key_value_hidden_state) + lora_scale * self.to_k_lora(key_value_hidden_state)
|
||||
value = attn.to_v(key_value_hidden_state) + lora_scale * self.to_v_lora(key_value_hidden_state)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states) + lora_scale * self.to_out_lora(hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,389 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.attention import FeedForward
|
||||
|
||||
from typing import Dict, Any
|
||||
from cameractrl.models.resnet import InflatedGroupNorm
|
||||
from cameractrl.models.attention_processor import PoseAdaptorAttnProcessor
|
||||
|
||||
from einops import rearrange
|
||||
import math
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
# Zero out the parameters of a module and return it.
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
@dataclass
|
||||
class TemporalTransformer3DModelOutput(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
def get_motion_module(
|
||||
in_channels,
|
||||
motion_module_type: str,
|
||||
motion_module_kwargs: dict
|
||||
):
|
||||
if motion_module_type == "Vanilla":
|
||||
return VanillaTemporalModule(in_channels=in_channels, **motion_module_kwargs)
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
|
||||
class VanillaTemporalModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads=8,
|
||||
num_transformer_block=2,
|
||||
attention_block_types=("Temporal_Self",),
|
||||
temporal_position_encoding=True,
|
||||
temporal_position_encoding_max_len=32,
|
||||
temporal_attention_dim_div=1,
|
||||
cross_attention_dim=320,
|
||||
zero_initialize=True,
|
||||
encoder_hidden_states_query=(False, False),
|
||||
attention_activation_scale=1.0,
|
||||
attention_processor_kwargs: Dict = {},
|
||||
causal_temporal_attention=False,
|
||||
causal_temporal_attention_mask_type="",
|
||||
rescale_output_factor=1.0
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.temporal_transformer = TemporalTransformer3DModel(
|
||||
in_channels=in_channels,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=in_channels // num_attention_heads // temporal_attention_dim_div,
|
||||
num_layers=num_transformer_block,
|
||||
attention_block_types=attention_block_types,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
encoder_hidden_states_query=encoder_hidden_states_query,
|
||||
attention_activation_scale=attention_activation_scale,
|
||||
attention_processor_kwargs=attention_processor_kwargs,
|
||||
causal_temporal_attention=causal_temporal_attention,
|
||||
causal_temporal_attention_mask_type=causal_temporal_attention_mask_type,
|
||||
rescale_output_factor=rescale_output_factor
|
||||
)
|
||||
|
||||
if zero_initialize:
|
||||
self.temporal_transformer.proj_out = zero_module(self.temporal_transformer.proj_out)
|
||||
|
||||
def forward(self, hidden_states, temb=None, encoder_hidden_states=None, attention_mask=None,
|
||||
cross_attention_kwargs: Dict[str, Any] = {}):
|
||||
hidden_states = self.temporal_transformer(hidden_states, encoder_hidden_states, attention_mask, cross_attention_kwargs=cross_attention_kwargs)
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformer3DModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
num_layers,
|
||||
attention_block_types=("Temporal_Self", "Temporal_Self",),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=320,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=32,
|
||||
encoder_hidden_states_query=(False, False),
|
||||
attention_activation_scale=1.0,
|
||||
attention_processor_kwargs: Dict = {},
|
||||
|
||||
causal_temporal_attention=None,
|
||||
causal_temporal_attention_mask_type="",
|
||||
rescale_output_factor=1.0
|
||||
):
|
||||
super().__init__()
|
||||
assert causal_temporal_attention is not None
|
||||
self.causal_temporal_attention = causal_temporal_attention
|
||||
|
||||
assert (not causal_temporal_attention) or (causal_temporal_attention_mask_type != "")
|
||||
self.causal_temporal_attention_mask_type = causal_temporal_attention_mask_type
|
||||
self.causal_temporal_attention_mask = None
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm = InflatedGroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
TemporalTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
attention_block_types=attention_block_types,
|
||||
dropout=dropout,
|
||||
norm_num_groups=norm_num_groups,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
encoder_hidden_states_query=encoder_hidden_states_query,
|
||||
attention_activation_scale=attention_activation_scale,
|
||||
attention_processor_kwargs=attention_processor_kwargs,
|
||||
rescale_output_factor=rescale_output_factor,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
def get_causal_temporal_attention_mask(self, hidden_states):
|
||||
batch_size, sequence_length, dim = hidden_states.shape
|
||||
|
||||
if self.causal_temporal_attention_mask is None or self.causal_temporal_attention_mask.shape != (
|
||||
batch_size, sequence_length, sequence_length):
|
||||
if self.causal_temporal_attention_mask_type == "causal":
|
||||
# 1. vanilla causal mask
|
||||
mask = torch.tril(torch.ones(sequence_length, sequence_length))
|
||||
|
||||
elif self.causal_temporal_attention_mask_type == "2-seq":
|
||||
# 2. 2-seq
|
||||
mask = torch.zeros(sequence_length, sequence_length)
|
||||
mask[:sequence_length // 2, :sequence_length // 2] = 1
|
||||
mask[-sequence_length // 2:, -sequence_length // 2:] = 1
|
||||
|
||||
elif self.causal_temporal_attention_mask_type == "0-prev":
|
||||
# attn to the previous frame
|
||||
indices = torch.arange(sequence_length)
|
||||
indices_prev = indices - 1
|
||||
indices_prev[0] = 0
|
||||
mask = torch.zeros(sequence_length, sequence_length)
|
||||
mask[:, 0] = 1.
|
||||
mask[indices, indices_prev] = 1.
|
||||
|
||||
elif self.causal_temporal_attention_mask_type == "0":
|
||||
# only attn to first frame
|
||||
mask = torch.zeros(sequence_length, sequence_length)
|
||||
mask[:, 0] = 1
|
||||
|
||||
elif self.causal_temporal_attention_mask_type == "wo-self":
|
||||
indices = torch.arange(sequence_length)
|
||||
mask = torch.ones(sequence_length, sequence_length)
|
||||
mask[indices, indices] = 0
|
||||
|
||||
elif self.causal_temporal_attention_mask_type == "circle":
|
||||
indices = torch.arange(sequence_length)
|
||||
indices_prev = indices - 1
|
||||
indices_prev[0] = 0
|
||||
|
||||
mask = torch.eye(sequence_length)
|
||||
mask[indices, indices_prev] = 1
|
||||
mask[0, -1] = 1
|
||||
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
# generate attention mask fron binary values
|
||||
mask = mask.masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))
|
||||
mask = mask.unsqueeze(0)
|
||||
mask = mask.repeat(batch_size, 1, 1)
|
||||
|
||||
self.causal_temporal_attention_mask = mask.to(hidden_states.device)
|
||||
|
||||
return self.causal_temporal_attention_mask
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None,
|
||||
cross_attention_kwargs: Dict[str, Any] = {},):
|
||||
residual = hidden_states
|
||||
|
||||
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||
height, width = hidden_states.shape[-2:]
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b h w) f c")
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
attention_mask = self.get_causal_temporal_attention_mask(
|
||||
hidden_states) if self.causal_temporal_attention else attention_mask
|
||||
|
||||
# Transformer Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask, cross_attention_kwargs=cross_attention_kwargs)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = rearrange(hidden_states, "(b h w) f c -> b c f h w", h=height, w=width)
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
attention_block_types=("Temporal_Self", "Temporal_Self",),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=32,
|
||||
encoder_hidden_states_query=(False, False),
|
||||
attention_activation_scale=1.0,
|
||||
attention_processor_kwargs: Dict = {},
|
||||
rescale_output_factor=1.0
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
attention_blocks = []
|
||||
norms = []
|
||||
self.attention_block_types = attention_block_types
|
||||
|
||||
for block_idx, block_name in enumerate(attention_block_types):
|
||||
attention_blocks.append(
|
||||
TemporalSelfAttention(
|
||||
attention_mode=block_name,
|
||||
cross_attention_dim=cross_attention_dim if block_name in ['Temporal_Cross', 'Temporal_Pose_Adaptor'] else None,
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
rescale_output_factor=rescale_output_factor,
|
||||
)
|
||||
)
|
||||
norms.append(nn.LayerNorm(dim))
|
||||
|
||||
self.attention_blocks = nn.ModuleList(attention_blocks)
|
||||
self.norms = nn.ModuleList(norms)
|
||||
|
||||
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||
self.ff_norm = nn.LayerNorm(dim)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, cross_attention_kwargs: Dict[str, Any] = {}):
|
||||
for attention_block, norm, attention_block_type in zip(self.attention_blocks, self.norms, self.attention_block_types):
|
||||
norm_hidden_states = norm(hidden_states)
|
||||
hidden_states = attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=norm_hidden_states if attention_block_type == 'Temporal_Self' else encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
**cross_attention_kwargs
|
||||
) + hidden_states
|
||||
|
||||
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model,
|
||||
dropout=0.,
|
||||
max_len=32,
|
||||
):
|
||||
super().__init__()
|
||||
self.dropout = nn.Dropout(p=dropout)
|
||||
position = torch.arange(max_len).unsqueeze(1)
|
||||
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
|
||||
pe = torch.zeros(1, max_len, d_model)
|
||||
pe[0, :, 0::2] = torch.sin(position * div_term)
|
||||
pe[0, :, 1::2] = torch.cos(position * div_term)
|
||||
self.register_buffer('pe', pe)
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.pe[:, :x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class TemporalSelfAttention(Attention):
|
||||
def __init__(
|
||||
self,
|
||||
attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=32,
|
||||
rescale_output_factor=1.0,
|
||||
*args, **kwargs
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
assert attention_mode == "Temporal_Self"
|
||||
|
||||
self.pos_encoder = PositionalEncoding(
|
||||
kwargs["query_dim"],
|
||||
max_len=temporal_position_encoding_max_len
|
||||
) if temporal_position_encoding else None
|
||||
self.rescale_output_factor = rescale_output_factor
|
||||
|
||||
def set_use_memory_efficient_attention_xformers(
|
||||
self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None
|
||||
):
|
||||
# disable motion module efficient xformers to avoid bad results, don't know why
|
||||
# TODO: fix this bug
|
||||
pass
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, **cross_attention_kwargs):
|
||||
# The `Attention` class can call different attention processors / attention functions
|
||||
# here we simply pass along all tensors to the selected processor class
|
||||
# For standard processors that are defined here, `**cross_attention_kwargs` is empty
|
||||
|
||||
# add position encoding
|
||||
if self.pos_encoder is not None:
|
||||
hidden_states = self.pos_encoder(hidden_states)
|
||||
if "pose_feature" in cross_attention_kwargs:
|
||||
pose_feature = cross_attention_kwargs["pose_feature"]
|
||||
if pose_feature.ndim == 5:
|
||||
pose_feature = rearrange(pose_feature, "b c f h w -> (b h w) f c")
|
||||
else:
|
||||
assert pose_feature.ndim == 3
|
||||
cross_attention_kwargs["pose_feature"] = pose_feature
|
||||
|
||||
if isinstance(self.processor, PoseAdaptorAttnProcessor):
|
||||
return self.processor(
|
||||
self,
|
||||
hidden_states,
|
||||
cross_attention_kwargs.pop('pose_feature'),
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=attention_mask,
|
||||
**cross_attention_kwargs,
|
||||
)
|
||||
elif hasattr(self.processor, "__call__"):
|
||||
return self.processor.__call__(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=attention_mask,
|
||||
**cross_attention_kwargs,
|
||||
)
|
||||
else:
|
||||
return self.processor(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=attention_mask,
|
||||
**cross_attention_kwargs,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from cameractrl.models.motion_module import TemporalTransformerBlock
|
||||
|
||||
|
||||
def get_parameter_dtype(parameter: torch.nn.Module):
|
||||
try:
|
||||
params = tuple(parameter.parameters())
|
||||
if len(params) > 0:
|
||||
return params[0].dtype
|
||||
|
||||
buffers = tuple(parameter.buffers())
|
||||
if len(buffers) > 0:
|
||||
return buffers[0].dtype
|
||||
|
||||
except StopIteration:
|
||||
# For torch.nn.DataParallel compatibility in PyTorch 1.5
|
||||
|
||||
def find_tensor_attributes(module: torch.nn.Module) -> List[Tuple[str, Tensor]]:
|
||||
tuples = [(k, v) for k, v in module.__dict__.items() if torch.is_tensor(v)]
|
||||
return tuples
|
||||
|
||||
gen = parameter._named_members(get_members_fn=find_tensor_attributes)
|
||||
first_tuple = next(gen)
|
||||
return first_tuple[1].dtype
|
||||
|
||||
|
||||
def conv_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D convolution module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.Conv1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return nn.Conv2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return nn.Conv3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
|
||||
|
||||
def avg_pool_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D average pooling module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.AvgPool1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return nn.AvgPool2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return nn.AvgPool3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
|
||||
|
||||
class PoseAdaptor(nn.Module):
|
||||
def __init__(self, unet, pose_encoder):
|
||||
super().__init__()
|
||||
self.unet = unet
|
||||
self.pose_encoder = pose_encoder
|
||||
|
||||
def forward(self, noisy_latents, timesteps, encoder_hidden_states, pose_embedding):
|
||||
assert pose_embedding.ndim == 5
|
||||
bs = pose_embedding.shape[0] # b c f h w
|
||||
pose_embedding_features = self.pose_encoder(pose_embedding) # bf c h w
|
||||
pose_embedding_features = [rearrange(x, '(b f) c h w -> b c f h w', b=bs)
|
||||
for x in pose_embedding_features]
|
||||
noise_pred = self.unet(noisy_latents,
|
||||
timesteps,
|
||||
encoder_hidden_states,
|
||||
pose_embedding_features=pose_embedding_features).sample
|
||||
return noise_pred
|
||||
|
||||
|
||||
class Downsample(nn.Module):
|
||||
"""
|
||||
A downsampling layer with an optional convolution.
|
||||
:param channels: channels in the inputs and outputs.
|
||||
:param use_conv: a bool determining if a convolution is applied.
|
||||
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
|
||||
downsampling occurs in the inner-two dimensions.
|
||||
"""
|
||||
|
||||
def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.dims = dims
|
||||
stride = 2 if dims != 3 else (1, 2, 2)
|
||||
if use_conv:
|
||||
self.op = conv_nd(dims, self.channels, self.out_channels, 3, stride=stride, padding=padding)
|
||||
else:
|
||||
assert self.channels == self.out_channels
|
||||
self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride)
|
||||
|
||||
def forward(self, x):
|
||||
assert x.shape[1] == self.channels
|
||||
return self.op(x)
|
||||
|
||||
|
||||
class ResnetBlock(nn.Module):
|
||||
|
||||
def __init__(self, in_c, out_c, down, ksize=3, sk=False, use_conv=True):
|
||||
super().__init__()
|
||||
ps = ksize // 2
|
||||
if in_c != out_c or sk == False:
|
||||
self.in_conv = nn.Conv2d(in_c, out_c, ksize, 1, ps)
|
||||
else:
|
||||
self.in_conv = None
|
||||
self.block1 = nn.Conv2d(out_c, out_c, 3, 1, 1)
|
||||
self.act = nn.ReLU()
|
||||
self.block2 = nn.Conv2d(out_c, out_c, ksize, 1, ps)
|
||||
if sk == False:
|
||||
self.skep = nn.Conv2d(in_c, out_c, ksize, 1, ps)
|
||||
else:
|
||||
self.skep = None
|
||||
|
||||
self.down = down
|
||||
if self.down == True:
|
||||
self.down_opt = Downsample(in_c, use_conv=use_conv)
|
||||
|
||||
def forward(self, x):
|
||||
if self.down == True:
|
||||
x = self.down_opt(x)
|
||||
if self.in_conv is not None: # edit
|
||||
x = self.in_conv(x)
|
||||
|
||||
h = self.block1(x)
|
||||
h = self.act(h)
|
||||
h = self.block2(h)
|
||||
if self.skep is not None:
|
||||
return h + self.skep(x)
|
||||
else:
|
||||
return h + x
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model,
|
||||
dropout=0.,
|
||||
max_len=32,
|
||||
):
|
||||
super().__init__()
|
||||
self.dropout = nn.Dropout(p=dropout)
|
||||
position = torch.arange(max_len).unsqueeze(1)
|
||||
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
|
||||
pe = torch.zeros(1, max_len, d_model)
|
||||
pe[0, :, 0::2, ...] = torch.sin(position * div_term)
|
||||
pe[0, :, 1::2, ...] = torch.cos(position * div_term)
|
||||
pe.unsqueeze_(-1).unsqueeze_(-1)
|
||||
self.register_buffer('pe', pe)
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.pe[:, :x.size(1), ...]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class CameraPoseEncoder(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
downscale_factor,
|
||||
channels=[320, 640, 1280, 1280],
|
||||
nums_rb=3,
|
||||
cin=64,
|
||||
ksize=3,
|
||||
sk=False,
|
||||
use_conv=True,
|
||||
compression_factor=1,
|
||||
temporal_attention_nhead=8,
|
||||
attention_block_types=("Temporal_Self", ),
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=16,
|
||||
rescale_output_factor=1.0):
|
||||
super(CameraPoseEncoder, self).__init__()
|
||||
self.unshuffle = nn.PixelUnshuffle(downscale_factor)
|
||||
self.channels = channels
|
||||
self.nums_rb = nums_rb
|
||||
self.encoder_down_conv_blocks = nn.ModuleList()
|
||||
self.encoder_down_attention_blocks = nn.ModuleList()
|
||||
for i in range(len(channels)):
|
||||
conv_layers = nn.ModuleList()
|
||||
temporal_attention_layers = nn.ModuleList()
|
||||
for j in range(nums_rb):
|
||||
if j == 0 and i != 0:
|
||||
in_dim = channels[i - 1]
|
||||
out_dim = int(channels[i] / compression_factor)
|
||||
conv_layer = ResnetBlock(in_dim, out_dim, down=True, ksize=ksize, sk=sk, use_conv=use_conv)
|
||||
elif j == 0:
|
||||
in_dim = channels[0]
|
||||
out_dim = int(channels[i] / compression_factor)
|
||||
conv_layer = ResnetBlock(in_dim, out_dim, down=False, ksize=ksize, sk=sk, use_conv=use_conv)
|
||||
elif j == nums_rb - 1:
|
||||
in_dim = channels[i] / compression_factor
|
||||
out_dim = channels[i]
|
||||
conv_layer = ResnetBlock(in_dim, out_dim, down=False, ksize=ksize, sk=sk, use_conv=use_conv)
|
||||
else:
|
||||
in_dim = int(channels[i] / compression_factor)
|
||||
out_dim = int(channels[i] / compression_factor)
|
||||
conv_layer = ResnetBlock(in_dim, out_dim, down=False, ksize=ksize, sk=sk, use_conv=use_conv)
|
||||
temporal_attention_layer = TemporalTransformerBlock(dim=out_dim,
|
||||
num_attention_heads=temporal_attention_nhead,
|
||||
attention_head_dim=int(out_dim / temporal_attention_nhead),
|
||||
attention_block_types=attention_block_types,
|
||||
dropout=0.0,
|
||||
cross_attention_dim=None,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
rescale_output_factor=rescale_output_factor)
|
||||
conv_layers.append(conv_layer)
|
||||
temporal_attention_layers.append(temporal_attention_layer)
|
||||
self.encoder_down_conv_blocks.append(conv_layers)
|
||||
self.encoder_down_attention_blocks.append(temporal_attention_layers)
|
||||
|
||||
self.encoder_conv_in = nn.Conv2d(cin, channels[0], 3, 1, 1)
|
||||
|
||||
@property
|
||||
def dtype(self) -> torch.dtype:
|
||||
"""
|
||||
`torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype).
|
||||
"""
|
||||
return get_parameter_dtype(self)
|
||||
|
||||
def forward(self, x):
|
||||
# unshuffle
|
||||
bs = x.shape[0]
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = self.unshuffle(x)
|
||||
# extract features
|
||||
features = []
|
||||
x = self.encoder_conv_in(x)
|
||||
for res_block, attention_block in zip(self.encoder_down_conv_blocks, self.encoder_down_attention_blocks):
|
||||
for res_layer, attention_layer in zip(res_block, attention_block):
|
||||
x = res_layer(x)
|
||||
h, w = x.shape[-2:]
|
||||
x = rearrange(x, '(b f) c h w -> (b h w) f c', b=bs)
|
||||
x = attention_layer(x)
|
||||
x = rearrange(x, '(b h w) f c -> (b f) c h w', h=h, w=w)
|
||||
features.append(x)
|
||||
return features
|
||||
@@ -0,0 +1,440 @@
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/resnet.py
|
||||
|
||||
from einops import rearrange, repeat
|
||||
from functools import partial
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from diffusers.models.activations import get_activation
|
||||
from diffusers.models.normalization import AdaGroupNorm
|
||||
from diffusers.models.attention_processor import SpatialNorm
|
||||
|
||||
|
||||
class InflatedConv3d(nn.Conv2d):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class InflatedGroupNorm(nn.GroupNorm):
|
||||
def forward(self, x):
|
||||
# return super().forward(x)
|
||||
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
def zero_module(module):
|
||||
# Zero out the parameters of a module and return it.
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
class FusionBlock2D(nn.Module):
|
||||
r"""
|
||||
A Resnet block.
|
||||
|
||||
Parameters:
|
||||
in_channels (`int`): The number of channels in the input.
|
||||
out_channels (`int`, *optional*, default to be `None`):
|
||||
The number of output channels for the first conv2d layer. If None, same as `in_channels`.
|
||||
dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use.
|
||||
temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding.
|
||||
groups (`int`, *optional*, default to `32`): The number of groups to use for the first normalization layer.
|
||||
groups_out (`int`, *optional*, default to None):
|
||||
The number of groups to use for the second normalization layer. if set to None, same as `groups`.
|
||||
eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization.
|
||||
non_linearity (`str`, *optional*, default to `"swish"`): the activation function to use.
|
||||
time_embedding_norm (`str`, *optional*, default to `"default"` ): Time scale shift config.
|
||||
By default, apply timestep embedding conditioning with a simple shift mechanism. Choose "scale_shift" or
|
||||
"ada_group" for a stronger conditioning with scale and shift.
|
||||
kernel (`torch.FloatTensor`, optional, default to None): FIR filter, see
|
||||
[`~models.resnet.FirUpsample2D`] and [`~models.resnet.FirDownsample2D`].
|
||||
output_scale_factor (`float`, *optional*, default to be `1.0`): the scale factor to use for the output.
|
||||
use_in_shortcut (`bool`, *optional*, default to `True`):
|
||||
If `True`, add a 1x1 nn.conv2d layer for skip-connection.
|
||||
up (`bool`, *optional*, default to `False`): If `True`, add an upsample layer.
|
||||
down (`bool`, *optional*, default to `False`): If `True`, add a downsample layer.
|
||||
conv_shortcut_bias (`bool`, *optional*, default to `True`): If `True`, adds a learnable bias to the
|
||||
`conv_shortcut` output.
|
||||
conv_2d_out_channels (`int`, *optional*, default to `None`): the number of channels in the output.
|
||||
If None, same as `out_channels`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels,
|
||||
out_channels=None,
|
||||
conv_shortcut=False,
|
||||
dropout=0.0,
|
||||
temb_channels=512,
|
||||
groups=32,
|
||||
groups_out=None,
|
||||
pre_norm=True,
|
||||
eps=1e-6,
|
||||
non_linearity="swish",
|
||||
skip_time_act=False,
|
||||
time_embedding_norm="default", # default, scale_shift, ada_group, spatial
|
||||
kernel=None,
|
||||
output_scale_factor=1.0,
|
||||
use_in_shortcut=None,
|
||||
up=False,
|
||||
down=False,
|
||||
conv_shortcut_bias: bool = True,
|
||||
conv_2d_out_channels: Optional[int] = None,
|
||||
|
||||
zero_init=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.pre_norm = pre_norm
|
||||
self.pre_norm = True
|
||||
|
||||
in_channels = in_channels * 2
|
||||
self.in_channels = in_channels
|
||||
|
||||
out_channels = in_channels * 3 if out_channels is None else out_channels * 3
|
||||
self.out_channels = out_channels
|
||||
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.up = up
|
||||
self.down = down
|
||||
self.output_scale_factor = output_scale_factor
|
||||
self.time_embedding_norm = time_embedding_norm
|
||||
self.skip_time_act = skip_time_act
|
||||
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
|
||||
if self.time_embedding_norm == "ada_group":
|
||||
self.norm1 = AdaGroupNorm(temb_channels, in_channels, groups, eps=eps)
|
||||
elif self.time_embedding_norm == "spatial":
|
||||
self.norm1 = SpatialNorm(in_channels, temb_channels)
|
||||
else:
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
|
||||
self.conv1 = torch.nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
if temb_channels is not None:
|
||||
if self.time_embedding_norm == "default":
|
||||
self.time_emb_proj = torch.nn.Linear(temb_channels, out_channels)
|
||||
elif self.time_embedding_norm == "scale_shift":
|
||||
self.time_emb_proj = torch.nn.Linear(temb_channels, 2 * out_channels)
|
||||
elif self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
|
||||
self.time_emb_proj = None
|
||||
else:
|
||||
raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
|
||||
if self.time_embedding_norm == "ada_group":
|
||||
self.norm2 = AdaGroupNorm(temb_channels, out_channels, groups_out, eps=eps)
|
||||
elif self.time_embedding_norm == "spatial":
|
||||
self.norm2 = SpatialNorm(out_channels, temb_channels)
|
||||
else:
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
conv_2d_out_channels = conv_2d_out_channels or out_channels
|
||||
self.conv2 = torch.nn.Conv2d(out_channels, conv_2d_out_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
|
||||
self.upsample = self.downsample = None
|
||||
if self.up:
|
||||
if kernel == "fir":
|
||||
fir_kernel = (1, 3, 3, 1)
|
||||
self.upsample = lambda x: upsample_2d(x, kernel=fir_kernel)
|
||||
elif kernel == "sde_vp":
|
||||
self.upsample = partial(F.interpolate, scale_factor=2.0, mode="nearest")
|
||||
else:
|
||||
self.upsample = Upsample2D(in_channels, use_conv=False)
|
||||
elif self.down:
|
||||
if kernel == "fir":
|
||||
fir_kernel = (1, 3, 3, 1)
|
||||
self.downsample = lambda x: downsample_2d(x, kernel=fir_kernel)
|
||||
elif kernel == "sde_vp":
|
||||
self.downsample = partial(F.avg_pool2d, kernel_size=2, stride=2)
|
||||
else:
|
||||
self.downsample = Downsample2D(in_channels, use_conv=False, padding=1, name="op")
|
||||
|
||||
self.use_in_shortcut = self.in_channels != conv_2d_out_channels if use_in_shortcut is None else use_in_shortcut
|
||||
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
self.conv_shortcut = torch.nn.Conv2d(
|
||||
in_channels, conv_2d_out_channels, kernel_size=1, stride=1, padding=0, bias=conv_shortcut_bias
|
||||
)
|
||||
|
||||
conv_out = torch.nn.Conv2d(
|
||||
conv_2d_out_channels, conv_2d_out_channels, kernel_size=1, stride=1, padding=0,
|
||||
)
|
||||
self.conv_out = zero_module(conv_out) if zero_init else conv_out
|
||||
|
||||
def forward(self, init_hidden_state, post_hidden_states, temb):
|
||||
# init_hidden_state: b c 1 h w
|
||||
# post_hidden_states: b c (f-1) h w
|
||||
|
||||
video_length = post_hidden_states.shape[2]
|
||||
repeated_init_hidden_state = repeat(init_hidden_state, "b c f h w -> b c (n f) h w", n=video_length)
|
||||
|
||||
hidden_states = torch.cat([repeated_init_hidden_state, post_hidden_states], dim=1)
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
input_tensor = hidden_states
|
||||
|
||||
if temb.shape[0] != input_tensor.shape[0]:
|
||||
temb = repeat(temb, "b c -> (b n) c", n=input_tensor.shape[0] // temb.shape[0])
|
||||
assert temb.shape[0] == input_tensor.shape[0], f"{temb.shape}, {input_tensor.shape}"
|
||||
|
||||
if self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
|
||||
hidden_states = self.norm1(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
if self.upsample is not None:
|
||||
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||
if hidden_states.shape[0] >= 64:
|
||||
input_tensor = input_tensor.contiguous()
|
||||
hidden_states = hidden_states.contiguous()
|
||||
input_tensor = self.upsample(input_tensor)
|
||||
hidden_states = self.upsample(hidden_states)
|
||||
elif self.downsample is not None:
|
||||
input_tensor = self.downsample(input_tensor)
|
||||
hidden_states = self.downsample(hidden_states)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if self.time_emb_proj is not None:
|
||||
if not self.skip_time_act:
|
||||
temb = self.nonlinearity(temb)
|
||||
temb = self.time_emb_proj(temb)[:, :, None, None]
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
if self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
|
||||
hidden_states = self.norm2(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "scale_shift":
|
||||
scale, shift = torch.chunk(temb, 2, dim=1)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
output_tensor = self.conv_out(output_tensor)
|
||||
|
||||
output_tensor = rearrange(output_tensor, "(b f) c h w -> b c f h w", f=video_length)
|
||||
scale_1, scale_2, shift = output_tensor.chunk(3, dim=1)
|
||||
|
||||
# output_tensor = (1 + scale_1) * repeated_init_hidden_state + scale_2 * post_hidden_states + shift
|
||||
output_tensor = scale_1 * repeated_init_hidden_state + (1 + scale_2) * post_hidden_states + shift
|
||||
|
||||
return output_tensor
|
||||
|
||||
class Upsample3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, use_conv_transpose=False, out_channels=None, name="conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.use_conv_transpose = use_conv_transpose
|
||||
self.name = name
|
||||
|
||||
conv = None
|
||||
if use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
elif use_conv:
|
||||
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, padding=1)
|
||||
|
||||
def forward(self, hidden_states, output_size=None):
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
if self.use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
|
||||
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
||||
dtype = hidden_states.dtype
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
|
||||
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||
if hidden_states.shape[0] >= 64:
|
||||
hidden_states = hidden_states.contiguous()
|
||||
|
||||
# if `output_size` is passed we force the interpolation output
|
||||
# size and do not make use of `scale_factor=2`
|
||||
if output_size is None:
|
||||
hidden_states = F.interpolate(hidden_states, scale_factor=[1.0, 2.0, 2.0], mode="nearest")
|
||||
else:
|
||||
hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest")
|
||||
|
||||
# If the input is bfloat16, we cast back to bfloat16
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(dtype)
|
||||
|
||||
# if self.use_conv:
|
||||
# if self.name == "conv":
|
||||
# hidden_states = self.conv(hidden_states)
|
||||
# else:
|
||||
# hidden_states = self.Conv2d_0(hidden_states)
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Downsample3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.padding = padding
|
||||
stride = 2
|
||||
self.name = name
|
||||
|
||||
if use_conv:
|
||||
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, stride=stride, padding=padding)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, hidden_states):
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
if self.use_conv and self.padding == 0:
|
||||
raise NotImplementedError
|
||||
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ResnetBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels,
|
||||
out_channels=None,
|
||||
conv_shortcut=False,
|
||||
dropout=0.0,
|
||||
temb_channels=512,
|
||||
groups=32,
|
||||
groups_out=None,
|
||||
pre_norm=True,
|
||||
eps=1e-6,
|
||||
non_linearity="swish",
|
||||
time_embedding_norm="default",
|
||||
output_scale_factor=1.0,
|
||||
use_in_shortcut=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.pre_norm = pre_norm
|
||||
self.pre_norm = True
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.time_embedding_norm = time_embedding_norm
|
||||
self.output_scale_factor = output_scale_factor
|
||||
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
|
||||
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if temb_channels is not None:
|
||||
if self.time_embedding_norm == "default":
|
||||
time_emb_proj_out_channels = out_channels
|
||||
elif self.time_embedding_norm == "scale_shift":
|
||||
time_emb_proj_out_channels = out_channels * 2
|
||||
else:
|
||||
raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||
|
||||
self.time_emb_proj = torch.nn.Linear(temb_channels, time_emb_proj_out_channels)
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if non_linearity == "swish":
|
||||
self.nonlinearity = lambda x: F.silu(x)
|
||||
elif non_linearity == "mish":
|
||||
self.nonlinearity = Mish()
|
||||
elif non_linearity == "silu":
|
||||
self.nonlinearity = nn.SiLU()
|
||||
|
||||
self.use_in_shortcut = self.in_channels != self.out_channels if use_in_shortcut is None else use_in_shortcut
|
||||
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
self.conv_shortcut = InflatedConv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
def forward(self, input_tensor, temb):
|
||||
# input: b c f h w
|
||||
|
||||
hidden_states = input_tensor
|
||||
|
||||
video_length = hidden_states.shape[2]
|
||||
emb = repeat(emb, "b c -> (b f) c", f=video_length)
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if temb is not None:
|
||||
temb = self.time_emb_proj(self.nonlinearity(temb))[:, :, None, None, None]
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "scale_shift":
|
||||
scale, shift = torch.chunk(temb, 2, dim=1)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
return output_tensor
|
||||
|
||||
|
||||
class Mish(torch.nn.Module):
|
||||
def forward(self, hidden_states):
|
||||
return hidden_states * torch.tanh(torch.nn.functional.softplus(hidden_states))
|
||||
@@ -0,0 +1,818 @@
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/unet_2d_blocks.py
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from einops import rearrange, repeat
|
||||
from diffusers.models.resnet import Downsample2D, Upsample2D, ResnetBlock2D
|
||||
from diffusers.models.transformer_2d import Transformer2DModel
|
||||
|
||||
from cameractrl.models.motion_module import get_motion_module
|
||||
|
||||
|
||||
def get_down_block(
|
||||
down_block_type,
|
||||
num_layers,
|
||||
in_channels,
|
||||
out_channels,
|
||||
temb_channels,
|
||||
add_downsample,
|
||||
resnet_eps,
|
||||
resnet_act_fn,
|
||||
attn_num_head_channels,
|
||||
resnet_groups=None,
|
||||
cross_attention_dim=None,
|
||||
downsample_padding=None,
|
||||
dual_cross_attention=False,
|
||||
use_linear_projection=False,
|
||||
only_cross_attention=False,
|
||||
upcast_attention=False,
|
||||
resnet_time_scale_shift="default",
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
down_block_type = down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type
|
||||
if down_block_type == "DownBlock3D":
|
||||
return DownBlock3D(
|
||||
num_layers=num_layers,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
add_downsample=add_downsample,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
downsample_padding=downsample_padding,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
use_motion_module=use_motion_module,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
)
|
||||
elif down_block_type == "CrossAttnDownBlock3D":
|
||||
if cross_attention_dim is None:
|
||||
raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlock3D")
|
||||
return CrossAttnDownBlock3D(
|
||||
num_layers=num_layers,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
add_downsample=add_downsample,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
downsample_padding=downsample_padding,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
attn_num_head_channels=attn_num_head_channels,
|
||||
dual_cross_attention=dual_cross_attention,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
use_motion_module=use_motion_module,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
)
|
||||
raise ValueError(f"{down_block_type} does not exist.")
|
||||
|
||||
|
||||
def get_up_block(
|
||||
up_block_type,
|
||||
num_layers,
|
||||
in_channels,
|
||||
out_channels,
|
||||
prev_output_channel,
|
||||
temb_channels,
|
||||
add_upsample,
|
||||
resnet_eps,
|
||||
resnet_act_fn,
|
||||
attn_num_head_channels,
|
||||
resnet_groups=None,
|
||||
cross_attention_dim=None,
|
||||
dual_cross_attention=False,
|
||||
use_linear_projection=False,
|
||||
only_cross_attention=False,
|
||||
upcast_attention=False,
|
||||
resnet_time_scale_shift="default",
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
up_block_type = up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
|
||||
if up_block_type == "UpBlock3D":
|
||||
return UpBlock3D(
|
||||
num_layers=num_layers,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
prev_output_channel=prev_output_channel,
|
||||
temb_channels=temb_channels,
|
||||
add_upsample=add_upsample,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
use_motion_module=use_motion_module,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
)
|
||||
elif up_block_type == "CrossAttnUpBlock3D":
|
||||
if cross_attention_dim is None:
|
||||
raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlock3D")
|
||||
return CrossAttnUpBlock3D(
|
||||
num_layers=num_layers,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
prev_output_channel=prev_output_channel,
|
||||
temb_channels=temb_channels,
|
||||
add_upsample=add_upsample,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
attn_num_head_channels=attn_num_head_channels,
|
||||
dual_cross_attention=dual_cross_attention,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
use_motion_module=use_motion_module,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
)
|
||||
raise ValueError(f"{up_block_type} does not exist.")
|
||||
|
||||
|
||||
class UNetMidBlock3DCrossAttn(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
temb_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
attn_num_head_channels=1,
|
||||
output_scale_factor=1.0,
|
||||
cross_attention_dim=1280,
|
||||
dual_cross_attention=False,
|
||||
use_linear_projection=False,
|
||||
upcast_attention=False,
|
||||
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.has_cross_attention = True
|
||||
self.attn_num_head_channels = attn_num_head_channels
|
||||
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||
|
||||
# there is always at least one resnet
|
||||
resnets = [
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
]
|
||||
attentions = []
|
||||
motion_modules = []
|
||||
|
||||
for _ in range(num_layers):
|
||||
if dual_cross_attention: raise NotImplementedError
|
||||
attentions.append(
|
||||
Transformer2DModel(
|
||||
attn_num_head_channels,
|
||||
in_channels // attn_num_head_channels,
|
||||
in_channels=in_channels,
|
||||
num_layers=1,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
norm_num_groups=resnet_groups,
|
||||
use_linear_projection=use_linear_projection,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
)
|
||||
motion_modules.append(
|
||||
get_motion_module(
|
||||
in_channels=in_channels,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
) if use_motion_module else None
|
||||
)
|
||||
resnets.append(
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
)
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.motion_modules = nn.ModuleList(motion_modules) if use_motion_module else motion_modules
|
||||
|
||||
def forward(self, hidden_states, temb=None, encoder_hidden_states=None, attention_mask=None,
|
||||
motion_module_alpha=1., cross_attention_kwargs=None, motion_cross_attention_kwargs=None):
|
||||
video_length = hidden_states.shape[2]
|
||||
temb_repeated = repeat(temb, "b c -> (b f) c", f=video_length)
|
||||
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = self.resnets[0](hidden_states, temb_repeated)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
lora_scale = getattr(self, "lora_scale", None)
|
||||
if lora_scale != None:
|
||||
cross_attention_kwargs = {"scale": lora_scale}
|
||||
motion_lora_scale = getattr(self, "motion_lora_scale", None)
|
||||
if motion_lora_scale != None:
|
||||
if motion_cross_attention_kwargs is None:
|
||||
motion_cross_attention_kwargs = {"scale": motion_lora_scale}
|
||||
else:
|
||||
motion_cross_attention_kwargs.update({"scale": motion_lora_scale})
|
||||
|
||||
for attn, resnet, motion_module in zip(self.attentions, self.resnets[1:], self.motion_modules):
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = attn(hidden_states, encoder_hidden_states=encoder_hidden_states,
|
||||
cross_attention_kwargs=cross_attention_kwargs).sample
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
# motion module
|
||||
if motion_module is not None:
|
||||
# hidden_states = motion_module_alpha * motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states) + hidden_states
|
||||
hidden_states = motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states, cross_attention_kwargs=motion_cross_attention_kwargs)
|
||||
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = resnet(hidden_states, temb_repeated)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CrossAttnDownBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
temb_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
attn_num_head_channels=1,
|
||||
cross_attention_dim=1280,
|
||||
output_scale_factor=1.0,
|
||||
downsample_padding=1,
|
||||
add_downsample=True,
|
||||
dual_cross_attention=False,
|
||||
use_linear_projection=False,
|
||||
only_cross_attention=False,
|
||||
upcast_attention=False,
|
||||
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
attentions = []
|
||||
motion_modules = []
|
||||
|
||||
self.has_cross_attention = True
|
||||
self.attn_num_head_channels = attn_num_head_channels
|
||||
|
||||
for i in range(num_layers):
|
||||
in_channels = in_channels if i == 0 else out_channels
|
||||
resnets.append(
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
)
|
||||
|
||||
if dual_cross_attention:
|
||||
raise NotImplementedError
|
||||
attentions.append(
|
||||
Transformer2DModel(
|
||||
attn_num_head_channels,
|
||||
out_channels // attn_num_head_channels,
|
||||
in_channels=out_channels,
|
||||
num_layers=1,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
norm_num_groups=resnet_groups,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
)
|
||||
motion_modules.append(
|
||||
get_motion_module(
|
||||
in_channels=out_channels,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
) if use_motion_module else None
|
||||
)
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.motion_modules = nn.ModuleList(motion_modules) if use_motion_module else motion_modules
|
||||
|
||||
if add_downsample:
|
||||
self.downsamplers = nn.ModuleList(
|
||||
[
|
||||
Downsample2D(
|
||||
out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states, temb=None, encoder_hidden_states=None, attention_mask=None,
|
||||
motion_module_alpha=1., cross_attention_kwargs={}, motion_cross_attention_kwargs={}):
|
||||
video_length = hidden_states.shape[2]
|
||||
temb_repeated = repeat(temb, "b c -> (b f) c", f=video_length)
|
||||
|
||||
output_states = ()
|
||||
|
||||
lora_scale = getattr(self, "lora_scale", None)
|
||||
if lora_scale != None:
|
||||
cross_attention_kwargs["scale"] = lora_scale
|
||||
motion_lora_scale = getattr(self, "motion_lora_scale", None)
|
||||
if motion_lora_scale != None:
|
||||
if motion_cross_attention_kwargs is None:
|
||||
motion_cross_attention_kwargs = {"scale": motion_lora_scale}
|
||||
else:
|
||||
motion_cross_attention_kwargs.update({"scale": motion_lora_scale})
|
||||
|
||||
for resnet, attn, motion_module in zip(self.resnets, self.attentions, self.motion_modules):
|
||||
if self.training and self.gradient_checkpointing:
|
||||
raise NotImplementedError
|
||||
|
||||
def create_custom_forward(module, return_dict=None):
|
||||
def custom_forward(*inputs):
|
||||
if return_dict is not None:
|
||||
return module(*inputs, return_dict=return_dict)
|
||||
else:
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(resnet), hidden_states, temb)
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(attn, return_dict=False),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
)[0]
|
||||
if motion_module is not None:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(motion_module),
|
||||
hidden_states.requires_grad_(), temb,
|
||||
encoder_hidden_states)
|
||||
|
||||
else:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = resnet(hidden_states, temb_repeated)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = attn(hidden_states, encoder_hidden_states=encoder_hidden_states,
|
||||
cross_attention_kwargs=cross_attention_kwargs).sample
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
# motion module
|
||||
if motion_module is not None:
|
||||
# hidden_states = motion_module_alpha * motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states) + hidden_states
|
||||
hidden_states = motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states, cross_attention_kwargs=motion_cross_attention_kwargs)
|
||||
|
||||
output_states += (hidden_states,)
|
||||
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = downsampler(hidden_states)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
output_states += (hidden_states,)
|
||||
|
||||
return hidden_states, output_states
|
||||
|
||||
|
||||
class DownBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
temb_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
output_scale_factor=1.0,
|
||||
add_downsample=True,
|
||||
downsample_padding=1,
|
||||
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
motion_modules = []
|
||||
|
||||
for i in range(num_layers):
|
||||
in_channels = in_channels if i == 0 else out_channels
|
||||
resnets.append(
|
||||
ResnetBlock2D(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
)
|
||||
motion_modules.append(
|
||||
get_motion_module(
|
||||
in_channels=out_channels,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
) if use_motion_module else None
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.motion_modules = nn.ModuleList(motion_modules) if use_motion_module else motion_modules
|
||||
|
||||
if add_downsample:
|
||||
self.downsamplers = nn.ModuleList(
|
||||
[
|
||||
Downsample2D(
|
||||
out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states, temb=None, encoder_hidden_states=None, motion_module_alpha=1.,
|
||||
motion_cross_attention_kwargs={}, **kwargs):
|
||||
video_length = hidden_states.shape[2]
|
||||
temb_repeated = repeat(temb, "b c -> (b f) c", f=video_length)
|
||||
output_states = ()
|
||||
motion_lora_scale = getattr(self, "motion_lora_scale", None)
|
||||
if motion_lora_scale != None:
|
||||
if motion_cross_attention_kwargs is None:
|
||||
motion_cross_attention_kwargs = {"scale": motion_lora_scale}
|
||||
else:
|
||||
motion_cross_attention_kwargs.update({"scale": motion_lora_scale})
|
||||
|
||||
for resnet, motion_module in zip(self.resnets, self.motion_modules):
|
||||
if self.training and self.gradient_checkpointing:
|
||||
raise NotImplementedError
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(resnet), hidden_states, temb)
|
||||
if motion_module is not None:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(motion_module),
|
||||
hidden_states.requires_grad_(), temb,
|
||||
encoder_hidden_states)
|
||||
else:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = resnet(hidden_states, temb_repeated)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
# motion module
|
||||
if motion_module is not None:
|
||||
hidden_states = motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states, cross_attention_kwargs=motion_cross_attention_kwargs)
|
||||
|
||||
output_states += (hidden_states,)
|
||||
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = downsampler(hidden_states)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
output_states += (hidden_states,)
|
||||
|
||||
return hidden_states, output_states
|
||||
|
||||
|
||||
class CrossAttnUpBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
prev_output_channel: int,
|
||||
temb_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
attn_num_head_channels=1,
|
||||
cross_attention_dim=1280,
|
||||
output_scale_factor=1.0,
|
||||
add_upsample=True,
|
||||
dual_cross_attention=False,
|
||||
use_linear_projection=False,
|
||||
only_cross_attention=False,
|
||||
upcast_attention=False,
|
||||
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
attentions = []
|
||||
motion_modules = []
|
||||
|
||||
self.has_cross_attention = True
|
||||
self.attn_num_head_channels = attn_num_head_channels
|
||||
|
||||
for i in range(num_layers):
|
||||
res_skip_channels = in_channels if (i == num_layers - 1) else out_channels
|
||||
resnet_in_channels = prev_output_channel if i == 0 else out_channels
|
||||
|
||||
resnets.append(
|
||||
ResnetBlock2D(
|
||||
in_channels=resnet_in_channels + res_skip_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
)
|
||||
|
||||
if dual_cross_attention:
|
||||
raise NotImplementedError
|
||||
attentions.append(
|
||||
Transformer2DModel(
|
||||
attn_num_head_channels,
|
||||
out_channels // attn_num_head_channels,
|
||||
in_channels=out_channels,
|
||||
num_layers=1,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
norm_num_groups=resnet_groups,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
)
|
||||
motion_modules.append(
|
||||
get_motion_module(
|
||||
in_channels=out_channels,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
) if use_motion_module else None
|
||||
)
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.motion_modules = nn.ModuleList(motion_modules) if use_motion_module else motion_modules
|
||||
|
||||
if add_upsample:
|
||||
self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)])
|
||||
else:
|
||||
self.upsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states,
|
||||
res_hidden_states_tuple,
|
||||
temb=None,
|
||||
encoder_hidden_states=None,
|
||||
upsample_size=None,
|
||||
attention_mask=None,
|
||||
motion_module_alpha=1.,
|
||||
cross_attention_kwargs=None,
|
||||
motion_cross_attention_kwargs={}
|
||||
):
|
||||
video_length = hidden_states.shape[2]
|
||||
temb_repeated = repeat(temb, "b c -> (b f) c", f=video_length)
|
||||
|
||||
lora_scale = getattr(self, "lora_scale", None)
|
||||
if lora_scale != None:
|
||||
cross_attention_kwargs = {"scale": lora_scale}
|
||||
motion_lora_scale = getattr(self, "motion_lora_scale", None)
|
||||
if motion_lora_scale != None:
|
||||
if motion_cross_attention_kwargs is None:
|
||||
motion_cross_attention_kwargs = {"scale": motion_lora_scale}
|
||||
else:
|
||||
motion_cross_attention_kwargs.update({"scale": motion_lora_scale})
|
||||
|
||||
for resnet, attn, motion_module in zip(self.resnets, self.attentions, self.motion_modules):
|
||||
# pop res hidden states
|
||||
res_hidden_states = res_hidden_states_tuple[-1]
|
||||
res_hidden_states_tuple = res_hidden_states_tuple[:-1]
|
||||
hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)
|
||||
|
||||
if self.training and self.gradient_checkpointing:
|
||||
raise NotImplementedError
|
||||
|
||||
def create_custom_forward(module, return_dict=None):
|
||||
def custom_forward(*inputs):
|
||||
if return_dict is not None:
|
||||
return module(*inputs, return_dict=return_dict)
|
||||
else:
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(resnet), hidden_states, temb)
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(attn, return_dict=False),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
)[0]
|
||||
if motion_module is not None:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(motion_module),
|
||||
hidden_states.requires_grad_(), temb,
|
||||
encoder_hidden_states)
|
||||
|
||||
else:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = resnet(hidden_states, temb_repeated)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = attn(hidden_states, encoder_hidden_states=encoder_hidden_states,
|
||||
cross_attention_kwargs=cross_attention_kwargs).sample
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
# motion module
|
||||
if motion_module is not None:
|
||||
# hidden_states = motion_module_alpha * motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states) + hidden_states
|
||||
hidden_states = motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states, cross_attention_kwargs=motion_cross_attention_kwargs)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = upsampler(hidden_states, upsample_size)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class UpBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
prev_output_channel: int,
|
||||
out_channels: int,
|
||||
temb_channels: int,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
output_scale_factor=1.0,
|
||||
add_upsample=True,
|
||||
|
||||
use_motion_module=None,
|
||||
motion_module_type=None,
|
||||
motion_module_kwargs=None,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
motion_modules = []
|
||||
|
||||
for i in range(num_layers):
|
||||
res_skip_channels = in_channels if (i == num_layers - 1) else out_channels
|
||||
resnet_in_channels = prev_output_channel if i == 0 else out_channels
|
||||
|
||||
resnets.append(
|
||||
ResnetBlock2D(
|
||||
in_channels=resnet_in_channels + res_skip_channels,
|
||||
out_channels=out_channels,
|
||||
temb_channels=temb_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
dropout=dropout,
|
||||
time_embedding_norm=resnet_time_scale_shift,
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
)
|
||||
)
|
||||
motion_modules.append(
|
||||
get_motion_module(
|
||||
in_channels=out_channels,
|
||||
motion_module_type=motion_module_type,
|
||||
motion_module_kwargs=motion_module_kwargs,
|
||||
) if use_motion_module else None
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.motion_modules = nn.ModuleList(motion_modules) if use_motion_module else motion_modules
|
||||
|
||||
if add_upsample:
|
||||
self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)])
|
||||
else:
|
||||
self.upsamplers = None
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None, encoder_hidden_states=None,
|
||||
motion_module_alpha=1., motion_cross_attention_kwargs={}, **kwargs):
|
||||
video_length = hidden_states.shape[2]
|
||||
temb_repeated = repeat(temb, "b c -> (b f) c", f=video_length)
|
||||
|
||||
motion_lora_scale = getattr(self, "motion_lora_scale", None)
|
||||
if motion_lora_scale != None:
|
||||
if motion_cross_attention_kwargs is None:
|
||||
motion_cross_attention_kwargs = {"scale": motion_lora_scale}
|
||||
else:
|
||||
motion_cross_attention_kwargs.update({"scale": motion_lora_scale})
|
||||
|
||||
for resnet, motion_module in zip(self.resnets, self.motion_modules):
|
||||
# pop res hidden states
|
||||
res_hidden_states = res_hidden_states_tuple[-1]
|
||||
res_hidden_states_tuple = res_hidden_states_tuple[:-1]
|
||||
hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)
|
||||
|
||||
if self.training and self.gradient_checkpointing:
|
||||
raise NotImplementedError
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(resnet), hidden_states, temb)
|
||||
if motion_module is not None:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(motion_module),
|
||||
hidden_states.requires_grad_(), temb,
|
||||
encoder_hidden_states)
|
||||
else:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = resnet(hidden_states, temb_repeated)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
# motion module
|
||||
if motion_module is not None:
|
||||
hidden_states = motion_module(hidden_states, temb=temb, encoder_hidden_states=encoder_hidden_states, cross_attention_kwargs=motion_cross_attention_kwargs)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
for upsampler in self.upsamplers:
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
hidden_states = upsampler(hidden_states, upsample_size)
|
||||
hidden_states = rearrange(hidden_states, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -0,0 +1,722 @@
|
||||
# Adapted from https://github.com/showlab/Tune-A-Video/blob/main/tuneavideo/pipelines/pipeline_tuneavideo.py
|
||||
|
||||
import inspect
|
||||
import torch
|
||||
|
||||
import numpy as np
|
||||
|
||||
from typing import Callable, List, Optional, Union
|
||||
from dataclasses import dataclass
|
||||
from diffusers.utils import is_accelerate_available
|
||||
from packaging import version
|
||||
from einops import rearrange
|
||||
from transformers import CLIPTextModel, CLIPTokenizer
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.models import AutoencoderKL
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import (
|
||||
DDIMScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
PNDMScheduler,
|
||||
)
|
||||
from diffusers.loaders import LoraLoaderMixin
|
||||
from diffusers.utils import deprecate, logging, BaseOutput
|
||||
|
||||
from cameractrl.models.pose_adaptor import CameraPoseEncoder
|
||||
from cameractrl.models.unet import UNet3DConditionModel
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnimationPipelineOutput(BaseOutput):
|
||||
videos: Union[torch.Tensor, np.ndarray]
|
||||
|
||||
|
||||
class AnimationPipeline(DiffusionPipeline, LoraLoaderMixin):
|
||||
_optional_components = []
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
unet: UNet3DConditionModel,
|
||||
scheduler: Union[
|
||||
DDIMScheduler,
|
||||
PNDMScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
],
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
|
||||
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
|
||||
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
|
||||
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
|
||||
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
|
||||
" file"
|
||||
)
|
||||
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["steps_offset"] = 1
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
|
||||
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
|
||||
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
|
||||
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
|
||||
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
|
||||
)
|
||||
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["clip_sample"] = False
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
is_unet_version_less_0_9_0 = hasattr(unet.config, "_diffusers_version") and version.parse(
|
||||
version.parse(unet.config._diffusers_version).base_version
|
||||
) < version.parse("0.9.0.dev0")
|
||||
is_unet_sample_size_less_64 = hasattr(unet.config, "sample_size") and unet.config.sample_size < 64
|
||||
if is_unet_version_less_0_9_0 and is_unet_sample_size_less_64:
|
||||
deprecation_message = (
|
||||
"The configuration file of the unet has set the default `sample_size` to smaller than"
|
||||
" 64 which seems highly unlikely. If your checkpoint is a fine-tuned version of any of the"
|
||||
" following: \n- CompVis/stable-diffusion-v1-4 \n- CompVis/stable-diffusion-v1-3 \n-"
|
||||
" CompVis/stable-diffusion-v1-2 \n- CompVis/stable-diffusion-v1-1 \n- runwayml/stable-diffusion-v1-5"
|
||||
" \n- runwayml/stable-diffusion-inpainting \n you should change 'sample_size' to 64 in the"
|
||||
" configuration file. Please make sure to update the config accordingly as leaving `sample_size=32`"
|
||||
" in the config might lead to incorrect results in future versions. If you have downloaded this"
|
||||
" checkpoint from the Hugging Face Hub, it would be very nice if you could open a Pull request for"
|
||||
" the `unet/config.json` file"
|
||||
)
|
||||
deprecate("sample_size<64", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(unet.config)
|
||||
new_config["sample_size"] = 64
|
||||
unet._internal_dict = FrozenDict(new_config)
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_sequential_cpu_offload(self, gpu_id=0):
|
||||
if is_accelerate_available():
|
||||
from accelerate import cpu_offload
|
||||
else:
|
||||
raise ImportError("Please install accelerate via `pip install accelerate`")
|
||||
|
||||
device = torch.device(f"cuda:{gpu_id}")
|
||||
|
||||
for cpu_offloaded_model in [self.unet, self.text_encoder, self.vae]:
|
||||
if cpu_offloaded_model is not None:
|
||||
cpu_offload(cpu_offloaded_model, device)
|
||||
|
||||
|
||||
@property
|
||||
def _execution_device(self):
|
||||
if self.device != torch.device("meta") or not hasattr(self.unet, "_hf_hook"):
|
||||
return self.device
|
||||
for module in self.unet.modules():
|
||||
if (
|
||||
hasattr(module, "_hf_hook")
|
||||
and hasattr(module._hf_hook, "execution_device")
|
||||
and module._hf_hook.execution_device is not None
|
||||
):
|
||||
return torch.device(module._hf_hook.execution_device)
|
||||
return self.device
|
||||
|
||||
def _encode_prompt(self, prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt):
|
||||
batch_size = len(prompt) if isinstance(prompt, list) else 1
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {self.tokenizer.model_max_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||
attention_mask = text_inputs.attention_mask.to(device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
text_embeddings = self.text_encoder(
|
||||
text_input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
text_embeddings = text_embeddings[0]
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
bs_embed, seq_len, _ = text_embeddings.shape
|
||||
text_embeddings = text_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
text_embeddings = text_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
# get unconditional embeddings for classifier free guidance
|
||||
if do_classifier_free_guidance:
|
||||
uncond_tokens: List[str]
|
||||
if negative_prompt is None:
|
||||
uncond_tokens = [""] * batch_size
|
||||
elif type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif isinstance(negative_prompt, str):
|
||||
uncond_tokens = [negative_prompt]
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
else:
|
||||
uncond_tokens = negative_prompt
|
||||
|
||||
max_length = text_input_ids.shape[-1]
|
||||
uncond_input = self.tokenizer(
|
||||
uncond_tokens,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||
attention_mask = uncond_input.attention_mask.to(device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
uncond_embeddings = self.text_encoder(
|
||||
uncond_input.input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
uncond_embeddings = uncond_embeddings[0]
|
||||
|
||||
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
|
||||
seq_len = uncond_embeddings.shape[1]
|
||||
uncond_embeddings = uncond_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
uncond_embeddings = uncond_embeddings.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
text_embeddings = torch.cat([uncond_embeddings, text_embeddings])
|
||||
|
||||
return text_embeddings
|
||||
|
||||
def decode_latents(self, latents):
|
||||
video_length = latents.shape[2]
|
||||
latents = 1 / 0.18215 * latents
|
||||
latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
# video = self.vae.decode(latents).sample
|
||||
video = []
|
||||
for frame_idx in range(latents.shape[0]):
|
||||
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
|
||||
video = torch.cat(video)
|
||||
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
|
||||
video = video.cpu().float().numpy()
|
||||
return video
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
def check_inputs(self, prompt, height, width, callback_steps):
|
||||
if not isinstance(prompt, str) and not isinstance(prompt, list):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
if (callback_steps is None) or (
|
||||
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
|
||||
def prepare_latents(self, batch_size, num_channels_latents, video_length, height, width, dtype, device, generator, latents=None):
|
||||
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
if latents is None:
|
||||
rand_device = "cpu" if device.type == "mps" else device
|
||||
|
||||
if isinstance(generator, list):
|
||||
shape = shape
|
||||
# shape = (1,) + shape[1:]
|
||||
latents = [
|
||||
torch.randn(shape, generator=generator[i], device=rand_device, dtype=dtype)
|
||||
for i in range(batch_size)
|
||||
]
|
||||
latents = torch.cat(latents, dim=0).to(device)
|
||||
else:
|
||||
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype).to(device)
|
||||
else:
|
||||
if latents.shape != shape:
|
||||
raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")
|
||||
latents = latents.to(device)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
video_length: Optional[int],
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "tensor",
|
||||
return_dict: bool = True,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: Optional[int] = 1,
|
||||
|
||||
multidiff_total_steps: int = 1,
|
||||
multidiff_overlaps: int = 12,
|
||||
**kwargs,
|
||||
):
|
||||
# Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
# Check inputs. Raise error if not correct
|
||||
self.check_inputs(prompt, height, width, callback_steps)
|
||||
|
||||
# Define call parameters
|
||||
# batch_size = 1 if isinstance(prompt, str) else len(prompt)
|
||||
batch_size = 1
|
||||
if latents is not None:
|
||||
batch_size = latents.shape[0]
|
||||
if isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# Encode input prompt
|
||||
prompt = prompt if isinstance(prompt, list) else [prompt] * batch_size
|
||||
if negative_prompt is not None:
|
||||
negative_prompt = negative_prompt if isinstance(negative_prompt, list) else [negative_prompt] * batch_size
|
||||
text_embeddings = self._encode_prompt(
|
||||
prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt
|
||||
)
|
||||
|
||||
# Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# Prepare latent variables
|
||||
single_model_length = video_length
|
||||
video_length = multidiff_total_steps * (video_length - multidiff_overlaps) + multidiff_overlaps
|
||||
num_channels_latents = self.unet.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
video_length,
|
||||
height,
|
||||
width,
|
||||
text_embeddings.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
latents_dtype = latents.dtype
|
||||
|
||||
# Prepare extra step kwargs.
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
# Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
noise_pred_full = torch.zeros_like(latents).to(latents.device)
|
||||
mask_full = torch.zeros_like(latents).to(latents.device)
|
||||
noise_preds = []
|
||||
|
||||
for multidiff_step in range(multidiff_total_steps):
|
||||
start_idx = multidiff_step * (single_model_length - multidiff_overlaps)
|
||||
latent_partial = latents[:, :, start_idx: start_idx + single_model_length].contiguous()
|
||||
mask_full[:, :, start_idx: start_idx + single_model_length] += 1
|
||||
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latent_partial] * 2) if do_classifier_free_guidance else latent_partial
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=text_embeddings).sample.to(dtype=latents_dtype)
|
||||
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
noise_preds.append(noise_pred)
|
||||
|
||||
for pred_idx, noise_pred in enumerate(noise_preds):
|
||||
start_idx = pred_idx * (single_model_length - multidiff_overlaps)
|
||||
noise_pred_full[:, :, start_idx: start_idx + single_model_length] += noise_pred / mask_full[:, :, start_idx: start_idx + single_model_length]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred_full, t, latents, **extra_step_kwargs).prev_sample
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
callback(i, t, latents)
|
||||
|
||||
# Post-processing
|
||||
video = self.decode_latents(latents)
|
||||
|
||||
# Convert to tensor
|
||||
if output_type == "tensor":
|
||||
video = torch.from_numpy(video)
|
||||
|
||||
if not return_dict:
|
||||
return video
|
||||
|
||||
return AnimationPipelineOutput(videos=video)
|
||||
|
||||
|
||||
class CameraCtrlPipeline(AnimationPipeline):
|
||||
_optional_components = []
|
||||
|
||||
def __init__(self,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
unet: UNet3DConditionModel,
|
||||
scheduler: Union[
|
||||
DDIMScheduler,
|
||||
PNDMScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
DPMSolverMultistepScheduler],
|
||||
pose_encoder: CameraPoseEncoder):
|
||||
|
||||
super().__init__(vae, text_encoder, tokenizer, unet, scheduler)
|
||||
|
||||
self.register_modules(
|
||||
pose_encoder=pose_encoder
|
||||
)
|
||||
|
||||
def decode_latents(self, latents):
|
||||
video_length = latents.shape[2]
|
||||
latents = 1 / 0.18215 * latents
|
||||
latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
# video = self.vae.decode(latents).sample
|
||||
video = []
|
||||
for frame_idx in range(latents.shape[0]):
|
||||
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
|
||||
video = torch.cat(video)
|
||||
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
|
||||
video = video.cpu().float().numpy()
|
||||
return video
|
||||
|
||||
def _encode_prompt(self, prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt):
|
||||
batch_size = len(prompt) if isinstance(prompt, list) else 1
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {self.tokenizer.model_max_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||
attention_mask = text_inputs.attention_mask.to(device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
text_embeddings = self.text_encoder(
|
||||
text_input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
text_embeddings = text_embeddings[0]
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
bs_embed, seq_len, _ = text_embeddings.shape
|
||||
text_embeddings = text_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
text_embeddings = text_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
# get unconditional embeddings for classifier free guidance
|
||||
if do_classifier_free_guidance:
|
||||
uncond_tokens: List[str]
|
||||
if negative_prompt is None:
|
||||
uncond_tokens = [""] * batch_size
|
||||
elif type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif isinstance(negative_prompt, str):
|
||||
uncond_tokens = [negative_prompt]
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
else:
|
||||
uncond_tokens = negative_prompt
|
||||
|
||||
max_length = text_input_ids.shape[-1]
|
||||
uncond_input = self.tokenizer(
|
||||
uncond_tokens,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||
attention_mask = uncond_input.attention_mask.to(device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
uncond_embeddings = self.text_encoder(
|
||||
uncond_input.input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
uncond_embeddings = uncond_embeddings[0]
|
||||
|
||||
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
|
||||
seq_len = uncond_embeddings.shape[1]
|
||||
uncond_embeddings = uncond_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
uncond_embeddings = uncond_embeddings.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
text_embeddings = torch.cat([uncond_embeddings, text_embeddings])
|
||||
|
||||
return text_embeddings
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
pose_embedding: torch.FloatTensor,
|
||||
video_length: Optional[int],
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "tensor",
|
||||
return_dict: bool = True,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: Optional[int] = 1,
|
||||
multidiff_total_steps: int = 1,
|
||||
multidiff_overlaps: int = 12,
|
||||
**kwargs,
|
||||
):
|
||||
# Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
# Check inputs. Raise error if not correct
|
||||
self.check_inputs(prompt, height, width, callback_steps)
|
||||
|
||||
# Define call parameters
|
||||
# batch_size = 1 if isinstance(prompt, str) else len(prompt)
|
||||
batch_size = 1
|
||||
if latents is not None:
|
||||
batch_size = latents.shape[0]
|
||||
if isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
|
||||
device = pose_embedding[0].device if isinstance(pose_embedding, list) else pose_embedding.device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# Encode input prompt
|
||||
prompt = prompt if isinstance(prompt, list) else [prompt] * batch_size
|
||||
if negative_prompt is not None:
|
||||
negative_prompt = negative_prompt if isinstance(negative_prompt, list) else [negative_prompt] * batch_size
|
||||
text_embeddings = self._encode_prompt(
|
||||
prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt
|
||||
) # [2bf, l, c]
|
||||
|
||||
# Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# Prepare latent variables
|
||||
single_model_length = video_length
|
||||
video_length = multidiff_total_steps * (video_length - multidiff_overlaps) + multidiff_overlaps
|
||||
num_channels_latents = self.unet.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
video_length,
|
||||
height,
|
||||
width,
|
||||
text_embeddings.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
) # b c f h w
|
||||
latents_dtype = latents.dtype
|
||||
|
||||
# Prepare extra step kwargs.
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
if isinstance(pose_embedding, list):
|
||||
assert all([x.ndim == 5 for x in pose_embedding])
|
||||
bs = pose_embedding[0].shape[0]
|
||||
pose_embedding_features = []
|
||||
for pe in pose_embedding:
|
||||
pose_embedding_feature = self.pose_encoder(pe)
|
||||
pose_embedding_feature = [rearrange(x, '(b f) c h w -> b c f h w', b=bs) for x in pose_embedding_feature]
|
||||
pose_embedding_features.append(pose_embedding_feature)
|
||||
else:
|
||||
bs = pose_embedding.shape[0]
|
||||
assert pose_embedding.ndim == 5
|
||||
pose_embedding_features = self.pose_encoder(pose_embedding) # bf, c, h, w
|
||||
pose_embedding_features = [rearrange(x, '(b f) c h w -> b c f h w', b=bs)
|
||||
for x in pose_embedding_features]
|
||||
|
||||
# Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
if isinstance(pose_embedding_features[0], list):
|
||||
pose_embedding_features = [[torch.cat([x, x], dim=0) for x in pose_embedding_feature]
|
||||
for pose_embedding_feature in pose_embedding_features] \
|
||||
if do_classifier_free_guidance else pose_embedding_features
|
||||
else:
|
||||
pose_embedding_features = [torch.cat([x, x], dim=0) for x in pose_embedding_features] \
|
||||
if do_classifier_free_guidance else pose_embedding_features # [2b c f h w]
|
||||
import comfy.utils
|
||||
pbar = comfy.utils.ProgressBar(num_inference_steps)
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
noise_pred_full = torch.zeros_like(latents).to(latents.device)
|
||||
mask_full = torch.zeros_like(latents).to(latents.device)
|
||||
noise_preds = []
|
||||
for multidiff_step in range(multidiff_total_steps):
|
||||
start_idx = multidiff_step * (single_model_length - multidiff_overlaps)
|
||||
latent_partial = latents[:, :, start_idx: start_idx + single_model_length].contiguous()
|
||||
mask_full[:, :, start_idx: start_idx + single_model_length] += 1
|
||||
|
||||
if isinstance(pose_embedding, list):
|
||||
pose_embedding_features_input = pose_embedding_features[multidiff_step]
|
||||
else:
|
||||
pose_embedding_features_input = [x[:, :, start_idx: start_idx + single_model_length]
|
||||
for x in pose_embedding_features]
|
||||
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latent_partial] * 2) if do_classifier_free_guidance else latent_partial # [2b c f h w]
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=text_embeddings,
|
||||
pose_embedding_features=pose_embedding_features_input).sample.to(dtype=latents_dtype)
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
noise_preds.append(noise_pred)
|
||||
for pred_idx, noise_pred in enumerate(noise_preds):
|
||||
start_idx = pred_idx * (single_model_length - multidiff_overlaps)
|
||||
noise_pred_full[:, :, start_idx: start_idx + single_model_length] += noise_pred / mask_full[:, :, start_idx: start_idx + single_model_length]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1 b c f h w
|
||||
latents = self.scheduler.step(noise_pred_full, t, latents, **extra_step_kwargs).prev_sample
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
pbar.update(1)
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
callback(i, t, latents)
|
||||
|
||||
# Post-processing
|
||||
video = self.decode_latents(latents)
|
||||
|
||||
# Convert to tensor
|
||||
if output_type == "tensor":
|
||||
video = torch.from_numpy(video)
|
||||
|
||||
if not return_dict:
|
||||
return video
|
||||
|
||||
return AnimationPipelineOutput(videos=video)
|
||||
@@ -0,0 +1,556 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2023 The HuggingFace Inc. team.
|
||||
#
|
||||
# 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.
|
||||
""" Conversion script for the Stable Diffusion checkpoints."""
|
||||
|
||||
import re
|
||||
from transformers import CLIPTextModel
|
||||
|
||||
def shave_segments(path, n_shave_prefix_segments=1):
|
||||
"""
|
||||
Removes segments. Positive values shave the first segments, negative shave the last segments.
|
||||
"""
|
||||
if n_shave_prefix_segments >= 0:
|
||||
return ".".join(path.split(".")[n_shave_prefix_segments:])
|
||||
else:
|
||||
return ".".join(path.split(".")[:n_shave_prefix_segments])
|
||||
|
||||
|
||||
def renew_resnet_paths(old_list, n_shave_prefix_segments=0):
|
||||
"""
|
||||
Updates paths inside resnets to the new naming scheme (local renaming)
|
||||
"""
|
||||
mapping = []
|
||||
for old_item in old_list:
|
||||
new_item = old_item.replace("in_layers.0", "norm1")
|
||||
new_item = new_item.replace("in_layers.2", "conv1")
|
||||
|
||||
new_item = new_item.replace("out_layers.0", "norm2")
|
||||
new_item = new_item.replace("out_layers.3", "conv2")
|
||||
|
||||
new_item = new_item.replace("emb_layers.1", "time_emb_proj")
|
||||
new_item = new_item.replace("skip_connection", "conv_shortcut")
|
||||
|
||||
new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||
|
||||
mapping.append({"old": old_item, "new": new_item})
|
||||
|
||||
return mapping
|
||||
|
||||
|
||||
def renew_vae_resnet_paths(old_list, n_shave_prefix_segments=0):
|
||||
"""
|
||||
Updates paths inside resnets to the new naming scheme (local renaming)
|
||||
"""
|
||||
mapping = []
|
||||
for old_item in old_list:
|
||||
new_item = old_item
|
||||
|
||||
new_item = new_item.replace("nin_shortcut", "conv_shortcut")
|
||||
new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||
|
||||
mapping.append({"old": old_item, "new": new_item})
|
||||
|
||||
return mapping
|
||||
|
||||
|
||||
def renew_attention_paths(old_list, n_shave_prefix_segments=0):
|
||||
"""
|
||||
Updates paths inside attentions to the new naming scheme (local renaming)
|
||||
"""
|
||||
mapping = []
|
||||
for old_item in old_list:
|
||||
new_item = old_item
|
||||
|
||||
# new_item = new_item.replace('norm.weight', 'group_norm.weight')
|
||||
# new_item = new_item.replace('norm.bias', 'group_norm.bias')
|
||||
|
||||
# new_item = new_item.replace('proj_out.weight', 'proj_attn.weight')
|
||||
# new_item = new_item.replace('proj_out.bias', 'proj_attn.bias')
|
||||
|
||||
# new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||
|
||||
mapping.append({"old": old_item, "new": new_item})
|
||||
|
||||
return mapping
|
||||
|
||||
|
||||
def renew_vae_attention_paths(old_list, n_shave_prefix_segments=0):
|
||||
"""
|
||||
Updates paths inside attentions to the new naming scheme (local renaming)
|
||||
"""
|
||||
mapping = []
|
||||
for old_item in old_list:
|
||||
new_item = old_item
|
||||
|
||||
new_item = new_item.replace("norm.weight", "group_norm.weight")
|
||||
new_item = new_item.replace("norm.bias", "group_norm.bias")
|
||||
|
||||
new_item = new_item.replace("q.weight", "query.weight")
|
||||
new_item = new_item.replace("q.bias", "query.bias")
|
||||
|
||||
new_item = new_item.replace("k.weight", "key.weight")
|
||||
new_item = new_item.replace("k.bias", "key.bias")
|
||||
|
||||
new_item = new_item.replace("v.weight", "value.weight")
|
||||
new_item = new_item.replace("v.bias", "value.bias")
|
||||
|
||||
new_item = new_item.replace("proj_out.weight", "proj_attn.weight")
|
||||
new_item = new_item.replace("proj_out.bias", "proj_attn.bias")
|
||||
|
||||
new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||
|
||||
mapping.append({"old": old_item, "new": new_item})
|
||||
|
||||
return mapping
|
||||
|
||||
|
||||
def assign_to_checkpoint(
|
||||
paths, checkpoint, old_checkpoint, attention_paths_to_split=None, additional_replacements=None, config=None
|
||||
):
|
||||
"""
|
||||
This does the final conversion step: take locally converted weights and apply a global renaming to them. It splits
|
||||
attention layers, and takes into account additional replacements that may arise.
|
||||
|
||||
Assigns the weights to the new checkpoint.
|
||||
"""
|
||||
assert isinstance(paths, list), "Paths should be a list of dicts containing 'old' and 'new' keys."
|
||||
|
||||
# Splits the attention layers into three variables.
|
||||
if attention_paths_to_split is not None:
|
||||
for path, path_map in attention_paths_to_split.items():
|
||||
old_tensor = old_checkpoint[path]
|
||||
channels = old_tensor.shape[0] // 3
|
||||
|
||||
target_shape = (-1, channels) if len(old_tensor.shape) == 3 else (-1)
|
||||
|
||||
num_heads = old_tensor.shape[0] // config["num_head_channels"] // 3
|
||||
|
||||
old_tensor = old_tensor.reshape((num_heads, 3 * channels // num_heads) + old_tensor.shape[1:])
|
||||
query, key, value = old_tensor.split(channels // num_heads, dim=1)
|
||||
|
||||
checkpoint[path_map["query"]] = query.reshape(target_shape)
|
||||
checkpoint[path_map["key"]] = key.reshape(target_shape)
|
||||
checkpoint[path_map["value"]] = value.reshape(target_shape)
|
||||
|
||||
for path in paths:
|
||||
new_path = path["new"]
|
||||
|
||||
# These have already been assigned
|
||||
if attention_paths_to_split is not None and new_path in attention_paths_to_split:
|
||||
continue
|
||||
|
||||
# Global renaming happens here
|
||||
new_path = new_path.replace("middle_block.0", "mid_block.resnets.0")
|
||||
new_path = new_path.replace("middle_block.1", "mid_block.attentions.0")
|
||||
new_path = new_path.replace("middle_block.2", "mid_block.resnets.1")
|
||||
|
||||
if additional_replacements is not None:
|
||||
for replacement in additional_replacements:
|
||||
new_path = new_path.replace(replacement["old"], replacement["new"])
|
||||
|
||||
# proj_attn.weight has to be converted from conv 1D to linear
|
||||
if "proj_attn.weight" in new_path:
|
||||
checkpoint[new_path] = old_checkpoint[path["old"]][:, :, 0]
|
||||
else:
|
||||
checkpoint[new_path] = old_checkpoint[path["old"]]
|
||||
|
||||
|
||||
def conv_attn_to_linear(checkpoint):
|
||||
keys = list(checkpoint.keys())
|
||||
attn_keys = ["query.weight", "key.weight", "value.weight"]
|
||||
for key in keys:
|
||||
if ".".join(key.split(".")[-2:]) in attn_keys:
|
||||
if checkpoint[key].ndim > 2:
|
||||
checkpoint[key] = checkpoint[key][:, :, 0, 0]
|
||||
elif "proj_attn.weight" in key:
|
||||
if checkpoint[key].ndim > 2:
|
||||
checkpoint[key] = checkpoint[key][:, :, 0]
|
||||
|
||||
|
||||
def convert_ldm_unet_checkpoint(checkpoint, config, path=None, extract_ema=False, controlnet=False):
|
||||
"""
|
||||
Takes a state dict and a config, and returns a converted checkpoint.
|
||||
"""
|
||||
|
||||
# extract state_dict for UNet
|
||||
unet_state_dict = {}
|
||||
keys = list(checkpoint.keys())
|
||||
|
||||
if controlnet:
|
||||
unet_key = "control_model."
|
||||
else:
|
||||
unet_key = "model.diffusion_model."
|
||||
|
||||
# at least a 100 parameters have to start with `model_ema` in order for the checkpoint to be EMA
|
||||
if sum(k.startswith("model_ema") for k in keys) > 100 and extract_ema:
|
||||
print(f"Checkpoint {path} has both EMA and non-EMA weights.")
|
||||
print(
|
||||
"In this conversion only the EMA weights are extracted. If you want to instead extract the non-EMA"
|
||||
" weights (useful to continue fine-tuning), please make sure to remove the `--extract_ema` flag."
|
||||
)
|
||||
for key in keys:
|
||||
if key.startswith("model.diffusion_model"):
|
||||
flat_ema_key = "model_ema." + "".join(key.split(".")[1:])
|
||||
unet_state_dict[key.replace(unet_key, "")] = checkpoint.pop(flat_ema_key)
|
||||
else:
|
||||
if sum(k.startswith("model_ema") for k in keys) > 100:
|
||||
print(
|
||||
"In this conversion only the non-EMA weights are extracted. If you want to instead extract the EMA"
|
||||
" weights (usually better for inference), please make sure to add the `--extract_ema` flag."
|
||||
)
|
||||
|
||||
for key in keys:
|
||||
if key.startswith(unet_key):
|
||||
unet_state_dict[key.replace(unet_key, "")] = checkpoint.pop(key)
|
||||
|
||||
new_checkpoint = {}
|
||||
|
||||
new_checkpoint["time_embedding.linear_1.weight"] = unet_state_dict["time_embed.0.weight"]
|
||||
new_checkpoint["time_embedding.linear_1.bias"] = unet_state_dict["time_embed.0.bias"]
|
||||
new_checkpoint["time_embedding.linear_2.weight"] = unet_state_dict["time_embed.2.weight"]
|
||||
new_checkpoint["time_embedding.linear_2.bias"] = unet_state_dict["time_embed.2.bias"]
|
||||
|
||||
if config["class_embed_type"] is None:
|
||||
# No parameters to port
|
||||
...
|
||||
elif config["class_embed_type"] == "timestep" or config["class_embed_type"] == "projection":
|
||||
new_checkpoint["class_embedding.linear_1.weight"] = unet_state_dict["label_emb.0.0.weight"]
|
||||
new_checkpoint["class_embedding.linear_1.bias"] = unet_state_dict["label_emb.0.0.bias"]
|
||||
new_checkpoint["class_embedding.linear_2.weight"] = unet_state_dict["label_emb.0.2.weight"]
|
||||
new_checkpoint["class_embedding.linear_2.bias"] = unet_state_dict["label_emb.0.2.bias"]
|
||||
else:
|
||||
raise NotImplementedError(f"Not implemented `class_embed_type`: {config['class_embed_type']}")
|
||||
|
||||
new_checkpoint["conv_in.weight"] = unet_state_dict["input_blocks.0.0.weight"]
|
||||
new_checkpoint["conv_in.bias"] = unet_state_dict["input_blocks.0.0.bias"]
|
||||
|
||||
if not controlnet:
|
||||
new_checkpoint["conv_norm_out.weight"] = unet_state_dict["out.0.weight"]
|
||||
new_checkpoint["conv_norm_out.bias"] = unet_state_dict["out.0.bias"]
|
||||
new_checkpoint["conv_out.weight"] = unet_state_dict["out.2.weight"]
|
||||
new_checkpoint["conv_out.bias"] = unet_state_dict["out.2.bias"]
|
||||
|
||||
# Retrieves the keys for the input blocks only
|
||||
num_input_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "input_blocks" in layer})
|
||||
input_blocks = {
|
||||
layer_id: [key for key in unet_state_dict if f"input_blocks.{layer_id}" in key]
|
||||
for layer_id in range(num_input_blocks)
|
||||
}
|
||||
|
||||
# Retrieves the keys for the middle blocks only
|
||||
num_middle_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "middle_block" in layer})
|
||||
middle_blocks = {
|
||||
layer_id: [key for key in unet_state_dict if f"middle_block.{layer_id}" in key]
|
||||
for layer_id in range(num_middle_blocks)
|
||||
}
|
||||
|
||||
# Retrieves the keys for the output blocks only
|
||||
num_output_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "output_blocks" in layer})
|
||||
output_blocks = {
|
||||
layer_id: [key for key in unet_state_dict if f"output_blocks.{layer_id}" in key]
|
||||
for layer_id in range(num_output_blocks)
|
||||
}
|
||||
|
||||
for i in range(1, num_input_blocks):
|
||||
block_id = (i - 1) // (config["layers_per_block"] + 1)
|
||||
layer_in_block_id = (i - 1) % (config["layers_per_block"] + 1)
|
||||
|
||||
resnets = [
|
||||
key for key in input_blocks[i] if f"input_blocks.{i}.0" in key and f"input_blocks.{i}.0.op" not in key
|
||||
]
|
||||
attentions = [key for key in input_blocks[i] if f"input_blocks.{i}.1" in key]
|
||||
|
||||
if f"input_blocks.{i}.0.op.weight" in unet_state_dict:
|
||||
new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.weight"] = unet_state_dict.pop(
|
||||
f"input_blocks.{i}.0.op.weight"
|
||||
)
|
||||
new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.bias"] = unet_state_dict.pop(
|
||||
f"input_blocks.{i}.0.op.bias"
|
||||
)
|
||||
|
||||
paths = renew_resnet_paths(resnets)
|
||||
meta_path = {"old": f"input_blocks.{i}.0", "new": f"down_blocks.{block_id}.resnets.{layer_in_block_id}"}
|
||||
assign_to_checkpoint(
|
||||
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||
)
|
||||
|
||||
if len(attentions):
|
||||
paths = renew_attention_paths(attentions)
|
||||
meta_path = {"old": f"input_blocks.{i}.1", "new": f"down_blocks.{block_id}.attentions.{layer_in_block_id}"}
|
||||
assign_to_checkpoint(
|
||||
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||
)
|
||||
|
||||
resnet_0 = middle_blocks[0]
|
||||
attentions = middle_blocks[1]
|
||||
resnet_1 = middle_blocks[2]
|
||||
|
||||
resnet_0_paths = renew_resnet_paths(resnet_0)
|
||||
assign_to_checkpoint(resnet_0_paths, new_checkpoint, unet_state_dict, config=config)
|
||||
|
||||
resnet_1_paths = renew_resnet_paths(resnet_1)
|
||||
assign_to_checkpoint(resnet_1_paths, new_checkpoint, unet_state_dict, config=config)
|
||||
|
||||
attentions_paths = renew_attention_paths(attentions)
|
||||
meta_path = {"old": "middle_block.1", "new": "mid_block.attentions.0"}
|
||||
assign_to_checkpoint(
|
||||
attentions_paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||
)
|
||||
|
||||
for i in range(num_output_blocks):
|
||||
block_id = i // (config["layers_per_block"] + 1)
|
||||
layer_in_block_id = i % (config["layers_per_block"] + 1)
|
||||
output_block_layers = [shave_segments(name, 2) for name in output_blocks[i]]
|
||||
output_block_list = {}
|
||||
|
||||
for layer in output_block_layers:
|
||||
layer_id, layer_name = layer.split(".")[0], shave_segments(layer, 1)
|
||||
if layer_id in output_block_list:
|
||||
output_block_list[layer_id].append(layer_name)
|
||||
else:
|
||||
output_block_list[layer_id] = [layer_name]
|
||||
|
||||
if len(output_block_list) > 1:
|
||||
resnets = [key for key in output_blocks[i] if f"output_blocks.{i}.0" in key]
|
||||
attentions = [key for key in output_blocks[i] if f"output_blocks.{i}.1" in key]
|
||||
|
||||
resnet_0_paths = renew_resnet_paths(resnets)
|
||||
paths = renew_resnet_paths(resnets)
|
||||
|
||||
meta_path = {"old": f"output_blocks.{i}.0", "new": f"up_blocks.{block_id}.resnets.{layer_in_block_id}"}
|
||||
assign_to_checkpoint(
|
||||
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||
)
|
||||
|
||||
output_block_list = {k: sorted(v) for k, v in output_block_list.items()}
|
||||
if ["conv.bias", "conv.weight"] in output_block_list.values():
|
||||
index = list(output_block_list.values()).index(["conv.bias", "conv.weight"])
|
||||
new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.weight"] = unet_state_dict[
|
||||
f"output_blocks.{i}.{index}.conv.weight"
|
||||
]
|
||||
new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.bias"] = unet_state_dict[
|
||||
f"output_blocks.{i}.{index}.conv.bias"
|
||||
]
|
||||
|
||||
# Clear attentions as they have been attributed above.
|
||||
if len(attentions) == 2:
|
||||
attentions = []
|
||||
|
||||
if len(attentions):
|
||||
paths = renew_attention_paths(attentions)
|
||||
meta_path = {
|
||||
"old": f"output_blocks.{i}.1",
|
||||
"new": f"up_blocks.{block_id}.attentions.{layer_in_block_id}",
|
||||
}
|
||||
assign_to_checkpoint(
|
||||
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||
)
|
||||
else:
|
||||
resnet_0_paths = renew_resnet_paths(output_block_layers, n_shave_prefix_segments=1)
|
||||
for path in resnet_0_paths:
|
||||
old_path = ".".join(["output_blocks", str(i), path["old"]])
|
||||
new_path = ".".join(["up_blocks", str(block_id), "resnets", str(layer_in_block_id), path["new"]])
|
||||
|
||||
new_checkpoint[new_path] = unet_state_dict[old_path]
|
||||
|
||||
if controlnet:
|
||||
# conditioning embedding
|
||||
|
||||
orig_index = 0
|
||||
|
||||
new_checkpoint["controlnet_cond_embedding.conv_in.weight"] = unet_state_dict.pop(
|
||||
f"input_hint_block.{orig_index}.weight"
|
||||
)
|
||||
new_checkpoint["controlnet_cond_embedding.conv_in.bias"] = unet_state_dict.pop(
|
||||
f"input_hint_block.{orig_index}.bias"
|
||||
)
|
||||
|
||||
orig_index += 2
|
||||
|
||||
diffusers_index = 0
|
||||
|
||||
while diffusers_index < 6:
|
||||
new_checkpoint[f"controlnet_cond_embedding.blocks.{diffusers_index}.weight"] = unet_state_dict.pop(
|
||||
f"input_hint_block.{orig_index}.weight"
|
||||
)
|
||||
new_checkpoint[f"controlnet_cond_embedding.blocks.{diffusers_index}.bias"] = unet_state_dict.pop(
|
||||
f"input_hint_block.{orig_index}.bias"
|
||||
)
|
||||
diffusers_index += 1
|
||||
orig_index += 2
|
||||
|
||||
new_checkpoint["controlnet_cond_embedding.conv_out.weight"] = unet_state_dict.pop(
|
||||
f"input_hint_block.{orig_index}.weight"
|
||||
)
|
||||
new_checkpoint["controlnet_cond_embedding.conv_out.bias"] = unet_state_dict.pop(
|
||||
f"input_hint_block.{orig_index}.bias"
|
||||
)
|
||||
|
||||
# down blocks
|
||||
for i in range(num_input_blocks):
|
||||
new_checkpoint[f"controlnet_down_blocks.{i}.weight"] = unet_state_dict.pop(f"zero_convs.{i}.0.weight")
|
||||
new_checkpoint[f"controlnet_down_blocks.{i}.bias"] = unet_state_dict.pop(f"zero_convs.{i}.0.bias")
|
||||
|
||||
# mid block
|
||||
new_checkpoint["controlnet_mid_block.weight"] = unet_state_dict.pop("middle_block_out.0.weight")
|
||||
new_checkpoint["controlnet_mid_block.bias"] = unet_state_dict.pop("middle_block_out.0.bias")
|
||||
|
||||
return new_checkpoint
|
||||
|
||||
|
||||
def convert_ldm_vae_checkpoint(checkpoint, config):
|
||||
# extract state dict for VAE
|
||||
vae_state_dict = {}
|
||||
keys = list(checkpoint.keys())
|
||||
vae_key = "first_stage_model." if any(k.startswith("first_stage_model.") for k in keys) else ""
|
||||
for key in keys:
|
||||
if key.startswith(vae_key):
|
||||
vae_state_dict[key.replace(vae_key, "")] = checkpoint.get(key)
|
||||
|
||||
new_checkpoint = {}
|
||||
|
||||
new_checkpoint["encoder.conv_in.weight"] = vae_state_dict["encoder.conv_in.weight"]
|
||||
new_checkpoint["encoder.conv_in.bias"] = vae_state_dict["encoder.conv_in.bias"]
|
||||
new_checkpoint["encoder.conv_out.weight"] = vae_state_dict["encoder.conv_out.weight"]
|
||||
new_checkpoint["encoder.conv_out.bias"] = vae_state_dict["encoder.conv_out.bias"]
|
||||
new_checkpoint["encoder.conv_norm_out.weight"] = vae_state_dict["encoder.norm_out.weight"]
|
||||
new_checkpoint["encoder.conv_norm_out.bias"] = vae_state_dict["encoder.norm_out.bias"]
|
||||
|
||||
new_checkpoint["decoder.conv_in.weight"] = vae_state_dict["decoder.conv_in.weight"]
|
||||
new_checkpoint["decoder.conv_in.bias"] = vae_state_dict["decoder.conv_in.bias"]
|
||||
new_checkpoint["decoder.conv_out.weight"] = vae_state_dict["decoder.conv_out.weight"]
|
||||
new_checkpoint["decoder.conv_out.bias"] = vae_state_dict["decoder.conv_out.bias"]
|
||||
new_checkpoint["decoder.conv_norm_out.weight"] = vae_state_dict["decoder.norm_out.weight"]
|
||||
new_checkpoint["decoder.conv_norm_out.bias"] = vae_state_dict["decoder.norm_out.bias"]
|
||||
|
||||
new_checkpoint["quant_conv.weight"] = vae_state_dict["quant_conv.weight"]
|
||||
new_checkpoint["quant_conv.bias"] = vae_state_dict["quant_conv.bias"]
|
||||
new_checkpoint["post_quant_conv.weight"] = vae_state_dict["post_quant_conv.weight"]
|
||||
new_checkpoint["post_quant_conv.bias"] = vae_state_dict["post_quant_conv.bias"]
|
||||
|
||||
# Retrieves the keys for the encoder down blocks only
|
||||
num_down_blocks = len({".".join(layer.split(".")[:3]) for layer in vae_state_dict if "encoder.down" in layer})
|
||||
down_blocks = {
|
||||
layer_id: [key for key in vae_state_dict if f"down.{layer_id}" in key] for layer_id in range(num_down_blocks)
|
||||
}
|
||||
|
||||
# Retrieves the keys for the decoder up blocks only
|
||||
num_up_blocks = len({".".join(layer.split(".")[:3]) for layer in vae_state_dict if "decoder.up" in layer})
|
||||
up_blocks = {
|
||||
layer_id: [key for key in vae_state_dict if f"up.{layer_id}" in key] for layer_id in range(num_up_blocks)
|
||||
}
|
||||
|
||||
for i in range(num_down_blocks):
|
||||
resnets = [key for key in down_blocks[i] if f"down.{i}" in key and f"down.{i}.downsample" not in key]
|
||||
|
||||
if f"encoder.down.{i}.downsample.conv.weight" in vae_state_dict:
|
||||
new_checkpoint[f"encoder.down_blocks.{i}.downsamplers.0.conv.weight"] = vae_state_dict.pop(
|
||||
f"encoder.down.{i}.downsample.conv.weight"
|
||||
)
|
||||
new_checkpoint[f"encoder.down_blocks.{i}.downsamplers.0.conv.bias"] = vae_state_dict.pop(
|
||||
f"encoder.down.{i}.downsample.conv.bias"
|
||||
)
|
||||
|
||||
paths = renew_vae_resnet_paths(resnets)
|
||||
meta_path = {"old": f"down.{i}.block", "new": f"down_blocks.{i}.resnets"}
|
||||
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||
|
||||
mid_resnets = [key for key in vae_state_dict if "encoder.mid.block" in key]
|
||||
num_mid_res_blocks = 2
|
||||
for i in range(1, num_mid_res_blocks + 1):
|
||||
resnets = [key for key in mid_resnets if f"encoder.mid.block_{i}" in key]
|
||||
|
||||
paths = renew_vae_resnet_paths(resnets)
|
||||
meta_path = {"old": f"mid.block_{i}", "new": f"mid_block.resnets.{i - 1}"}
|
||||
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||
|
||||
mid_attentions = [key for key in vae_state_dict if "encoder.mid.attn" in key]
|
||||
paths = renew_vae_attention_paths(mid_attentions)
|
||||
meta_path = {"old": "mid.attn_1", "new": "mid_block.attentions.0"}
|
||||
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||
conv_attn_to_linear(new_checkpoint)
|
||||
|
||||
for i in range(num_up_blocks):
|
||||
block_id = num_up_blocks - 1 - i
|
||||
resnets = [
|
||||
key for key in up_blocks[block_id] if f"up.{block_id}" in key and f"up.{block_id}.upsample" not in key
|
||||
]
|
||||
|
||||
if f"decoder.up.{block_id}.upsample.conv.weight" in vae_state_dict:
|
||||
new_checkpoint[f"decoder.up_blocks.{i}.upsamplers.0.conv.weight"] = vae_state_dict[
|
||||
f"decoder.up.{block_id}.upsample.conv.weight"
|
||||
]
|
||||
new_checkpoint[f"decoder.up_blocks.{i}.upsamplers.0.conv.bias"] = vae_state_dict[
|
||||
f"decoder.up.{block_id}.upsample.conv.bias"
|
||||
]
|
||||
|
||||
paths = renew_vae_resnet_paths(resnets)
|
||||
meta_path = {"old": f"up.{block_id}.block", "new": f"up_blocks.{i}.resnets"}
|
||||
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||
|
||||
mid_resnets = [key for key in vae_state_dict if "decoder.mid.block" in key]
|
||||
num_mid_res_blocks = 2
|
||||
for i in range(1, num_mid_res_blocks + 1):
|
||||
resnets = [key for key in mid_resnets if f"decoder.mid.block_{i}" in key]
|
||||
|
||||
paths = renew_vae_resnet_paths(resnets)
|
||||
meta_path = {"old": f"mid.block_{i}", "new": f"mid_block.resnets.{i - 1}"}
|
||||
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||
|
||||
mid_attentions = [key for key in vae_state_dict if "decoder.mid.attn" in key]
|
||||
paths = renew_vae_attention_paths(mid_attentions)
|
||||
meta_path = {"old": "mid.attn_1", "new": "mid_block.attentions.0"}
|
||||
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||
conv_attn_to_linear(new_checkpoint)
|
||||
return new_checkpoint
|
||||
|
||||
|
||||
def convert_ldm_clip_checkpoint(checkpoint):
|
||||
text_model = CLIPTextModel.from_pretrained("openai/clip-vit-large-patch14")
|
||||
keys = list(checkpoint.keys())
|
||||
|
||||
text_model_dict = {}
|
||||
|
||||
for key in keys:
|
||||
if key.startswith("cond_stage_model.transformer"):
|
||||
text_model_dict[key[len("cond_stage_model.transformer.") :]] = checkpoint[key]
|
||||
|
||||
text_model.load_state_dict(text_model_dict)
|
||||
|
||||
return text_model
|
||||
|
||||
|
||||
textenc_conversion_lst = [
|
||||
("cond_stage_model.model.positional_embedding", "text_model.embeddings.position_embedding.weight"),
|
||||
("cond_stage_model.model.token_embedding.weight", "text_model.embeddings.token_embedding.weight"),
|
||||
("cond_stage_model.model.ln_final.weight", "text_model.final_layer_norm.weight"),
|
||||
("cond_stage_model.model.ln_final.bias", "text_model.final_layer_norm.bias"),
|
||||
]
|
||||
textenc_conversion_map = {x[0]: x[1] for x in textenc_conversion_lst}
|
||||
|
||||
textenc_transformer_conversion_lst = [
|
||||
# (stable-diffusion, HF Diffusers)
|
||||
("resblocks.", "text_model.encoder.layers."),
|
||||
("ln_1", "layer_norm1"),
|
||||
("ln_2", "layer_norm2"),
|
||||
(".c_fc.", ".fc1."),
|
||||
(".c_proj.", ".fc2."),
|
||||
(".attn", ".self_attn"),
|
||||
("ln_final.", "transformer.text_model.final_layer_norm."),
|
||||
("token_embedding.weight", "transformer.text_model.embeddings.token_embedding.weight"),
|
||||
("positional_embedding", "transformer.text_model.embeddings.position_embedding.weight"),
|
||||
]
|
||||
protected = {re.escape(x[0]): x[1] for x in textenc_transformer_conversion_lst}
|
||||
textenc_pattern = re.compile("|".join(protected.keys()))
|
||||
@@ -0,0 +1,154 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2023, Haofan Wang, Qixun Wang, All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
""" Conversion script for the LoRA's safetensors checkpoints. """
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from diffusers import StableDiffusionPipeline
|
||||
import pdb
|
||||
|
||||
|
||||
|
||||
def convert_motion_lora_ckpt_to_diffusers(pipeline, state_dict, alpha=1.0):
|
||||
# directly update weight in diffusers model
|
||||
for key in state_dict:
|
||||
# only process lora down key
|
||||
if "up." in key: continue
|
||||
|
||||
up_key = key.replace(".down.", ".up.")
|
||||
model_key = key.replace("processor.", "").replace("_lora", "").replace("down.", "").replace("up.", "")
|
||||
model_key = model_key.replace("to_out.", "to_out.0.")
|
||||
layer_infos = model_key.split(".")[:-1]
|
||||
|
||||
curr_layer = pipeline.unet
|
||||
while len(layer_infos) > 0:
|
||||
temp_name = layer_infos.pop(0)
|
||||
curr_layer = curr_layer.__getattr__(temp_name)
|
||||
|
||||
weight_down = state_dict[key]
|
||||
weight_up = state_dict[up_key]
|
||||
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).to(curr_layer.weight.data.device)
|
||||
|
||||
return pipeline
|
||||
|
||||
|
||||
|
||||
def convert_lora(pipeline, state_dict, LORA_PREFIX_UNET="lora_unet", LORA_PREFIX_TEXT_ENCODER="lora_te", alpha=0.6):
|
||||
# load base model
|
||||
# pipeline = StableDiffusionPipeline.from_pretrained(base_model_path, torch_dtype=torch.float32)
|
||||
|
||||
# load LoRA weight from .safetensors
|
||||
# state_dict = load_file(checkpoint_path)
|
||||
|
||||
visited = []
|
||||
|
||||
# directly update weight in diffusers model
|
||||
for key in state_dict:
|
||||
# it is suggested to print out the key, it usually will be something like below
|
||||
# "lora_te_text_model_encoder_layers_0_self_attn_k_proj.lora_down.weight"
|
||||
|
||||
# as we have set the alpha beforehand, so just skip
|
||||
if ".alpha" in key or key in visited:
|
||||
continue
|
||||
|
||||
if "text" in key:
|
||||
layer_infos = key.split(".")[0].split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
|
||||
curr_layer = pipeline.text_encoder
|
||||
else:
|
||||
layer_infos = key.split(".")[0].split(LORA_PREFIX_UNET + "_")[-1].split("_")
|
||||
curr_layer = pipeline.unet
|
||||
|
||||
# find the target layer
|
||||
temp_name = layer_infos.pop(0)
|
||||
while len(layer_infos) > -1:
|
||||
try:
|
||||
curr_layer = curr_layer.__getattr__(temp_name)
|
||||
if len(layer_infos) > 0:
|
||||
temp_name = layer_infos.pop(0)
|
||||
elif len(layer_infos) == 0:
|
||||
break
|
||||
except Exception:
|
||||
if len(temp_name) > 0:
|
||||
temp_name += "_" + layer_infos.pop(0)
|
||||
else:
|
||||
temp_name = layer_infos.pop(0)
|
||||
|
||||
pair_keys = []
|
||||
if "lora_down" in key:
|
||||
pair_keys.append(key.replace("lora_down", "lora_up"))
|
||||
pair_keys.append(key)
|
||||
else:
|
||||
pair_keys.append(key)
|
||||
pair_keys.append(key.replace("lora_up", "lora_down"))
|
||||
|
||||
# update weight
|
||||
if len(state_dict[pair_keys[0]].shape) == 4:
|
||||
weight_up = state_dict[pair_keys[0]].squeeze(3).squeeze(2).to(torch.float32)
|
||||
weight_down = state_dict[pair_keys[1]].squeeze(3).squeeze(2).to(torch.float32)
|
||||
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3).to(curr_layer.weight.data.device)
|
||||
else:
|
||||
weight_up = state_dict[pair_keys[0]].to(torch.float32)
|
||||
weight_down = state_dict[pair_keys[1]].to(torch.float32)
|
||||
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).to(curr_layer.weight.data.device)
|
||||
|
||||
# update visited list
|
||||
for item in pair_keys:
|
||||
visited.append(item)
|
||||
|
||||
return pipeline
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--base_model_path", default=None, type=str, required=True, help="Path to the base model in diffusers format."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint_path", default=None, type=str, required=True, help="Path to the checkpoint to convert."
|
||||
)
|
||||
parser.add_argument("--dump_path", default=None, type=str, required=True, help="Path to the output model.")
|
||||
parser.add_argument(
|
||||
"--lora_prefix_unet", default="lora_unet", type=str, help="The prefix of UNet weight in safetensors"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora_prefix_text_encoder",
|
||||
default="lora_te",
|
||||
type=str,
|
||||
help="The prefix of text encoder weight in safetensors",
|
||||
)
|
||||
parser.add_argument("--alpha", default=0.75, type=float, help="The merging ratio in W = W0 + alpha * deltaW")
|
||||
parser.add_argument(
|
||||
"--to_safetensors", action="store_true", help="Whether to store pipeline in safetensors format or not."
|
||||
)
|
||||
parser.add_argument("--device", type=str, help="Device to use (e.g. cpu, cuda:0, cuda:1, etc.)")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
base_model_path = args.base_model_path
|
||||
checkpoint_path = args.checkpoint_path
|
||||
dump_path = args.dump_path
|
||||
lora_prefix_unet = args.lora_prefix_unet
|
||||
lora_prefix_text_encoder = args.lora_prefix_text_encoder
|
||||
alpha = args.alpha
|
||||
|
||||
pipe = convert(base_model_path, checkpoint_path, lora_prefix_unet, lora_prefix_text_encoder, alpha)
|
||||
|
||||
pipe = pipe.to(args.device)
|
||||
pipe.save_pretrained(args.dump_path, safe_serialization=args.to_safetensors)
|
||||
@@ -0,0 +1,148 @@
|
||||
import os
|
||||
import functools
|
||||
import logging
|
||||
import sys
|
||||
import imageio
|
||||
import atexit
|
||||
import importlib
|
||||
import torch
|
||||
import torchvision
|
||||
import numpy as np
|
||||
from termcolor import colored
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
def instantiate_from_config(config, **additional_kwargs):
|
||||
if not "target" in config:
|
||||
if config == '__is_first_stage__':
|
||||
return None
|
||||
elif config == "__is_unconditional__":
|
||||
return None
|
||||
raise KeyError("Expected key `target` to instantiate.")
|
||||
|
||||
additional_kwargs.update(config.get("kwargs", dict()))
|
||||
return get_obj_from_str(config["target"])(**additional_kwargs)
|
||||
|
||||
|
||||
def get_obj_from_str(string, reload=False):
|
||||
module, cls = string.rsplit(".", 1)
|
||||
if reload:
|
||||
module_imp = importlib.import_module(module)
|
||||
importlib.reload(module_imp)
|
||||
return getattr(importlib.import_module(module, package=None), cls)
|
||||
|
||||
|
||||
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=8):
|
||||
videos = rearrange(videos, "b c t h w -> t b c h w")
|
||||
outputs = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=n_rows)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
if rescale:
|
||||
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
|
||||
x = (x * 255).numpy().astype(np.uint8)
|
||||
outputs.append(x)
|
||||
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
imageio.mimsave(path, outputs, fps=fps)
|
||||
|
||||
|
||||
# Logger utils are copied from detectron2
|
||||
class _ColorfulFormatter(logging.Formatter):
|
||||
def __init__(self, *args, **kwargs):
|
||||
self._root_name = kwargs.pop("root_name") + "."
|
||||
self._abbrev_name = kwargs.pop("abbrev_name", "")
|
||||
if len(self._abbrev_name):
|
||||
self._abbrev_name = self._abbrev_name + "."
|
||||
super(_ColorfulFormatter, self).__init__(*args, **kwargs)
|
||||
|
||||
def formatMessage(self, record):
|
||||
record.name = record.name.replace(self._root_name, self._abbrev_name)
|
||||
log = super(_ColorfulFormatter, self).formatMessage(record)
|
||||
if record.levelno == logging.WARNING:
|
||||
prefix = colored("WARNING", "red", attrs=["blink"])
|
||||
elif record.levelno == logging.ERROR or record.levelno == logging.CRITICAL:
|
||||
prefix = colored("ERROR", "red", attrs=["blink", "underline"])
|
||||
else:
|
||||
return log
|
||||
return prefix + " " + log
|
||||
|
||||
|
||||
# cache the opened file object, so that different calls to `setup_logger`
|
||||
# with the same file name can safely write to the same file.
|
||||
@functools.lru_cache(maxsize=None)
|
||||
def _cached_log_stream(filename):
|
||||
# use 1K buffer if writing to cloud storage
|
||||
io = open(filename, "a", buffering=1024 if "://" in filename else -1)
|
||||
atexit.register(io.close)
|
||||
return io
|
||||
|
||||
@functools.lru_cache()
|
||||
def setup_logger(output, distributed_rank, color=True, name='AnimateDiff', abbrev_name=None):
|
||||
logger = logging.getLogger(name)
|
||||
logger.setLevel(logging.DEBUG)
|
||||
logger.propagate = False
|
||||
|
||||
if abbrev_name is None:
|
||||
abbrev_name = 'AD'
|
||||
plain_formatter = logging.Formatter(
|
||||
"[%(asctime)s] %(name)s:%(lineno)d %(levelname)s: %(message)s", datefmt="%m/%d %H:%M:%S"
|
||||
)
|
||||
|
||||
# stdout logging: master only
|
||||
if distributed_rank == 0:
|
||||
ch = logging.StreamHandler(stream=sys.stdout)
|
||||
ch.setLevel(logging.DEBUG)
|
||||
if color:
|
||||
formatter = _ColorfulFormatter(
|
||||
colored("[%(asctime)s %(name)s:%(lineno)d]: ", "green") + "%(message)s",
|
||||
datefmt="%m/%d %H:%M:%S",
|
||||
root_name=name,
|
||||
abbrev_name=str(abbrev_name),
|
||||
)
|
||||
else:
|
||||
formatter = plain_formatter
|
||||
ch.setFormatter(formatter)
|
||||
logger.addHandler(ch)
|
||||
|
||||
# file logging: all workers
|
||||
if output is not None:
|
||||
if output.endswith(".txt") or output.endswith(".log"):
|
||||
filename = output
|
||||
else:
|
||||
filename = os.path.join(output, "log.txt")
|
||||
if distributed_rank > 0:
|
||||
filename = filename + ".rank{}".format(distributed_rank)
|
||||
os.makedirs(os.path.dirname(filename), exist_ok=True)
|
||||
|
||||
fh = logging.StreamHandler(_cached_log_stream(filename))
|
||||
fh.setLevel(logging.DEBUG)
|
||||
fh.setFormatter(plain_formatter)
|
||||
logger.addHandler(fh)
|
||||
|
||||
return logger
|
||||
|
||||
|
||||
def format_time(elapsed_time):
|
||||
# Time thresholds
|
||||
minute = 60
|
||||
hour = 60 * minute
|
||||
day = 24 * hour
|
||||
|
||||
days, remainder = divmod(elapsed_time, day)
|
||||
hours, remainder = divmod(remainder, hour)
|
||||
minutes, seconds = divmod(remainder, minute)
|
||||
|
||||
formatted_time = ""
|
||||
|
||||
if days > 0:
|
||||
formatted_time += f"{int(days)} days "
|
||||
if hours > 0:
|
||||
formatted_time += f"{int(hours)} hours "
|
||||
if minutes > 0:
|
||||
formatted_time += f"{int(minutes)} minutes "
|
||||
if seconds > 0:
|
||||
formatted_time += f"{seconds:.2f} seconds"
|
||||
|
||||
return formatted_time.strip()
|
||||
@@ -0,0 +1,99 @@
|
||||
|
||||
output_dir: "output/cameractrl_model"
|
||||
pretrained_model_path: "[replace with SD1.5 root path]"
|
||||
unet_subfolder: "unet_webvidlora_v3"
|
||||
|
||||
train_data:
|
||||
root_path: "[replace RealEstate10K root path]"
|
||||
annotation_json: "annotations/train.json"
|
||||
sample_stride: 8
|
||||
sample_n_frames: 16
|
||||
relative_pose: true
|
||||
zero_t_first_frame: true
|
||||
sample_size: [256, 384]
|
||||
rescale_fxy: true
|
||||
shuffle_frames: true
|
||||
use_flip: true
|
||||
|
||||
validation_data:
|
||||
root_path: "[replace RealEstate10K root path]"
|
||||
annotation_json: "annotations/validation.json"
|
||||
sample_stride: 8
|
||||
sample_n_frames: 16
|
||||
relative_pose: true
|
||||
zero_t_first_frame: true
|
||||
sample_size: [256, 384]
|
||||
rescale_fxy: true
|
||||
shuffle_frames: false
|
||||
use_flip: false
|
||||
return_clip_name: true
|
||||
|
||||
unet_additional_kwargs:
|
||||
use_motion_module : true
|
||||
motion_module_resolutions : [ 1,2,4,8 ]
|
||||
unet_use_cross_frame_attention : false
|
||||
unet_use_temporal_attention : false
|
||||
motion_module_mid_block: false
|
||||
motion_module_type: Vanilla
|
||||
motion_module_kwargs:
|
||||
num_attention_heads : 8
|
||||
num_transformer_block : 1
|
||||
attention_block_types : [ "Temporal_Self", "Temporal_Self" ]
|
||||
temporal_position_encoding : true
|
||||
temporal_position_encoding_max_len : 32
|
||||
temporal_attention_dim_div : 1
|
||||
zero_initialize : false
|
||||
|
||||
lora_rank: 2
|
||||
lora_scale: 1.0
|
||||
lora_ckpt: "[Replace with RealEstate10k image LoRA ckpt]"
|
||||
motion_module_ckpt: "[Replace with ADV3 motion module]"
|
||||
|
||||
pose_encoder_kwargs:
|
||||
downscale_factor: 8
|
||||
channels: [320, 640, 1280, 1280]
|
||||
nums_rb: 2
|
||||
cin: 384
|
||||
ksize: 1
|
||||
sk: true
|
||||
use_conv: false
|
||||
compression_factor: 1
|
||||
temporal_attention_nhead: 8
|
||||
attention_block_types: ["Temporal_Self", ]
|
||||
temporal_position_encoding: true
|
||||
temporal_position_encoding_max_len: 16
|
||||
attention_processor_kwargs:
|
||||
add_spatial: false
|
||||
spatial_attn_names: 'attn1'
|
||||
add_temporal: true
|
||||
temporal_attn_names: '0'
|
||||
pose_feature_dimensions: [320, 640, 1280, 1280]
|
||||
query_condition: true
|
||||
key_value_condition: true
|
||||
scale: 1.0
|
||||
noise_scheduler_kwargs:
|
||||
num_train_timesteps: 1000
|
||||
beta_start: 0.00085
|
||||
beta_end: 0.012
|
||||
beta_schedule: "linear"
|
||||
steps_offset: 1
|
||||
clip_sample: false
|
||||
|
||||
do_sanity_check: true
|
||||
|
||||
max_train_epoch: -1
|
||||
max_train_steps: 25000
|
||||
validation_steps: 1000
|
||||
validation_steps_tuple: [2, ]
|
||||
|
||||
learning_rate: 1.e-4
|
||||
|
||||
num_workers: 8
|
||||
train_batch_size: 2
|
||||
checkpointing_epochs: -1
|
||||
checkpointing_steps: 1000
|
||||
|
||||
mixed_precision_training: true
|
||||
global_seed: 42
|
||||
logger_interval: 10
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
output_dir: "output/image_lora"
|
||||
pretrained_model_path: "[replace with SD1.5 root path]"
|
||||
unet_subfolder: "unet_webvidlora_v3"
|
||||
|
||||
train_data:
|
||||
root_path: "[replace RealEstate10K root path]"
|
||||
annotation_json: "annotations/train.json"
|
||||
sample_size: [256, 384]
|
||||
is_image: true
|
||||
|
||||
validation_data:
|
||||
prompts:
|
||||
- "a kitchen with large windows overlooking a lake"
|
||||
- "a hallway leading to a laundry room with a washer and dryer"
|
||||
- "a bedroom with a bed, chair, and window blinds"
|
||||
- "a living room with a couch, piano and tv"
|
||||
- "a bathroom with a walk in shower and a bathtub"
|
||||
- "a bedroom with an exercise bike and a bed"
|
||||
- "a deck with chairs overlooking a wooded area"
|
||||
- "the porch of a house is decorated with christmas lights"
|
||||
- "a hallway with a staircase and a painting on the wall"
|
||||
- "a large brown house with green grass and bushes"
|
||||
- "a kitchen with white appliances and wooden cabinets"
|
||||
- "a kitchen with wooden cabinets and counter tops"
|
||||
- "a bedroom with a large bed and a chandelier"
|
||||
- "a kitchen and dining room with hardwood floors"
|
||||
- "a hallway leading to a bedroom and bathroom"
|
||||
- "a dining room with a chandelier hanging from the ceiling"
|
||||
num_inference_steps: 50
|
||||
guidance_scale: 8.
|
||||
|
||||
noise_scheduler_kwargs:
|
||||
num_train_timesteps: 1000
|
||||
beta_start: 0.00085
|
||||
beta_end: 0.012
|
||||
beta_schedule: "scaled_linear"
|
||||
steps_offset: 1
|
||||
clip_sample: false
|
||||
|
||||
do_sanity_check: true
|
||||
max_train_epoch: -1
|
||||
max_train_steps: 10000
|
||||
validation_steps: 2000
|
||||
validation_steps_tuple: [2,]
|
||||
|
||||
learning_rate: 1.e-4
|
||||
|
||||
lora_rank: 2
|
||||
|
||||
num_workers: 8
|
||||
train_batch_size: 32
|
||||
|
||||
checkpointing_epochs: -1
|
||||
checkpointing_steps: 2000
|
||||
|
||||
mixed_precision_training: true
|
||||
enable_xformers_memory_efficient_attention: false
|
||||
|
||||
global_seed: 42
|
||||
logger_interval: 10
|
||||
@@ -0,0 +1,16 @@
|
||||
#!/bin/bash
|
||||
|
||||
set -x
|
||||
|
||||
CONFIG=$1
|
||||
GPUS=$2
|
||||
PT_SCRIPT=$3
|
||||
RANDOM_PORT=$((49152 + RANDOM % 16384))
|
||||
|
||||
python -m torch.distributed.launch \
|
||||
--nproc_per_node=$GPUS \
|
||||
--master_port=$RANDOM_PORT \
|
||||
${PT_SCRIPT} \
|
||||
--config=${CONFIG} \
|
||||
--launcher=pytorch \
|
||||
--port=${RANDOM_PORT}
|
||||
@@ -0,0 +1,315 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
from packaging import version as pver
|
||||
from einops import rearrange
|
||||
from safetensors import safe_open
|
||||
|
||||
from omegaconf import OmegaConf
|
||||
from diffusers import (
|
||||
AutoencoderKL,
|
||||
DDIMScheduler
|
||||
)
|
||||
from transformers import CLIPTextModel, CLIPTokenizer
|
||||
from diffusers.pipelines.stable_diffusion.convert_from_ckpt import convert_ldm_vae_checkpoint, \
|
||||
convert_ldm_clip_checkpoint
|
||||
|
||||
from cameractrl.utils.util import save_videos_grid
|
||||
from cameractrl.models.unet import UNet3DConditionModelPoseCond
|
||||
from cameractrl.models.pose_adaptor import CameraPoseEncoder
|
||||
from cameractrl.pipelines.pipeline_animation import CameraCtrlPipeline
|
||||
from cameractrl.utils.convert_from_ckpt import convert_ldm_unet_checkpoint
|
||||
from cameractrl.data.dataset import Camera
|
||||
|
||||
|
||||
def setup_for_distributed(is_master):
|
||||
"""
|
||||
This function disables printing when not in master process
|
||||
"""
|
||||
import builtins as __builtin__
|
||||
builtin_print = __builtin__.print
|
||||
|
||||
def print(*args, **kwargs):
|
||||
force = kwargs.pop('force', False)
|
||||
if is_master or force:
|
||||
builtin_print(*args, **kwargs)
|
||||
|
||||
__builtin__.print = print
|
||||
|
||||
|
||||
def custom_meshgrid(*args):
|
||||
# ref: https://pytorch.org/docs/stable/generated/torch.meshgrid.html?highlight=meshgrid#torch.meshgrid
|
||||
if pver.parse(torch.__version__) < pver.parse('1.10'):
|
||||
return torch.meshgrid(*args)
|
||||
else:
|
||||
return torch.meshgrid(*args, indexing='ij')
|
||||
|
||||
|
||||
def get_relative_pose(cam_params):
|
||||
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
|
||||
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
|
||||
cam_to_origin = 0
|
||||
target_cam_c2w = np.array([
|
||||
[1, 0, 0, 0],
|
||||
[0, 1, 0, -cam_to_origin],
|
||||
[0, 0, 1, 0],
|
||||
[0, 0, 0, 1]
|
||||
])
|
||||
abs2rel = target_cam_c2w @ abs_w2cs[0]
|
||||
ret_poses = [target_cam_c2w, ] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
|
||||
ret_poses = np.array(ret_poses, dtype=np.float32)
|
||||
return ret_poses
|
||||
|
||||
|
||||
def ray_condition(K, c2w, H, W, device):
|
||||
# c2w: B, V, 4, 4
|
||||
# K: B, V, 4
|
||||
|
||||
B = K.shape[0]
|
||||
|
||||
j, i = custom_meshgrid(
|
||||
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
|
||||
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
|
||||
)
|
||||
i = i.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
j = j.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
|
||||
fx, fy, cx, cy = K.chunk(4, dim=-1) # B,V, 1
|
||||
|
||||
zs = torch.ones_like(i) # [B, HxW]
|
||||
xs = (i - cx) / fx * zs
|
||||
ys = (j - cy) / fy * zs
|
||||
zs = zs.expand_as(ys)
|
||||
|
||||
directions = torch.stack((xs, ys, zs), dim=-1) # B, V, HW, 3
|
||||
directions = directions / directions.norm(dim=-1, keepdim=True) # B, V, HW, 3
|
||||
|
||||
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2) # B, V, 3, HW
|
||||
rays_o = c2w[..., :3, 3] # B, V, 3
|
||||
rays_o = rays_o[:, :, None].expand_as(rays_d) # B, V, 3, HW
|
||||
# c2w @ dirctions
|
||||
rays_dxo = torch.cross(rays_o, rays_d)
|
||||
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
|
||||
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6) # B, V, H, W, 6
|
||||
# plucker = plucker.permute(0, 1, 4, 2, 3)
|
||||
return plucker
|
||||
|
||||
|
||||
def load_personalized_base_model(pipeline, personalized_base_model):
|
||||
print(f'Load civitai base model from {personalized_base_model}')
|
||||
if personalized_base_model.endswith(".safetensors"):
|
||||
dreambooth_state_dict = {}
|
||||
with safe_open(personalized_base_model, framework="pt", device="cpu") as f:
|
||||
for key in f.keys():
|
||||
dreambooth_state_dict[key] = f.get_tensor(key)
|
||||
elif personalized_base_model.endswith(".ckpt"):
|
||||
dreambooth_state_dict = torch.load(personalized_base_model, map_location="cpu")
|
||||
|
||||
# 1. vae
|
||||
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, pipeline.vae.config)
|
||||
pipeline.vae.load_state_dict(converted_vae_checkpoint)
|
||||
# 2. unet
|
||||
converted_unet_checkpoint = convert_ldm_unet_checkpoint(dreambooth_state_dict, pipeline.unet.config)
|
||||
_, unetu = pipeline.unet.load_state_dict(converted_unet_checkpoint, strict=False)
|
||||
assert len(unetu) == 0
|
||||
# 3. text_model
|
||||
pipeline.text_encoder = convert_ldm_clip_checkpoint(dreambooth_state_dict, text_encoder=pipeline.text_encoder)
|
||||
del dreambooth_state_dict
|
||||
return pipeline
|
||||
|
||||
|
||||
def get_pipeline(ori_model_path, unet_subfolder, image_lora_rank, image_lora_ckpt, unet_additional_kwargs,
|
||||
unet_mm_ckpt, pose_encoder_kwargs, attention_processor_kwargs,
|
||||
noise_scheduler_kwargs, pose_adaptor_ckpt, personalized_base_model, gpu_id):
|
||||
vae = AutoencoderKL.from_pretrained(ori_model_path, subfolder="vae")
|
||||
tokenizer = CLIPTokenizer.from_pretrained(ori_model_path, subfolder="tokenizer")
|
||||
text_encoder = CLIPTextModel.from_pretrained(ori_model_path, subfolder="text_encoder")
|
||||
unet = UNet3DConditionModelPoseCond.from_pretrained_2d(ori_model_path, subfolder=unet_subfolder,
|
||||
unet_additional_kwargs=unet_additional_kwargs)
|
||||
pose_encoder = CameraPoseEncoder(**pose_encoder_kwargs)
|
||||
print(f"Setting the attention processors")
|
||||
unet.set_all_attn_processor(add_spatial_lora=image_lora_ckpt is not None,
|
||||
add_motion_lora=False,
|
||||
lora_kwargs={"lora_rank": image_lora_rank, "lora_scale": 1.0},
|
||||
motion_lora_kwargs={"lora_rank": -1, "lora_scale": 1.0},
|
||||
**attention_processor_kwargs)
|
||||
|
||||
if image_lora_ckpt is not None:
|
||||
print(f"Loading the lora checkpoint from {image_lora_ckpt}")
|
||||
lora_checkpoints = torch.load(image_lora_ckpt, map_location=unet.device)
|
||||
if 'lora_state_dict' in lora_checkpoints.keys():
|
||||
lora_checkpoints = lora_checkpoints['lora_state_dict']
|
||||
_, lora_u = unet.load_state_dict(lora_checkpoints, strict=False)
|
||||
assert len(lora_u) == 0
|
||||
print(f'Loading done')
|
||||
|
||||
if unet_mm_ckpt is not None:
|
||||
print(f"Loading the motion module checkpoint from {unet_mm_ckpt}")
|
||||
mm_checkpoints = torch.load(unet_mm_ckpt, map_location=unet.device)
|
||||
_, mm_u = unet.load_state_dict(mm_checkpoints, strict=False)
|
||||
assert len(mm_u) == 0
|
||||
print("Loading done")
|
||||
|
||||
print(f"Loading pose adaptor")
|
||||
pose_adaptor_checkpoint = torch.load(pose_adaptor_ckpt, map_location='cpu')
|
||||
pose_encoder_state_dict = pose_adaptor_checkpoint['pose_encoder_state_dict']
|
||||
pose_encoder_m, pose_encoder_u = pose_encoder.load_state_dict(pose_encoder_state_dict)
|
||||
assert len(pose_encoder_u) == 0 and len(pose_encoder_m) == 0
|
||||
attention_processor_state_dict = pose_adaptor_checkpoint['attention_processor_state_dict']
|
||||
_, attn_proc_u = unet.load_state_dict(attention_processor_state_dict, strict=False)
|
||||
assert len(attn_proc_u) == 0
|
||||
print(f"Loading done")
|
||||
|
||||
noise_scheduler = DDIMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
|
||||
vae.to(gpu_id)
|
||||
text_encoder.to(gpu_id)
|
||||
unet.to(gpu_id)
|
||||
pose_encoder.to(gpu_id)
|
||||
pipe = CameraCtrlPipeline(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
unet=unet,
|
||||
scheduler=noise_scheduler,
|
||||
pose_encoder=pose_encoder)
|
||||
if personalized_base_model is not None:
|
||||
load_personalized_base_model(pipeline=pipe, personalized_base_model=personalized_base_model)
|
||||
pipe.enable_vae_slicing()
|
||||
pipe = pipe.to(gpu_id)
|
||||
|
||||
return pipe
|
||||
|
||||
|
||||
def main(args):
|
||||
os.makedirs(args.out_root, exist_ok=True)
|
||||
rank = args.local_rank
|
||||
setup_for_distributed(rank == 0)
|
||||
gpu_id = rank % torch.cuda.device_count()
|
||||
model_configs = OmegaConf.load(args.model_config)
|
||||
unet_additional_kwargs = model_configs[
|
||||
'unet_additional_kwargs'] if 'unet_additional_kwargs' in model_configs else None
|
||||
noise_scheduler_kwargs = model_configs['noise_scheduler_kwargs']
|
||||
pose_encoder_kwargs = model_configs['pose_encoder_kwargs']
|
||||
attention_processor_kwargs = model_configs['attention_processor_kwargs']
|
||||
|
||||
print(f'Constructing pipeline')
|
||||
pipeline = get_pipeline(args.ori_model_path, args.unet_subfolder, args.image_lora_rank, args.image_lora_ckpt,
|
||||
unet_additional_kwargs, args.motion_module_ckpt, pose_encoder_kwargs, attention_processor_kwargs,
|
||||
noise_scheduler_kwargs, args.pose_adaptor_ckpt,
|
||||
args.personalized_base_model, f"cuda:{gpu_id}")
|
||||
device = torch.device(f"cuda:{gpu_id}")
|
||||
print('Done')
|
||||
print('Loading K, R, t matrix')
|
||||
with open(args.trajectory_file, 'r') as f:
|
||||
poses = f.readlines()
|
||||
poses = [pose.strip().split(' ') for pose in poses[1:]]
|
||||
|
||||
cam_params = [[float(x) for x in pose] for pose in poses]
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
sample_wh_ratio = args.image_width / args.image_height
|
||||
pose_wh_ratio = args.original_pose_width / args.original_pose_height
|
||||
if pose_wh_ratio > sample_wh_ratio:
|
||||
resized_ori_w = args.image_height * pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / args.image_width
|
||||
else:
|
||||
resized_ori_h = args.image_width / pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / args.image_height
|
||||
intrinsic = np.asarray([[cam_param.fx * args.image_width,
|
||||
cam_param.fy * args.image_height,
|
||||
cam_param.cx * args.image_width,
|
||||
cam_param.cy * args.image_height]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
|
||||
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
|
||||
c2ws = get_relative_pose(cam_params)
|
||||
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
|
||||
plucker_embedding = ray_condition(K, c2ws, args.image_height, args.image_width, device='cpu')[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
|
||||
plucker_embedding = plucker_embedding[None].to(device) # B V 6 H W
|
||||
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b c f h w")
|
||||
if args.visualization_captions.endswith('.json'):
|
||||
json_file = json.load(open(args.visualization_captions, 'r'))
|
||||
captions = json_file['captions'] if 'captions' in json_file else json_file['prompts']
|
||||
if args.use_negative_prompt:
|
||||
negative_prompts = json_file['negative_prompts']
|
||||
else:
|
||||
negative_prompts = None
|
||||
if isinstance(captions[0], dict):
|
||||
captions = [cap['caption'] for cap in captions]
|
||||
if args.use_specific_seeds:
|
||||
specific_seeds = json_file['seeds']
|
||||
else:
|
||||
specific_seeds = None
|
||||
elif args.visualization_captions.endswith('.txt'):
|
||||
with open(args.visualization_captions, 'r') as f:
|
||||
captions = f.readlines()
|
||||
captions = [cap.strip() for cap in captions]
|
||||
negative_prompts = None
|
||||
specific_seeds = None
|
||||
N = int(len(captions) // args.n_procs)
|
||||
remainder = int(len(captions) % args.n_procs)
|
||||
prompts_per_gpu = [N + 1 if gpu_id < remainder else N for gpu_id in range(args.n_procs)]
|
||||
low_idx = sum(prompts_per_gpu[:gpu_id])
|
||||
high_idx = low_idx + prompts_per_gpu[gpu_id]
|
||||
prompts = captions[low_idx: high_idx]
|
||||
negative_prompts = negative_prompts[low_idx: high_idx] if negative_prompts is not None else None
|
||||
specific_seeds = specific_seeds[low_idx: high_idx] if specific_seeds is not None else None
|
||||
print(f"rank {rank} / {torch.cuda.device_count()}, number of prompts: {len(prompts)}")
|
||||
generator = torch.Generator(device=device)
|
||||
generator.manual_seed(42)
|
||||
for local_idx, caption in tqdm(enumerate(prompts)):
|
||||
if specific_seeds is not None:
|
||||
specific_seed = specific_seeds[local_idx]
|
||||
generator.manual_seed(specific_seed)
|
||||
sample = pipeline(
|
||||
prompt=caption,
|
||||
negative_prompt=negative_prompts[local_idx] if negative_prompts is not None else None,
|
||||
pose_embedding=plucker_embedding,
|
||||
video_length=args.video_length,
|
||||
height=args.image_height,
|
||||
width=args.image_width,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).videos # [1, 3, f, h, w]
|
||||
save_name = "_".join(caption.split(" "))
|
||||
save_videos_grid(sample, f"{args.out_root}/{save_name}.mp4")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--out_root", type=str)
|
||||
parser.add_argument("--image_height", type=int, default=256)
|
||||
parser.add_argument("--image_width", type=int, default=384)
|
||||
parser.add_argument("--video_length", type=int, default=16)
|
||||
parser.add_argument("--ori_model_path", type=str, help='path to the sd model folder')
|
||||
parser.add_argument("--unet_subfolder", type=str, help='subfolder name of unet ckpt')
|
||||
parser.add_argument("--motion_module_ckpt", type=str, help='path to the animatediff motion module ckpt')
|
||||
parser.add_argument("--image_lora_rank", type=int, default=2)
|
||||
parser.add_argument("--image_lora_ckpt", default=None)
|
||||
parser.add_argument("--personalized_base_model", default=None)
|
||||
parser.add_argument("--pose_adaptor_ckpt", default=None, help='path to the camera control model ckpt')
|
||||
parser.add_argument("--model_config", type=str)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=25)
|
||||
parser.add_argument("--guidance_scale", type=float, default=14.0)
|
||||
parser.add_argument("--visualization_captions", required=True, help='prompts path, json or txt')
|
||||
parser.add_argument("--use_negative_prompt", action='store_true', help='whether to use negative prompts')
|
||||
parser.add_argument("--use_specific_seeds", action='store_true', help='whether to use specific seeds for each prompt')
|
||||
parser.add_argument("--trajectory_file", required=True, help='txt file')
|
||||
parser.add_argument("--original_pose_width", type=int, default=1280, help='the width of the video used to extract camera trajectory')
|
||||
parser.add_argument("--original_pose_height", type=int, default=720, help='the height of the video used to extract camera trajectory')
|
||||
parser.add_argument("--n_procs", type=int, default=8)
|
||||
|
||||
# DDP args
|
||||
parser.add_argument("--world_size", default=1, type=int,
|
||||
help="number of the distributed processes.")
|
||||
parser.add_argument('--local_rank', type=int, default=-1,
|
||||
help='Replica rank on the current node. This field is required '
|
||||
'by `torch.distributed.launch`.')
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,398 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import folder_paths
|
||||
comfy_path = os.path.dirname(folder_paths.__file__)
|
||||
|
||||
import sys
|
||||
cameractrl_path=f'{comfy_path}/custom_nodes/ComfyUI-CameraCtrl'
|
||||
sys.path.insert(0,cameractrl_path)
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
from packaging import version as pver
|
||||
from einops import rearrange
|
||||
from safetensors import safe_open
|
||||
|
||||
from omegaconf import OmegaConf
|
||||
from diffusers import (
|
||||
AutoencoderKL,
|
||||
DDIMScheduler
|
||||
)
|
||||
from transformers import CLIPTextModel, CLIPTokenizer
|
||||
from diffusers.pipelines.stable_diffusion.convert_from_ckpt import convert_ldm_vae_checkpoint, \
|
||||
convert_ldm_clip_checkpoint
|
||||
|
||||
from cameractrl.utils.util import save_videos_grid
|
||||
from cameractrl.models.unet import UNet3DConditionModelPoseCond
|
||||
from cameractrl.models.pose_adaptor import CameraPoseEncoder
|
||||
from cameractrl.pipelines.pipeline_animation import CameraCtrlPipeline
|
||||
from cameractrl.utils.convert_from_ckpt import convert_ldm_unet_checkpoint
|
||||
|
||||
class Camera(object):
|
||||
def __init__(self, entry):
|
||||
fx, fy, cx, cy = entry[0:4]
|
||||
self.fx = fx
|
||||
self.fy = fy
|
||||
self.cx = cx
|
||||
self.cy = cy
|
||||
w2c_mat = np.array(entry[6:]).reshape(3, 4)
|
||||
w2c_mat_4x4 = np.eye(4)
|
||||
w2c_mat_4x4[:3, :] = w2c_mat
|
||||
self.w2c_mat = w2c_mat_4x4
|
||||
self.c2w_mat = np.linalg.inv(w2c_mat_4x4)
|
||||
|
||||
|
||||
def setup_for_distributed(is_master):
|
||||
"""
|
||||
This function disables printing when not in master process
|
||||
"""
|
||||
import builtins as __builtin__
|
||||
builtin_print = __builtin__.print
|
||||
|
||||
def print(*args, **kwargs):
|
||||
force = kwargs.pop('force', False)
|
||||
if is_master or force:
|
||||
builtin_print(*args, **kwargs)
|
||||
|
||||
__builtin__.print = print
|
||||
|
||||
|
||||
def custom_meshgrid(*args):
|
||||
# ref: https://pytorch.org/docs/stable/generated/torch.meshgrid.html?highlight=meshgrid#torch.meshgrid
|
||||
if pver.parse(torch.__version__) < pver.parse('1.10'):
|
||||
return torch.meshgrid(*args)
|
||||
else:
|
||||
return torch.meshgrid(*args, indexing='ij')
|
||||
|
||||
|
||||
def get_relative_pose(cam_params):
|
||||
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
|
||||
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
|
||||
cam_to_origin = 0
|
||||
target_cam_c2w = np.array([
|
||||
[1, 0, 0, 0],
|
||||
[0, 1, 0, -cam_to_origin],
|
||||
[0, 0, 1, 0],
|
||||
[0, 0, 0, 1]
|
||||
])
|
||||
abs2rel = target_cam_c2w @ abs_w2cs[0]
|
||||
ret_poses = [target_cam_c2w, ] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
|
||||
ret_poses = np.array(ret_poses, dtype=np.float32)
|
||||
return ret_poses
|
||||
|
||||
|
||||
def ray_condition(K, c2w, H, W, device):
|
||||
# c2w: B, V, 4, 4
|
||||
# K: B, V, 4
|
||||
|
||||
B = K.shape[0]
|
||||
|
||||
j, i = custom_meshgrid(
|
||||
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
|
||||
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
|
||||
)
|
||||
i = i.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
j = j.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
|
||||
fx, fy, cx, cy = K.chunk(4, dim=-1) # B,V, 1
|
||||
|
||||
zs = torch.ones_like(i) # [B, HxW]
|
||||
xs = (i - cx) / fx * zs
|
||||
ys = (j - cy) / fy * zs
|
||||
zs = zs.expand_as(ys)
|
||||
|
||||
directions = torch.stack((xs, ys, zs), dim=-1) # B, V, HW, 3
|
||||
directions = directions / directions.norm(dim=-1, keepdim=True) # B, V, HW, 3
|
||||
|
||||
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2) # B, V, 3, HW
|
||||
rays_o = c2w[..., :3, 3] # B, V, 3
|
||||
rays_o = rays_o[:, :, None].expand_as(rays_d) # B, V, 3, HW
|
||||
# c2w @ dirctions
|
||||
rays_dxo = torch.cross(rays_o, rays_d)
|
||||
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
|
||||
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6) # B, V, H, W, 6
|
||||
# plucker = plucker.permute(0, 1, 4, 2, 3)
|
||||
return plucker
|
||||
|
||||
|
||||
def load_personalized_base_model(pipeline, personalized_base_model):
|
||||
print(f'Load civitai base model from {personalized_base_model}')
|
||||
if personalized_base_model.endswith(".safetensors"):
|
||||
dreambooth_state_dict = {}
|
||||
with safe_open(personalized_base_model, framework="pt", device="cpu") as f:
|
||||
for key in f.keys():
|
||||
dreambooth_state_dict[key] = f.get_tensor(key)
|
||||
elif personalized_base_model.endswith(".ckpt"):
|
||||
dreambooth_state_dict = torch.load(personalized_base_model, map_location="cpu")
|
||||
|
||||
# 1. vae
|
||||
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, pipeline.vae.config)
|
||||
pipeline.vae.load_state_dict(converted_vae_checkpoint)
|
||||
# 2. unet
|
||||
converted_unet_checkpoint = convert_ldm_unet_checkpoint(dreambooth_state_dict, pipeline.unet.config)
|
||||
_, unetu = pipeline.unet.load_state_dict(converted_unet_checkpoint, strict=False)
|
||||
assert len(unetu) == 0
|
||||
# 3. text_model
|
||||
pipeline.text_encoder = convert_ldm_clip_checkpoint(dreambooth_state_dict, text_encoder=pipeline.text_encoder)
|
||||
del dreambooth_state_dict
|
||||
return pipeline
|
||||
|
||||
|
||||
def get_pipeline(vae,tokenizer,text_encoder,unet, image_lora_rank, image_lora_ckpt, unet_additional_kwargs,
|
||||
unet_mm_ckpt, pose_encoder_kwargs, attention_processor_kwargs,
|
||||
noise_scheduler_kwargs, pose_adaptor_ckpt, personalized_base_model, gpu_id):
|
||||
#vae = AutoencoderKL.from_pretrained(ori_model_path, subfolder="vae")
|
||||
#tokenizer = CLIPTokenizer.from_pretrained(ori_model_path, subfolder="tokenizer")
|
||||
#text_encoder = CLIPTextModel.from_pretrained(ori_model_path, subfolder="text_encoder")
|
||||
#unet = UNet3DConditionModelPoseCond.from_pretrained_2d(ori_model_path, subfolder=unet_subfolder,unet_additional_kwargs=unet_additional_kwargs)
|
||||
'''
|
||||
unet=UNet3DConditionModelPoseCond(unet)
|
||||
unet.down_block_types = [
|
||||
"CrossAttnDownBlock3D",
|
||||
"CrossAttnDownBlock3D",
|
||||
"CrossAttnDownBlock3D",
|
||||
"DownBlock3D"
|
||||
]
|
||||
unet.up_block_types = [
|
||||
"UpBlock3D",
|
||||
"CrossAttnUpBlock3D",
|
||||
"CrossAttnUpBlock3D",
|
||||
"CrossAttnUpBlock3D"
|
||||
]
|
||||
'''
|
||||
|
||||
pose_encoder = CameraPoseEncoder(**pose_encoder_kwargs)
|
||||
print(f"Setting the attention processors")
|
||||
unet.set_all_attn_processor(add_spatial_lora=image_lora_ckpt is not None,
|
||||
add_motion_lora=False,
|
||||
lora_kwargs={"lora_rank": image_lora_rank, "lora_scale": 1.0},
|
||||
motion_lora_kwargs={"lora_rank": -1, "lora_scale": 1.0},
|
||||
**attention_processor_kwargs)
|
||||
|
||||
if image_lora_ckpt is not None:
|
||||
print(f"Loading the lora checkpoint from {image_lora_ckpt}")
|
||||
lora_checkpoints = torch.load(image_lora_ckpt, map_location=unet.device)
|
||||
if 'lora_state_dict' in lora_checkpoints.keys():
|
||||
lora_checkpoints = lora_checkpoints['lora_state_dict']
|
||||
_, lora_u = unet.load_state_dict(lora_checkpoints, strict=False)
|
||||
assert len(lora_u) == 0
|
||||
print(f'Loading done')
|
||||
|
||||
if unet_mm_ckpt is not None:
|
||||
print(f"Loading the motion module checkpoint from {unet_mm_ckpt}")
|
||||
mm_checkpoints = torch.load(unet_mm_ckpt, map_location=unet.device)
|
||||
_, mm_u = unet.load_state_dict(mm_checkpoints, strict=False)
|
||||
assert len(mm_u) == 0
|
||||
print("Loading done")
|
||||
|
||||
print(f"Loading pose adaptor")
|
||||
pose_adaptor_checkpoint = torch.load(pose_adaptor_ckpt, map_location='cpu')
|
||||
pose_encoder_state_dict = pose_adaptor_checkpoint['pose_encoder_state_dict']
|
||||
pose_encoder_m, pose_encoder_u = pose_encoder.load_state_dict(pose_encoder_state_dict)
|
||||
assert len(pose_encoder_u) == 0 and len(pose_encoder_m) == 0
|
||||
attention_processor_state_dict = pose_adaptor_checkpoint['attention_processor_state_dict']
|
||||
_, attn_proc_u = unet.load_state_dict(attention_processor_state_dict, strict=False)
|
||||
assert len(attn_proc_u) == 0
|
||||
print(f"Loading done")
|
||||
|
||||
noise_scheduler = DDIMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
|
||||
vae.to(gpu_id)
|
||||
text_encoder.to(gpu_id)
|
||||
unet.to(gpu_id)
|
||||
pose_encoder.to(gpu_id)
|
||||
pipe = CameraCtrlPipeline(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
unet=unet,
|
||||
scheduler=noise_scheduler,
|
||||
pose_encoder=pose_encoder)
|
||||
if personalized_base_model is not None:
|
||||
load_personalized_base_model(pipeline=pipe, personalized_base_model=personalized_base_model)
|
||||
pipe.enable_vae_slicing()
|
||||
pipe = pipe.to(gpu_id)
|
||||
|
||||
return pipe
|
||||
|
||||
class CameraCtrlLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"sd15_ckpt": (folder_paths.get_filename_list("checkpoints"),),
|
||||
"ad_v3_sd15_adapter_ckpt": (folder_paths.get_filename_list("loras"), {"default": "v3_sd15_adapter.ckpt"}),
|
||||
"ad_v3_ckpt": (os.listdir(os.path.join(folder_paths.models_dir,"animatediff_models")), {"default": "v3_sd15_mm.ckpt"}),
|
||||
"cameractrl_ckpt": (folder_paths.get_filename_list("checkpoints"), {"default": "CameraCtrl.ckpt"}),
|
||||
},
|
||||
"optional": {
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CameraCtrlPipeline",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "CameraCtrl"
|
||||
|
||||
def run(self,sd15_ckpt,ad_v3_sd15_adapter_ckpt,ad_v3_ckpt,cameractrl_ckpt):
|
||||
sd15_ckpt = folder_paths.get_full_path("checkpoints", sd15_ckpt)
|
||||
ad_v3_sd15_adapter_ckpt = folder_paths.get_full_path("loras", ad_v3_sd15_adapter_ckpt)
|
||||
ad_v3_ckpt=os.path.join(os.path.join(folder_paths.models_dir,"animatediff_models"),ad_v3_ckpt)
|
||||
cameractrl_ckpt = folder_paths.get_full_path("checkpoints", cameractrl_ckpt)
|
||||
|
||||
from diffusers import StableDiffusionPipeline
|
||||
pipe = StableDiffusionPipeline.from_single_file(sd15_ckpt).to("cuda")
|
||||
|
||||
rank = -1
|
||||
setup_for_distributed(rank == 0)
|
||||
gpu_id = rank % torch.cuda.device_count()
|
||||
|
||||
model_configs = OmegaConf.load(f'{cameractrl_path}/configs/train_cameractrl/adv3_256_384_cameractrl_relora.yaml')
|
||||
unet_additional_kwargs = model_configs[
|
||||
'unet_additional_kwargs'] if 'unet_additional_kwargs' in model_configs else None
|
||||
noise_scheduler_kwargs = model_configs['noise_scheduler_kwargs']
|
||||
pose_encoder_kwargs = model_configs['pose_encoder_kwargs']
|
||||
attention_processor_kwargs = model_configs['attention_processor_kwargs']
|
||||
|
||||
unet=pipe.unet
|
||||
fused_state_dict = unet.state_dict()
|
||||
lora_state_dict = torch.load(ad_v3_sd15_adapter_ckpt, map_location='cuda')
|
||||
if 'state_dict' in lora_state_dict:
|
||||
lora_state_dict = lora_state_dict['state_dict']
|
||||
print(f'Loading done')
|
||||
print(f'Fusing the lora weight to unet weight')
|
||||
used_lora_key = []
|
||||
for lora_key in ['to_q', 'to_k', 'to_v', 'to_out']:
|
||||
unet_keys = [x for x in fused_state_dict.keys() if lora_key in x and "bias" not in x]
|
||||
print(f'There are {len(unet_keys)} unet keys for lora key: {lora_key}')
|
||||
for unet_key in unet_keys:
|
||||
prefixes = unet_key.split('.')
|
||||
idx = prefixes.index(lora_key)
|
||||
lora_down_key = ".".join(prefixes[:idx]) + f".processor.{lora_key}_lora.down" + f".{prefixes[-1]}"
|
||||
lora_up_key = ".".join(prefixes[:idx]) + f".processor.{lora_key}_lora.up" + f".{prefixes[-1]}"
|
||||
assert lora_down_key in lora_state_dict and lora_up_key in lora_state_dict
|
||||
print(f'Fusing lora weight for {unet_key}')
|
||||
fused_state_dict[unet_key] = fused_state_dict[unet_key] + torch.bmm(lora_state_dict[lora_up_key][None, ...], lora_state_dict[lora_down_key][None, ...])[0] * 1.0
|
||||
used_lora_key.append(lora_down_key)
|
||||
used_lora_key.append(lora_up_key)
|
||||
assert len(set(used_lora_key) - set(lora_state_dict.keys())) == 0
|
||||
print(f'Fusing done')
|
||||
from diffusers.utils import SAFETENSORS_WEIGHTS_NAME
|
||||
save_path = os.path.join(os.path.join(cameractrl_path, 'unet_webvidlora_v3'),SAFETENSORS_WEIGHTS_NAME)
|
||||
print(f'Saving the fused state dict to {save_path}')
|
||||
|
||||
from safetensors.torch import save_file
|
||||
save_file(fused_state_dict, save_path)
|
||||
unet = UNet3DConditionModelPoseCond.from_pretrained_2d(cameractrl_path, subfolder='unet_webvidlora_v3',unet_additional_kwargs=unet_additional_kwargs)
|
||||
|
||||
print(f'Constructing pipeline')
|
||||
image_lora_rank=2
|
||||
image_lora_ckpt=None
|
||||
personalized_base_model=None
|
||||
pipeline = get_pipeline(pipe.vae,pipe.tokenizer,pipe.text_encoder,unet, image_lora_rank, image_lora_ckpt,
|
||||
unet_additional_kwargs, ad_v3_ckpt, pose_encoder_kwargs, attention_processor_kwargs,
|
||||
noise_scheduler_kwargs, cameractrl_ckpt,
|
||||
personalized_base_model, f"cuda:{gpu_id}")
|
||||
return (pipeline,)
|
||||
|
||||
class CameraCtrlRun:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"pipeline":("CameraCtrlPipeline",),
|
||||
"camera_trajectory":("STRING",{"multiline": True, "default":"[[0.474812461,0.844111024,0.5,0.5,0,0,0.780003667,0.059620168,-0.622928321,0.726968666,-0.062449891,0.997897983,0.017311305,0.217967188,0.622651041,0.025398925,0.782087326,-1.002211444],[0.474812461,0.844111024,0.5,0.5,0,0,0.743836701,0.064830206,-0.66520977,0.951841944,-0.068305343,0.997446954,0.020830527,0.206496789,0.664861917,0.029942872,0.746365905,-1.084913992],[0.474812461,0.844111024,0.5,0.5,0,0,0.697046876,0.070604131,-0.713540971,1.208789672,-0.074218854,0.996899366,0.026138915,0.196421447,0.713174045,0.034738146,0.700125754,-1.130142078],[0.474812461,0.844111024,0.5,0.5,0,0,0.635762572,0.077846259,-0.767949164,1.465161122,-0.080595709,0.996158004,0.034256749,0.157107229,0.767665446,0.040114246,0.639594078,-1.13689307],[0.474812461,0.844111024,0.5,0.5,0,0,0.593250692,0.083153486,-0.800711632,1.635091834,-0.085384794,0.995539784,0.040124334,0.135863998,0.800476789,0.04456481,0.597704709,-1.166997229],[0.474812461,0.844111024,0.5,0.5,0,0,0.555486798,0.087166689,-0.826943994,1.803789619,-0.089439675,0.99498421,0.044799786,0.145490422,0.826701283,0.049075913,0.560496747,-1.24382735],[0.474812461,0.844111024,0.5,0.5,0,0,0.523399472,0.09026666,-0.847292721,1.945815368,-0.093254104,0.994468153,0.048340045,0.174777447,0.846969128,0.053712368,0.528921843,-1.336914479],[0.474812461,0.844111024,0.5,0.5,0,0,0.491546303,0.09212707,-0.865964711,2.093852892,-0.095617607,0.994085968,0.051482171,0.196702533,0.865586221,0.057495601,0.497448236,-1.43970938],[0.474812461,0.844111024,0.5,0.5,0,0,0.475284129,0.093297184,-0.87487179,2.200792438,-0.096743606,0.993874133,0.053430639,0.209217395,0.874497354,0.059243519,0.481398523,-1.547068315],[0.474812461,0.844111024,0.5,0.5,0,0,0.464444131,0.093880348,-0.880612373,2.324141986,-0.097857766,0.993716478,0.054326952,0.220651207,0.880179226,0.060942926,0.470712721,-1.712512928],[0.474812461,0.844111024,0.5,0.5,0,0,0.458157241,0.093640216,-0.883925021,2.44310089,-0.098046601,0.993691206,0.054448847,0.257385043,0.883447111,0.061719712,0.464447916,-1.885672329],[0.474812461,0.844111024,0.5,0.5,0,0,0.457354397,0.09350872,-0.884354591,2.543246338,-0.097820736,0.993711591,0.054482624,0.281562244,0.883888066,0.061590351,0.463625461,-2.094829165],[0.474812461,0.844111024,0.5,0.5,0,0,0.465170115,0.093944497,-0.880222261,2.606377358,-0.097235762,0.99375838,0.054675922,0.277376127,0.879864752,0.060155477,0.471401453,-2.299280675],[0.474812461,0.844111024,0.5,0.5,0,0,0.511845231,0.090872414,-0.854257941,2.5767741,-0.093636356,0.994366586,0.049672548,0.270516319,0.853959382,0.054564942,0.517470777,-2.624374352],[0.474812461,0.844111024,0.5,0.5,0,0,0.590568483,0.083218277,-0.802685261,2.398318316,-0.085610889,0.995516419,0.04022257,0.282138215,0.80243355,0.044964414,0.59504503,-3.012309268],[0.474812461,0.844111024,0.5,0.5,0,0,0.684302032,0.072693504,-0.725566208,2.086323553,-0.074529484,0.996780157,0.029575195,0.310959312,0.725379944,0.03383771,0.68751651,-3.456740526]]"}),
|
||||
"image_width":("INT",{"default":384}),
|
||||
"image_height":("INT",{"default":256}),
|
||||
"original_pose_width":("INT",{"default":1280}),
|
||||
"original_pose_height":("INT",{"default":720}),
|
||||
"prompt":("STRING",{"multiline": True, "default":"A serene mountain lake at sunrise, with mist hovering over the water."}),
|
||||
"negative_prompt":("STRING",{"multiline": True, "default":"Strange motion trajectory, a poor composition and deformed video, worst quality, normal quality, low quality, low resolution, duplicate and ugly"}),
|
||||
"video_length":("INT",{"default":16}),
|
||||
"num_inference_steps":("INT",{"default":30}),
|
||||
"guidance_scale":("FLOAT",{"default":6.0}),
|
||||
"seed":("INT",{"default":1234}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "CameraCtrl"
|
||||
|
||||
def run(self,pipeline,camera_trajectory,image_width,image_height,original_pose_width,original_pose_height,prompt,negative_prompt,video_length,num_inference_steps,guidance_scale,seed):
|
||||
rank = -1
|
||||
setup_for_distributed(rank == 0)
|
||||
gpu_id = rank % torch.cuda.device_count()
|
||||
|
||||
generator = torch.Generator(device='cuda').manual_seed(seed)
|
||||
device = torch.device(f"cuda:{gpu_id}")
|
||||
print('Done')
|
||||
print('Loading K, R, t matrix')
|
||||
#with open(args.trajectory_file, 'r') as f:
|
||||
# poses = f.readlines()
|
||||
#poses = [pose.strip().split(' ') for pose in poses[1:]]
|
||||
|
||||
poses=json.loads(camera_trajectory)
|
||||
cam_params = [[float(x) for x in pose] for pose in poses]
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
sample_wh_ratio = image_width / image_height
|
||||
pose_wh_ratio = original_pose_width / original_pose_height
|
||||
if pose_wh_ratio > sample_wh_ratio:
|
||||
resized_ori_w = image_height * pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / image_width
|
||||
else:
|
||||
resized_ori_h = image_width / pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / image_height
|
||||
intrinsic = np.asarray([[cam_param.fx * image_width,
|
||||
cam_param.fy * image_height,
|
||||
cam_param.cx * image_width,
|
||||
cam_param.cy * image_height]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
|
||||
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
|
||||
c2ws = get_relative_pose(cam_params)
|
||||
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
|
||||
plucker_embedding = ray_condition(K, c2ws, image_height, image_width, device='cpu')[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
|
||||
plucker_embedding = plucker_embedding[None].to(device) # B V 6 H W
|
||||
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b c f h w")
|
||||
captions=[prompt]
|
||||
negative_prompts=[negative_prompt]
|
||||
specific_seeds=[seed]
|
||||
n_procs=8
|
||||
N = int(len(captions) // n_procs)
|
||||
remainder = int(len(captions) % n_procs)
|
||||
prompts_per_gpu = [N + 1 if gpu_id < remainder else N for gpu_id in range(n_procs)]
|
||||
low_idx = sum(prompts_per_gpu[:gpu_id])
|
||||
high_idx = low_idx + prompts_per_gpu[gpu_id]
|
||||
prompts = captions[low_idx: high_idx]
|
||||
negative_prompts = negative_prompts[low_idx: high_idx] if negative_prompts is not None else None
|
||||
specific_seeds = specific_seeds[low_idx: high_idx] if specific_seeds is not None else None
|
||||
print(f"rank {rank} / {torch.cuda.device_count()}, number of prompts: {len(prompts)}")
|
||||
generator = torch.Generator(device=device)
|
||||
generator.manual_seed(seed)
|
||||
for local_idx, caption in tqdm(enumerate(prompts)):
|
||||
if specific_seeds is not None:
|
||||
specific_seed = specific_seeds[local_idx]
|
||||
generator.manual_seed(specific_seed)
|
||||
videos = pipeline(
|
||||
prompt=caption,
|
||||
negative_prompt=negative_prompts[local_idx] if negative_prompts is not None else None,
|
||||
pose_embedding=plucker_embedding,
|
||||
video_length=video_length,
|
||||
height=image_height,
|
||||
width=image_width,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
generator=generator,
|
||||
).videos # [1, 3, f, h, w]
|
||||
videos=videos.permute(0,2,3,4,1)
|
||||
return videos
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"CameraCtrlLoader":CameraCtrlLoader,
|
||||
"CameraCtrlRun":CameraCtrlRun,
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
diffusers
|
||||
xformers
|
||||
imageio
|
||||
imageio[ffmpeg]
|
||||
opencv-python
|
||||
transformers
|
||||
gdown
|
||||
einops
|
||||
decord
|
||||
omegaconf
|
||||
safetensors
|
||||
gradio
|
||||
wandb
|
||||
termcolor
|
||||
@@ -0,0 +1,28 @@
|
||||
#!/bin/bash
|
||||
|
||||
set -x
|
||||
|
||||
PARTITION=$1
|
||||
JOB_NAME=$2
|
||||
GPUS=$3
|
||||
CONFIG=$4
|
||||
PT_SCRIPT=${5}
|
||||
CPUS_PER_TASK=${CPUS_PER_TASK:-10}
|
||||
RANDOM_PORT=$((49152 + RANDOM % 16384))
|
||||
if [ $GPUS -lt 8 ]; then
|
||||
GPUS_PER_NODE=${GPUS_PER_NODE:-$GPUS}
|
||||
else
|
||||
GPUS_PER_NODE=${GPUS_PER_NODE:-8}
|
||||
fi
|
||||
|
||||
srun -p ${PARTITION} \
|
||||
--job-name=${JOB_NAME} \
|
||||
--gres=gpu:${GPUS_PER_NODE} \
|
||||
--ntasks=${GPUS} \
|
||||
--ntasks-per-node=${GPUS_PER_NODE} \
|
||||
--cpus-per-task=${CPUS_PER_TASK} \
|
||||
--kill-on-bad-exit=1 \
|
||||
python -u ${PT_SCRIPT} \
|
||||
--config=${CONFIG} \
|
||||
--launcher=slurm \
|
||||
--port=${RANDOM_PORT}
|
||||
@@ -0,0 +1,30 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import os.path as osp
|
||||
from collections import defaultdict
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--video_folder', required=True, help='Path to the down loaded realestate10k txt files')
|
||||
parser.add_argument('--save_path', required=True)
|
||||
parser.add_argument('--save_name', required=True)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = get_args()
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
all_txts = os.listdir(args.video_folder)
|
||||
print(f'There are {len(all_txts)} video clips in the folder {args.video_folder}')
|
||||
video_paths = defaultdict(list)
|
||||
for txt in tqdm(all_txts):
|
||||
with open(osp.join(args.video_folder, txt), 'r') as f:
|
||||
lines = f.readlines()
|
||||
video_name = lines[0].strip().split('=')[-1]
|
||||
video_paths[video_name].append(txt.split('.')[0])
|
||||
print(f'There are {len(video_paths)} videos in the folder {args.video_folder}')
|
||||
with open(osp.join(args.save_path, args.save_name), 'w') as f:
|
||||
json.dump(video_paths, fp=f)
|
||||
@@ -0,0 +1,43 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import os.path as osp
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--root_path', required=True)
|
||||
parser.add_argument('--caption_json', required=True)
|
||||
parser.add_argument('--save_path', required=True)
|
||||
parser.add_argument('--save_name', required=True)
|
||||
parser.add_argument('--video_suffix', default='.mp4')
|
||||
parser.add_argument('--video_folder', default='video_clips')
|
||||
parser.add_argument('--pose_suffix', default='.txt')
|
||||
parser.add_argument('--pose_folder', default='pose_files')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
args = get_args()
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
save_root = args.root_path
|
||||
captions = json.load(open(args.caption_json, 'r'))
|
||||
captions = {k: v[0] for k, v in captions.items()}
|
||||
all_results = []
|
||||
for clip_path, caption in tqdm(captions.items()):
|
||||
clip_path = '/'.join(clip_path.split('/')[-2:])
|
||||
clip_name = clip_path.split('/')[-1].replace(args.video_suffix, '')
|
||||
clip_relative_path = osp.join(args.video_folder, clip_path)
|
||||
if not osp.exists(osp.join(save_root, clip_relative_path)):
|
||||
continue
|
||||
pose_file = args.pose_folder + '/' + clip_name + args.pose_suffix
|
||||
if not osp.exists(osp.join(save_root, pose_file)):
|
||||
continue
|
||||
|
||||
all_results.append({"clip_name": clip_name, "clip_path": clip_relative_path,
|
||||
"pose_file": pose_file, "caption": caption})
|
||||
print(f'There are {len(all_results)} clips after the processing')
|
||||
with open(osp.join(args.save_path, args.save_name), 'w') as f:
|
||||
json.dump(all_results, fp=f)
|
||||
print(f'Saved the generated json file to {osp.join(args.save_path, args.save_name)}')
|
||||
@@ -0,0 +1,50 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import os.path as osp
|
||||
from tqdm import tqdm
|
||||
from moviepy.editor import VideoFileClip
|
||||
import imageio
|
||||
from decord import VideoReader
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--video_root', required=True)
|
||||
parser.add_argument('--save_path', required=True)
|
||||
parser.add_argument('--video2clip_json', required=True, help='generated by gather_realestate.py')
|
||||
parser.add_argument('--clip_txt_path', required=True, help='path to the downloaded realestate txt files')
|
||||
parser.add_argument('--low_idx', type=int, default=0, help='used for parallel processing')
|
||||
parser.add_argument('--high_idx', type=int, default=-1, help='used for parallel processing')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
args = get_args()
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
video2clips = json.load(open(args.video2clip_json, 'r'))
|
||||
video_names = list(video2clips.keys())[args.low_idx: args.high_idx] if args.high_idx != -1 else list(video2clips.keys())
|
||||
video2clips = {k: v for k, v in video2clips.items() if k in video_names}
|
||||
|
||||
for video_name, clip_list in tqdm(video2clips.items()):
|
||||
video_path = osp.join(args.video_root, video_name + '.mp4')
|
||||
if not osp.exists(video_path):
|
||||
continue
|
||||
video = VideoFileClip(video_path)
|
||||
clip_save_path = osp.join(args.save_path, video_name)
|
||||
os.makedirs(clip_save_path, exist_ok=True)
|
||||
for clip in tqdm(clip_list):
|
||||
clip_save_name = clip + '.mp4'
|
||||
if osp.exists(osp.join(clip_save_path, clip_save_name)):
|
||||
continue
|
||||
with open(osp.join(args.clip_txt_path, clip + '.txt'), 'r') as f:
|
||||
lines = f.readlines()
|
||||
frames = [x for x in lines[1: ]]
|
||||
timesteps = [int(x.split(' ')[0]) for x in frames]
|
||||
if timesteps[-1] <= timesteps[0]:
|
||||
continue
|
||||
timestamps_seconds = [x / 1000000.0 for x in timesteps]
|
||||
frames = [video.get_frame(t) for t in timestamps_seconds]
|
||||
imageio.mimsave(osp.join(clip_save_path, clip_save_name), frames, fps=video.fps)
|
||||
video_reader = VideoReader(osp.join(clip_save_path, clip_save_name))
|
||||
assert len(video_reader) == len(timesteps)
|
||||
@@ -0,0 +1,57 @@
|
||||
import argparse
|
||||
import torch
|
||||
import os
|
||||
import shutil
|
||||
from diffusers.models import UNet2DConditionModel
|
||||
from diffusers.utils import SAFETENSORS_WEIGHTS_NAME
|
||||
from safetensors.torch import save_file
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--lora_scale', type=float, default=1.0)
|
||||
parser.add_argument('--lora_ckpt_path', type=str, required=True)
|
||||
parser.add_argument('--unet_ckpt_path', type=str, required=True, help='root path of the sd1.5 model')
|
||||
parser.add_argument('--save_path', type=str, required=True, help='args.unet_ckpt_path + a new subfolder name')
|
||||
parser.add_argument('--unet_config_path', type=str, required=True, help='path to unet config, in the `unet` subfolder of args.unet_ckpt_path')
|
||||
parser.add_argument('--lora_keys', nargs='*', type=str, default=['to_q', 'to_k', 'to_v', 'to_out'])
|
||||
parser.add_argument('--negative_lora_keys', type=str, default="bias")
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
args = get_args()
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
unet = UNet2DConditionModel.from_pretrained(args.unet_ckpt_path, subfolder='unet')
|
||||
fused_state_dict = unet.state_dict()
|
||||
|
||||
print(f'Loading the lora weights from {args.lora_ckpt_path}')
|
||||
lora_state_dict = torch.load(args.lora_ckpt_path, map_location='cpu')
|
||||
if 'state_dict' in lora_state_dict:
|
||||
lora_state_dict = lora_state_dict['state_dict']
|
||||
print(f'Loading done')
|
||||
print(f'Fusing the lora weight to unet weight')
|
||||
used_lora_key = []
|
||||
for lora_key in args.lora_keys:
|
||||
unet_keys = [x for x in fused_state_dict.keys() if lora_key in x and args.negative_lora_keys not in x]
|
||||
print(f'There are {len(unet_keys)} unet keys for lora key: {lora_key}')
|
||||
for unet_key in unet_keys:
|
||||
prefixes = unet_key.split('.')
|
||||
idx = prefixes.index(lora_key)
|
||||
lora_down_key = ".".join(prefixes[:idx]) + f".processor.{lora_key}_lora.down" + f".{prefixes[-1]}"
|
||||
lora_up_key = ".".join(prefixes[:idx]) + f".processor.{lora_key}_lora.up" + f".{prefixes[-1]}"
|
||||
assert lora_down_key in lora_state_dict and lora_up_key in lora_state_dict
|
||||
print(f'Fusing lora weight for {unet_key}')
|
||||
fused_state_dict[unet_key] = fused_state_dict[unet_key] + torch.bmm(lora_state_dict[lora_up_key][None, ...], lora_state_dict[lora_down_key][None, ...])[0] * args.lora_scale
|
||||
used_lora_key.append(lora_down_key)
|
||||
used_lora_key.append(lora_up_key)
|
||||
assert len(set(used_lora_key) - set(lora_state_dict.keys())) == 0
|
||||
print(f'Fusing done')
|
||||
save_path = os.path.join(args.save_path, SAFETENSORS_WEIGHTS_NAME)
|
||||
print(f'Saving the fused state dict to {save_path}')
|
||||
save_file(fused_state_dict, save_path)
|
||||
config_dst_path = os.path.join(args.save_path, 'config.json')
|
||||
print(f'Copying the unet config to {config_dst_path}')
|
||||
shutil.copy(args.unet_config_path, config_dst_path)
|
||||
print('Done!')
|
||||
@@ -0,0 +1,74 @@
|
||||
import argparse
|
||||
import json
|
||||
import random
|
||||
import os
|
||||
import os.path as osp
|
||||
import imageio
|
||||
import cv2
|
||||
import numpy as np
|
||||
from decord import VideoReader
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--save_path', required=True)
|
||||
parser.add_argument('--clip_names', nargs='*', help='the txt file names')
|
||||
parser.add_argument('--clip_txt_path', required=True, help='root path of downloaded realestate10k txt files')
|
||||
parser.add_argument('--json_file', required=True, help='json file generated using generate_realestate_json.py')
|
||||
parser.add_argument('--trajectory_names', nargs='*', help='saving names for each trajectory')
|
||||
parser.add_argument('--sample_stride', type=int, default=4)
|
||||
parser.add_argument('--num_frames', type=int, default=16)
|
||||
parser.add_argument('--video_width', type=int, default=384)
|
||||
parser.add_argument('--video_height', type=int, default=256)
|
||||
parser.add_argument('--save_images', action='store_true')
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
args = get_args()
|
||||
os.makedirs(args.save_path, exist_ok=True)
|
||||
os.makedirs(osp.join(args.save_path, 'selected_pose_files'), exist_ok=True)
|
||||
os.makedirs(osp.join(args.save_path, 'selected_clips'), exist_ok=True)
|
||||
if args.save_images:
|
||||
os.makedirs(osp.join(args.save_path, 'selected_images'), exist_ok=True)
|
||||
clip_infos = json.load(open(args.json_file, 'r'))
|
||||
clip_name2clip_info = {x['clip_name']: x for x in clip_infos}
|
||||
clip_name2clip_info = {x: clip_name2clip_info[x] for x in args.clip_names}
|
||||
selected_clip_infos = []
|
||||
trajectory_names = args.clip_names if args.trajectory_names is None else args.trajectory_names
|
||||
for clip_info, trajectory_name in zip(clip_name2clip_info.values(), trajectory_names):
|
||||
pose_file = osp.join(args.clip_txt_path, clip_info['pose_file'])
|
||||
with open(pose_file, 'r') as f:
|
||||
poses = f.readlines()
|
||||
html = poses[0].strip()
|
||||
poses = [x.strip() for x in poses[1:]]
|
||||
total_frames = len(poses)
|
||||
cropped_length = args.num_frames * args.sample_stride
|
||||
start_frame_ind = random.randint(0, max(0, total_frames - cropped_length - 1))
|
||||
end_frame_ind = min(start_frame_ind + cropped_length, total_frames)
|
||||
assert end_frame_ind - start_frame_ind >= args.num_frames
|
||||
frame_ind = np.linspace(start_frame_ind, end_frame_ind - 1, args.num_frames, dtype=int)
|
||||
poses = [html, ] + [poses[ind] for ind in frame_ind]
|
||||
pose_save_file = osp.join(args.save_path, 'selected_pose_files', trajectory_name + '.txt')
|
||||
with open(pose_save_file, 'w') as f:
|
||||
for pose in poses:
|
||||
f.write(pose + '\n')
|
||||
clip_file = osp.join(args.clip_txt_path, clip_info['clip_path'])
|
||||
video_reader = VideoReader(clip_file)
|
||||
video_batch = video_reader.get_batch(frame_ind).asnumpy()
|
||||
video_batch = [cv2.resize(x, dsize=(args.video_width, args.video_height)) for x in video_batch]
|
||||
clip_save_file = osp.join(args.save_path, 'selected_clips', trajectory_name + '.mp4')
|
||||
imageio.mimsave(clip_save_file, video_batch, fps=8)
|
||||
selected_clip_infos.append({'clip_name': clip_info['clip_name'], 'caption': clip_info['caption'],
|
||||
'clip_path': clip_save_file, 'pose_file': pose_save_file,
|
||||
'trajectory_name': trajectory_name})
|
||||
if args.save_images:
|
||||
images_save_path = osp.join(args.save_path, 'selected_images', trajectory_name)
|
||||
os.makedirs(images_save_path, exist_ok=True)
|
||||
for image_idx, image in zip(frame_ind, video_batch):
|
||||
image_save_path = osp.join(images_save_path, f'{image_idx}.jpg')
|
||||
cv2.imwrite(image_save_path, cv2.cvtColor(image, cv2.COLOR_RGB2BGR))
|
||||
selected_clip_infos[-1].update({'images_save_path': images_save_path})
|
||||
with open(osp.join(args.save_path, 'selected_clip_infos.json'), 'w') as f:
|
||||
json.dump(selected_clip_infos, fp=f)
|
||||