mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
add rerender lora
This commit is contained in:
@@ -0,0 +1 @@
|
||||
animation/rerender-lora/Anaconda3-2023.09-0-Linux-x86_64.sh filter=lfs diff=lfs merge=lfs -text
|
||||
@@ -0,0 +1,24 @@
|
||||
FROM nvidia/cuda:11.8.0-devel-ubuntu22.04
|
||||
|
||||
# Set the working directory in the container
|
||||
WORKDIR /usr/src/app
|
||||
|
||||
# Copy the current directory contents into the container at /usr/src/app
|
||||
COPY . /usr/src/app
|
||||
|
||||
RUN sh Anaconda3-2023.09-0-Linux-x86_64.sh -b -p /usr/local/miniconda
|
||||
ENV PATH /usr/local/miniconda/bin:$PATH
|
||||
|
||||
RUN apt update
|
||||
RUN apt-get install -y libglib2.0-0 libsm6 libxrender1 libxext6
|
||||
|
||||
# Install any needed packages specified in requirements.txt
|
||||
RUN conda env create -f environment.yml
|
||||
|
||||
# Make RUN commands use the new environment:
|
||||
SHELL ["conda", "run", "-n", "rerender", "/bin/bash", "-c"]
|
||||
|
||||
RUN conda run -n rerender python install.py
|
||||
|
||||
# Run python script when the container launches
|
||||
ENTRYPOINT ["conda", "run", "-n", "rerender", "python", "rerender.py", "--cfg", "config/real2sculpture.json"]
|
||||
@@ -0,0 +1,14 @@
|
||||
# S-Lab License 1.0
|
||||
|
||||
Copyright 2023 S-Lab
|
||||
|
||||
Redistribution and use for non-commercial purpose in source and binary forms, with or without modification, are permitted provided that the following conditions are met:
|
||||
1. Redistributions of source code must retain the above copyright notice, this list of conditions and the following disclaimer.
|
||||
2. Redistributions in binary form must reproduce the above copyright notice, this list of conditions and the following disclaimer in the documentation and/or other materials provided with the distribution.
|
||||
3. Neither the name of the copyright holder nor the names of its contributors may be used to endorse or promote products derived from this software without specific prior written permission.\
|
||||
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
4. In the event that redistribution and/or use for commercial purpose in source or binary forms, with or without modification is required, please contact the contributor(s) of the work.
|
||||
|
||||
|
||||
---
|
||||
For the commercial use of the code, please consult Prof. Chen Change Loy (ccloy@ntu.edu.sg)
|
||||
@@ -0,0 +1,371 @@
|
||||
# Rerender A Video - Official PyTorch Implementation
|
||||
|
||||

|
||||
|
||||
<!--https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/82c35efb-e86b-4376-bfbe-6b69159b8879-->
|
||||
|
||||
|
||||
**Rerender A Video: Zero-Shot Text-Guided Video-to-Video Translation**<br>
|
||||
[Shuai Yang](https://williamyang1991.github.io/), [Yifan Zhou](https://zhouyifan.net/), [Ziwei Liu](https://liuziwei7.github.io/) and [Chen Change Loy](https://www.mmlab-ntu.com/person/ccloy/)<br>
|
||||
in SIGGRAPH Asia 2023 Conference Proceedings <br>
|
||||
[**Project Page**](https://www.mmlab-ntu.com/project/rerender/) | [**Paper**](https://arxiv.org/abs/2306.07954) | [**Supplementary Video**](https://youtu.be/cxfxdepKVaM) | [**Input Data and Video Results**](https://drive.google.com/file/d/1HkxG5eiLM_TQbbMZYOwjDbd5gWisOy4m/view?usp=sharing) <br>
|
||||
|
||||
<a href="https://huggingface.co/spaces/Anonymous-sub/Rerender"><img src="https://huggingface.co/datasets/huggingface/badges/raw/main/open-in-hf-spaces-sm-dark.svg" alt="Web Demo"></a> 
|
||||
|
||||
> **Abstract:** *Large text-to-image diffusion models have exhibited impressive proficiency in generating high-quality images. However, when applying these models to video domain, ensuring temporal consistency across video frames remains a formidable challenge. This paper proposes a novel zero-shot text-guided video-to-video translation framework to adapt image models to videos. The framework includes two parts: key frame translation and full video translation. The first part uses an adapted diffusion model to generate key frames, with hierarchical cross-frame constraints applied to enforce coherence in shapes, textures and colors. The second part propagates the key frames to other frames with temporal-aware patch matching and frame blending. Our framework achieves global style and local texture temporal consistency at a low cost (without re-training or optimization). The adaptation is compatible with existing image diffusion techniques, allowing our framework to take advantage of them, such as customizing a specific subject with LoRA, and introducing extra spatial guidance with ControlNet. Extensive experimental results demonstrate the effectiveness of our proposed framework over existing methods in rendering high-quality and temporally-coherent videos.*
|
||||
|
||||
**Features**:<br>
|
||||
- **Temporal consistency**: cross-frame constraints for low-level temporal consistency.
|
||||
- **Zero-shot**: no training or fine-tuning required.
|
||||
- **Flexibility**: compatible with off-the-shelf models (e.g., [ControlNet](https://github.com/lllyasviel/ControlNet), [LoRA](https://civitai.com/)) for customized translation.
|
||||
|
||||
https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/811fdea3-f0da-49c9-92b8-2d2ad360f0d6
|
||||
|
||||
## Updates
|
||||
- [10/2023] New features: [Loose cross-frame attention](#loose-cross-frame-attention) and [FreeU](#freeu).
|
||||
- [09/2023] Code is released.
|
||||
- [09/2023] Accepted to SIGGRAPH Asia 2023 Conference Proceedings!
|
||||
- [06/2023] Integrated to 🤗 [Hugging Face](https://huggingface.co/spaces/Anonymous-sub/Rerender). Enjoy the web demo!
|
||||
- [05/2023] This website is created.
|
||||
|
||||
### TODO
|
||||
- [x] Integrate into Diffusers.
|
||||
- [x] ~~Integrate [FreeU](https://github.com/ChenyangSi/FreeU) into Rerender~~
|
||||
- [x] ~~Add Inference instructions in README.md.~~
|
||||
- [x] ~~Add Examples to webUI.~~
|
||||
- [x] ~~Add optional poisson fusion to the pipeline.~~
|
||||
- [x] ~~Add Installation instructions for Windows~~
|
||||
|
||||
## Installation
|
||||
|
||||
*Please make sure your installation path only contain English letters or _*
|
||||
|
||||
1. Clone the repository. (Don't forget --recursive. Otherwise, please run `git submodule update --init --recursive`)
|
||||
|
||||
```shell
|
||||
git clone git@github.com:williamyang1991/Rerender_A_Video.git --recursive
|
||||
cd Rerender_A_Video
|
||||
```
|
||||
|
||||
2. If you have installed PyTorch CUDA, you can simply set up the environment with pip.
|
||||
|
||||
```shell
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
You can also create a new conda environment from scratch.
|
||||
|
||||
```shell
|
||||
conda env create -f environment.yml
|
||||
conda activate rerender
|
||||
```
|
||||
24GB VRAM is required. Please refer to https://github.com/williamyang1991/Rerender_A_Video/pull/23#issue-1900789461 to reduce memory consumption.
|
||||
|
||||
3. Run the installation script. The required models will be downloaded in `./models`.
|
||||
|
||||
```shell
|
||||
python install.py
|
||||
```
|
||||
|
||||
4. You can run the demo with `rerender.py`
|
||||
|
||||
```shell
|
||||
python rerender.py --cfg config/real2sculpture.json
|
||||
```
|
||||
|
||||
<details>
|
||||
<summary>Installation on Windows</summary>
|
||||
|
||||
Before running the above 1-4 steps, you need prepare:
|
||||
1. Install [CUDA](https://developer.nvidia.com/cuda-toolkit-archive)
|
||||
2. Install [git](https://git-scm.com/download/win)
|
||||
3. Install [VS](https://visualstudio.microsoft.com/) with Windows 10/11 SDK (for building deps/ebsynth/bin/ebsynth.exe)
|
||||
4. [Here](https://github.com/williamyang1991/Rerender_A_Video/issues/18#issuecomment-1752712233) are more information. If building ebsynth fails, we provides our complied [ebsynth](https://drive.google.com/drive/folders/1oSB3imKwZGz69q2unBUfcgmQpzwccoyD?usp=sharing).
|
||||
</details>
|
||||
|
||||
<details id="issues">
|
||||
<summary>🔥🔥🔥 <b>Installation or Running Fails?</b> 🔥🔥🔥</summary>
|
||||
|
||||
1. In case building ebsynth fails, we provides our complied [ebsynth](https://drive.google.com/drive/folders/1oSB3imKwZGz69q2unBUfcgmQpzwccoyD?usp=sharing)
|
||||
2. `FileNotFoundError: [Errno 2] No such file or directory: 'xxxx.bin' or 'xxxx.jpg'`:
|
||||
- make sure your path only contains English letters or _ (https://github.com/williamyang1991/Rerender_A_Video/issues/18#issuecomment-1723361433)
|
||||
- find the code `python video_blend.py ...` in the error log and use it to manually run the ebsynth part, which is more stable than WebUI.
|
||||
- if some non-keyframes are generated but somes are not, rather than missing all non-keyframes in '/out_xx/', you may refer to https://github.com/williamyang1991/Rerender_A_Video/issues/38#issuecomment-1730668991
|
||||
5. `KeyError: 'dataset'`: upgrade Gradio to the latest version (https://github.com/williamyang1991/Rerender_A_Video/issues/14#issuecomment-1722778672, https://github.com/AUTOMATIC1111/stable-diffusion-webui/issues/11855)
|
||||
6. Error when processing videos: manually install ffmpeg (https://github.com/williamyang1991/Rerender_A_Video/issues/19#issuecomment-1723685825, https://github.com/williamyang1991/Rerender_A_Video/issues/29#issuecomment-1726091112)
|
||||
7. `ERR_ADDRESS_INVALID` Cannot open the webUI in browser: replace 0.0.0.0 with 127.0.0.1 in webUI.py (https://github.com/williamyang1991/Rerender_A_Video/issues/19#issuecomment-1723685825)
|
||||
8. `CUDA out of memory`:
|
||||
- Using xformers (https://github.com/williamyang1991/Rerender_A_Video/pull/23#issue-1900789461)
|
||||
- Set `"use_limit_device_resolution"` to `true` in the config to resize the video according to your VRAM (https://github.com/williamyang1991/Rerender_A_Video/issues/79). An example config `config/van_gogh_man_dynamic_resolution.json` is provided.
|
||||
10. `AttributeError: module 'keras.backend' has no attribute 'is_tensor'`: update einops (https://github.com/williamyang1991/Rerender_A_Video/issues/26#issuecomment-1726682446)
|
||||
11. `IndexError: list index out of range`: use the original DDIM steps of 20 (https://github.com/williamyang1991/Rerender_A_Video/issues/30#issuecomment-1729039779)
|
||||
12. One-click installation https://github.com/williamyang1991/Rerender_A_Video/issues/99
|
||||
|
||||
</details>
|
||||
|
||||
|
||||
## (1) Inference
|
||||
|
||||
### WebUI (recommended)
|
||||
|
||||
```
|
||||
python webUI.py
|
||||
```
|
||||
The Gradio app also allows you to flexibly change the inference options. Just try it for more details. (For WebUI, you need to download [revAnimated_v11](https://civitai.com/models/7371/rev-animated?modelVersionId=19575) and [realisticVisionV20_v20](https://civitai.com/models/4201?modelVersionId=29460) to `./models/` after Installation)
|
||||
|
||||
Upload your video, input the prompt, select the seed, and hit:
|
||||
- **Run 1st Key Frame**: only translate the first frame, so you can adjust the prompts/models/parameters to find your ideal output appearance before running the whole video.
|
||||
- **Run Key Frames**: translate all the key frames based on the settings of the first frame, so you can adjust the temporal-related parameters for better temporal consistency before running the whole video.
|
||||
- **Run Propagation**: propagate the key frames to other frames for full video translation
|
||||
- **Run All**: **Run 1st Key Frame**, **Run Key Frames** and **Run Propagation**
|
||||
|
||||

|
||||
|
||||
|
||||
We provide abundant advanced options to play with
|
||||
|
||||
<details id="option0">
|
||||
<summary> <b>Using customized models</b></summary>
|
||||
|
||||
- Using LoRA/Dreambooth/Finetuned/Mixed SD models
|
||||
- Modify `sd_model_cfg.py` to add paths to the saved SD models
|
||||
- How to use LoRA: https://github.com/williamyang1991/Rerender_A_Video/issues/39#issuecomment-1730678296
|
||||
- Using other controls from ControlNet (e.g., Depth, Pose)
|
||||
- Add more options like `control_type = gr.Dropdown(['HED', 'canny', 'depth']` here https://github.com/williamyang1991/Rerender_A_Video/blob/b6cafb5d80a79a3ef831c689ffad92ec095f2794/webUI.py#L690
|
||||
- Add model loading options like `elif control_type == 'depth':` following https://github.com/williamyang1991/Rerender_A_Video/blob/b6cafb5d80a79a3ef831c689ffad92ec095f2794/webUI.py#L88
|
||||
- Add model detectors like `elif control_type == 'depth':` following https://github.com/williamyang1991/Rerender_A_Video/blob/b6cafb5d80a79a3ef831c689ffad92ec095f2794/webUI.py#L122
|
||||
- One example is given [here](https://huggingface.co/spaces/Anonymous-sub/Rerender/discussions/10/files)
|
||||
|
||||
</details>
|
||||
|
||||
<details id="option1">
|
||||
<summary> <b>Advanced options for the 1st frame translation</b></summary>
|
||||
|
||||
1. Resolution related (**Frame resolution**, **left/top/right/bottom crop length**): crop the frame and resize its short side to 512.
|
||||
2. ControlNet related:
|
||||
- **ControlNet strength**: how well the output matches the input control edges
|
||||
- **Control type**: HED edge or Canny edge
|
||||
- **Canny low/high threshold**: low values for more edge details
|
||||
3. SDEdit related:
|
||||
- **Denoising strength**: repaint degree (low value to make the output look more like the original video)
|
||||
- **Preserve color**: preserve the color of the original video
|
||||
4. SD related:
|
||||
- **Steps**: denoising step
|
||||
- **CFG scale**: how well the output matches the prompt
|
||||
- **Base model**: base Stable Diffusion model (SD 1.5)
|
||||
- Stable Diffusion 1.5: official model
|
||||
- [revAnimated_v11](https://civitai.com/models/7371/rev-animated?modelVersionId=19575): a semi-realistic (2.5D) model
|
||||
- [realisticVisionV20_v20](https://civitai.com/models/4201?modelVersionId=29460): a photo-realistic model
|
||||
- **Added prompt/Negative prompt**: supplementary prompts
|
||||
5. FreeU related:
|
||||
- **FreeU first/second-stage backbone factor**: =1 do nothing; >1 enhance output color and details
|
||||
- **FreeU first/second-stage skip factor**: =1 do nothing; <1 enhance output color and details
|
||||
|
||||
</details>
|
||||
|
||||
<details id="option2">
|
||||
<summary> <b>Advanced options for the key frame translation</b></summary>
|
||||
|
||||
1. Key frame related
|
||||
- **Key frame frequency (K)**: Uniformly sample the key frame every K frames. Small value for large or fast motions.
|
||||
- **Number of key frames (M)**: The final output video will have K*M+1 frames with M+1 key frames.
|
||||
2. Temporal consistency related
|
||||
- Cross-frame attention:
|
||||
- **Cross-frame attention start/end**: When applying cross-frame attention for global style consistency
|
||||
- **Cross-frame attention update frequency (N)**: Update the reference style frame every N key frames. Should be large for long videos to avoid error accumulation.
|
||||
- **Loose Cross-frame attention**: Using cross-frame attention in fewer layers to better match the input video (for video with large motions)
|
||||
- **Shape-aware fusion** Check to use this feature
|
||||
- **Shape-aware fusion start/end**: When applying shape-aware fusion for local shape consistency
|
||||
- **Pixel-aware fusion** Check to use this feature
|
||||
- **Pixel-aware fusion start/end**: When applying pixel-aware fusion for pixel-level temporal consistency
|
||||
- **Pixel-aware fusion strength**: The strength to preserve the non-inpainting region. Small to avoid error accumulation. Large to avoid burry textures.
|
||||
- **Pixel-aware fusion detail level**: The strength to sharpen the inpainting region. Small to avoid error accumulation. Large to avoid burry textures.
|
||||
- **Smooth fusion boundary**: Check to smooth the inpainting boundary (avoid error accumulation).
|
||||
- **Color-aware AdaIN** Check to use this feature
|
||||
- **Color-aware AdaIN start/end**: When applying AdaIN to make the video color consistent with the first frame
|
||||
|
||||
</details>
|
||||
|
||||
<details id="option3">
|
||||
<summary> <b>Advanced options for the full video translation</b></summary>
|
||||
|
||||
1. **Gradient blending**: apply Poisson Blending to reduce ghosting artifacts. May slow the process and increase flickers.
|
||||
2. **Number of parallel processes**: multiprocessing to speed up the process. Large value (8) is recommended.
|
||||
</details>
|
||||
|
||||

|
||||
|
||||
|
||||
### Command Line
|
||||
|
||||
We also provide a flexible script `rerender.py` to run our method.
|
||||
|
||||
#### Simple mode
|
||||
|
||||
Set the options via command line. For example,
|
||||
|
||||
```shell
|
||||
python rerender.py --input videos/pexels-antoni-shkraba-8048492-540x960-25fps.mp4 --output result/man/man.mp4 --prompt "a handsome man in van gogh painting"
|
||||
```
|
||||
|
||||
The script will run the full pipeline. A work directory will be created at `result/man` and the result video will be saved as `result/man/man.mp4`
|
||||
|
||||
#### Advanced mode
|
||||
|
||||
Set the options via a config file. For example,
|
||||
|
||||
```shell
|
||||
python rerender.py --cfg config/van_gogh_man.json
|
||||
```
|
||||
|
||||
The script will run the full pipeline.
|
||||
We provide some examples of the config in `config` directory.
|
||||
Most options in the config is the same as those in WebUI.
|
||||
Please check the explanations in the WebUI section.
|
||||
|
||||
Specifying customized models by setting `sd_model` in config. For example:
|
||||
```json
|
||||
{
|
||||
"sd_model": "models/realisticVisionV20_v20.safetensors",
|
||||
}
|
||||
```
|
||||
|
||||
#### Customize the pipeline
|
||||
|
||||
Similar to WebUI, we provide three-step workflow: Rerender the first key frame, then rerender the full key frames, finally rerender the full video with propagation. To run only a single step, specify options `-one`, `-nb` and `-nr`:
|
||||
|
||||
1. Rerender the first key frame
|
||||
```shell
|
||||
python rerender.py --cfg config/van_gogh_man.json -one -nb
|
||||
```
|
||||
2. Rerender the full key frames
|
||||
```shell
|
||||
python rerender.py --cfg config/van_gogh_man.json -nb
|
||||
```
|
||||
3. Rerender the full video with propagation
|
||||
```shell
|
||||
python rerender.py --cfg config/van_gogh_man.json -nr
|
||||
```
|
||||
|
||||
#### Our Ebsynth implementation
|
||||
|
||||
We provide a separate Ebsynth python script `video_blend.py` with the temporal blending algorithm introduced in
|
||||
[Stylizing Video by Example](https://dcgi.fel.cvut.cz/home/sykorad/ebsynth.html) for interpolating style between key frames.
|
||||
It can work on your own stylized key frames independently of our Rerender algorithm.
|
||||
|
||||
Usage:
|
||||
```shell
|
||||
video_blend.py [-h] [--output OUTPUT] [--fps FPS] [--beg BEG] [--end END] [--itv ITV] [--key KEY]
|
||||
[--n_proc N_PROC] [-ps] [-ne] [-tmp]
|
||||
name
|
||||
|
||||
positional arguments:
|
||||
name Path to input video
|
||||
|
||||
optional arguments:
|
||||
-h, --help show this help message and exit
|
||||
--output OUTPUT Path to output video
|
||||
--fps FPS The FPS of output video
|
||||
--beg BEG The index of the first frame to be stylized
|
||||
--end END The index of the last frame to be stylized
|
||||
--itv ITV The interval of key frame
|
||||
--key KEY The subfolder name of stylized key frames
|
||||
--n_proc N_PROC The max process count
|
||||
-ps Use poisson gradient blending
|
||||
-ne Do not run ebsynth (use previous ebsynth output)
|
||||
-tmp Keep temporary output
|
||||
```
|
||||
For example, to run Ebsynth on video `man.mp4`,
|
||||
1. Put the stylized key frames to `videos/man/keys` for every 10 frames (named as `0001.png`, `0011.png`, ...)
|
||||
2. Put the original video frames in `videos/man/video` (named as `0001.png`, `0002.png`, ...).
|
||||
3. Run Ebsynth on the first 101 frames of the video with poisson gradient blending and save the result to `videos/man/blend.mp4` under FPS 25 with the following command:
|
||||
```shell
|
||||
python video_blend.py videos/man \
|
||||
--beg 1 \
|
||||
--end 101 \
|
||||
--itv 10 \
|
||||
--key keys \
|
||||
--output videos/man/blend.mp4 \
|
||||
--fps 25.0 \
|
||||
-ps
|
||||
```
|
||||
|
||||
## (2) Results
|
||||
|
||||
### Key frame translation
|
||||
|
||||
|
||||
<table class="center">
|
||||
<tr>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/18666871-f273-44b2-ae67-7be85d43e2f6" raw=true></td>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/61f59540-f06e-4e5a-86b6-1d7cb8ed6300" raw=true></td>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/8e8ad51a-6a71-4b34-8633-382192d0f17c" raw=true></td>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/b03cd35f-5d90-471a-9aa9-5c7773d7ac39" raw=true></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=27.5% align="center">white ancient Greek sculpture, Venus de Milo, light pink and blue background</td>
|
||||
<td width=27.5% align="center">a handsome Greek man</td>
|
||||
<td width=21.5% align="center">a traditional mountain in chinese ink wash painting</td>
|
||||
<td width=23.5% align="center">a cartoon tiger</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
<table class="center">
|
||||
<tr>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/649a789e-0c41-41cf-94a4-0d524dcfb282" raw=true></td>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/73590c16-916f-4ee6-881a-44a201dd85dd" raw=true></td>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/fbdc0b8e-6046-414f-a37e-3cd9dd0adf5d" raw=true></td>
|
||||
<td><img src="https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/eb11d807-2afa-4609-a074-34300b67e6aa" raw=true></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td width=26.0% align="center">a swan in chinese ink wash painting, monochrome</td>
|
||||
<td width=29.0% align="center">a beautiful woman in CG style</td>
|
||||
<td width=21.5% align="center">a clean simple white jade sculpture</td>
|
||||
<td width=24.0% align="center">a fluorescent jellyfish in the deep dark blue sea</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### Full video translation
|
||||
|
||||
Text-guided virtual character generation.
|
||||
|
||||
|
||||
https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/1405b257-e59a-427f-890d-7652e6bed0a4
|
||||
|
||||
|
||||
https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/efee8cc6-9708-4124-bf6a-49baf91349fc
|
||||
|
||||
|
||||
Video stylization and video editing.
|
||||
|
||||
|
||||
https://github.com/williamyang1991/Rerender_A_Video/assets/18130694/1b72585c-99c0-401d-b240-5b8016df7a3f
|
||||
|
||||
## New Features
|
||||
|
||||
Compared to the conference version, we are keeping adding new features.
|
||||
|
||||

|
||||
|
||||
#### Loose cross-frame attention
|
||||
By using cross-frame attention in less layers, our results will better match the input video, thus reducing ghosting artifacts caused by inconsistencies. This feature can be activated by checking `Loose Cross-frame attention` in the <a href="#option2">Advanced options for the key frame translation</a> for WebUI or setting `loose_cfattn` for script (see `config/real2sculpture_loose_cfattn.json`).
|
||||
|
||||
#### FreeU
|
||||
[FreeU](https://github.com/ChenyangSi/FreeU) is a method that improves diffusion model sample quality at no costs. We find featured with FreeU, our results will have higher contrast and saturation, richer details, and more vivid colors. This feature can be used by setting FreeU backbone factors and skip factors in the <a href="#option1">Advanced options for the 1st frame translation</a> for WebUI or setting `freeu_args` for script (see `config/real2sculpture_freeu.json`).
|
||||
|
||||
## Citation
|
||||
|
||||
If you find this work useful for your research, please consider citing our paper:
|
||||
|
||||
```bibtex
|
||||
@inproceedings{yang2023rerender,
|
||||
title = {Rerender A Video: Zero-Shot Text-Guided Video-to-Video Translation},
|
||||
author = {Yang, Shuai and Zhou, Yifan and Liu, Ziwei and and Loy, Chen Change},
|
||||
booktitle = {ACM SIGGRAPH Asia Conference Proceedings},
|
||||
year = {2023},
|
||||
}
|
||||
```
|
||||
|
||||
## Acknowledgments
|
||||
|
||||
The code is mainly developed based on [ControlNet](https://github.com/lllyasviel/ControlNet), [Stable Diffusion](https://github.com/Stability-AI/stablediffusion), [GMFlow](https://github.com/haofeixu/gmflow) and [Ebsynth](https://github.com/jamriska/ebsynth).
|
||||
@@ -0,0 +1,105 @@
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from flow.flow_utils import flow_calc, read_flow, read_mask
|
||||
|
||||
|
||||
class BaseGuide:
|
||||
|
||||
def __init__(self):
|
||||
...
|
||||
|
||||
def get_cmd(self, i, weight) -> str:
|
||||
return (f'-guide {os.path.abspath(self.imgs[0])} '
|
||||
f'{os.path.abspath(self.imgs[i])} -weight {weight}')
|
||||
|
||||
|
||||
class ColorGuide(BaseGuide):
|
||||
|
||||
def __init__(self, imgs):
|
||||
super().__init__()
|
||||
self.imgs = imgs
|
||||
|
||||
|
||||
class PositionalGuide(BaseGuide):
|
||||
|
||||
def __init__(self, flow_paths, save_paths):
|
||||
super().__init__()
|
||||
flows = [read_flow(f) for f in flow_paths]
|
||||
masks = [read_mask(f) for f in flow_paths]
|
||||
# TODO: modify the format of flow to numpy
|
||||
H, W = flows[0].shape[2:]
|
||||
first_img = PositionalGuide.__generate_first_img(H, W)
|
||||
prev_img = first_img
|
||||
imgs = [first_img]
|
||||
cid = 0
|
||||
for flow, mask in zip(flows, masks):
|
||||
cur_img = flow_calc.warp(prev_img, flow,
|
||||
'nearest').astype(np.uint8)
|
||||
cur_img = cv2.inpaint(cur_img, mask, 30, cv2.INPAINT_TELEA)
|
||||
prev_img = cur_img
|
||||
imgs.append(cur_img)
|
||||
cid += 1
|
||||
cv2.imwrite(f'guide/{cid}.jpg', mask)
|
||||
|
||||
for path, img in zip(save_paths, imgs):
|
||||
cv2.imwrite(path, img)
|
||||
self.imgs = save_paths
|
||||
|
||||
@staticmethod
|
||||
def __generate_first_img(H, W):
|
||||
Hs = np.linspace(0, 1, H)
|
||||
Ws = np.linspace(0, 1, W)
|
||||
i, j = np.meshgrid(Hs, Ws, indexing='ij')
|
||||
r = (i * 255).astype(np.uint8)
|
||||
g = (j * 255).astype(np.uint8)
|
||||
b = np.zeros(r.shape)
|
||||
res = np.stack((b, g, r), 2)
|
||||
return res
|
||||
|
||||
|
||||
class EdgeGuide(BaseGuide):
|
||||
|
||||
def __init__(self, imgs, save_paths):
|
||||
super().__init__()
|
||||
edges = [EdgeGuide.__generate_edge(cv2.imread(img)) for img in imgs]
|
||||
for path, img in zip(save_paths, edges):
|
||||
cv2.imwrite(path, img)
|
||||
self.imgs = save_paths
|
||||
|
||||
@staticmethod
|
||||
def __generate_edge(img):
|
||||
filter = np.array([[0, -1, 0], [-1, 4, -1], [0, -1, 0]])
|
||||
res = cv2.filter2D(img, -1, filter)
|
||||
return res
|
||||
|
||||
|
||||
class TemporalGuide(BaseGuide):
|
||||
|
||||
def __init__(self, key_img, stylized_imgs, flow_paths, save_paths):
|
||||
super().__init__()
|
||||
self.flows = [read_flow(f) for f in flow_paths]
|
||||
self.masks = [read_mask(f) for f in flow_paths]
|
||||
self.stylized_imgs = stylized_imgs
|
||||
self.imgs = save_paths
|
||||
|
||||
first_img = cv2.imread(key_img)
|
||||
cv2.imwrite(self.imgs[0], first_img)
|
||||
|
||||
def get_cmd(self, i, weight) -> str:
|
||||
if i == 0:
|
||||
warped_img = self.stylized_imgs[0]
|
||||
else:
|
||||
prev_img = cv2.imread(self.stylized_imgs[i - 1])
|
||||
warped_img = flow_calc.warp(prev_img, self.flows[i - 1],
|
||||
'nearest').astype(np.uint8)
|
||||
|
||||
warped_img = cv2.inpaint(warped_img, self.masks[i - 1], 30,
|
||||
cv2.INPAINT_TELEA)
|
||||
|
||||
cv2.imwrite(self.imgs[i], warped_img)
|
||||
|
||||
return super().get_cmd(i, weight)
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def histogram_transform(img: np.ndarray, means: np.ndarray, stds: np.ndarray,
|
||||
target_means: np.ndarray, target_stds: np.ndarray):
|
||||
means = means.reshape((1, 1, 3))
|
||||
stds = stds.reshape((1, 1, 3))
|
||||
target_means = target_means.reshape((1, 1, 3))
|
||||
target_stds = target_stds.reshape((1, 1, 3))
|
||||
x = img.astype(np.float32)
|
||||
x = (x - means) * target_stds / stds + target_means
|
||||
# x = np.round(x)
|
||||
# x = np.clip(x, 0, 255)
|
||||
# x = x.astype(np.uint8)
|
||||
return x
|
||||
|
||||
|
||||
def blend(a: np.ndarray,
|
||||
b: np.ndarray,
|
||||
min_error: np.ndarray,
|
||||
weight1=0.5,
|
||||
weight2=0.5):
|
||||
a = cv2.cvtColor(a, cv2.COLOR_BGR2Lab)
|
||||
b = cv2.cvtColor(b, cv2.COLOR_BGR2Lab)
|
||||
min_error = cv2.cvtColor(min_error, cv2.COLOR_BGR2Lab)
|
||||
a_mean = np.mean(a, axis=(0, 1))
|
||||
a_std = np.std(a, axis=(0, 1))
|
||||
b_mean = np.mean(b, axis=(0, 1))
|
||||
b_std = np.std(b, axis=(0, 1))
|
||||
min_error_mean = np.mean(min_error, axis=(0, 1))
|
||||
min_error_std = np.std(min_error, axis=(0, 1))
|
||||
|
||||
t_mean_val = 0.5 * 256
|
||||
t_std_val = (1 / 36) * 256
|
||||
t_mean = np.ones([3], dtype=np.float32) * t_mean_val
|
||||
t_std = np.ones([3], dtype=np.float32) * t_std_val
|
||||
a = histogram_transform(a, a_mean, a_std, t_mean, t_std)
|
||||
|
||||
b = histogram_transform(b, b_mean, b_std, t_mean, t_std)
|
||||
ab = (a * weight1 + b * weight2 - t_mean_val) / 0.5 + t_mean_val
|
||||
ab_mean = np.mean(ab, axis=(0, 1))
|
||||
ab_std = np.std(ab, axis=(0, 1))
|
||||
ab = histogram_transform(ab, ab_mean, ab_std, min_error_mean,
|
||||
min_error_std)
|
||||
ab = np.round(ab)
|
||||
ab = np.clip(ab, 0, 255)
|
||||
ab = ab.astype(np.uint8)
|
||||
ab = cv2.cvtColor(ab, cv2.COLOR_Lab2BGR)
|
||||
return ab
|
||||
@@ -0,0 +1,93 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
import scipy
|
||||
|
||||
As = None
|
||||
prev_states = None
|
||||
|
||||
|
||||
def construct_A(h, w, grad_weight):
|
||||
indgx_x = []
|
||||
indgx_y = []
|
||||
indgy_x = []
|
||||
indgy_y = []
|
||||
vdx = []
|
||||
vdy = []
|
||||
for i in range(h):
|
||||
for j in range(w):
|
||||
if i < h - 1:
|
||||
indgx_x += [i * w + j]
|
||||
indgx_y += [i * w + j]
|
||||
vdx += [1]
|
||||
indgx_x += [i * w + j]
|
||||
indgx_y += [(i + 1) * w + j]
|
||||
vdx += [-1]
|
||||
if j < w - 1:
|
||||
indgy_x += [i * w + j]
|
||||
indgy_y += [i * w + j]
|
||||
vdy += [1]
|
||||
indgy_x += [i * w + j]
|
||||
indgy_y += [i * w + j + 1]
|
||||
vdy += [-1]
|
||||
Ix = scipy.sparse.coo_array(
|
||||
(np.ones(h * w), (np.arange(h * w), np.arange(h * w))),
|
||||
shape=(h * w, h * w)).tocsc()
|
||||
Gx = scipy.sparse.coo_array(
|
||||
(np.array(vdx), (np.array(indgx_x), np.array(indgx_y))),
|
||||
shape=(h * w, h * w)).tocsc()
|
||||
Gy = scipy.sparse.coo_array(
|
||||
(np.array(vdy), (np.array(indgy_x), np.array(indgy_y))),
|
||||
shape=(h * w, h * w)).tocsc()
|
||||
As = []
|
||||
for i in range(3):
|
||||
As += [
|
||||
scipy.sparse.vstack([Gx * grad_weight[i], Gy * grad_weight[i], Ix])
|
||||
]
|
||||
return As
|
||||
|
||||
|
||||
# blendI, I1, I2, mask should be RGB unit8 type
|
||||
# return poissson fusion result (RGB unit8 type)
|
||||
# I1 and I2: propagated results from previous and subsequent key frames
|
||||
# mask: pixel selection mask
|
||||
# blendI: contrastive-preserving blending results of I1 and I2
|
||||
def poisson_fusion(blendI, I1, I2, mask, grad_weight=[2.5, 0.5, 0.5]):
|
||||
global As
|
||||
global prev_states
|
||||
|
||||
Iab = cv2.cvtColor(blendI, cv2.COLOR_BGR2LAB).astype(float)
|
||||
Ia = cv2.cvtColor(I1, cv2.COLOR_BGR2LAB).astype(float)
|
||||
Ib = cv2.cvtColor(I2, cv2.COLOR_BGR2LAB).astype(float)
|
||||
m = (mask > 0).astype(float)[:, :, np.newaxis]
|
||||
h, w, c = Iab.shape
|
||||
|
||||
# fuse the gradient of I1 and I2 with mask
|
||||
gx = np.zeros_like(Ia)
|
||||
gy = np.zeros_like(Ia)
|
||||
gx[:-1, :, :] = (Ia[:-1, :, :] - Ia[1:, :, :]) * (1 - m[:-1, :, :]) + (
|
||||
Ib[:-1, :, :] - Ib[1:, :, :]) * m[:-1, :, :]
|
||||
gy[:, :-1, :] = (Ia[:, :-1, :] - Ia[:, 1:, :]) * (1 - m[:, :-1, :]) + (
|
||||
Ib[:, :-1, :] - Ib[:, 1:, :]) * m[:, :-1, :]
|
||||
|
||||
# construct A for solving Ax=b
|
||||
crt_states = (h, w, grad_weight)
|
||||
if As is None or crt_states != prev_states:
|
||||
As = construct_A(*crt_states)
|
||||
prev_states = crt_states
|
||||
|
||||
final = []
|
||||
for i in range(3):
|
||||
weight = grad_weight[i]
|
||||
im_dx = np.clip(gx[:, :, i].reshape(h * w, 1), -100, 100)
|
||||
im_dy = np.clip(gy[:, :, i].reshape(h * w, 1), -100, 100)
|
||||
im = Iab[:, :, i].reshape(h * w, 1)
|
||||
im_mean = im.mean()
|
||||
im = im - im_mean
|
||||
A = As[i]
|
||||
b = np.vstack([im_dx * weight, im_dy * weight, im])
|
||||
out = scipy.sparse.linalg.lsqr(A, b)
|
||||
out_im = (out[0] + im_mean).reshape(h, w, 1)
|
||||
final += [out_im]
|
||||
|
||||
final = np.clip(np.concatenate(final, axis=2), 0, 255)
|
||||
return cv2.cvtColor(final.astype(np.uint8), cv2.COLOR_LAB2BGR)
|
||||
@@ -0,0 +1,189 @@
|
||||
import os
|
||||
import shutil
|
||||
|
||||
|
||||
class VideoSequence:
|
||||
|
||||
def __init__(self,
|
||||
base_dir,
|
||||
beg_frame,
|
||||
end_frame,
|
||||
interval,
|
||||
input_subdir='videos',
|
||||
key_subdir='keys0',
|
||||
tmp_subdir='tmp',
|
||||
input_format='frame%04d.jpg',
|
||||
key_format='%04d.jpg',
|
||||
out_subdir_format='out_%d',
|
||||
blending_out_subdir='blend',
|
||||
output_format='%04d.jpg'):
|
||||
if (end_frame - beg_frame) % interval != 0:
|
||||
end_frame -= (end_frame - beg_frame) % interval
|
||||
|
||||
self.__base_dir = base_dir
|
||||
self.__input_dir = os.path.join(base_dir, input_subdir)
|
||||
self.__key_dir = os.path.join(base_dir, key_subdir)
|
||||
self.__tmp_dir = os.path.join(base_dir, tmp_subdir)
|
||||
self.__input_format = input_format
|
||||
self.__blending_out_dir = os.path.join(base_dir, blending_out_subdir)
|
||||
self.__key_format = key_format
|
||||
self.__out_subdir_format = out_subdir_format
|
||||
self.__output_format = output_format
|
||||
self.__beg_frame = beg_frame
|
||||
self.__end_frame = end_frame
|
||||
self.__interval = interval
|
||||
self.__n_seq = (end_frame - beg_frame) // interval
|
||||
self.__make_out_dirs()
|
||||
os.makedirs(self.__tmp_dir, exist_ok=True)
|
||||
|
||||
@property
|
||||
def beg_frame(self):
|
||||
return self.__beg_frame
|
||||
|
||||
@property
|
||||
def end_frame(self):
|
||||
return self.__end_frame
|
||||
|
||||
@property
|
||||
def n_seq(self):
|
||||
return self.__n_seq
|
||||
|
||||
@property
|
||||
def interval(self):
|
||||
return self.__interval
|
||||
|
||||
@property
|
||||
def blending_dir(self):
|
||||
return os.path.abspath(self.__blending_out_dir)
|
||||
|
||||
def remove_out_and_tmp(self):
|
||||
for i in range(self.n_seq + 1):
|
||||
out_dir = self.__get_out_subdir(i)
|
||||
shutil.rmtree(out_dir)
|
||||
shutil.rmtree(self.__tmp_dir)
|
||||
|
||||
def get_input_sequence(self, i, is_forward=True):
|
||||
beg_id = self.get_sequence_beg_id(i)
|
||||
end_id = self.get_sequence_beg_id(i + 1)
|
||||
if is_forward:
|
||||
id_list = list(range(beg_id, end_id))
|
||||
else:
|
||||
id_list = list(range(end_id, beg_id, -1))
|
||||
path_dir = [
|
||||
os.path.join(self.__input_dir, self.__input_format % id)
|
||||
for id in id_list
|
||||
]
|
||||
return path_dir
|
||||
|
||||
def get_output_sequence(self, i, is_forward=True):
|
||||
beg_id = self.get_sequence_beg_id(i)
|
||||
end_id = self.get_sequence_beg_id(i + 1)
|
||||
if is_forward:
|
||||
id_list = list(range(beg_id, end_id))
|
||||
else:
|
||||
i += 1
|
||||
id_list = list(range(end_id, beg_id, -1))
|
||||
out_subdir = self.__get_out_subdir(i)
|
||||
path_dir = [
|
||||
os.path.join(out_subdir, self.__output_format % id)
|
||||
for id in id_list
|
||||
]
|
||||
return path_dir
|
||||
|
||||
def get_temporal_sequence(self, i, is_forward=True):
|
||||
beg_id = self.get_sequence_beg_id(i)
|
||||
end_id = self.get_sequence_beg_id(i + 1)
|
||||
if is_forward:
|
||||
id_list = list(range(beg_id, end_id))
|
||||
else:
|
||||
i += 1
|
||||
id_list = list(range(end_id, beg_id, -1))
|
||||
tmp_dir = self.__get_tmp_out_subdir(i)
|
||||
path_dir = [
|
||||
os.path.join(tmp_dir, 'temporal_' + self.__output_format % id)
|
||||
for id in id_list
|
||||
]
|
||||
return path_dir
|
||||
|
||||
def get_edge_sequence(self, i, is_forward=True):
|
||||
beg_id = self.get_sequence_beg_id(i)
|
||||
end_id = self.get_sequence_beg_id(i + 1)
|
||||
if is_forward:
|
||||
id_list = list(range(beg_id, end_id))
|
||||
else:
|
||||
i += 1
|
||||
id_list = list(range(end_id, beg_id, -1))
|
||||
tmp_dir = self.__get_tmp_out_subdir(i)
|
||||
path_dir = [
|
||||
os.path.join(tmp_dir, 'edge_' + self.__output_format % id)
|
||||
for id in id_list
|
||||
]
|
||||
return path_dir
|
||||
|
||||
def get_pos_sequence(self, i, is_forward=True):
|
||||
beg_id = self.get_sequence_beg_id(i)
|
||||
end_id = self.get_sequence_beg_id(i + 1)
|
||||
if is_forward:
|
||||
id_list = list(range(beg_id, end_id))
|
||||
else:
|
||||
i += 1
|
||||
id_list = list(range(end_id, beg_id, -1))
|
||||
tmp_dir = self.__get_tmp_out_subdir(i)
|
||||
path_dir = [
|
||||
os.path.join(tmp_dir, 'pos_' + self.__output_format % id)
|
||||
for id in id_list
|
||||
]
|
||||
return path_dir
|
||||
|
||||
def get_flow_sequence(self, i, is_forward=True):
|
||||
beg_id = self.get_sequence_beg_id(i)
|
||||
end_id = self.get_sequence_beg_id(i + 1)
|
||||
if is_forward:
|
||||
id_list = list(range(beg_id, end_id - 1))
|
||||
path_dir = [
|
||||
os.path.join(self.__tmp_dir, 'flow_f_%04d.npy' % id)
|
||||
for id in id_list
|
||||
]
|
||||
else:
|
||||
id_list = list(range(end_id, beg_id + 1, -1))
|
||||
path_dir = [
|
||||
os.path.join(self.__tmp_dir, 'flow_b_%04d.npy' % id)
|
||||
for id in id_list
|
||||
]
|
||||
|
||||
return path_dir
|
||||
|
||||
def get_input_img(self, i):
|
||||
return os.path.join(self.__input_dir, self.__input_format % i)
|
||||
|
||||
def get_key_img(self, i):
|
||||
sequence_beg_id = self.get_sequence_beg_id(i)
|
||||
return os.path.join(self.__key_dir,
|
||||
self.__key_format % sequence_beg_id)
|
||||
|
||||
def get_blending_img(self, i):
|
||||
return os.path.join(self.__blending_out_dir, self.__output_format % i)
|
||||
|
||||
def get_sequence_beg_id(self, i):
|
||||
return i * self.__interval + self.__beg_frame
|
||||
|
||||
def __get_out_subdir(self, i):
|
||||
dir_id = self.get_sequence_beg_id(i)
|
||||
out_subdir = os.path.join(self.__base_dir,
|
||||
self.__out_subdir_format % dir_id)
|
||||
return out_subdir
|
||||
|
||||
def __get_tmp_out_subdir(self, i):
|
||||
dir_id = self.get_sequence_beg_id(i)
|
||||
tmp_out_subdir = os.path.join(self.__tmp_dir,
|
||||
self.__out_subdir_format % dir_id)
|
||||
return tmp_out_subdir
|
||||
|
||||
def __make_out_dirs(self):
|
||||
os.makedirs(self.__base_dir, exist_ok=True)
|
||||
os.makedirs(self.__blending_out_dir, exist_ok=True)
|
||||
for i in range(self.__n_seq + 1):
|
||||
out_subdir = self.__get_out_subdir(i)
|
||||
tmp_subdir = self.__get_tmp_out_subdir(i)
|
||||
os.makedirs(out_subdir, exist_ok=True)
|
||||
os.makedirs(tmp_subdir, exist_ok=True)
|
||||
@@ -0,0 +1,32 @@
|
||||
{
|
||||
"input": "videos/pexels-cottonbro-studio-6649832-960x506-25fps.mp4",
|
||||
"output": "videos/pexels-cottonbro-studio-6649832-960x506-25fps/blend.mp4",
|
||||
"work_dir": "videos/pexels-cottonbro-studio-6649832-960x506-25fps",
|
||||
"key_subdir": "keys",
|
||||
"sd_model": "models/realisticVisionV20_v20.safetensors",
|
||||
"frame_count": 102,
|
||||
"interval": 10,
|
||||
"crop": [
|
||||
0,
|
||||
180,
|
||||
0,
|
||||
0
|
||||
],
|
||||
"prompt": "white ancient Greek sculpture, Venus de Milo, light pink and blue background",
|
||||
"a_prompt": "RAW photo, subject, (high detailed skin:1.2), 8k uhd, dslr, soft lighting, high quality, film grain, Fujifilm XT3",
|
||||
"n_prompt": "(deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, mutated hands and fingers:1.4), (deformed, distorted, disfigured:1.3), poorly drawn, bad anatomy, wrong anatomy, extra limb, missing limb, floating limbs, disconnected limbs, mutation, mutated, ugly, disgusting, amputation",
|
||||
"x0_strength": 0.95,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
{
|
||||
"input": "videos/pexels-cottonbro-studio-6649832-960x506-25fps.mp4",
|
||||
"output": "videos/pexels-cottonbro-studio-6649832-960x506-25fps/blend.mp4",
|
||||
"work_dir": "videos/pexels-cottonbro-studio-6649832-960x506-25fps",
|
||||
"key_subdir": "keys",
|
||||
"sd_model": "models/realisticVisionV20_v20.safetensors",
|
||||
"frame_count": 102,
|
||||
"interval": 10,
|
||||
"crop": [
|
||||
0,
|
||||
180,
|
||||
0,
|
||||
0
|
||||
],
|
||||
"prompt": "white ancient Greek sculpture, Venus de Milo, light pink and blue background",
|
||||
"a_prompt": "RAW photo, subject, (high detailed skin:1.2), 8k uhd, dslr, soft lighting, high quality, film grain, Fujifilm XT3",
|
||||
"n_prompt": "(deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, mutated hands and fingers:1.4), (deformed, distorted, disfigured:1.3), poorly drawn, bad anatomy, wrong anatomy, extra limb, missing limb, floating limbs, disconnected limbs, mutation, mutated, ugly, disgusting, amputation",
|
||||
"x0_strength": 0.95,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
],
|
||||
"freeu_args": [
|
||||
1.1,
|
||||
1.2,
|
||||
1.0,
|
||||
0.2
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"input": "videos/pexels-cottonbro-studio-6649832-960x506-25fps.mp4",
|
||||
"output": "videos/pexels-cottonbro-studio-6649832-960x506-25fps/blend.mp4",
|
||||
"work_dir": "videos/pexels-cottonbro-studio-6649832-960x506-25fps",
|
||||
"key_subdir": "keys",
|
||||
"sd_model": "models/realisticVisionV20_v20.safetensors",
|
||||
"frame_count": 102,
|
||||
"interval": 10,
|
||||
"crop": [
|
||||
0,
|
||||
180,
|
||||
0,
|
||||
0
|
||||
],
|
||||
"prompt": "white ancient Greek sculpture, Venus de Milo, light pink and blue background",
|
||||
"a_prompt": "RAW photo, subject, (high detailed skin:1.2), 8k uhd, dslr, soft lighting, high quality, film grain, Fujifilm XT3",
|
||||
"n_prompt": "(deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, mutated hands and fingers:1.4), (deformed, distorted, disfigured:1.3), poorly drawn, bad anatomy, wrong anatomy, extra limb, missing limb, floating limbs, disconnected limbs, mutation, mutated, ugly, disgusting, amputation",
|
||||
"x0_strength": 0.95,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
],
|
||||
"loose_cfattn": true
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
{
|
||||
"input": "videos/testanime.mp4",
|
||||
"output": "videos/testanime/blend.mp4",
|
||||
"work_dir": "videos/testanime",
|
||||
"key_subdir": "keys",
|
||||
"sd_model": "/home/salt/clone/ComfyUI/models/checkpoints/animePastelDream_softBakedVae.safetensors",
|
||||
"lora_path": "/home/salt/Downloads/ya-v100-000029.safetensors",
|
||||
"interval": 8,
|
||||
"crop": [
|
||||
0,
|
||||
180,
|
||||
0,
|
||||
0
|
||||
],
|
||||
"prompt": "1girl, ya, def clothe, 1girl, long hair, skirt, animal ears, pantyhose, blunt bangs, bangs, gloves",
|
||||
"a_prompt": "",
|
||||
"n_prompt": "(worst quality, low quality:1.1), (bdpro:0.65), Low quality, Photo, Artifacts, Table, Paper, Pencils, Pages, Wall, watermark, signature, EasyNegative, censored, nipples, faceless",
|
||||
"x0_strength": 0.95,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"input": "videos/testing123/RunTurnaround10240000-0320__v2.mp4",
|
||||
"output": "videos/testing123/RunTurnaround10240000-0320_out.mp4",
|
||||
"work_dir": "videos/testing123/out",
|
||||
"key_subdir": "keys",
|
||||
"sd_model": "/home/salt/clone/ComfyUI/models/checkpoints/animePastelDream_softBakedVae.safetensors",
|
||||
"interval": 10,
|
||||
"crop": [
|
||||
0,
|
||||
180,
|
||||
0,
|
||||
0
|
||||
],
|
||||
"prompt": "masterpiece, best quality, 1girl, running, white tank top, black jogging pants, two tone hair, red hair, black hair, full body, outdoors, mountains, snowy mountains, top of mountain, cloudy sky",
|
||||
"a_prompt": "1girl, anime girl, semi realistic, 3d art, 3d render, outside, running, red and black hair",
|
||||
"n_prompt": "(deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, mutated hands and fingers:1.4), (deformed, distorted, disfigured:1.3), poorly drawn, bad anatomy, wrong anatomy, extra limb, missing limb, floating limbs, disconnected limbs, mutation, mutated, ugly, disgusting, amputation",
|
||||
"x0_strength": 0.95,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
{
|
||||
"input": "videos/pexels-antoni-shkraba-8048492-540x960-25fps.mp4",
|
||||
"output": "videos/pexels-antoni-shkraba-8048492-540x960-25fps/blend.mp4",
|
||||
"work_dir": "videos/pexels-antoni-shkraba-8048492-540x960-25fps",
|
||||
"key_subdir": "keys",
|
||||
"frame_count": 102,
|
||||
"interval": 10,
|
||||
"crop": [
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
280
|
||||
],
|
||||
"prompt": "a handsome man in van gogh painting",
|
||||
"a_prompt": "best quality, extremely detailed",
|
||||
"n_prompt": "longbody, lowres, bad anatomy, bad hands, missing fingers, extra digit, fewer digits, cropped, worst quality, low quality",
|
||||
"x0_strength": 1.05,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
],
|
||||
"image_resolution": 512
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"input": "videos/pexels-antoni-shkraba-8048492-540x960-25fps.mp4",
|
||||
"output": "videos/pexels-antoni-shkraba-8048492-540x960-25fps/blend.mp4",
|
||||
"work_dir": "videos/pexels-antoni-shkraba-8048492-540x960-25fps",
|
||||
"key_subdir": "keys",
|
||||
"frame_count": 102,
|
||||
"interval": 10,
|
||||
"crop": [
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
280
|
||||
],
|
||||
"prompt": "a handsome man in van gogh painting",
|
||||
"a_prompt": "best quality, extremely detailed",
|
||||
"n_prompt": "longbody, lowres, bad anatomy, bad hands, missing fingers, extra digit, fewer digits, cropped, worst quality, low quality",
|
||||
"x0_strength": 1.05,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
0.8,
|
||||
1
|
||||
],
|
||||
"image_resolution": 512,
|
||||
"use_limit_device_resolution": true
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
{
|
||||
"input": "videos/pexels-koolshooters-7322716.mp4",
|
||||
"output": "videos/pexels-koolshooters-7322716/blend.mp4",
|
||||
"work_dir": "videos/pexels-koolshooters-7322716",
|
||||
"key_subdir": "keys",
|
||||
"frame_count": 102,
|
||||
"interval": 10,
|
||||
"sd_model": "models/revAnimated_v11.safetensors",
|
||||
"prompt": "a beautiful woman in CG style",
|
||||
"a_prompt": "best quality, extremely detailed",
|
||||
"n_prompt": "longbody, lowres, bad anatomy, bad hands, missing fingers, extra digit, fewer digits, cropped, worst quality, low quality",
|
||||
"x0_strength": 0.75,
|
||||
"control_type": "canny",
|
||||
"canny_low": 50,
|
||||
"canny_high": 100,
|
||||
"control_strength": 0.7,
|
||||
"seed": 0,
|
||||
"warp_period": [
|
||||
0,
|
||||
0.1
|
||||
],
|
||||
"ada_period": [
|
||||
1,
|
||||
1
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
name: rerender
|
||||
channels:
|
||||
- pytorch
|
||||
- defaults
|
||||
dependencies:
|
||||
- python=3.8.5
|
||||
- pip=20.3
|
||||
- cudatoolkit=11.3
|
||||
- pytorch=1.12.1
|
||||
- torchvision=0.13.1
|
||||
- numpy=1.23.1
|
||||
- pip:
|
||||
- gradio==3.44.4
|
||||
- albumentations==1.3.0
|
||||
- opencv-contrib-python==4.3.0.36
|
||||
- imageio==2.9.0
|
||||
- imageio-ffmpeg==0.4.2
|
||||
- pytorch-lightning==1.5.0
|
||||
- omegaconf==2.1.1
|
||||
- test-tube>=0.7.5
|
||||
- streamlit==1.12.1
|
||||
- einops==0.3.0
|
||||
- transformers==4.19.2
|
||||
- webdataset==0.2.5
|
||||
- kornia==0.6
|
||||
- open_clip_torch==2.0.2
|
||||
- invisible-watermark>=0.1.5
|
||||
- streamlit-drawable-canvas==0.8.0
|
||||
- torchmetrics==0.6.0
|
||||
- timm==0.6.12
|
||||
- addict==2.4.0
|
||||
- yapf==0.32.0
|
||||
- prettytable==3.6.0
|
||||
- safetensors==0.2.7
|
||||
- basicsr==1.4.2
|
||||
- blendmodes
|
||||
- numba==0.57.0
|
||||
@@ -0,0 +1,260 @@
|
||||
import os
|
||||
import sys
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
parent_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
gmflow_dir = os.path.join(parent_dir, 'deps/gmflow')
|
||||
sys.path.insert(0, gmflow_dir)
|
||||
|
||||
from gmflow.gmflow import GMFlow # noqa: E702 E402 F401
|
||||
from utils.utils import InputPadder # noqa: E702 E402
|
||||
|
||||
|
||||
def coords_grid(b, h, w, homogeneous=False, device=None):
|
||||
y, x = torch.meshgrid(torch.arange(h), torch.arange(w), indexing='ij') # [H, W]
|
||||
|
||||
stacks = [x, y]
|
||||
|
||||
if homogeneous:
|
||||
ones = torch.ones_like(x) # [H, W]
|
||||
stacks.append(ones)
|
||||
|
||||
grid = torch.stack(stacks, dim=0).float() # [2, H, W] or [3, H, W]
|
||||
|
||||
grid = grid[None].repeat(b, 1, 1, 1) # [B, 2, H, W] or [B, 3, H, W]
|
||||
|
||||
if device is not None:
|
||||
grid = grid.to(device)
|
||||
|
||||
return grid
|
||||
|
||||
|
||||
def bilinear_sample(img,
|
||||
sample_coords,
|
||||
mode='bilinear',
|
||||
padding_mode='zeros',
|
||||
return_mask=False):
|
||||
# img: [B, C, H, W]
|
||||
# sample_coords: [B, 2, H, W] in image scale
|
||||
if sample_coords.size(1) != 2: # [B, H, W, 2]
|
||||
sample_coords = sample_coords.permute(0, 3, 1, 2)
|
||||
|
||||
b, _, h, w = sample_coords.shape
|
||||
|
||||
# Normalize to [-1, 1]
|
||||
x_grid = 2 * sample_coords[:, 0] / (w - 1) - 1
|
||||
y_grid = 2 * sample_coords[:, 1] / (h - 1) - 1
|
||||
|
||||
grid = torch.stack([x_grid, y_grid], dim=-1) # [B, H, W, 2]
|
||||
|
||||
img = F.grid_sample(img,
|
||||
grid,
|
||||
mode=mode,
|
||||
padding_mode=padding_mode,
|
||||
align_corners=True)
|
||||
|
||||
if return_mask:
|
||||
mask = (x_grid >= -1) & (y_grid >= -1) & (x_grid <= 1) & (
|
||||
y_grid <= 1) # [B, H, W]
|
||||
|
||||
return img, mask
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def flow_warp(feature,
|
||||
flow,
|
||||
mask=False,
|
||||
mode='bilinear',
|
||||
padding_mode='zeros'):
|
||||
b, c, h, w = feature.size()
|
||||
assert flow.size(1) == 2
|
||||
grid = coords_grid(b, h, w).to(flow.device)
|
||||
if grid.shape != flow.shape:
|
||||
print()
|
||||
grid += flow # [B, 2, H, W]
|
||||
|
||||
return bilinear_sample(feature,
|
||||
grid,
|
||||
mode=mode,
|
||||
padding_mode=padding_mode,
|
||||
return_mask=mask)
|
||||
|
||||
|
||||
def forward_backward_consistency_check(fwd_flow,
|
||||
bwd_flow,
|
||||
alpha=0.01,
|
||||
beta=0.5):
|
||||
# fwd_flow, bwd_flow: [B, 2, H, W]
|
||||
# alpha and beta values are following UnFlow
|
||||
# (https://arxiv.org/abs/1711.07837)
|
||||
assert fwd_flow.dim() == 4 and bwd_flow.dim() == 4
|
||||
assert fwd_flow.size(1) == 2 and bwd_flow.size(1) == 2
|
||||
flow_mag = torch.norm(fwd_flow, dim=1) + torch.norm(bwd_flow,
|
||||
dim=1) # [B, H, W]
|
||||
|
||||
warped_bwd_flow = flow_warp(bwd_flow, fwd_flow) # [B, 2, H, W]
|
||||
warped_fwd_flow = flow_warp(fwd_flow, bwd_flow) # [B, 2, H, W]
|
||||
|
||||
diff_fwd = torch.norm(fwd_flow + warped_bwd_flow, dim=1) # [B, H, W]
|
||||
diff_bwd = torch.norm(bwd_flow + warped_fwd_flow, dim=1)
|
||||
|
||||
threshold = alpha * flow_mag + beta
|
||||
|
||||
fwd_occ = (diff_fwd > threshold).float() # [B, H, W]
|
||||
bwd_occ = (diff_bwd > threshold).float()
|
||||
|
||||
return fwd_occ, bwd_occ
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def get_warped_and_mask(flow_model,
|
||||
image1,
|
||||
image2,
|
||||
image3=None,
|
||||
pixel_consistency=False):
|
||||
if image3 is None:
|
||||
image3 = image1
|
||||
padder = InputPadder(image1.shape, padding_factor=8)
|
||||
image1, image2 = padder.pad(image1[None].cuda(), image2[None].cuda())
|
||||
results_dict = flow_model(image1,
|
||||
image2,
|
||||
attn_splits_list=[2],
|
||||
corr_radius_list=[-1],
|
||||
prop_radius_list=[-1],
|
||||
pred_bidir_flow=True)
|
||||
flow_pr = results_dict['flow_preds'][-1] # [B, 2, H, W]
|
||||
fwd_flow = padder.unpad(flow_pr[0]).unsqueeze(0) # [1, 2, H, W]
|
||||
bwd_flow = padder.unpad(flow_pr[1]).unsqueeze(0) # [1, 2, H, W]
|
||||
fwd_occ, bwd_occ = forward_backward_consistency_check(
|
||||
fwd_flow, bwd_flow) # [1, H, W] float
|
||||
if pixel_consistency:
|
||||
warped_image1 = flow_warp(image1, bwd_flow)
|
||||
bwd_occ = torch.clamp(
|
||||
bwd_occ +
|
||||
(abs(image2 - warped_image1).mean(dim=1) > 255 * 0.25).float(), 0,
|
||||
1).unsqueeze(0)
|
||||
warped_results = flow_warp(image3, bwd_flow)
|
||||
return warped_results, bwd_occ, bwd_flow
|
||||
|
||||
|
||||
class FlowCalc():
|
||||
|
||||
def __init__(self, model_path='./models/gmflow_sintel-0c07dcb3.pth'):
|
||||
flow_model = GMFlow(
|
||||
feature_channels=128,
|
||||
num_scales=1,
|
||||
upsample_factor=8,
|
||||
num_head=1,
|
||||
attention_type='swin',
|
||||
ffn_dim_expansion=4,
|
||||
num_transformer_layers=6,
|
||||
).to('cuda')
|
||||
|
||||
checkpoint = torch.load(model_path,
|
||||
map_location=lambda storage, loc: storage)
|
||||
weights = checkpoint['model'] if 'model' in checkpoint else checkpoint
|
||||
flow_model.load_state_dict(weights, strict=False)
|
||||
flow_model.eval()
|
||||
self.model = flow_model
|
||||
|
||||
@torch.no_grad()
|
||||
def get_flow(self, image1, image2, save_path=None):
|
||||
|
||||
if save_path is not None and os.path.exists(save_path):
|
||||
bwd_flow = read_flow(save_path)
|
||||
return bwd_flow
|
||||
|
||||
image1 = torch.from_numpy(image1).permute(2, 0, 1).float()
|
||||
image2 = torch.from_numpy(image2).permute(2, 0, 1).float()
|
||||
padder = InputPadder(image1.shape, padding_factor=8)
|
||||
image1, image2 = padder.pad(image1[None].cuda(), image2[None].cuda())
|
||||
results_dict = self.model(image1,
|
||||
image2,
|
||||
attn_splits_list=[2],
|
||||
corr_radius_list=[-1],
|
||||
prop_radius_list=[-1],
|
||||
pred_bidir_flow=True)
|
||||
flow_pr = results_dict['flow_preds'][-1] # [B, 2, H, W]
|
||||
fwd_flow = padder.unpad(flow_pr[0]).unsqueeze(0) # [1, 2, H, W]
|
||||
bwd_flow = padder.unpad(flow_pr[1]).unsqueeze(0) # [1, 2, H, W]
|
||||
fwd_occ, bwd_occ = forward_backward_consistency_check(
|
||||
fwd_flow, bwd_flow) # [1, H, W] float
|
||||
if save_path is not None:
|
||||
flow_np = bwd_flow.cpu().numpy()
|
||||
np.save(save_path, flow_np)
|
||||
mask_path = os.path.splitext(save_path)[0] + '.png'
|
||||
bwd_occ = bwd_occ.cpu().permute(1, 2, 0).to(
|
||||
torch.long).numpy() * 255
|
||||
cv2.imwrite(mask_path, bwd_occ)
|
||||
|
||||
return bwd_flow
|
||||
|
||||
@torch.no_grad()
|
||||
def get_mask(self, image1, image2, save_path=None):
|
||||
|
||||
if save_path is not None:
|
||||
mask_path = os.path.splitext(save_path)[0] + '.png'
|
||||
if os.path.exists(mask_path):
|
||||
return read_mask(mask_path)
|
||||
|
||||
image1 = torch.from_numpy(image1).permute(2, 0, 1).float()
|
||||
image2 = torch.from_numpy(image2).permute(2, 0, 1).float()
|
||||
padder = InputPadder(image1.shape, padding_factor=8)
|
||||
image1, image2 = padder.pad(image1[None].cuda(), image2[None].cuda())
|
||||
results_dict = self.model(image1,
|
||||
image2,
|
||||
attn_splits_list=[2],
|
||||
corr_radius_list=[-1],
|
||||
prop_radius_list=[-1],
|
||||
pred_bidir_flow=True)
|
||||
flow_pr = results_dict['flow_preds'][-1] # [B, 2, H, W]
|
||||
fwd_flow = padder.unpad(flow_pr[0]).unsqueeze(0) # [1, 2, H, W]
|
||||
bwd_flow = padder.unpad(flow_pr[1]).unsqueeze(0) # [1, 2, H, W]
|
||||
fwd_occ, bwd_occ = forward_backward_consistency_check(
|
||||
fwd_flow, bwd_flow) # [1, H, W] float
|
||||
if save_path is not None:
|
||||
flow_np = bwd_flow.cpu().numpy()
|
||||
np.save(save_path, flow_np)
|
||||
mask_path = os.path.splitext(save_path)[0] + '.png'
|
||||
bwd_occ = bwd_occ.cpu().permute(1, 2, 0).to(
|
||||
torch.long).numpy() * 255
|
||||
cv2.imwrite(mask_path, bwd_occ)
|
||||
|
||||
return bwd_occ
|
||||
|
||||
def warp(self, img, flow, mode='bilinear'):
|
||||
expand = False
|
||||
if len(img.shape) == 2:
|
||||
expand = True
|
||||
img = np.expand_dims(img, 2)
|
||||
|
||||
img = torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0)
|
||||
dtype = img.dtype
|
||||
img = img.to(torch.float)
|
||||
res = flow_warp(img, flow, mode=mode)
|
||||
res = res.to(dtype)
|
||||
res = res[0].cpu().permute(1, 2, 0).numpy()
|
||||
if expand:
|
||||
res = res[:, :, 0]
|
||||
return res
|
||||
|
||||
|
||||
def read_flow(save_path):
|
||||
flow_np = np.load(save_path)
|
||||
bwd_flow = torch.from_numpy(flow_np)
|
||||
return bwd_flow
|
||||
|
||||
|
||||
def read_mask(save_path):
|
||||
mask_path = os.path.splitext(save_path)[0] + '.png'
|
||||
mask = cv2.imread(mask_path)
|
||||
mask = cv2.cvtColor(mask, cv2.COLOR_BGR2GRAY)
|
||||
return mask
|
||||
|
||||
|
||||
flow_calc = FlowCalc()
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,78 @@
|
||||
import os
|
||||
import platform
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def build_ebsynth():
|
||||
if os.path.exists('deps/ebsynth/bin/ebsynth'):
|
||||
print('Ebsynth has been built.')
|
||||
return
|
||||
|
||||
os_str = platform.system()
|
||||
|
||||
if os_str == 'Windows':
|
||||
print('Build Ebsynth Windows 64 bit.',
|
||||
'If you want to build for 32 bit, please modify install.py.')
|
||||
cmd = '.\\build-win64-cpu+cuda.bat'
|
||||
exe_file = 'deps/ebsynth/bin/ebsynth.exe'
|
||||
elif os_str == 'Linux':
|
||||
cmd = 'bash build-linux-cpu+cuda.sh'
|
||||
exe_file = 'deps/ebsynth/bin/ebsynth'
|
||||
elif os_str == 'Darwin':
|
||||
cmd = 'sh build-macos-cpu_only.sh'
|
||||
exe_file = 'deps/ebsynth/bin/ebsynth.app'
|
||||
else:
|
||||
print('Cannot recognize OS. Ebsynth installation stopped.')
|
||||
return
|
||||
|
||||
os.chdir('deps/ebsynth')
|
||||
print(cmd)
|
||||
os.system(cmd)
|
||||
os.chdir('../..')
|
||||
if os.path.exists(exe_file):
|
||||
print('Ebsynth installed successfully.')
|
||||
else:
|
||||
print('Failed to install Ebsynth.')
|
||||
|
||||
|
||||
def download(url, dir, name=None):
|
||||
os.makedirs(dir, exist_ok=True)
|
||||
if name is None:
|
||||
name = url.split('/')[-1]
|
||||
path = os.path.join(dir, name)
|
||||
if not os.path.exists(path):
|
||||
print(f'Install {name} ...')
|
||||
open(path, 'wb').write(requests.get(url).content)
|
||||
print('Install successfully.')
|
||||
|
||||
|
||||
def download_gmflow_ckpt():
|
||||
url = ('https://huggingface.co/PKUWilliamYang/Rerender/'
|
||||
'resolve/main/models/gmflow_sintel-0c07dcb3.pth')
|
||||
download(url, 'models')
|
||||
|
||||
|
||||
def download_controlnet_canny():
|
||||
url = ('https://huggingface.co/lllyasviel/ControlNet/'
|
||||
'resolve/main/models/control_sd15_canny.pth')
|
||||
download(url, 'models')
|
||||
|
||||
|
||||
def download_controlnet_hed():
|
||||
url = ('https://huggingface.co/lllyasviel/ControlNet/'
|
||||
'resolve/main/models/control_sd15_hed.pth')
|
||||
download(url, 'models')
|
||||
|
||||
|
||||
def download_vae():
|
||||
url = ('https://huggingface.co/stabilityai/sd-vae-ft-mse-original'
|
||||
'/resolve/main/vae-ft-mse-840000-ema-pruned.ckpt')
|
||||
download(url, 'models')
|
||||
|
||||
|
||||
build_ebsynth()
|
||||
download_gmflow_ckpt()
|
||||
download_controlnet_canny()
|
||||
download_controlnet_hed()
|
||||
download_vae()
|
||||
@@ -0,0 +1,25 @@
|
||||
addict==2.4.0
|
||||
albumentations==1.3.0
|
||||
basicsr==1.4.2
|
||||
blendmodes
|
||||
einops==0.3.0
|
||||
gradio==3.44.4
|
||||
imageio==2.9.0
|
||||
imageio-ffmpeg==0.4.2
|
||||
invisible-watermark==0.1.5
|
||||
kornia==0.6
|
||||
numba==0.57.0
|
||||
omegaconf==2.1.1
|
||||
open_clip_torch==2.0.2
|
||||
prettytable==3.6.0
|
||||
pytorch-lightning==1.5.0
|
||||
safetensors==0.2.7
|
||||
streamlit==1.12.1
|
||||
streamlit-drawable-canvas==0.8.0
|
||||
test-tube==0.7.5
|
||||
timm==0.6.12
|
||||
torchmetrics==0.6.0
|
||||
transformers==4.19.2
|
||||
webdataset==0.2.5
|
||||
yapf==0.32.0
|
||||
xformers
|
||||
@@ -0,0 +1,474 @@
|
||||
import argparse
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as T
|
||||
from blendmodes.blend import BlendType, blendLayers
|
||||
from PIL import Image
|
||||
from pytorch_lightning import seed_everything
|
||||
from safetensors.torch import load_file
|
||||
from skimage import exposure
|
||||
|
||||
import src.import_util # noqa: F401
|
||||
from deps.ControlNet.annotator.canny import CannyDetector
|
||||
from deps.ControlNet.annotator.hed import HEDdetector
|
||||
from deps.ControlNet.annotator.util import HWC3
|
||||
from deps.ControlNet.cldm.cldm import ControlLDM
|
||||
from deps.ControlNet.cldm.model import create_model, load_state_dict
|
||||
from deps.gmflow.gmflow.gmflow import GMFlow
|
||||
from flow.flow_utils import get_warped_and_mask
|
||||
from src.config import RerenderConfig
|
||||
from src.controller import AttentionControl
|
||||
from src.ddim_v_hacked import DDIMVSampler
|
||||
from src.freeu import freeu_forward
|
||||
from src.img_util import find_flat_region, numpy2tensor
|
||||
from src.video_util import frame_to_video, get_fps, prepare_frames
|
||||
from src.lora import load_lora, apply_lora
|
||||
|
||||
blur = T.GaussianBlur(kernel_size=(9, 9), sigma=(18, 18))
|
||||
totensor = T.PILToTensor()
|
||||
|
||||
|
||||
def setup_color_correction(image):
|
||||
correction_target = cv2.cvtColor(np.asarray(image.copy()),
|
||||
cv2.COLOR_RGB2LAB)
|
||||
return correction_target
|
||||
|
||||
|
||||
def apply_color_correction(correction, original_image):
|
||||
image = Image.fromarray(
|
||||
cv2.cvtColor(
|
||||
exposure.match_histograms(cv2.cvtColor(np.asarray(original_image),
|
||||
cv2.COLOR_RGB2LAB),
|
||||
correction,
|
||||
channel_axis=2),
|
||||
cv2.COLOR_LAB2RGB).astype('uint8'))
|
||||
|
||||
image = blendLayers(image, original_image, BlendType.LUMINOSITY)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
def rerender(cfg: RerenderConfig, first_img_only: bool, key_video_path: str):
|
||||
# Preprocess input
|
||||
prepare_frames(cfg.input_path, cfg.input_dir, cfg.image_resolution, cfg.crop, cfg.use_limit_device_resolution)
|
||||
|
||||
# Load models
|
||||
if cfg.control_type == 'HED':
|
||||
detector = HEDdetector()
|
||||
elif cfg.control_type == 'canny':
|
||||
canny_detector = CannyDetector()
|
||||
low_threshold = cfg.canny_low
|
||||
high_threshold = cfg.canny_high
|
||||
|
||||
def apply_canny(x):
|
||||
return canny_detector(x, low_threshold, high_threshold)
|
||||
|
||||
detector = apply_canny
|
||||
|
||||
model: ControlLDM = create_model(
|
||||
'./deps/ControlNet/models/cldm_v15.yaml').cpu()
|
||||
if cfg.control_type == 'HED':
|
||||
model.load_state_dict(
|
||||
load_state_dict('./models/control_sd15_hed.pth', location='cuda'))
|
||||
elif cfg.control_type == 'canny':
|
||||
model.load_state_dict(
|
||||
load_state_dict('./models/control_sd15_canny.pth',
|
||||
location='cuda'))
|
||||
model = model.cuda()
|
||||
model.control_scales = [cfg.control_strength] * 13
|
||||
|
||||
if cfg.sd_model is not None:
|
||||
model_ext = os.path.splitext(cfg.sd_model)[1]
|
||||
if model_ext == '.safetensors':
|
||||
model.load_state_dict(load_file(cfg.sd_model), strict=False)
|
||||
elif model_ext == '.ckpt' or model_ext == '.pth':
|
||||
model.load_state_dict(torch.load(cfg.sd_model)['state_dict'],
|
||||
strict=False)
|
||||
|
||||
# apply lora if exists
|
||||
if cfg.lora_path:
|
||||
lora_weights = load_lora(cfg.lora_path)
|
||||
apply_lora(model.model, model.cond_stage_model, lora_weights)
|
||||
|
||||
try:
|
||||
model.first_stage_model.load_state_dict(torch.load(
|
||||
'./models/vae-ft-mse-840000-ema-pruned.ckpt')['state_dict'],
|
||||
strict=False)
|
||||
except Exception:
|
||||
print('Warning: We suggest you download the fine-tuned VAE',
|
||||
'otherwise the generation quality will be degraded')
|
||||
|
||||
model.model.diffusion_model.forward = \
|
||||
freeu_forward(model.model.diffusion_model, *cfg.freeu_args)
|
||||
ddim_v_sampler = DDIMVSampler(model)
|
||||
|
||||
flow_model = GMFlow(
|
||||
feature_channels=128,
|
||||
num_scales=1,
|
||||
upsample_factor=8,
|
||||
num_head=1,
|
||||
attention_type='swin',
|
||||
ffn_dim_expansion=4,
|
||||
num_transformer_layers=6,
|
||||
).to('cuda')
|
||||
|
||||
checkpoint = torch.load('models/gmflow_sintel-0c07dcb3.pth',
|
||||
map_location=lambda storage, loc: storage)
|
||||
weights = checkpoint['model'] if 'model' in checkpoint else checkpoint
|
||||
flow_model.load_state_dict(weights, strict=False)
|
||||
flow_model.eval()
|
||||
|
||||
num_samples = 1
|
||||
ddim_steps = 20
|
||||
scale = 7.5
|
||||
|
||||
seed = cfg.seed
|
||||
if seed == -1:
|
||||
seed = random.randint(0, 65535)
|
||||
eta = 0.0
|
||||
|
||||
prompt = cfg.prompt
|
||||
a_prompt = cfg.a_prompt
|
||||
n_prompt = cfg.n_prompt
|
||||
prompt = prompt + ', ' + a_prompt
|
||||
|
||||
style_update_freq = cfg.style_update_freq
|
||||
pixelfusion = True
|
||||
color_preserve = cfg.color_preserve
|
||||
|
||||
x0_strength = 1 - cfg.x0_strength
|
||||
mask_period = cfg.mask_period
|
||||
firstx0 = True
|
||||
controller = AttentionControl(cfg.inner_strength, cfg.mask_period,
|
||||
cfg.cross_period, cfg.ada_period,
|
||||
cfg.warp_period, cfg.loose_cfattn)
|
||||
|
||||
imgs = sorted(os.listdir(cfg.input_dir))
|
||||
imgs = [os.path.join(cfg.input_dir, img) for img in imgs]
|
||||
if cfg.frame_count >= 0:
|
||||
imgs = imgs[:cfg.frame_count]
|
||||
|
||||
with torch.no_grad():
|
||||
frame = cv2.imread(imgs[0])
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
img = HWC3(frame)
|
||||
H, W, C = img.shape
|
||||
|
||||
img_ = numpy2tensor(img)
|
||||
# if color_preserve:
|
||||
# img_ = numpy2tensor(img)
|
||||
# else:
|
||||
# img_ = apply_color_correction(color_corrections,
|
||||
# Image.fromarray(img))
|
||||
# img_ = totensor(img_).unsqueeze(0)[:, :3] / 127.5 - 1
|
||||
encoder_posterior = model.encode_first_stage(img_.cuda())
|
||||
x0 = model.get_first_stage_encoding(encoder_posterior).detach()
|
||||
|
||||
detected_map = detector(img)
|
||||
detected_map = HWC3(detected_map)
|
||||
# For visualization
|
||||
detected_img = 255 - detected_map
|
||||
|
||||
control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0
|
||||
control = torch.stack([control for _ in range(num_samples)], dim=0)
|
||||
control = einops.rearrange(control, 'b h w c -> b c h w').clone()
|
||||
cond = {
|
||||
'c_concat': [control],
|
||||
'c_crossattn':
|
||||
[model.get_learned_conditioning([prompt] * num_samples)]
|
||||
}
|
||||
un_cond = {
|
||||
'c_concat': [control],
|
||||
'c_crossattn':
|
||||
[model.get_learned_conditioning([n_prompt] * num_samples)]
|
||||
}
|
||||
shape = (4, H // 8, W // 8)
|
||||
|
||||
controller.set_task('initfirst')
|
||||
seed_everything(seed)
|
||||
samples, _ = ddim_v_sampler.sample(ddim_steps,
|
||||
num_samples,
|
||||
shape,
|
||||
cond,
|
||||
verbose=False,
|
||||
eta=eta,
|
||||
unconditional_guidance_scale=scale,
|
||||
unconditional_conditioning=un_cond,
|
||||
controller=controller,
|
||||
x0=x0,
|
||||
strength=x0_strength)
|
||||
x_samples = model.decode_first_stage(samples)
|
||||
pre_result = x_samples
|
||||
pre_img = img
|
||||
first_result = pre_result
|
||||
first_img = pre_img
|
||||
|
||||
x_samples = (
|
||||
einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +
|
||||
127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
color_corrections = setup_color_correction(Image.fromarray(x_samples[0]))
|
||||
Image.fromarray(x_samples[0]).save(os.path.join(cfg.first_dir,
|
||||
'first.jpg'))
|
||||
cv2.imwrite(os.path.join(cfg.first_dir, 'first_edge.jpg'), detected_img)
|
||||
|
||||
if first_img_only:
|
||||
exit(0)
|
||||
|
||||
for i in range(0, min(len(imgs), cfg.frame_count) - 1, cfg.interval):
|
||||
cid = i + 1
|
||||
print(cid)
|
||||
if cid <= (len(imgs) - 1):
|
||||
frame = cv2.imread(imgs[cid])
|
||||
else:
|
||||
frame = cv2.imread(imgs[len(imgs) - 1])
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
img = HWC3(frame)
|
||||
|
||||
if color_preserve:
|
||||
img_ = numpy2tensor(img)
|
||||
else:
|
||||
img_ = apply_color_correction(color_corrections,
|
||||
Image.fromarray(img))
|
||||
img_ = totensor(img_).unsqueeze(0)[:, :3] / 127.5 - 1
|
||||
encoder_posterior = model.encode_first_stage(img_.cuda())
|
||||
x0 = model.get_first_stage_encoding(encoder_posterior).detach()
|
||||
|
||||
detected_map = detector(img)
|
||||
detected_map = HWC3(detected_map)
|
||||
|
||||
control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0
|
||||
control = torch.stack([control for _ in range(num_samples)], dim=0)
|
||||
control = einops.rearrange(control, 'b h w c -> b c h w').clone()
|
||||
cond['c_concat'] = [control]
|
||||
un_cond['c_concat'] = [control]
|
||||
|
||||
image1 = torch.from_numpy(pre_img).permute(2, 0, 1).float()
|
||||
image2 = torch.from_numpy(img).permute(2, 0, 1).float()
|
||||
warped_pre, bwd_occ_pre, bwd_flow_pre = get_warped_and_mask(
|
||||
flow_model, image1, image2, pre_result, False)
|
||||
blend_mask_pre = blur(
|
||||
F.max_pool2d(bwd_occ_pre, kernel_size=9, stride=1, padding=4))
|
||||
blend_mask_pre = torch.clamp(blend_mask_pre + bwd_occ_pre, 0, 1)
|
||||
|
||||
image1 = torch.from_numpy(first_img).permute(2, 0, 1).float()
|
||||
warped_0, bwd_occ_0, bwd_flow_0 = get_warped_and_mask(
|
||||
flow_model, image1, image2, first_result, False)
|
||||
blend_mask_0 = blur(
|
||||
F.max_pool2d(bwd_occ_0, kernel_size=9, stride=1, padding=4))
|
||||
blend_mask_0 = torch.clamp(blend_mask_0 + bwd_occ_0, 0, 1)
|
||||
|
||||
if firstx0:
|
||||
mask = 1 - F.max_pool2d(blend_mask_0, kernel_size=8)
|
||||
controller.set_warp(
|
||||
F.interpolate(bwd_flow_0 / 8.0,
|
||||
scale_factor=1. / 8,
|
||||
mode='bilinear'), mask)
|
||||
else:
|
||||
mask = 1 - F.max_pool2d(blend_mask_pre, kernel_size=8)
|
||||
controller.set_warp(
|
||||
F.interpolate(bwd_flow_pre / 8.0,
|
||||
scale_factor=1. / 8,
|
||||
mode='bilinear'), mask)
|
||||
|
||||
controller.set_task('keepx0, keepstyle')
|
||||
seed_everything(seed)
|
||||
samples, intermediates = ddim_v_sampler.sample(
|
||||
ddim_steps,
|
||||
num_samples,
|
||||
shape,
|
||||
cond,
|
||||
verbose=False,
|
||||
eta=eta,
|
||||
unconditional_guidance_scale=scale,
|
||||
unconditional_conditioning=un_cond,
|
||||
controller=controller,
|
||||
x0=x0,
|
||||
strength=x0_strength)
|
||||
direct_result = model.decode_first_stage(samples)
|
||||
|
||||
if not pixelfusion:
|
||||
pre_result = direct_result
|
||||
pre_img = img
|
||||
viz = (
|
||||
einops.rearrange(direct_result, 'b c h w -> b h w c') * 127.5 +
|
||||
127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
|
||||
else:
|
||||
|
||||
blend_results = (1 - blend_mask_pre
|
||||
) * warped_pre + blend_mask_pre * direct_result
|
||||
blend_results = (
|
||||
1 - blend_mask_0) * warped_0 + blend_mask_0 * blend_results
|
||||
|
||||
bwd_occ = 1 - torch.clamp(1 - bwd_occ_pre + 1 - bwd_occ_0, 0, 1)
|
||||
blend_mask = blur(
|
||||
F.max_pool2d(bwd_occ, kernel_size=9, stride=1, padding=4))
|
||||
blend_mask = 1 - torch.clamp(blend_mask + bwd_occ, 0, 1)
|
||||
|
||||
encoder_posterior = model.encode_first_stage(blend_results)
|
||||
xtrg = model.get_first_stage_encoding(
|
||||
encoder_posterior).detach() # * mask
|
||||
blend_results_rec = model.decode_first_stage(xtrg)
|
||||
encoder_posterior = model.encode_first_stage(blend_results_rec)
|
||||
xtrg_rec = model.get_first_stage_encoding(
|
||||
encoder_posterior).detach()
|
||||
xtrg_ = (xtrg + 1 * (xtrg - xtrg_rec)) # * mask
|
||||
blend_results_rec_new = model.decode_first_stage(xtrg_)
|
||||
tmp = (abs(blend_results_rec_new - blend_results).mean(
|
||||
dim=1, keepdims=True) > 0.25).float()
|
||||
mask_x = F.max_pool2d((F.interpolate(
|
||||
tmp, scale_factor=1 / 8., mode='bilinear') > 0).float(),
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
mask = (1 - F.max_pool2d(1 - blend_mask, kernel_size=8)
|
||||
) # * (1-mask_x)
|
||||
|
||||
if cfg.smooth_boundary:
|
||||
noise_rescale = find_flat_region(mask)
|
||||
else:
|
||||
noise_rescale = torch.ones_like(mask)
|
||||
masks = []
|
||||
for j in range(ddim_steps):
|
||||
if j <= ddim_steps * mask_period[
|
||||
0] or j >= ddim_steps * mask_period[1]:
|
||||
masks += [None]
|
||||
else:
|
||||
masks += [mask * cfg.mask_strength]
|
||||
|
||||
# mask 3
|
||||
# xtrg = ((1-mask_x) *
|
||||
# (xtrg + xtrg - xtrg_rec) + mask_x * samples) * mask
|
||||
# mask 2
|
||||
# xtrg = (xtrg + 1 * (xtrg - xtrg_rec)) * mask
|
||||
xtrg = (xtrg + (1 - mask_x) * (xtrg - xtrg_rec)) * mask # mask 1
|
||||
|
||||
tasks = 'keepstyle, keepx0'
|
||||
if not firstx0:
|
||||
tasks += ', updatex0'
|
||||
if i % style_update_freq == 0:
|
||||
tasks += ', updatestyle'
|
||||
controller.set_task(tasks, 1.0)
|
||||
|
||||
seed_everything(seed)
|
||||
samples, _ = ddim_v_sampler.sample(
|
||||
ddim_steps,
|
||||
num_samples,
|
||||
shape,
|
||||
cond,
|
||||
verbose=False,
|
||||
eta=eta,
|
||||
unconditional_guidance_scale=scale,
|
||||
unconditional_conditioning=un_cond,
|
||||
controller=controller,
|
||||
x0=x0,
|
||||
strength=x0_strength,
|
||||
xtrg=xtrg,
|
||||
mask=masks,
|
||||
noise_rescale=noise_rescale)
|
||||
x_samples = model.decode_first_stage(samples)
|
||||
pre_result = x_samples
|
||||
pre_img = img
|
||||
|
||||
viz = (einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +
|
||||
127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
|
||||
Image.fromarray(viz[0]).save(
|
||||
os.path.join(cfg.key_dir, f'{cid:04d}.png'))
|
||||
if key_video_path is not None:
|
||||
fps = get_fps(cfg.input_path)
|
||||
fps //= cfg.interval
|
||||
frame_to_video(key_video_path, cfg.key_dir, fps, False)
|
||||
|
||||
|
||||
def postprocess(cfg: RerenderConfig, ne: bool, max_process: int, tmp: bool,
|
||||
ps: bool):
|
||||
video_base_dir = cfg.work_dir
|
||||
o_video = cfg.output_path
|
||||
fps = get_fps(cfg.input_path)
|
||||
|
||||
end_frame = cfg.frame_count - 1
|
||||
interval = cfg.interval
|
||||
key_dir = os.path.split(cfg.key_dir)[-1]
|
||||
use_e = '-ne' if ne else ''
|
||||
use_tmp = '-tmp' if tmp else ''
|
||||
use_ps = '-ps' if ps else ''
|
||||
o_video_cmd = f'--output {o_video}'
|
||||
|
||||
cmd = (
|
||||
f'python video_blend.py {video_base_dir} --beg 1 --end {end_frame} '
|
||||
f'--itv {interval} --key {key_dir} {use_e} {o_video_cmd} --fps {fps} '
|
||||
f'--n_proc {max_process} {use_tmp} {use_ps}')
|
||||
print(cmd)
|
||||
os.system(cmd)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--cfg', type=str, default=None)
|
||||
parser.add_argument('--input',
|
||||
type=str,
|
||||
default=None,
|
||||
help='The input path to video.')
|
||||
parser.add_argument('--output', type=str, default=None)
|
||||
parser.add_argument('--prompt', type=str, default=None)
|
||||
parser.add_argument('--key_video_path', type=str, default=None)
|
||||
parser.add_argument('-one',
|
||||
action='store_true',
|
||||
help='Run the first frame with ControlNet only')
|
||||
parser.add_argument('-nr',
|
||||
action='store_true',
|
||||
help='Do not run rerender and do postprocessing only')
|
||||
parser.add_argument('-nb',
|
||||
action='store_true',
|
||||
help='Do not run postprocessing and run rerender only')
|
||||
parser.add_argument(
|
||||
'-ne',
|
||||
action='store_true',
|
||||
help='Do not run ebsynth (use previous ebsynth temporary output)')
|
||||
parser.add_argument('-nps',
|
||||
action='store_true',
|
||||
help='Do not run poisson gradient blending')
|
||||
parser.add_argument('--n_proc',
|
||||
type=int,
|
||||
default=4,
|
||||
help='The max process count')
|
||||
parser.add_argument('--tmp',
|
||||
action='store_true',
|
||||
help='Keep ebsynth temporary output')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
cfg = RerenderConfig()
|
||||
if args.cfg is not None:
|
||||
cfg.create_from_path(args.cfg)
|
||||
if args.input is not None:
|
||||
print('Config has been loaded. --input is ignored.')
|
||||
if args.output is not None:
|
||||
print('Config has been loaded. --output is ignored.')
|
||||
if args.prompt is not None:
|
||||
print('Config has been loaded. --prompt is ignored.')
|
||||
else:
|
||||
if args.input is None:
|
||||
print('Config not found. --input is required.')
|
||||
exit(0)
|
||||
if args.output is None:
|
||||
print('Config not found. --output is required.')
|
||||
exit(0)
|
||||
if args.prompt is None:
|
||||
print('Config not found. --prompt is required.')
|
||||
exit(0)
|
||||
cfg.create_from_parameters(args.input, args.output, args.prompt)
|
||||
|
||||
if not args.nr:
|
||||
rerender(cfg, args.one, args.key_video_path)
|
||||
torch.cuda.empty_cache()
|
||||
if not args.nb:
|
||||
postprocess(cfg, args.ne, args.n_proc, args.tmp, not args.nps)
|
||||
@@ -0,0 +1,7 @@
|
||||
# The model dict is used for webUI only
|
||||
|
||||
model_dict = {
|
||||
'Stable Diffusion 1.5': '',
|
||||
'revAnimated_v11': 'models/revAnimated_v11.safetensors',
|
||||
'realisticVisionV20_v20': 'models/realisticVisionV20_v20.safetensors'
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Optional, Sequence, Tuple
|
||||
|
||||
from src.video_util import get_frame_count
|
||||
|
||||
|
||||
class RerenderConfig:
|
||||
|
||||
def __init__(self):
|
||||
...
|
||||
|
||||
def create_from_parameters(self,
|
||||
input_path: str,
|
||||
output_path: str,
|
||||
prompt: str,
|
||||
work_dir: Optional[str] = None,
|
||||
key_subdir: str = 'keys',
|
||||
frame_count: Optional[int] = None,
|
||||
interval: int = 10,
|
||||
crop: Sequence[int] = (0, 0, 0, 0),
|
||||
sd_model: Optional[str] = None,
|
||||
a_prompt: str = '',
|
||||
n_prompt: str = '',
|
||||
ddim_steps=20,
|
||||
scale=7.5,
|
||||
control_type: str = 'HED',
|
||||
control_strength=1,
|
||||
seed: int = -1,
|
||||
image_resolution: int = 512,
|
||||
use_limit_device_resolution: bool = False,
|
||||
x0_strength: float = -1,
|
||||
style_update_freq: int = 10,
|
||||
cross_period: Tuple[float, float] = (0, 1),
|
||||
warp_period: Tuple[float, float] = (0, 0.1),
|
||||
mask_period: Tuple[float, float] = (0.5, 0.8),
|
||||
ada_period: Tuple[float, float] = (1.0, 1.0),
|
||||
mask_strength: float = 0.5,
|
||||
inner_strength: float = 0.9,
|
||||
smooth_boundary: bool = True,
|
||||
color_preserve: bool = True,
|
||||
loose_cfattn: bool = False,
|
||||
freeu_args: Tuple[int] = (1, 1, 1, 1),
|
||||
lora_path: Optional[str] = None,
|
||||
**kwargs):
|
||||
self.input_path = input_path
|
||||
self.output_path = output_path
|
||||
self.prompt = prompt
|
||||
self.work_dir = work_dir
|
||||
if work_dir is None:
|
||||
self.work_dir = os.path.dirname(output_path)
|
||||
self.key_dir = os.path.join(self.work_dir, key_subdir)
|
||||
self.first_dir = os.path.join(self.work_dir, 'first')
|
||||
|
||||
# Split video into frames
|
||||
if not os.path.isfile(input_path):
|
||||
raise FileNotFoundError(f'Cannot find video file {input_path}')
|
||||
self.input_dir = os.path.join(self.work_dir, 'video')
|
||||
|
||||
self.frame_count = frame_count
|
||||
if frame_count is None:
|
||||
self.frame_count = get_frame_count(self.input_path)
|
||||
self.interval = interval
|
||||
self.crop = crop
|
||||
self.sd_model = sd_model
|
||||
self.a_prompt = a_prompt
|
||||
self.n_prompt = n_prompt
|
||||
self.ddim_steps = ddim_steps
|
||||
self.scale = scale
|
||||
self.control_type = control_type
|
||||
if self.control_type == 'canny':
|
||||
self.canny_low = kwargs.get('canny_low', 100)
|
||||
self.canny_high = kwargs.get('canny_high', 200)
|
||||
else:
|
||||
self.canny_low = None
|
||||
self.canny_high = None
|
||||
self.control_strength = control_strength
|
||||
self.seed = seed
|
||||
self.image_resolution = image_resolution
|
||||
self.use_limit_device_resolution = use_limit_device_resolution
|
||||
self.x0_strength = x0_strength
|
||||
self.style_update_freq = style_update_freq
|
||||
self.cross_period = cross_period
|
||||
self.mask_period = mask_period
|
||||
self.warp_period = warp_period
|
||||
self.ada_period = ada_period
|
||||
self.mask_strength = mask_strength
|
||||
self.inner_strength = inner_strength
|
||||
self.smooth_boundary = smooth_boundary
|
||||
self.color_preserve = color_preserve
|
||||
self.loose_cfattn = loose_cfattn
|
||||
self.freeu_args = freeu_args
|
||||
self.lora_path = lora_path
|
||||
|
||||
os.makedirs(self.input_dir, exist_ok=True)
|
||||
os.makedirs(self.work_dir, exist_ok=True)
|
||||
os.makedirs(self.key_dir, exist_ok=True)
|
||||
os.makedirs(self.first_dir, exist_ok=True)
|
||||
|
||||
def create_from_path(self, cfg_path: str):
|
||||
with open(cfg_path, 'r') as fp:
|
||||
cfg = json.load(fp)
|
||||
kwargs = dict()
|
||||
|
||||
def append_if_not_none(key):
|
||||
value = cfg.get(key, None)
|
||||
if value is not None:
|
||||
kwargs[key] = value
|
||||
|
||||
kwargs['input_path'] = cfg['input']
|
||||
kwargs['output_path'] = cfg['output']
|
||||
kwargs['prompt'] = cfg['prompt']
|
||||
append_if_not_none('work_dir')
|
||||
append_if_not_none('key_subdir')
|
||||
append_if_not_none('frame_count')
|
||||
append_if_not_none('interval')
|
||||
append_if_not_none('crop')
|
||||
append_if_not_none('sd_model')
|
||||
append_if_not_none('a_prompt')
|
||||
append_if_not_none('n_prompt')
|
||||
append_if_not_none('ddim_steps')
|
||||
append_if_not_none('scale')
|
||||
append_if_not_none('control_type')
|
||||
if kwargs.get('control_type', '') == 'canny':
|
||||
append_if_not_none('canny_low')
|
||||
append_if_not_none('canny_high')
|
||||
append_if_not_none('control_strength')
|
||||
append_if_not_none('seed')
|
||||
append_if_not_none('image_resolution')
|
||||
append_if_not_none('use_limit_device_resolution')
|
||||
append_if_not_none('x0_strength')
|
||||
append_if_not_none('style_update_freq')
|
||||
append_if_not_none('cross_period')
|
||||
append_if_not_none('warp_period')
|
||||
append_if_not_none('mask_period')
|
||||
append_if_not_none('ada_period')
|
||||
append_if_not_none('mask_strength')
|
||||
append_if_not_none('inner_strength')
|
||||
append_if_not_none('smooth_boundary')
|
||||
append_if_not_none('color_perserve')
|
||||
append_if_not_none('loose_cfattn')
|
||||
append_if_not_none('freeu_args')
|
||||
append_if_not_none('lora_path')
|
||||
self.create_from_parameters(**kwargs)
|
||||
|
||||
def create_from_json(self, json_dict):
|
||||
kwargs = dict()
|
||||
def append_if_not_none(key):
|
||||
value = json_dict.get(key, None)
|
||||
if value is not None:
|
||||
kwargs[key] = value
|
||||
|
||||
kwargs['input_path'] = json_dict['input']
|
||||
kwargs['output_path'] = json_dict['output']
|
||||
kwargs['prompt'] = json_dict['prompt']
|
||||
append_if_not_none('work_dir')
|
||||
append_if_not_none('key_subdir')
|
||||
append_if_not_none('frame_count')
|
||||
append_if_not_none('interval')
|
||||
append_if_not_none('crop')
|
||||
append_if_not_none('sd_model')
|
||||
append_if_not_none('a_prompt')
|
||||
append_if_not_none('n_prompt')
|
||||
append_if_not_none('ddim_steps')
|
||||
append_if_not_none('scale')
|
||||
append_if_not_none('control_type')
|
||||
if kwargs.get('control_type', '') == 'canny':
|
||||
append_if_not_none('canny_low')
|
||||
append_if_not_none('canny_high')
|
||||
append_if_not_none('control_strength')
|
||||
append_if_not_none('seed')
|
||||
append_if_not_none('image_resolution')
|
||||
append_if_not_none('use_limit_device_resolution')
|
||||
append_if_not_none('x0_strength')
|
||||
append_if_not_none('style_update_freq')
|
||||
append_if_not_none('cross_period')
|
||||
append_if_not_none('warp_period')
|
||||
append_if_not_none('mask_period')
|
||||
append_if_not_none('ada_period')
|
||||
append_if_not_none('mask_strength')
|
||||
append_if_not_none('inner_strength')
|
||||
append_if_not_none('smooth_boundary')
|
||||
append_if_not_none('color_perserve')
|
||||
append_if_not_none('loose_cfattn')
|
||||
append_if_not_none('freeu_args')
|
||||
self.create_from_parameters(**kwargs)
|
||||
|
||||
@property
|
||||
def use_warp(self):
|
||||
return self.warp_period[0] <= self.warp_period[1]
|
||||
|
||||
@property
|
||||
def use_mask(self):
|
||||
return self.mask_period[0] <= self.mask_period[1]
|
||||
|
||||
@property
|
||||
def use_ada(self):
|
||||
return self.ada_period[0] <= self.ada_period[1]
|
||||
@@ -0,0 +1,143 @@
|
||||
import gc
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from flow.flow_utils import flow_warp
|
||||
|
||||
# AdaIn
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps=1e-5):
|
||||
# eps is a small value added to the variance to avoid divide-by-zero.
|
||||
size = feat.size()
|
||||
assert (len(size) == 4)
|
||||
N, C = size[:2]
|
||||
feat_var = feat.view(N, C, -1).var(dim=2) + eps
|
||||
feat_std = feat_var.sqrt().view(N, C, 1, 1)
|
||||
feat_mean = feat.view(N, C, -1).mean(dim=2).view(N, C, 1, 1)
|
||||
return feat_mean, feat_std
|
||||
|
||||
|
||||
class AttentionControl():
|
||||
|
||||
def __init__(self,
|
||||
inner_strength,
|
||||
mask_period,
|
||||
cross_period,
|
||||
ada_period,
|
||||
warp_period,
|
||||
loose_cfatnn=False):
|
||||
self.step_store = self.get_empty_store()
|
||||
self.cur_step = 0
|
||||
self.total_step = 0
|
||||
self.cur_index = 0
|
||||
self.init_store = False
|
||||
self.restore = False
|
||||
self.update = False
|
||||
self.flow = None
|
||||
self.mask = None
|
||||
self.restorex0 = False
|
||||
self.updatex0 = False
|
||||
self.inner_strength = inner_strength
|
||||
self.cross_period = cross_period
|
||||
self.mask_period = mask_period
|
||||
self.ada_period = ada_period
|
||||
self.warp_period = warp_period
|
||||
self.up_resolution = 1280 if loose_cfatnn else 1281
|
||||
|
||||
@staticmethod
|
||||
def get_empty_store():
|
||||
return {
|
||||
'first': [],
|
||||
'previous': [],
|
||||
'x0_previous': [],
|
||||
'first_ada': []
|
||||
}
|
||||
|
||||
def forward(self, context, is_cross: bool, place_in_unet: str):
|
||||
cross_period = (self.total_step * self.cross_period[0],
|
||||
self.total_step * self.cross_period[1])
|
||||
if not is_cross and place_in_unet == 'up' and context.shape[
|
||||
2] < self.up_resolution:
|
||||
if self.init_store:
|
||||
self.step_store['first'].append(context.detach())
|
||||
self.step_store['previous'].append(context.detach())
|
||||
if self.update:
|
||||
tmp = context.clone().detach()
|
||||
if self.restore and self.cur_step >= cross_period[0] and \
|
||||
self.cur_step <= cross_period[1]:
|
||||
context = torch.cat(
|
||||
(self.step_store['first'][self.cur_index],
|
||||
self.step_store['previous'][self.cur_index]),
|
||||
dim=1).clone()
|
||||
if self.update:
|
||||
self.step_store['previous'][self.cur_index] = tmp
|
||||
self.cur_index += 1
|
||||
return context
|
||||
|
||||
def update_x0(self, x0):
|
||||
if self.init_store:
|
||||
self.step_store['x0_previous'].append(x0.detach())
|
||||
style_mean, style_std = calc_mean_std(x0.detach())
|
||||
self.step_store['first_ada'].append(style_mean.detach())
|
||||
self.step_store['first_ada'].append(style_std.detach())
|
||||
if self.updatex0:
|
||||
tmp = x0.clone().detach()
|
||||
if self.restorex0:
|
||||
if self.cur_step >= self.total_step * self.ada_period[
|
||||
0] and self.cur_step <= self.total_step * self.ada_period[
|
||||
1]:
|
||||
x0 = F.instance_norm(x0) * self.step_store['first_ada'][
|
||||
2 * self.cur_step +
|
||||
1] + self.step_store['first_ada'][2 * self.cur_step]
|
||||
if self.cur_step >= self.total_step * self.warp_period[
|
||||
0] and self.cur_step <= self.total_step * self.warp_period[
|
||||
1]:
|
||||
pre = self.step_store['x0_previous'][self.cur_step]
|
||||
x0 = flow_warp(pre, self.flow, mode='nearest') * self.mask + (
|
||||
1 - self.mask) * x0
|
||||
if self.updatex0:
|
||||
self.step_store['x0_previous'][self.cur_step] = tmp
|
||||
return x0
|
||||
|
||||
def set_warp(self, flow, mask):
|
||||
self.flow = flow.clone()
|
||||
self.mask = mask.clone()
|
||||
|
||||
def __call__(self, context, is_cross: bool, place_in_unet: str):
|
||||
context = self.forward(context, is_cross, place_in_unet)
|
||||
return context
|
||||
|
||||
def set_step(self, step):
|
||||
self.cur_step = step
|
||||
|
||||
def set_total_step(self, total_step):
|
||||
self.total_step = total_step
|
||||
self.cur_index = 0
|
||||
|
||||
def clear_store(self):
|
||||
del self.step_store
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
self.step_store = self.get_empty_store()
|
||||
|
||||
def set_task(self, task, restore_step=1.0):
|
||||
self.init_store = False
|
||||
self.restore = False
|
||||
self.update = False
|
||||
self.cur_index = 0
|
||||
self.restore_step = restore_step
|
||||
self.updatex0 = False
|
||||
self.restorex0 = False
|
||||
if 'initfirst' in task:
|
||||
self.init_store = True
|
||||
self.clear_store()
|
||||
if 'updatestyle' in task:
|
||||
self.update = True
|
||||
if 'keepstyle' in task:
|
||||
self.restore = True
|
||||
if 'updatex0' in task:
|
||||
self.updatex0 = True
|
||||
if 'keepx0' in task:
|
||||
self.restorex0 = True
|
||||
@@ -0,0 +1,587 @@
|
||||
"""SAMPLING ONLY."""
|
||||
|
||||
# CrossAttn precision handling
|
||||
import os
|
||||
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from deps.ControlNet.ldm.modules.diffusionmodules.util import (
|
||||
extract_into_tensor, make_ddim_sampling_parameters, make_ddim_timesteps,
|
||||
noise_like)
|
||||
|
||||
_ATTN_PRECISION = os.environ.get('ATTN_PRECISION', 'fp32')
|
||||
|
||||
|
||||
def register_attention_control(model, controller=None):
|
||||
|
||||
def ca_forward(self, place_in_unet):
|
||||
|
||||
def forward(x, context=None, mask=None):
|
||||
h = self.heads
|
||||
|
||||
q = self.to_q(x)
|
||||
is_cross = context is not None
|
||||
context = context if is_cross else x
|
||||
context = controller(context, is_cross, place_in_unet)
|
||||
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
q, k, v = map(
|
||||
lambda t: einops.rearrange(t, 'b n (h d) -> (b h) n d', h=h),
|
||||
(q, k, v))
|
||||
|
||||
# force cast to fp32 to avoid overflowing
|
||||
if _ATTN_PRECISION == 'fp32':
|
||||
with torch.autocast(enabled=False, device_type='cuda'):
|
||||
q, k = q.float(), k.float()
|
||||
sim = torch.einsum('b i d, b j d -> b i j', q,
|
||||
k) * self.scale
|
||||
else:
|
||||
sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale
|
||||
|
||||
del q, k
|
||||
|
||||
if mask is not None:
|
||||
mask = einops.rearrange(mask, 'b ... -> b (...)')
|
||||
max_neg_value = -torch.finfo(sim.dtype).max
|
||||
mask = einops.repeat(mask, 'b j -> (b h) () j', h=h)
|
||||
sim.masked_fill_(~mask, max_neg_value)
|
||||
|
||||
# attention, what we cannot get enough of
|
||||
sim = sim.softmax(dim=-1)
|
||||
|
||||
out = torch.einsum('b i j, b j d -> b i d', sim, v)
|
||||
out = einops.rearrange(out, '(b h) n d -> b n (h d)', h=h)
|
||||
return self.to_out(out)
|
||||
|
||||
return forward
|
||||
|
||||
class DummyController:
|
||||
|
||||
def __call__(self, *args):
|
||||
return args[0]
|
||||
|
||||
def __init__(self):
|
||||
self.cur_step = 0
|
||||
|
||||
if controller is None:
|
||||
controller = DummyController()
|
||||
|
||||
def register_recr(net_, place_in_unet):
|
||||
if net_.__class__.__name__ == 'CrossAttention':
|
||||
net_.forward = ca_forward(net_, place_in_unet)
|
||||
elif hasattr(net_, 'children'):
|
||||
for net__ in net_.children():
|
||||
register_recr(net__, place_in_unet)
|
||||
|
||||
sub_nets = model.named_children()
|
||||
for net in sub_nets:
|
||||
if 'input_blocks' in net[0]:
|
||||
register_recr(net[1], 'down')
|
||||
elif 'output_blocks' in net[0]:
|
||||
register_recr(net[1], 'up')
|
||||
elif 'middle_block' in net[0]:
|
||||
register_recr(net[1], 'mid')
|
||||
|
||||
|
||||
class DDIMVSampler(object):
|
||||
|
||||
def __init__(self, model, schedule='linear', **kwargs):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self.ddpm_num_timesteps = model.num_timesteps
|
||||
self.schedule = schedule
|
||||
|
||||
def register_buffer(self, name, attr):
|
||||
if type(attr) == torch.Tensor:
|
||||
if attr.device != torch.device('cuda'):
|
||||
attr = attr.to(torch.device('cuda'))
|
||||
setattr(self, name, attr)
|
||||
|
||||
def make_schedule(self,
|
||||
ddim_num_steps,
|
||||
ddim_discretize='uniform',
|
||||
ddim_eta=0.,
|
||||
verbose=True):
|
||||
self.ddim_timesteps = make_ddim_timesteps(
|
||||
ddim_discr_method=ddim_discretize,
|
||||
num_ddim_timesteps=ddim_num_steps,
|
||||
num_ddpm_timesteps=self.ddpm_num_timesteps,
|
||||
verbose=verbose)
|
||||
alphas_cumprod = self.model.alphas_cumprod
|
||||
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, \
|
||||
'alphas have to be defined for each timestep'
|
||||
|
||||
def to_torch(x):
|
||||
return x.clone().detach().to(torch.float32).to(self.model.device)
|
||||
|
||||
self.register_buffer('betas', to_torch(self.model.betas))
|
||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
||||
self.register_buffer('alphas_cumprod_prev',
|
||||
to_torch(self.model.alphas_cumprod_prev))
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.register_buffer('sqrt_alphas_cumprod',
|
||||
to_torch(np.sqrt(alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_one_minus_alphas_cumprod',
|
||||
to_torch(np.sqrt(1. - alphas_cumprod.cpu())))
|
||||
self.register_buffer('log_one_minus_alphas_cumprod',
|
||||
to_torch(np.log(1. - alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod',
|
||||
to_torch(np.sqrt(1. / alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod',
|
||||
to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1)))
|
||||
|
||||
# ddim sampling parameters
|
||||
ddim_sigmas, ddim_alphas, ddim_alphas_prev = \
|
||||
make_ddim_sampling_parameters(
|
||||
alphacums=alphas_cumprod.cpu(),
|
||||
ddim_timesteps=self.ddim_timesteps,
|
||||
eta=ddim_eta,
|
||||
verbose=verbose)
|
||||
self.register_buffer('ddim_sigmas', ddim_sigmas)
|
||||
self.register_buffer('ddim_alphas', ddim_alphas)
|
||||
self.register_buffer('ddim_alphas_prev', ddim_alphas_prev)
|
||||
self.register_buffer('ddim_sqrt_one_minus_alphas',
|
||||
np.sqrt(1. - ddim_alphas))
|
||||
sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt(
|
||||
(1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) *
|
||||
(1 - self.alphas_cumprod / self.alphas_cumprod_prev))
|
||||
self.register_buffer('ddim_sigmas_for_original_num_steps',
|
||||
sigmas_for_original_sampling_steps)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self,
|
||||
S,
|
||||
batch_size,
|
||||
shape,
|
||||
conditioning=None,
|
||||
callback=None,
|
||||
img_callback=None,
|
||||
quantize_x0=False,
|
||||
eta=0.,
|
||||
mask=None,
|
||||
x0=None,
|
||||
xtrg=None,
|
||||
noise_rescale=None,
|
||||
temperature=1.,
|
||||
noise_dropout=0.,
|
||||
score_corrector=None,
|
||||
corrector_kwargs=None,
|
||||
verbose=True,
|
||||
x_T=None,
|
||||
log_every_t=100,
|
||||
unconditional_guidance_scale=1.,
|
||||
unconditional_conditioning=None,
|
||||
dynamic_threshold=None,
|
||||
ucg_schedule=None,
|
||||
controller=None,
|
||||
strength=0.0,
|
||||
**kwargs):
|
||||
if conditioning is not None:
|
||||
if isinstance(conditioning, dict):
|
||||
ctmp = conditioning[list(conditioning.keys())[0]]
|
||||
while isinstance(ctmp, list):
|
||||
ctmp = ctmp[0]
|
||||
cbs = ctmp.shape[0]
|
||||
if cbs != batch_size:
|
||||
print(f'Warning: Got {cbs} conditionings'
|
||||
f'but batch-size is {batch_size}')
|
||||
|
||||
elif isinstance(conditioning, list):
|
||||
for ctmp in conditioning:
|
||||
if ctmp.shape[0] != batch_size:
|
||||
print(f'Warning: Got {cbs} conditionings'
|
||||
f'but batch-size is {batch_size}')
|
||||
|
||||
else:
|
||||
if conditioning.shape[0] != batch_size:
|
||||
print(f'Warning: Got {conditioning.shape[0]}'
|
||||
f'conditionings but batch-size is {batch_size}')
|
||||
|
||||
self.make_schedule(ddim_num_steps=S, ddim_eta=eta, verbose=verbose)
|
||||
# sampling
|
||||
C, H, W = shape
|
||||
size = (batch_size, C, H, W)
|
||||
print(f'Data shape for DDIM sampling is {size}, eta {eta}')
|
||||
|
||||
samples, intermediates = self.ddim_sampling(
|
||||
conditioning,
|
||||
size,
|
||||
callback=callback,
|
||||
img_callback=img_callback,
|
||||
quantize_denoised=quantize_x0,
|
||||
mask=mask,
|
||||
x0=x0,
|
||||
xtrg=xtrg,
|
||||
noise_rescale=noise_rescale,
|
||||
ddim_use_original_steps=False,
|
||||
noise_dropout=noise_dropout,
|
||||
temperature=temperature,
|
||||
score_corrector=score_corrector,
|
||||
corrector_kwargs=corrector_kwargs,
|
||||
x_T=x_T,
|
||||
log_every_t=log_every_t,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning,
|
||||
dynamic_threshold=dynamic_threshold,
|
||||
ucg_schedule=ucg_schedule,
|
||||
controller=controller,
|
||||
strength=strength,
|
||||
)
|
||||
return samples, intermediates
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sampling(self,
|
||||
cond,
|
||||
shape,
|
||||
x_T=None,
|
||||
ddim_use_original_steps=False,
|
||||
callback=None,
|
||||
timesteps=None,
|
||||
quantize_denoised=False,
|
||||
mask=None,
|
||||
x0=None,
|
||||
xtrg=None,
|
||||
noise_rescale=None,
|
||||
img_callback=None,
|
||||
log_every_t=100,
|
||||
temperature=1.,
|
||||
noise_dropout=0.,
|
||||
score_corrector=None,
|
||||
corrector_kwargs=None,
|
||||
unconditional_guidance_scale=1.,
|
||||
unconditional_conditioning=None,
|
||||
dynamic_threshold=None,
|
||||
ucg_schedule=None,
|
||||
controller=None,
|
||||
strength=0.0):
|
||||
|
||||
if strength == 1 and x0 is not None:
|
||||
return x0, None
|
||||
|
||||
register_attention_control(self.model.model.diffusion_model,
|
||||
controller)
|
||||
|
||||
device = self.model.betas.device
|
||||
b = shape[0]
|
||||
if x_T is None:
|
||||
img = torch.randn(shape, device=device)
|
||||
else:
|
||||
img = x_T
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = self.ddpm_num_timesteps if ddim_use_original_steps \
|
||||
else self.ddim_timesteps
|
||||
elif timesteps is not None and not ddim_use_original_steps:
|
||||
subset_end = int(
|
||||
min(timesteps / self.ddim_timesteps.shape[0], 1) *
|
||||
self.ddim_timesteps.shape[0]) - 1
|
||||
timesteps = self.ddim_timesteps[:subset_end]
|
||||
|
||||
intermediates = {'x_inter': [img], 'pred_x0': [img]}
|
||||
time_range = reversed(range(
|
||||
0, timesteps)) if ddim_use_original_steps else np.flip(timesteps)
|
||||
total_steps = timesteps if ddim_use_original_steps \
|
||||
else timesteps.shape[0]
|
||||
print(f'Running DDIM Sampling with {total_steps} timesteps')
|
||||
|
||||
iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps)
|
||||
if controller is not None:
|
||||
controller.set_total_step(total_steps)
|
||||
if mask is None:
|
||||
mask = [None] * total_steps
|
||||
|
||||
dir_xt = 0
|
||||
for i, step in enumerate(iterator):
|
||||
if controller is not None:
|
||||
controller.set_step(i)
|
||||
index = total_steps - i - 1
|
||||
ts = torch.full((b, ), step, device=device, dtype=torch.long)
|
||||
|
||||
if strength >= 0 and i == int(
|
||||
total_steps * strength) and x0 is not None:
|
||||
img = self.model.q_sample(x0, ts)
|
||||
if mask is not None and xtrg is not None:
|
||||
# TODO: deterministic forward pass?
|
||||
if type(mask) == list:
|
||||
weight = mask[i]
|
||||
else:
|
||||
weight = mask
|
||||
if weight is not None:
|
||||
rescale = torch.maximum(1. - weight, (1 - weight**2)**0.5 *
|
||||
controller.inner_strength)
|
||||
if noise_rescale is not None:
|
||||
rescale = (1. - weight) * (
|
||||
1 - noise_rescale) + rescale * noise_rescale
|
||||
img_ref = self.model.q_sample(xtrg, ts)
|
||||
img = img_ref * weight + (1. - weight) * (
|
||||
img - dir_xt) + rescale * dir_xt
|
||||
|
||||
if ucg_schedule is not None:
|
||||
assert len(ucg_schedule) == len(time_range)
|
||||
unconditional_guidance_scale = ucg_schedule[i]
|
||||
|
||||
outs = self.p_sample_ddim(
|
||||
img,
|
||||
cond,
|
||||
ts,
|
||||
index=index,
|
||||
use_original_steps=ddim_use_original_steps,
|
||||
quantize_denoised=quantize_denoised,
|
||||
temperature=temperature,
|
||||
noise_dropout=noise_dropout,
|
||||
score_corrector=score_corrector,
|
||||
corrector_kwargs=corrector_kwargs,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning,
|
||||
dynamic_threshold=dynamic_threshold,
|
||||
controller=controller,
|
||||
return_dir=True)
|
||||
img, pred_x0, dir_xt = outs
|
||||
if callback:
|
||||
callback(i)
|
||||
if img_callback:
|
||||
img_callback(pred_x0, i)
|
||||
|
||||
if index % log_every_t == 0 or index == total_steps - 1:
|
||||
intermediates['x_inter'].append(img)
|
||||
intermediates['pred_x0'].append(pred_x0)
|
||||
|
||||
return img, intermediates
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_ddim(self,
|
||||
x,
|
||||
c,
|
||||
t,
|
||||
index,
|
||||
repeat_noise=False,
|
||||
use_original_steps=False,
|
||||
quantize_denoised=False,
|
||||
temperature=1.,
|
||||
noise_dropout=0.,
|
||||
score_corrector=None,
|
||||
corrector_kwargs=None,
|
||||
unconditional_guidance_scale=1.,
|
||||
unconditional_conditioning=None,
|
||||
dynamic_threshold=None,
|
||||
controller=None,
|
||||
return_dir=False):
|
||||
b, *_, device = *x.shape, x.device
|
||||
|
||||
if unconditional_conditioning is None or \
|
||||
unconditional_guidance_scale == 1.:
|
||||
model_output = self.model.apply_model(x, t, c)
|
||||
else:
|
||||
model_t = self.model.apply_model(x, t, c)
|
||||
model_uncond = self.model.apply_model(x, t,
|
||||
unconditional_conditioning)
|
||||
model_output = model_uncond + unconditional_guidance_scale * (
|
||||
model_t - model_uncond)
|
||||
|
||||
if self.model.parameterization == 'v':
|
||||
e_t = self.model.predict_eps_from_z_and_v(x, t, model_output)
|
||||
else:
|
||||
e_t = model_output
|
||||
|
||||
if score_corrector is not None:
|
||||
assert self.model.parameterization == 'eps', 'not implemented'
|
||||
e_t = score_corrector.modify_score(self.model, e_t, x, t, c,
|
||||
**corrector_kwargs)
|
||||
|
||||
if use_original_steps:
|
||||
alphas = self.model.alphas_cumprod
|
||||
alphas_prev = self.model.alphas_cumprod_prev
|
||||
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod
|
||||
sigmas = self.model.ddim_sigmas_for_original_num_steps
|
||||
else:
|
||||
alphas = self.ddim_alphas
|
||||
alphas_prev = self.ddim_alphas_prev
|
||||
sqrt_one_minus_alphas = self.ddim_sqrt_one_minus_alphas
|
||||
sigmas = self.ddim_sigmas
|
||||
|
||||
# select parameters corresponding to the currently considered timestep
|
||||
a_t = torch.full((b, 1, 1, 1), alphas[index], device=device)
|
||||
a_prev = torch.full((b, 1, 1, 1), alphas_prev[index], device=device)
|
||||
sigma_t = torch.full((b, 1, 1, 1), sigmas[index], device=device)
|
||||
sqrt_one_minus_at = torch.full((b, 1, 1, 1),
|
||||
sqrt_one_minus_alphas[index],
|
||||
device=device)
|
||||
|
||||
# current prediction for x_0
|
||||
if self.model.parameterization != 'v':
|
||||
pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt()
|
||||
else:
|
||||
pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output)
|
||||
|
||||
if quantize_denoised:
|
||||
pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0)
|
||||
|
||||
if dynamic_threshold is not None:
|
||||
raise NotImplementedError()
|
||||
'''
|
||||
if mask is not None and xtrg is not None:
|
||||
pred_x0 = xtrg * mask + (1. - mask) * pred_x0
|
||||
'''
|
||||
|
||||
if controller is not None:
|
||||
pred_x0 = controller.update_x0(pred_x0)
|
||||
|
||||
# direction pointing to x_t
|
||||
dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t
|
||||
noise = sigma_t * noise_like(x.shape, device,
|
||||
repeat_noise) * temperature
|
||||
if noise_dropout > 0.:
|
||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
||||
x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise
|
||||
|
||||
if return_dir:
|
||||
return x_prev, pred_x0, dir_xt
|
||||
return x_prev, pred_x0
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self,
|
||||
x0,
|
||||
c,
|
||||
t_enc,
|
||||
use_original_steps=False,
|
||||
return_intermediates=None,
|
||||
unconditional_guidance_scale=1.0,
|
||||
unconditional_conditioning=None,
|
||||
callback=None):
|
||||
timesteps = np.arange(self.ddpm_num_timesteps
|
||||
) if use_original_steps else self.ddim_timesteps
|
||||
num_reference_steps = timesteps.shape[0]
|
||||
|
||||
assert t_enc <= num_reference_steps
|
||||
num_steps = t_enc
|
||||
|
||||
if use_original_steps:
|
||||
alphas_next = self.alphas_cumprod[:num_steps]
|
||||
alphas = self.alphas_cumprod_prev[:num_steps]
|
||||
else:
|
||||
alphas_next = self.ddim_alphas[:num_steps]
|
||||
alphas = torch.tensor(self.ddim_alphas_prev[:num_steps])
|
||||
|
||||
x_next = x0
|
||||
intermediates = []
|
||||
inter_steps = []
|
||||
for i in tqdm(range(num_steps), desc='Encoding Image'):
|
||||
t = torch.full((x0.shape[0], ),
|
||||
timesteps[i],
|
||||
device=self.model.device,
|
||||
dtype=torch.long)
|
||||
if unconditional_guidance_scale == 1.:
|
||||
noise_pred = self.model.apply_model(x_next, t, c)
|
||||
else:
|
||||
assert unconditional_conditioning is not None
|
||||
e_t_uncond, noise_pred = torch.chunk(
|
||||
self.model.apply_model(
|
||||
torch.cat((x_next, x_next)), torch.cat((t, t)),
|
||||
torch.cat((unconditional_conditioning, c))), 2)
|
||||
noise_pred = e_t_uncond + unconditional_guidance_scale * (
|
||||
noise_pred - e_t_uncond)
|
||||
xt_weighted = (alphas_next[i] / alphas[i]).sqrt() * x_next
|
||||
weighted_noise_pred = alphas_next[i].sqrt() * (
|
||||
(1 / alphas_next[i] - 1).sqrt() -
|
||||
(1 / alphas[i] - 1).sqrt()) * noise_pred
|
||||
x_next = xt_weighted + weighted_noise_pred
|
||||
if return_intermediates and i % (num_steps // return_intermediates
|
||||
) == 0 and i < num_steps - 1:
|
||||
intermediates.append(x_next)
|
||||
inter_steps.append(i)
|
||||
elif return_intermediates and i >= num_steps - 2:
|
||||
intermediates.append(x_next)
|
||||
inter_steps.append(i)
|
||||
if callback:
|
||||
callback(i)
|
||||
|
||||
out = {'x_encoded': x_next, 'intermediate_steps': inter_steps}
|
||||
if return_intermediates:
|
||||
out.update({'intermediates': intermediates})
|
||||
return x_next, out
|
||||
|
||||
@torch.no_grad()
|
||||
def stochastic_encode(self, x0, t, use_original_steps=False, noise=None):
|
||||
# fast, but does not allow for exact reconstruction
|
||||
# t serves as an index to gather the correct alphas
|
||||
if use_original_steps:
|
||||
sqrt_alphas_cumprod = self.sqrt_alphas_cumprod
|
||||
sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod
|
||||
else:
|
||||
sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas)
|
||||
sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x0)
|
||||
if t >= len(sqrt_alphas_cumprod):
|
||||
return noise
|
||||
return (
|
||||
extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) *
|
||||
noise)
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self,
|
||||
x_latent,
|
||||
cond,
|
||||
t_start,
|
||||
unconditional_guidance_scale=1.0,
|
||||
unconditional_conditioning=None,
|
||||
use_original_steps=False,
|
||||
callback=None):
|
||||
|
||||
timesteps = np.arange(self.ddpm_num_timesteps
|
||||
) if use_original_steps else self.ddim_timesteps
|
||||
timesteps = timesteps[:t_start]
|
||||
|
||||
time_range = np.flip(timesteps)
|
||||
total_steps = timesteps.shape[0]
|
||||
print(f'Running DDIM Sampling with {total_steps} timesteps')
|
||||
|
||||
iterator = tqdm(time_range, desc='Decoding image', total=total_steps)
|
||||
x_dec = x_latent
|
||||
for i, step in enumerate(iterator):
|
||||
index = total_steps - i - 1
|
||||
ts = torch.full((x_latent.shape[0], ),
|
||||
step,
|
||||
device=x_latent.device,
|
||||
dtype=torch.long)
|
||||
x_dec, _ = self.p_sample_ddim(
|
||||
x_dec,
|
||||
cond,
|
||||
ts,
|
||||
index=index,
|
||||
use_original_steps=use_original_steps,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning)
|
||||
if callback:
|
||||
callback(i)
|
||||
return x_dec
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps=1e-5):
|
||||
# eps is a small value added to the variance to avoid divide-by-zero.
|
||||
size = feat.size()
|
||||
assert (len(size) == 4)
|
||||
N, C = size[:2]
|
||||
feat_var = feat.view(N, C, -1).var(dim=2) + eps
|
||||
feat_std = feat_var.sqrt().view(N, C, 1, 1)
|
||||
feat_mean = feat.view(N, C, -1).mean(dim=2).view(N, C, 1, 1)
|
||||
return feat_mean, feat_std
|
||||
|
||||
|
||||
def adaptive_instance_normalization(content_feat, style_feat):
|
||||
assert (content_feat.size()[:2] == style_feat.size()[:2])
|
||||
size = content_feat.size()
|
||||
style_mean, style_std = calc_mean_std(style_feat)
|
||||
content_mean, content_std = calc_mean_std(content_feat)
|
||||
|
||||
normalized_feat = (content_feat -
|
||||
content_mean.expand(size)) / content_std.expand(size)
|
||||
return normalized_feat * style_std.expand(size) + style_mean.expand(size)
|
||||
@@ -0,0 +1,94 @@
|
||||
import torch
|
||||
import torch.fft as fft
|
||||
|
||||
|
||||
def Fourier_filter(x, threshold, scale):
|
||||
|
||||
x_freq = fft.fftn(x, dim=(-2, -1))
|
||||
x_freq = fft.fftshift(x_freq, dim=(-2, -1))
|
||||
|
||||
B, C, H, W = x_freq.shape
|
||||
mask = torch.ones((B, C, H, W)).cuda()
|
||||
|
||||
crow, ccol = H // 2, W // 2
|
||||
mask[..., crow - threshold:crow + threshold,
|
||||
ccol - threshold:ccol + threshold] = scale
|
||||
x_freq = x_freq * mask
|
||||
|
||||
x_freq = fft.ifftshift(x_freq, dim=(-2, -1))
|
||||
|
||||
x_filtered = fft.ifftn(x_freq, dim=(-2, -1)).real
|
||||
|
||||
return x_filtered
|
||||
|
||||
from deps.ControlNet.ldm.modules.diffusionmodules.util import \
|
||||
timestep_embedding # noqa:E501
|
||||
|
||||
|
||||
# backbone_scale1=1.1, backbone_scale2=1.2, skip_scale1=1.0, skip_scale2=0.2
|
||||
def freeu_forward(self,
|
||||
backbone_scale1=1.,
|
||||
backbone_scale2=1.,
|
||||
skip_scale1=1.,
|
||||
skip_scale2=1.):
|
||||
|
||||
def forward(x,
|
||||
timesteps=None,
|
||||
context=None,
|
||||
control=None,
|
||||
only_mid_control=False,
|
||||
**kwargs):
|
||||
hs = []
|
||||
with torch.no_grad():
|
||||
t_emb = timestep_embedding(timesteps,
|
||||
self.model_channels,
|
||||
repeat_only=False)
|
||||
emb = self.time_embed(t_emb)
|
||||
h = x.type(self.dtype)
|
||||
for module in self.input_blocks:
|
||||
h = module(h, emb, context)
|
||||
hs.append(h)
|
||||
h = self.middle_block(h, emb, context)
|
||||
|
||||
if control is not None:
|
||||
h += control.pop()
|
||||
'''
|
||||
for i, module in enumerate(self.output_blocks):
|
||||
if only_mid_control or control is None:
|
||||
h = torch.cat([h, hs.pop()], dim=1)
|
||||
else:
|
||||
h = torch.cat([h, hs.pop() + control.pop()], dim=1)
|
||||
h = module(h, emb, context)
|
||||
'''
|
||||
for i, module in enumerate(self.output_blocks):
|
||||
hs_ = hs.pop()
|
||||
|
||||
if h.shape[1] == 1280:
|
||||
hidden_mean = h.mean(1).unsqueeze(1)
|
||||
B = hidden_mean.shape[0]
|
||||
hidden_max, _ = torch.max(hidden_mean.view(B, -1), dim=-1, keepdim=True)
|
||||
hidden_min, _ = torch.min(hidden_mean.view(B, -1), dim=-1, keepdim=True)
|
||||
hidden_mean = (hidden_mean - hidden_min.unsqueeze(2).unsqueeze(3)) / (hidden_max - hidden_min).unsqueeze(2).unsqueeze(3)
|
||||
h[:, :640] = h[:, :640] * ((backbone_scale1 - 1) * hidden_mean + 1)
|
||||
# h[:, :640] = h[:, :640] * backbone_scale1
|
||||
hs_ = Fourier_filter(hs_, threshold=1, scale=skip_scale1)
|
||||
if h.shape[1] == 640:
|
||||
hidden_mean = h.mean(1).unsqueeze(1)
|
||||
B = hidden_mean.shape[0]
|
||||
hidden_max, _ = torch.max(hidden_mean.view(B, -1), dim=-1, keepdim=True)
|
||||
hidden_min, _ = torch.min(hidden_mean.view(B, -1), dim=-1, keepdim=True)
|
||||
hidden_mean = (hidden_mean - hidden_min.unsqueeze(2).unsqueeze(3)) / (hidden_max - hidden_min).unsqueeze(2).unsqueeze(3)
|
||||
h[:, :320] = h[:, :320] * ((backbone_scale2 - 1) * hidden_mean + 1)
|
||||
# h[:, :320] = h[:, :320] * backbone_scale2
|
||||
hs_ = Fourier_filter(hs_, threshold=1, scale=skip_scale2)
|
||||
|
||||
if only_mid_control or control is None:
|
||||
h = torch.cat([h, hs_], dim=1)
|
||||
else:
|
||||
h = torch.cat([h, hs_ + control.pop()], dim=1)
|
||||
h = module(h, emb, context)
|
||||
|
||||
h = h.type(x.dtype)
|
||||
return self.out(h)
|
||||
|
||||
return forward
|
||||
@@ -0,0 +1,23 @@
|
||||
import einops
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def find_flat_region(mask):
|
||||
device = mask.device
|
||||
kernel_x = torch.Tensor([[-1, 0, 1], [-1, 0, 1],
|
||||
[-1, 0, 1]]).unsqueeze(0).unsqueeze(0).to(device)
|
||||
kernel_y = torch.Tensor([[-1, -1, -1], [0, 0, 0],
|
||||
[1, 1, 1]]).unsqueeze(0).unsqueeze(0).to(device)
|
||||
mask_ = F.pad(mask.unsqueeze(0), (1, 1, 1, 1), mode='replicate')
|
||||
|
||||
grad_x = torch.nn.functional.conv2d(mask_, kernel_x)
|
||||
grad_y = torch.nn.functional.conv2d(mask_, kernel_y)
|
||||
return ((abs(grad_x) + abs(grad_y)) == 0).float()[0]
|
||||
|
||||
|
||||
def numpy2tensor(img):
|
||||
x0 = torch.from_numpy(img.copy()).float().cuda() / 255.0 * 2.0 - 1.
|
||||
x0 = torch.stack([x0], dim=0)
|
||||
return einops.rearrange(x0, 'b h w c -> b c h w').clone()
|
||||
@@ -0,0 +1,10 @@
|
||||
import os
|
||||
import sys
|
||||
|
||||
cur_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
gmflow_dir = os.path.join(cur_dir, 'deps/gmflow')
|
||||
controlnet_dir = os.path.join(cur_dir, 'deps/ControlNet')
|
||||
sys.path.insert(0, gmflow_dir)
|
||||
sys.path.insert(0, controlnet_dir)
|
||||
|
||||
import deps.ControlNet.share # noqa: F401 E402
|
||||
@@ -0,0 +1,182 @@
|
||||
import json
|
||||
import pathlib
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from safetensors import safe_open
|
||||
|
||||
unet_mapping = json.load(open('unet_lora_mapping.json'))
|
||||
|
||||
|
||||
class LoraLinear(torch.nn.Module):
|
||||
def __init__(self, old_linear, up=None, down=None):
|
||||
super().__init__()
|
||||
self.old_linear = old_linear
|
||||
self.up = up
|
||||
self.down = down
|
||||
self.alpha = 1.0
|
||||
self.new_linear = None
|
||||
|
||||
def forward(self, x):
|
||||
if self.new_linear is None:
|
||||
self.new_linear = nn.Linear(self.down.shape[1], self.up.shape[0])
|
||||
self.new_linear.weight.data = self.up @ self.down
|
||||
self.new_linear = self.new_linear.to('cuda').to(torch.float32)
|
||||
x1 = self.new_linear(x) * (self.alpha / float(self.up[1].size()[0])) if self.alpha is not None else 1
|
||||
x2 = self.old_linear(x)
|
||||
return x1 + x2
|
||||
|
||||
|
||||
class LoraConv2d(torch.nn.Module):
|
||||
def __init__(self, old_conv2d, up=None, down=None):
|
||||
super().__init__()
|
||||
self.old_conv2d = old_conv2d
|
||||
self.up = up
|
||||
self.down = down
|
||||
self.alpha = None
|
||||
self.new_conv2d = None
|
||||
|
||||
def forward(self, x):
|
||||
if self.new_conv2d is None:
|
||||
self.new_conv2d = nn.Conv2d(self.down.shape[1], self.up.shape[0], 1)
|
||||
self.new_conv2d.weight.data = (self.up.squeeze() @ self.down.squeeze()).unsqueeze(-1).unsqueeze(-1)
|
||||
self.new_conv2d = self.new_conv2d.to('cuda').to(torch.float32)
|
||||
x1 = self.new_conv2d(x) * (self.alpha / float(self.up[1].size()[0])) if self.alpha is not None else 1
|
||||
x2 = self.old_conv2d(x)
|
||||
return x1 + x2
|
||||
|
||||
|
||||
def _get_nested_attr(obj, attr_path):
|
||||
elements = attr_path.split('.')
|
||||
for elem in elements:
|
||||
if elem.isdigit(): # If the element is a digit, access it as an index
|
||||
obj = obj[int(elem)]
|
||||
else: # Otherwise, access it as an attribute
|
||||
obj = getattr(obj, elem)
|
||||
return obj
|
||||
|
||||
|
||||
def _set_nested_attr(obj, attr_path, value):
|
||||
elements = attr_path.split('.')
|
||||
for elem in elements[:-1]:
|
||||
if elem.isdigit():
|
||||
obj = obj[int(elem)]
|
||||
else:
|
||||
obj = getattr(obj, elem)
|
||||
setattr(obj, elements[-1], value)
|
||||
|
||||
|
||||
def _replace_underscores(input_string, exceptions):
|
||||
# Split the input string by underscores
|
||||
parts = input_string.split('_')
|
||||
|
||||
# Reconstruct the string, replacing underscores with dots
|
||||
# except in the specified exceptions
|
||||
output_parts = []
|
||||
skip_next = False
|
||||
for i, part in enumerate(parts):
|
||||
if skip_next:
|
||||
skip_next = False
|
||||
continue
|
||||
|
||||
# Check if this part combined with the next part is in exceptions
|
||||
if i < len(parts) - 1 and '_'.join([part, parts[i + 1]]) in exceptions:
|
||||
output_parts.append('_'.join([part, parts[i + 1]]))
|
||||
skip_next = True
|
||||
else:
|
||||
output_parts.append(part)
|
||||
|
||||
return '.'.join(output_parts)
|
||||
|
||||
|
||||
def apply_lora(unet, text_encoder, lora_weights):
|
||||
for param_name, param_value in lora_weights.items():
|
||||
if 'unet' in param_name:
|
||||
# Determine layer type
|
||||
if 'lora_down' in param_name:
|
||||
layer_type = 'down'
|
||||
elif 'lora_up' in param_name:
|
||||
layer_type = 'up'
|
||||
elif 'alpha' in param_name:
|
||||
layer_type = 'alpha'
|
||||
else:
|
||||
raise ValueError(f'Unknown layer type for {param_name}')
|
||||
|
||||
base_param_name = param_name.split('.')[0]
|
||||
replaced_name = _replace_underscores(base_param_name, exceptions=[
|
||||
'up_blocks', 'mid_block', 'down_blocks', 'transformer_blocks',
|
||||
'to_k', 'to_q', 'to_v', 'to_out', 'to_in', 'proj_in', 'proj_out',
|
||||
])
|
||||
|
||||
for diffusers_name, sd_name in unet_mapping.items():
|
||||
base_diffusers = diffusers_name.rsplit('.', 1)[0]
|
||||
base_sd = sd_name.rsplit('.', 1)[0]
|
||||
if base_diffusers in replaced_name:
|
||||
replaced_name = replaced_name.replace(base_diffusers, base_sd)
|
||||
|
||||
replaced_name = replaced_name.replace('lora.unet', 'diffusion_model')
|
||||
param_value = param_value.to('cuda').to(torch.float32)
|
||||
curModule = _get_nested_attr(unet, replaced_name)
|
||||
|
||||
# Update module type if needed
|
||||
if isinstance(curModule, torch.nn.Linear):
|
||||
_set_nested_attr(unet, replaced_name, LoraLinear(curModule))
|
||||
elif isinstance(curModule, torch.nn.Conv2d):
|
||||
_set_nested_attr(unet, replaced_name, LoraConv2d(curModule))
|
||||
elif not isinstance(curModule, (LoraLinear, LoraConv2d)):
|
||||
raise ValueError(f'Unknown layer type: {type(curModule)}')
|
||||
|
||||
if layer_type == 'down':
|
||||
_set_nested_attr(unet, replaced_name + '.down', param_value)
|
||||
elif layer_type == 'up':
|
||||
_set_nested_attr(unet, replaced_name + '.up', param_value)
|
||||
elif layer_type == 'alpha':
|
||||
_set_nested_attr(unet, replaced_name + '.alpha', param_value)
|
||||
elif 'te_text_model' in param_name:
|
||||
# Determine layer type
|
||||
if 'lora_down' in param_name:
|
||||
layer_type = 'down'
|
||||
elif 'lora_up' in param_name:
|
||||
layer_type = 'up'
|
||||
elif 'alpha' in param_name:
|
||||
layer_type = 'alpha'
|
||||
else:
|
||||
raise ValueError(f'Unknown layer type for {param_name}')
|
||||
|
||||
base_param_name = param_name.split('.')[0]
|
||||
replaced_name = base_param_name.replace('lora_te', 'transformer')
|
||||
replaced_name = _replace_underscores(replaced_name, exceptions=[
|
||||
'text_model', 'self_attn', 'k_proj', 'v_proj', 'q_proj', 'out_proj',
|
||||
])
|
||||
curModule = _get_nested_attr(text_encoder, replaced_name)
|
||||
if isinstance(curModule, torch.nn.Linear):
|
||||
_set_nested_attr(text_encoder, replaced_name, LoraLinear(curModule))
|
||||
elif not isinstance(curModule, LoraLinear):
|
||||
raise ValueError(f'Unknown layer type: {type(curModule)}')
|
||||
param_value = param_value.to('cuda').to(torch.float32)
|
||||
if layer_type == 'down':
|
||||
_set_nested_attr(text_encoder, replaced_name + '.down', param_value)
|
||||
elif layer_type == 'up':
|
||||
_set_nested_attr(text_encoder, replaced_name + '.up', param_value)
|
||||
elif layer_type == 'alpha':
|
||||
_set_nested_attr(text_encoder, replaced_name + '.alpha', param_value)
|
||||
|
||||
|
||||
def load_lora(filepath):
|
||||
path = pathlib.Path(filepath)
|
||||
if path.suffix == '.safetensors':
|
||||
try:
|
||||
print(f'Loading LORA weights using safetensors from {filepath}')
|
||||
lora_weights = {}
|
||||
with safe_open(filepath, framework="pt", device='cuda') as f:
|
||||
for key in f.keys():
|
||||
lora_weights[key] = f.get_tensor(key)
|
||||
except Exception as e:
|
||||
raise Exception(f'Error loading LORA weights: {str(e)}')
|
||||
else:
|
||||
try:
|
||||
print(f'Loading LORA weights using torch from {filepath}')
|
||||
lora_weights = torch.load(filepath, map_location='cuda')
|
||||
except Exception as e:
|
||||
raise Exception(f'Error loading LORA weights from {filepath}: {str(e)}')
|
||||
return lora_weights
|
||||
@@ -0,0 +1,127 @@
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import torch
|
||||
import imageio
|
||||
import numpy as np
|
||||
|
||||
|
||||
def video_to_frame(video_path: str,
|
||||
frame_dir: str,
|
||||
filename_pattern: str = 'frame%03d.jpg',
|
||||
log: bool = True,
|
||||
frame_edit_func=None):
|
||||
os.makedirs(frame_dir, exist_ok=True)
|
||||
|
||||
vidcap = cv2.VideoCapture(video_path)
|
||||
success, image = vidcap.read()
|
||||
|
||||
if log:
|
||||
print('img shape: ', image.shape[0:2])
|
||||
|
||||
count = 0
|
||||
while success:
|
||||
if frame_edit_func is not None:
|
||||
image = frame_edit_func(image)
|
||||
|
||||
cv2.imwrite(os.path.join(frame_dir, filename_pattern % count), image)
|
||||
success, image = vidcap.read()
|
||||
if log:
|
||||
print('Read a new frame: ', success, count)
|
||||
count += 1
|
||||
|
||||
vidcap.release()
|
||||
|
||||
|
||||
def frame_to_video(video_path: str, frame_dir: str, fps=30, log=True):
|
||||
|
||||
first_img = True
|
||||
writer = imageio.get_writer(video_path, fps=fps)
|
||||
|
||||
file_list = sorted(os.listdir(frame_dir))
|
||||
for file_name in file_list:
|
||||
if not (file_name.endswith('jpg') or file_name.endswith('png')):
|
||||
continue
|
||||
|
||||
fn = os.path.join(frame_dir, file_name)
|
||||
curImg = imageio.imread(fn)
|
||||
|
||||
if first_img:
|
||||
H, W = curImg.shape[0:2]
|
||||
if log:
|
||||
print('img shape', (H, W))
|
||||
first_img = False
|
||||
|
||||
writer.append_data(curImg)
|
||||
|
||||
writer.close()
|
||||
|
||||
|
||||
def get_fps(video_path: str):
|
||||
video = cv2.VideoCapture(video_path)
|
||||
fps = video.get(cv2.CAP_PROP_FPS)
|
||||
video.release()
|
||||
return fps
|
||||
|
||||
|
||||
def get_frame_count(video_path: str):
|
||||
video = cv2.VideoCapture(video_path)
|
||||
frame_count = int(video.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
video.release()
|
||||
return frame_count
|
||||
|
||||
|
||||
def resize_image(input_image, resolution):
|
||||
H, W, C = input_image.shape
|
||||
H = float(H)
|
||||
W = float(W)
|
||||
aspect_ratio = W / H
|
||||
k = float(resolution) / min(H, W)
|
||||
H *= k
|
||||
W *= k
|
||||
if H < W:
|
||||
W = resolution
|
||||
H = int(resolution / aspect_ratio)
|
||||
else:
|
||||
H = resolution
|
||||
W = int(aspect_ratio * resolution)
|
||||
H = int(np.round(H / 64.0)) * 64
|
||||
W = int(np.round(W / 64.0)) * 64
|
||||
img = cv2.resize(
|
||||
input_image, (W, H),
|
||||
interpolation=cv2.INTER_LANCZOS4 if k > 1 else cv2.INTER_AREA)
|
||||
return img
|
||||
|
||||
|
||||
def prepare_frames(input_path: str, output_dir: str, resolution: int, crop, use_limit_device_resolution=False):
|
||||
l, r, t, b = crop
|
||||
|
||||
if use_limit_device_resolution:
|
||||
resolution = vram_limit_device_resolution(resolution)
|
||||
|
||||
def crop_func(frame):
|
||||
H, W, C = frame.shape
|
||||
left = np.clip(l, 0, W)
|
||||
right = np.clip(W - r, left, W)
|
||||
top = np.clip(t, 0, H)
|
||||
bottom = np.clip(H - b, top, H)
|
||||
frame = frame[top:bottom, left:right]
|
||||
return resize_image(frame, resolution)
|
||||
|
||||
video_to_frame(input_path, output_dir, '%04d.png', False, crop_func)
|
||||
|
||||
|
||||
def vram_limit_device_resolution(resolution, device="cuda"):
|
||||
# get max limit target size
|
||||
gpu_vram = torch.cuda.get_device_properties(device).total_memory / (1024 ** 3)
|
||||
# table of gpu memory limit
|
||||
gpu_table = {24: 1280, 18: 1024, 14: 768, 10: 640, 8: 576, 7: 512, 6: 448, 5: 320, 4: 192, 0: 0}
|
||||
# get user resize for gpu
|
||||
device_resolution = max(val for key, val in gpu_table.items() if key <= gpu_vram)
|
||||
print(f"Limit VRAM is {gpu_vram} Gb and size {device_resolution}.")
|
||||
if gpu_vram < 4:
|
||||
print(f"Small VRAM to use GPU. Configuration resolution will be used.")
|
||||
if resolution < device_resolution:
|
||||
print(f"Video will not resize")
|
||||
return resolution
|
||||
return device_resolution
|
||||
@@ -0,0 +1,688 @@
|
||||
{
|
||||
"conv_in.weight": "input_blocks.0.0.weight",
|
||||
"conv_in.bias": "input_blocks.0.0.bias",
|
||||
"time_embedding.linear_1.weight": "time_embed.0.weight",
|
||||
"time_embedding.linear_1.bias": "time_embed.0.bias",
|
||||
"time_embedding.linear_2.weight": "time_embed.2.weight",
|
||||
"time_embedding.linear_2.bias": "time_embed.2.bias",
|
||||
"down_blocks.0.attentions.0.norm.weight": "input_blocks.1.1.norm.weight",
|
||||
"down_blocks.0.attentions.0.norm.bias": "input_blocks.1.1.norm.bias",
|
||||
"down_blocks.0.attentions.0.proj_in.weight": "input_blocks.1.1.proj_in.weight",
|
||||
"down_blocks.0.attentions.0.proj_in.bias": "input_blocks.1.1.proj_in.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn1.to_q.weight": "input_blocks.1.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn1.to_k.weight": "input_blocks.1.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn1.to_v.weight": "input_blocks.1.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "input_blocks.1.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "input_blocks.1.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "input_blocks.1.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "input_blocks.1.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.ff.net.2.weight": "input_blocks.1.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.ff.net.2.bias": "input_blocks.1.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn2.to_q.weight": "input_blocks.1.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn2.to_k.weight": "input_blocks.1.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn2.to_v.weight": "input_blocks.1.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "input_blocks.1.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "input_blocks.1.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.norm1.weight": "input_blocks.1.1.transformer_blocks.0.norm1.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.norm1.bias": "input_blocks.1.1.transformer_blocks.0.norm1.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.norm2.weight": "input_blocks.1.1.transformer_blocks.0.norm2.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.norm2.bias": "input_blocks.1.1.transformer_blocks.0.norm2.bias",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.norm3.weight": "input_blocks.1.1.transformer_blocks.0.norm3.weight",
|
||||
"down_blocks.0.attentions.0.transformer_blocks.0.norm3.bias": "input_blocks.1.1.transformer_blocks.0.norm3.bias",
|
||||
"down_blocks.0.attentions.0.proj_out.weight": "input_blocks.1.1.proj_out.weight",
|
||||
"down_blocks.0.attentions.0.proj_out.bias": "input_blocks.1.1.proj_out.bias",
|
||||
"down_blocks.0.attentions.1.norm.weight": "input_blocks.2.1.norm.weight",
|
||||
"down_blocks.0.attentions.1.norm.bias": "input_blocks.2.1.norm.bias",
|
||||
"down_blocks.0.attentions.1.proj_in.weight": "input_blocks.2.1.proj_in.weight",
|
||||
"down_blocks.0.attentions.1.proj_in.bias": "input_blocks.2.1.proj_in.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn1.to_q.weight": "input_blocks.2.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn1.to_k.weight": "input_blocks.2.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn1.to_v.weight": "input_blocks.2.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn1.to_out.0.weight": "input_blocks.2.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn1.to_out.0.bias": "input_blocks.2.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.ff.net.0.proj.weight": "input_blocks.2.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.ff.net.0.proj.bias": "input_blocks.2.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.ff.net.2.weight": "input_blocks.2.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.ff.net.2.bias": "input_blocks.2.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn2.to_q.weight": "input_blocks.2.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn2.to_k.weight": "input_blocks.2.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn2.to_v.weight": "input_blocks.2.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn2.to_out.0.weight": "input_blocks.2.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.attn2.to_out.0.bias": "input_blocks.2.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.norm1.weight": "input_blocks.2.1.transformer_blocks.0.norm1.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.norm1.bias": "input_blocks.2.1.transformer_blocks.0.norm1.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.norm2.weight": "input_blocks.2.1.transformer_blocks.0.norm2.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.norm2.bias": "input_blocks.2.1.transformer_blocks.0.norm2.bias",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.norm3.weight": "input_blocks.2.1.transformer_blocks.0.norm3.weight",
|
||||
"down_blocks.0.attentions.1.transformer_blocks.0.norm3.bias": "input_blocks.2.1.transformer_blocks.0.norm3.bias",
|
||||
"down_blocks.0.attentions.1.proj_out.weight": "input_blocks.2.1.proj_out.weight",
|
||||
"down_blocks.0.attentions.1.proj_out.bias": "input_blocks.2.1.proj_out.bias",
|
||||
"down_blocks.0.resnets.0.norm1.weight": "input_blocks.1.0.in_layers.0.weight",
|
||||
"down_blocks.0.resnets.0.norm1.bias": "input_blocks.1.0.in_layers.0.bias",
|
||||
"down_blocks.0.resnets.0.conv1.weight": "input_blocks.1.0.in_layers.2.weight",
|
||||
"down_blocks.0.resnets.0.conv1.bias": "input_blocks.1.0.in_layers.2.bias",
|
||||
"down_blocks.0.resnets.0.time_emb_proj.weight": "input_blocks.1.0.emb_layers.1.weight",
|
||||
"down_blocks.0.resnets.0.time_emb_proj.bias": "input_blocks.1.0.emb_layers.1.bias",
|
||||
"down_blocks.0.resnets.0.norm2.weight": "input_blocks.1.0.out_layers.0.weight",
|
||||
"down_blocks.0.resnets.0.norm2.bias": "input_blocks.1.0.out_layers.0.bias",
|
||||
"down_blocks.0.resnets.0.conv2.weight": "input_blocks.1.0.out_layers.3.weight",
|
||||
"down_blocks.0.resnets.0.conv2.bias": "input_blocks.1.0.out_layers.3.bias",
|
||||
"down_blocks.0.resnets.1.norm1.weight": "input_blocks.2.0.in_layers.0.weight",
|
||||
"down_blocks.0.resnets.1.norm1.bias": "input_blocks.2.0.in_layers.0.bias",
|
||||
"down_blocks.0.resnets.1.conv1.weight": "input_blocks.2.0.in_layers.2.weight",
|
||||
"down_blocks.0.resnets.1.conv1.bias": "input_blocks.2.0.in_layers.2.bias",
|
||||
"down_blocks.0.resnets.1.time_emb_proj.weight": "input_blocks.2.0.emb_layers.1.weight",
|
||||
"down_blocks.0.resnets.1.time_emb_proj.bias": "input_blocks.2.0.emb_layers.1.bias",
|
||||
"down_blocks.0.resnets.1.norm2.weight": "input_blocks.2.0.out_layers.0.weight",
|
||||
"down_blocks.0.resnets.1.norm2.bias": "input_blocks.2.0.out_layers.0.bias",
|
||||
"down_blocks.0.resnets.1.conv2.weight": "input_blocks.2.0.out_layers.3.weight",
|
||||
"down_blocks.0.resnets.1.conv2.bias": "input_blocks.2.0.out_layers.3.bias",
|
||||
"down_blocks.0.downsamplers.0.conv.weight": "input_blocks.3.0.op.weight",
|
||||
"down_blocks.0.downsamplers.0.conv.bias": "input_blocks.3.0.op.bias",
|
||||
"down_blocks.1.attentions.0.norm.weight": "input_blocks.4.1.norm.weight",
|
||||
"down_blocks.1.attentions.0.norm.bias": "input_blocks.4.1.norm.bias",
|
||||
"down_blocks.1.attentions.0.proj_in.weight": "input_blocks.4.1.proj_in.weight",
|
||||
"down_blocks.1.attentions.0.proj_in.bias": "input_blocks.4.1.proj_in.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_q.weight": "input_blocks.4.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k.weight": "input_blocks.4.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_v.weight": "input_blocks.4.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "input_blocks.4.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "input_blocks.4.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "input_blocks.4.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "input_blocks.4.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.ff.net.2.weight": "input_blocks.4.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.ff.net.2.bias": "input_blocks.4.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_q.weight": "input_blocks.4.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_k.weight": "input_blocks.4.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_v.weight": "input_blocks.4.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "input_blocks.4.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "input_blocks.4.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.norm1.weight": "input_blocks.4.1.transformer_blocks.0.norm1.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.norm1.bias": "input_blocks.4.1.transformer_blocks.0.norm1.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.norm2.weight": "input_blocks.4.1.transformer_blocks.0.norm2.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.norm2.bias": "input_blocks.4.1.transformer_blocks.0.norm2.bias",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.norm3.weight": "input_blocks.4.1.transformer_blocks.0.norm3.weight",
|
||||
"down_blocks.1.attentions.0.transformer_blocks.0.norm3.bias": "input_blocks.4.1.transformer_blocks.0.norm3.bias",
|
||||
"down_blocks.1.attentions.0.proj_out.weight": "input_blocks.4.1.proj_out.weight",
|
||||
"down_blocks.1.attentions.0.proj_out.bias": "input_blocks.4.1.proj_out.bias",
|
||||
"down_blocks.1.attentions.1.norm.weight": "input_blocks.5.1.norm.weight",
|
||||
"down_blocks.1.attentions.1.norm.bias": "input_blocks.5.1.norm.bias",
|
||||
"down_blocks.1.attentions.1.proj_in.weight": "input_blocks.5.1.proj_in.weight",
|
||||
"down_blocks.1.attentions.1.proj_in.bias": "input_blocks.5.1.proj_in.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_q.weight": "input_blocks.5.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_k.weight": "input_blocks.5.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_v.weight": "input_blocks.5.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0.weight": "input_blocks.5.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0.bias": "input_blocks.5.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj.weight": "input_blocks.5.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj.bias": "input_blocks.5.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.ff.net.2.weight": "input_blocks.5.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.ff.net.2.bias": "input_blocks.5.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_q.weight": "input_blocks.5.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_k.weight": "input_blocks.5.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_v.weight": "input_blocks.5.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0.weight": "input_blocks.5.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0.bias": "input_blocks.5.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.norm1.weight": "input_blocks.5.1.transformer_blocks.0.norm1.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.norm1.bias": "input_blocks.5.1.transformer_blocks.0.norm1.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.norm2.weight": "input_blocks.5.1.transformer_blocks.0.norm2.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.norm2.bias": "input_blocks.5.1.transformer_blocks.0.norm2.bias",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.norm3.weight": "input_blocks.5.1.transformer_blocks.0.norm3.weight",
|
||||
"down_blocks.1.attentions.1.transformer_blocks.0.norm3.bias": "input_blocks.5.1.transformer_blocks.0.norm3.bias",
|
||||
"down_blocks.1.attentions.1.proj_out.weight": "input_blocks.5.1.proj_out.weight",
|
||||
"down_blocks.1.attentions.1.proj_out.bias": "input_blocks.5.1.proj_out.bias",
|
||||
"down_blocks.1.resnets.0.norm1.weight": "input_blocks.4.0.in_layers.0.weight",
|
||||
"down_blocks.1.resnets.0.norm1.bias": "input_blocks.4.0.in_layers.0.bias",
|
||||
"down_blocks.1.resnets.0.conv1.weight": "input_blocks.4.0.in_layers.2.weight",
|
||||
"down_blocks.1.resnets.0.conv1.bias": "input_blocks.4.0.in_layers.2.bias",
|
||||
"down_blocks.1.resnets.0.time_emb_proj.weight": "input_blocks.4.0.emb_layers.1.weight",
|
||||
"down_blocks.1.resnets.0.time_emb_proj.bias": "input_blocks.4.0.emb_layers.1.bias",
|
||||
"down_blocks.1.resnets.0.norm2.weight": "input_blocks.4.0.out_layers.0.weight",
|
||||
"down_blocks.1.resnets.0.norm2.bias": "input_blocks.4.0.out_layers.0.bias",
|
||||
"down_blocks.1.resnets.0.conv2.weight": "input_blocks.4.0.out_layers.3.weight",
|
||||
"down_blocks.1.resnets.0.conv2.bias": "input_blocks.4.0.out_layers.3.bias",
|
||||
"down_blocks.1.resnets.0.conv_shortcut.weight": "input_blocks.4.0.skip_connection.weight",
|
||||
"down_blocks.1.resnets.0.conv_shortcut.bias": "input_blocks.4.0.skip_connection.bias",
|
||||
"down_blocks.1.resnets.1.norm1.weight": "input_blocks.5.0.in_layers.0.weight",
|
||||
"down_blocks.1.resnets.1.norm1.bias": "input_blocks.5.0.in_layers.0.bias",
|
||||
"down_blocks.1.resnets.1.conv1.weight": "input_blocks.5.0.in_layers.2.weight",
|
||||
"down_blocks.1.resnets.1.conv1.bias": "input_blocks.5.0.in_layers.2.bias",
|
||||
"down_blocks.1.resnets.1.time_emb_proj.weight": "input_blocks.5.0.emb_layers.1.weight",
|
||||
"down_blocks.1.resnets.1.time_emb_proj.bias": "input_blocks.5.0.emb_layers.1.bias",
|
||||
"down_blocks.1.resnets.1.norm2.weight": "input_blocks.5.0.out_layers.0.weight",
|
||||
"down_blocks.1.resnets.1.norm2.bias": "input_blocks.5.0.out_layers.0.bias",
|
||||
"down_blocks.1.resnets.1.conv2.weight": "input_blocks.5.0.out_layers.3.weight",
|
||||
"down_blocks.1.resnets.1.conv2.bias": "input_blocks.5.0.out_layers.3.bias",
|
||||
"down_blocks.1.downsamplers.0.conv.weight": "input_blocks.6.0.op.weight",
|
||||
"down_blocks.1.downsamplers.0.conv.bias": "input_blocks.6.0.op.bias",
|
||||
"down_blocks.2.attentions.0.norm.weight": "input_blocks.7.1.norm.weight",
|
||||
"down_blocks.2.attentions.0.norm.bias": "input_blocks.7.1.norm.bias",
|
||||
"down_blocks.2.attentions.0.proj_in.weight": "input_blocks.7.1.proj_in.weight",
|
||||
"down_blocks.2.attentions.0.proj_in.bias": "input_blocks.7.1.proj_in.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_q.weight": "input_blocks.7.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_k.weight": "input_blocks.7.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_v.weight": "input_blocks.7.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "input_blocks.7.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "input_blocks.7.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "input_blocks.7.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "input_blocks.7.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.ff.net.2.weight": "input_blocks.7.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.ff.net.2.bias": "input_blocks.7.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_q.weight": "input_blocks.7.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_k.weight": "input_blocks.7.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_v.weight": "input_blocks.7.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "input_blocks.7.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "input_blocks.7.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.norm1.weight": "input_blocks.7.1.transformer_blocks.0.norm1.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.norm1.bias": "input_blocks.7.1.transformer_blocks.0.norm1.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.norm2.weight": "input_blocks.7.1.transformer_blocks.0.norm2.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.norm2.bias": "input_blocks.7.1.transformer_blocks.0.norm2.bias",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.norm3.weight": "input_blocks.7.1.transformer_blocks.0.norm3.weight",
|
||||
"down_blocks.2.attentions.0.transformer_blocks.0.norm3.bias": "input_blocks.7.1.transformer_blocks.0.norm3.bias",
|
||||
"down_blocks.2.attentions.0.proj_out.weight": "input_blocks.7.1.proj_out.weight",
|
||||
"down_blocks.2.attentions.0.proj_out.bias": "input_blocks.7.1.proj_out.bias",
|
||||
"down_blocks.2.attentions.1.norm.weight": "input_blocks.8.1.norm.weight",
|
||||
"down_blocks.2.attentions.1.norm.bias": "input_blocks.8.1.norm.bias",
|
||||
"down_blocks.2.attentions.1.proj_in.weight": "input_blocks.8.1.proj_in.weight",
|
||||
"down_blocks.2.attentions.1.proj_in.bias": "input_blocks.8.1.proj_in.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_q.weight": "input_blocks.8.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_k.weight": "input_blocks.8.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_v.weight": "input_blocks.8.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_out.0.weight": "input_blocks.8.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_out.0.bias": "input_blocks.8.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.ff.net.0.proj.weight": "input_blocks.8.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.ff.net.0.proj.bias": "input_blocks.8.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.ff.net.2.weight": "input_blocks.8.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.ff.net.2.bias": "input_blocks.8.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_q.weight": "input_blocks.8.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_k.weight": "input_blocks.8.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_v.weight": "input_blocks.8.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_out.0.weight": "input_blocks.8.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_out.0.bias": "input_blocks.8.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.norm1.weight": "input_blocks.8.1.transformer_blocks.0.norm1.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.norm1.bias": "input_blocks.8.1.transformer_blocks.0.norm1.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.norm2.weight": "input_blocks.8.1.transformer_blocks.0.norm2.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.norm2.bias": "input_blocks.8.1.transformer_blocks.0.norm2.bias",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.norm3.weight": "input_blocks.8.1.transformer_blocks.0.norm3.weight",
|
||||
"down_blocks.2.attentions.1.transformer_blocks.0.norm3.bias": "input_blocks.8.1.transformer_blocks.0.norm3.bias",
|
||||
"down_blocks.2.attentions.1.proj_out.weight": "input_blocks.8.1.proj_out.weight",
|
||||
"down_blocks.2.attentions.1.proj_out.bias": "input_blocks.8.1.proj_out.bias",
|
||||
"down_blocks.2.resnets.0.norm1.weight": "input_blocks.7.0.in_layers.0.weight",
|
||||
"down_blocks.2.resnets.0.norm1.bias": "input_blocks.7.0.in_layers.0.bias",
|
||||
"down_blocks.2.resnets.0.conv1.weight": "input_blocks.7.0.in_layers.2.weight",
|
||||
"down_blocks.2.resnets.0.conv1.bias": "input_blocks.7.0.in_layers.2.bias",
|
||||
"down_blocks.2.resnets.0.time_emb_proj.weight": "input_blocks.7.0.emb_layers.1.weight",
|
||||
"down_blocks.2.resnets.0.time_emb_proj.bias": "input_blocks.7.0.emb_layers.1.bias",
|
||||
"down_blocks.2.resnets.0.norm2.weight": "input_blocks.7.0.out_layers.0.weight",
|
||||
"down_blocks.2.resnets.0.norm2.bias": "input_blocks.7.0.out_layers.0.bias",
|
||||
"down_blocks.2.resnets.0.conv2.weight": "input_blocks.7.0.out_layers.3.weight",
|
||||
"down_blocks.2.resnets.0.conv2.bias": "input_blocks.7.0.out_layers.3.bias",
|
||||
"down_blocks.2.resnets.0.conv_shortcut.weight": "input_blocks.7.0.skip_connection.weight",
|
||||
"down_blocks.2.resnets.0.conv_shortcut.bias": "input_blocks.7.0.skip_connection.bias",
|
||||
"down_blocks.2.resnets.1.norm1.weight": "input_blocks.8.0.in_layers.0.weight",
|
||||
"down_blocks.2.resnets.1.norm1.bias": "input_blocks.8.0.in_layers.0.bias",
|
||||
"down_blocks.2.resnets.1.conv1.weight": "input_blocks.8.0.in_layers.2.weight",
|
||||
"down_blocks.2.resnets.1.conv1.bias": "input_blocks.8.0.in_layers.2.bias",
|
||||
"down_blocks.2.resnets.1.time_emb_proj.weight": "input_blocks.8.0.emb_layers.1.weight",
|
||||
"down_blocks.2.resnets.1.time_emb_proj.bias": "input_blocks.8.0.emb_layers.1.bias",
|
||||
"down_blocks.2.resnets.1.norm2.weight": "input_blocks.8.0.out_layers.0.weight",
|
||||
"down_blocks.2.resnets.1.norm2.bias": "input_blocks.8.0.out_layers.0.bias",
|
||||
"down_blocks.2.resnets.1.conv2.weight": "input_blocks.8.0.out_layers.3.weight",
|
||||
"down_blocks.2.resnets.1.conv2.bias": "input_blocks.8.0.out_layers.3.bias",
|
||||
"down_blocks.2.downsamplers.0.conv.weight": "input_blocks.9.0.op.weight",
|
||||
"down_blocks.2.downsamplers.0.conv.bias": "input_blocks.9.0.op.bias",
|
||||
"down_blocks.3.resnets.0.norm1.weight": "input_blocks.10.0.in_layers.0.weight",
|
||||
"down_blocks.3.resnets.0.norm1.bias": "input_blocks.10.0.in_layers.0.bias",
|
||||
"down_blocks.3.resnets.0.conv1.weight": "input_blocks.10.0.in_layers.2.weight",
|
||||
"down_blocks.3.resnets.0.conv1.bias": "input_blocks.10.0.in_layers.2.bias",
|
||||
"down_blocks.3.resnets.0.time_emb_proj.weight": "input_blocks.10.0.emb_layers.1.weight",
|
||||
"down_blocks.3.resnets.0.time_emb_proj.bias": "input_blocks.10.0.emb_layers.1.bias",
|
||||
"down_blocks.3.resnets.0.norm2.weight": "input_blocks.10.0.out_layers.0.weight",
|
||||
"down_blocks.3.resnets.0.norm2.bias": "input_blocks.10.0.out_layers.0.bias",
|
||||
"down_blocks.3.resnets.0.conv2.weight": "input_blocks.10.0.out_layers.3.weight",
|
||||
"down_blocks.3.resnets.0.conv2.bias": "input_blocks.10.0.out_layers.3.bias",
|
||||
"down_blocks.3.resnets.1.norm1.weight": "input_blocks.11.0.in_layers.0.weight",
|
||||
"down_blocks.3.resnets.1.norm1.bias": "input_blocks.11.0.in_layers.0.bias",
|
||||
"down_blocks.3.resnets.1.conv1.weight": "input_blocks.11.0.in_layers.2.weight",
|
||||
"down_blocks.3.resnets.1.conv1.bias": "input_blocks.11.0.in_layers.2.bias",
|
||||
"down_blocks.3.resnets.1.time_emb_proj.weight": "input_blocks.11.0.emb_layers.1.weight",
|
||||
"down_blocks.3.resnets.1.time_emb_proj.bias": "input_blocks.11.0.emb_layers.1.bias",
|
||||
"down_blocks.3.resnets.1.norm2.weight": "input_blocks.11.0.out_layers.0.weight",
|
||||
"down_blocks.3.resnets.1.norm2.bias": "input_blocks.11.0.out_layers.0.bias",
|
||||
"down_blocks.3.resnets.1.conv2.weight": "input_blocks.11.0.out_layers.3.weight",
|
||||
"down_blocks.3.resnets.1.conv2.bias": "input_blocks.11.0.out_layers.3.bias",
|
||||
"up_blocks.0.resnets.0.norm1.weight": "output_blocks.0.0.in_layers.0.weight",
|
||||
"up_blocks.0.resnets.0.norm1.bias": "output_blocks.0.0.in_layers.0.bias",
|
||||
"up_blocks.0.resnets.0.conv1.weight": "output_blocks.0.0.in_layers.2.weight",
|
||||
"up_blocks.0.resnets.0.conv1.bias": "output_blocks.0.0.in_layers.2.bias",
|
||||
"up_blocks.0.resnets.0.time_emb_proj.weight": "output_blocks.0.0.emb_layers.1.weight",
|
||||
"up_blocks.0.resnets.0.time_emb_proj.bias": "output_blocks.0.0.emb_layers.1.bias",
|
||||
"up_blocks.0.resnets.0.norm2.weight": "output_blocks.0.0.out_layers.0.weight",
|
||||
"up_blocks.0.resnets.0.norm2.bias": "output_blocks.0.0.out_layers.0.bias",
|
||||
"up_blocks.0.resnets.0.conv2.weight": "output_blocks.0.0.out_layers.3.weight",
|
||||
"up_blocks.0.resnets.0.conv2.bias": "output_blocks.0.0.out_layers.3.bias",
|
||||
"up_blocks.0.resnets.0.conv_shortcut.weight": "output_blocks.0.0.skip_connection.weight",
|
||||
"up_blocks.0.resnets.0.conv_shortcut.bias": "output_blocks.0.0.skip_connection.bias",
|
||||
"up_blocks.0.resnets.1.norm1.weight": "output_blocks.1.0.in_layers.0.weight",
|
||||
"up_blocks.0.resnets.1.norm1.bias": "output_blocks.1.0.in_layers.0.bias",
|
||||
"up_blocks.0.resnets.1.conv1.weight": "output_blocks.1.0.in_layers.2.weight",
|
||||
"up_blocks.0.resnets.1.conv1.bias": "output_blocks.1.0.in_layers.2.bias",
|
||||
"up_blocks.0.resnets.1.time_emb_proj.weight": "output_blocks.1.0.emb_layers.1.weight",
|
||||
"up_blocks.0.resnets.1.time_emb_proj.bias": "output_blocks.1.0.emb_layers.1.bias",
|
||||
"up_blocks.0.resnets.1.norm2.weight": "output_blocks.1.0.out_layers.0.weight",
|
||||
"up_blocks.0.resnets.1.norm2.bias": "output_blocks.1.0.out_layers.0.bias",
|
||||
"up_blocks.0.resnets.1.conv2.weight": "output_blocks.1.0.out_layers.3.weight",
|
||||
"up_blocks.0.resnets.1.conv2.bias": "output_blocks.1.0.out_layers.3.bias",
|
||||
"up_blocks.0.resnets.1.conv_shortcut.weight": "output_blocks.1.0.skip_connection.weight",
|
||||
"up_blocks.0.resnets.1.conv_shortcut.bias": "output_blocks.1.0.skip_connection.bias",
|
||||
"up_blocks.0.resnets.2.norm1.weight": "output_blocks.2.0.in_layers.0.weight",
|
||||
"up_blocks.0.resnets.2.norm1.bias": "output_blocks.2.0.in_layers.0.bias",
|
||||
"up_blocks.0.resnets.2.conv1.weight": "output_blocks.2.0.in_layers.2.weight",
|
||||
"up_blocks.0.resnets.2.conv1.bias": "output_blocks.2.0.in_layers.2.bias",
|
||||
"up_blocks.0.resnets.2.time_emb_proj.weight": "output_blocks.2.0.emb_layers.1.weight",
|
||||
"up_blocks.0.resnets.2.time_emb_proj.bias": "output_blocks.2.0.emb_layers.1.bias",
|
||||
"up_blocks.0.resnets.2.norm2.weight": "output_blocks.2.0.out_layers.0.weight",
|
||||
"up_blocks.0.resnets.2.norm2.bias": "output_blocks.2.0.out_layers.0.bias",
|
||||
"up_blocks.0.resnets.2.conv2.weight": "output_blocks.2.0.out_layers.3.weight",
|
||||
"up_blocks.0.resnets.2.conv2.bias": "output_blocks.2.0.out_layers.3.bias",
|
||||
"up_blocks.0.resnets.2.conv_shortcut.weight": "output_blocks.2.0.skip_connection.weight",
|
||||
"up_blocks.0.resnets.2.conv_shortcut.bias": "output_blocks.2.0.skip_connection.bias",
|
||||
"up_blocks.0.upsamplers.0.conv.weight": "output_blocks.2.1.conv.weight",
|
||||
"up_blocks.0.upsamplers.0.conv.bias": "output_blocks.2.1.conv.bias",
|
||||
"up_blocks.1.attentions.0.norm.weight": "output_blocks.3.1.norm.weight",
|
||||
"up_blocks.1.attentions.0.norm.bias": "output_blocks.3.1.norm.bias",
|
||||
"up_blocks.1.attentions.0.proj_in.weight": "output_blocks.3.1.proj_in.weight",
|
||||
"up_blocks.1.attentions.0.proj_in.bias": "output_blocks.3.1.proj_in.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_q.weight": "output_blocks.3.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k.weight": "output_blocks.3.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_v.weight": "output_blocks.3.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.3.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.3.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.3.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.3.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.ff.net.2.weight": "output_blocks.3.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.ff.net.2.bias": "output_blocks.3.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_q.weight": "output_blocks.3.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_k.weight": "output_blocks.3.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_v.weight": "output_blocks.3.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.3.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.3.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.norm1.weight": "output_blocks.3.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.norm1.bias": "output_blocks.3.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.norm2.weight": "output_blocks.3.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.norm2.bias": "output_blocks.3.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.norm3.weight": "output_blocks.3.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.1.attentions.0.transformer_blocks.0.norm3.bias": "output_blocks.3.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.1.attentions.0.proj_out.weight": "output_blocks.3.1.proj_out.weight",
|
||||
"up_blocks.1.attentions.0.proj_out.bias": "output_blocks.3.1.proj_out.bias",
|
||||
"up_blocks.1.attentions.1.norm.weight": "output_blocks.4.1.norm.weight",
|
||||
"up_blocks.1.attentions.1.norm.bias": "output_blocks.4.1.norm.bias",
|
||||
"up_blocks.1.attentions.1.proj_in.weight": "output_blocks.4.1.proj_in.weight",
|
||||
"up_blocks.1.attentions.1.proj_in.bias": "output_blocks.4.1.proj_in.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_q.weight": "output_blocks.4.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_k.weight": "output_blocks.4.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_v.weight": "output_blocks.4.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.4.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.4.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.4.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.4.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.ff.net.2.weight": "output_blocks.4.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.ff.net.2.bias": "output_blocks.4.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_q.weight": "output_blocks.4.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_k.weight": "output_blocks.4.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_v.weight": "output_blocks.4.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.4.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.4.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.norm1.weight": "output_blocks.4.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.norm1.bias": "output_blocks.4.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.norm2.weight": "output_blocks.4.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.norm2.bias": "output_blocks.4.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.norm3.weight": "output_blocks.4.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.1.attentions.1.transformer_blocks.0.norm3.bias": "output_blocks.4.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.1.attentions.1.proj_out.weight": "output_blocks.4.1.proj_out.weight",
|
||||
"up_blocks.1.attentions.1.proj_out.bias": "output_blocks.4.1.proj_out.bias",
|
||||
"up_blocks.1.attentions.2.norm.weight": "output_blocks.5.1.norm.weight",
|
||||
"up_blocks.1.attentions.2.norm.bias": "output_blocks.5.1.norm.bias",
|
||||
"up_blocks.1.attentions.2.proj_in.weight": "output_blocks.5.1.proj_in.weight",
|
||||
"up_blocks.1.attentions.2.proj_in.bias": "output_blocks.5.1.proj_in.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_q.weight": "output_blocks.5.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_k.weight": "output_blocks.5.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_v.weight": "output_blocks.5.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.5.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.5.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.5.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.5.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.ff.net.2.weight": "output_blocks.5.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.ff.net.2.bias": "output_blocks.5.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_q.weight": "output_blocks.5.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_k.weight": "output_blocks.5.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_v.weight": "output_blocks.5.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.5.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.5.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.norm1.weight": "output_blocks.5.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.norm1.bias": "output_blocks.5.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.norm2.weight": "output_blocks.5.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.norm2.bias": "output_blocks.5.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.norm3.weight": "output_blocks.5.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.1.attentions.2.transformer_blocks.0.norm3.bias": "output_blocks.5.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.1.attentions.2.proj_out.weight": "output_blocks.5.1.proj_out.weight",
|
||||
"up_blocks.1.attentions.2.proj_out.bias": "output_blocks.5.1.proj_out.bias",
|
||||
"up_blocks.1.resnets.0.norm1.weight": "output_blocks.3.0.in_layers.0.weight",
|
||||
"up_blocks.1.resnets.0.norm1.bias": "output_blocks.3.0.in_layers.0.bias",
|
||||
"up_blocks.1.resnets.0.conv1.weight": "output_blocks.3.0.in_layers.2.weight",
|
||||
"up_blocks.1.resnets.0.conv1.bias": "output_blocks.3.0.in_layers.2.bias",
|
||||
"up_blocks.1.resnets.0.time_emb_proj.weight": "output_blocks.3.0.emb_layers.1.weight",
|
||||
"up_blocks.1.resnets.0.time_emb_proj.bias": "output_blocks.3.0.emb_layers.1.bias",
|
||||
"up_blocks.1.resnets.0.norm2.weight": "output_blocks.3.0.out_layers.0.weight",
|
||||
"up_blocks.1.resnets.0.norm2.bias": "output_blocks.3.0.out_layers.0.bias",
|
||||
"up_blocks.1.resnets.0.conv2.weight": "output_blocks.3.0.out_layers.3.weight",
|
||||
"up_blocks.1.resnets.0.conv2.bias": "output_blocks.3.0.out_layers.3.bias",
|
||||
"up_blocks.1.resnets.0.conv_shortcut.weight": "output_blocks.3.0.skip_connection.weight",
|
||||
"up_blocks.1.resnets.0.conv_shortcut.bias": "output_blocks.3.0.skip_connection.bias",
|
||||
"up_blocks.1.resnets.1.norm1.weight": "output_blocks.4.0.in_layers.0.weight",
|
||||
"up_blocks.1.resnets.1.norm1.bias": "output_blocks.4.0.in_layers.0.bias",
|
||||
"up_blocks.1.resnets.1.conv1.weight": "output_blocks.4.0.in_layers.2.weight",
|
||||
"up_blocks.1.resnets.1.conv1.bias": "output_blocks.4.0.in_layers.2.bias",
|
||||
"up_blocks.1.resnets.1.time_emb_proj.weight": "output_blocks.4.0.emb_layers.1.weight",
|
||||
"up_blocks.1.resnets.1.time_emb_proj.bias": "output_blocks.4.0.emb_layers.1.bias",
|
||||
"up_blocks.1.resnets.1.norm2.weight": "output_blocks.4.0.out_layers.0.weight",
|
||||
"up_blocks.1.resnets.1.norm2.bias": "output_blocks.4.0.out_layers.0.bias",
|
||||
"up_blocks.1.resnets.1.conv2.weight": "output_blocks.4.0.out_layers.3.weight",
|
||||
"up_blocks.1.resnets.1.conv2.bias": "output_blocks.4.0.out_layers.3.bias",
|
||||
"up_blocks.1.resnets.1.conv_shortcut.weight": "output_blocks.4.0.skip_connection.weight",
|
||||
"up_blocks.1.resnets.1.conv_shortcut.bias": "output_blocks.4.0.skip_connection.bias",
|
||||
"up_blocks.1.resnets.2.norm1.weight": "output_blocks.5.0.in_layers.0.weight",
|
||||
"up_blocks.1.resnets.2.norm1.bias": "output_blocks.5.0.in_layers.0.bias",
|
||||
"up_blocks.1.resnets.2.conv1.weight": "output_blocks.5.0.in_layers.2.weight",
|
||||
"up_blocks.1.resnets.2.conv1.bias": "output_blocks.5.0.in_layers.2.bias",
|
||||
"up_blocks.1.resnets.2.time_emb_proj.weight": "output_blocks.5.0.emb_layers.1.weight",
|
||||
"up_blocks.1.resnets.2.time_emb_proj.bias": "output_blocks.5.0.emb_layers.1.bias",
|
||||
"up_blocks.1.resnets.2.norm2.weight": "output_blocks.5.0.out_layers.0.weight",
|
||||
"up_blocks.1.resnets.2.norm2.bias": "output_blocks.5.0.out_layers.0.bias",
|
||||
"up_blocks.1.resnets.2.conv2.weight": "output_blocks.5.0.out_layers.3.weight",
|
||||
"up_blocks.1.resnets.2.conv2.bias": "output_blocks.5.0.out_layers.3.bias",
|
||||
"up_blocks.1.resnets.2.conv_shortcut.weight": "output_blocks.5.0.skip_connection.weight",
|
||||
"up_blocks.1.resnets.2.conv_shortcut.bias": "output_blocks.5.0.skip_connection.bias",
|
||||
"up_blocks.1.upsamplers.0.conv.weight": "output_blocks.5.2.conv.weight",
|
||||
"up_blocks.1.upsamplers.0.conv.bias": "output_blocks.5.2.conv.bias",
|
||||
"up_blocks.2.attentions.0.norm.weight": "output_blocks.6.1.norm.weight",
|
||||
"up_blocks.2.attentions.0.norm.bias": "output_blocks.6.1.norm.bias",
|
||||
"up_blocks.2.attentions.0.proj_in.weight": "output_blocks.6.1.proj_in.weight",
|
||||
"up_blocks.2.attentions.0.proj_in.bias": "output_blocks.6.1.proj_in.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn1.to_q.weight": "output_blocks.6.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn1.to_k.weight": "output_blocks.6.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn1.to_v.weight": "output_blocks.6.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.6.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.6.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.6.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.6.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.ff.net.2.weight": "output_blocks.6.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.ff.net.2.bias": "output_blocks.6.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn2.to_q.weight": "output_blocks.6.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn2.to_k.weight": "output_blocks.6.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn2.to_v.weight": "output_blocks.6.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.6.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.6.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.norm1.weight": "output_blocks.6.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.norm1.bias": "output_blocks.6.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.norm2.weight": "output_blocks.6.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.norm2.bias": "output_blocks.6.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.norm3.weight": "output_blocks.6.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.2.attentions.0.transformer_blocks.0.norm3.bias": "output_blocks.6.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.2.attentions.0.proj_out.weight": "output_blocks.6.1.proj_out.weight",
|
||||
"up_blocks.2.attentions.0.proj_out.bias": "output_blocks.6.1.proj_out.bias",
|
||||
"up_blocks.2.attentions.1.norm.weight": "output_blocks.7.1.norm.weight",
|
||||
"up_blocks.2.attentions.1.norm.bias": "output_blocks.7.1.norm.bias",
|
||||
"up_blocks.2.attentions.1.proj_in.weight": "output_blocks.7.1.proj_in.weight",
|
||||
"up_blocks.2.attentions.1.proj_in.bias": "output_blocks.7.1.proj_in.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn1.to_q.weight": "output_blocks.7.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn1.to_k.weight": "output_blocks.7.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn1.to_v.weight": "output_blocks.7.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.7.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.7.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.7.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.7.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.ff.net.2.weight": "output_blocks.7.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.ff.net.2.bias": "output_blocks.7.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn2.to_q.weight": "output_blocks.7.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn2.to_k.weight": "output_blocks.7.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn2.to_v.weight": "output_blocks.7.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.7.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.7.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.norm1.weight": "output_blocks.7.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.norm1.bias": "output_blocks.7.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.norm2.weight": "output_blocks.7.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.norm2.bias": "output_blocks.7.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.norm3.weight": "output_blocks.7.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.2.attentions.1.transformer_blocks.0.norm3.bias": "output_blocks.7.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.2.attentions.1.proj_out.weight": "output_blocks.7.1.proj_out.weight",
|
||||
"up_blocks.2.attentions.1.proj_out.bias": "output_blocks.7.1.proj_out.bias",
|
||||
"up_blocks.2.attentions.2.norm.weight": "output_blocks.8.1.norm.weight",
|
||||
"up_blocks.2.attentions.2.norm.bias": "output_blocks.8.1.norm.bias",
|
||||
"up_blocks.2.attentions.2.proj_in.weight": "output_blocks.8.1.proj_in.weight",
|
||||
"up_blocks.2.attentions.2.proj_in.bias": "output_blocks.8.1.proj_in.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn1.to_q.weight": "output_blocks.8.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn1.to_k.weight": "output_blocks.8.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn1.to_v.weight": "output_blocks.8.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.8.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.8.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.8.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.8.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.ff.net.2.weight": "output_blocks.8.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.ff.net.2.bias": "output_blocks.8.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn2.to_q.weight": "output_blocks.8.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn2.to_k.weight": "output_blocks.8.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn2.to_v.weight": "output_blocks.8.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.8.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.8.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.norm1.weight": "output_blocks.8.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.norm1.bias": "output_blocks.8.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.norm2.weight": "output_blocks.8.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.norm2.bias": "output_blocks.8.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.norm3.weight": "output_blocks.8.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.2.attentions.2.transformer_blocks.0.norm3.bias": "output_blocks.8.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.2.attentions.2.proj_out.weight": "output_blocks.8.1.proj_out.weight",
|
||||
"up_blocks.2.attentions.2.proj_out.bias": "output_blocks.8.1.proj_out.bias",
|
||||
"up_blocks.2.resnets.0.norm1.weight": "output_blocks.6.0.in_layers.0.weight",
|
||||
"up_blocks.2.resnets.0.norm1.bias": "output_blocks.6.0.in_layers.0.bias",
|
||||
"up_blocks.2.resnets.0.conv1.weight": "output_blocks.6.0.in_layers.2.weight",
|
||||
"up_blocks.2.resnets.0.conv1.bias": "output_blocks.6.0.in_layers.2.bias",
|
||||
"up_blocks.2.resnets.0.time_emb_proj.weight": "output_blocks.6.0.emb_layers.1.weight",
|
||||
"up_blocks.2.resnets.0.time_emb_proj.bias": "output_blocks.6.0.emb_layers.1.bias",
|
||||
"up_blocks.2.resnets.0.norm2.weight": "output_blocks.6.0.out_layers.0.weight",
|
||||
"up_blocks.2.resnets.0.norm2.bias": "output_blocks.6.0.out_layers.0.bias",
|
||||
"up_blocks.2.resnets.0.conv2.weight": "output_blocks.6.0.out_layers.3.weight",
|
||||
"up_blocks.2.resnets.0.conv2.bias": "output_blocks.6.0.out_layers.3.bias",
|
||||
"up_blocks.2.resnets.0.conv_shortcut.weight": "output_blocks.6.0.skip_connection.weight",
|
||||
"up_blocks.2.resnets.0.conv_shortcut.bias": "output_blocks.6.0.skip_connection.bias",
|
||||
"up_blocks.2.resnets.1.norm1.weight": "output_blocks.7.0.in_layers.0.weight",
|
||||
"up_blocks.2.resnets.1.norm1.bias": "output_blocks.7.0.in_layers.0.bias",
|
||||
"up_blocks.2.resnets.1.conv1.weight": "output_blocks.7.0.in_layers.2.weight",
|
||||
"up_blocks.2.resnets.1.conv1.bias": "output_blocks.7.0.in_layers.2.bias",
|
||||
"up_blocks.2.resnets.1.time_emb_proj.weight": "output_blocks.7.0.emb_layers.1.weight",
|
||||
"up_blocks.2.resnets.1.time_emb_proj.bias": "output_blocks.7.0.emb_layers.1.bias",
|
||||
"up_blocks.2.resnets.1.norm2.weight": "output_blocks.7.0.out_layers.0.weight",
|
||||
"up_blocks.2.resnets.1.norm2.bias": "output_blocks.7.0.out_layers.0.bias",
|
||||
"up_blocks.2.resnets.1.conv2.weight": "output_blocks.7.0.out_layers.3.weight",
|
||||
"up_blocks.2.resnets.1.conv2.bias": "output_blocks.7.0.out_layers.3.bias",
|
||||
"up_blocks.2.resnets.1.conv_shortcut.weight": "output_blocks.7.0.skip_connection.weight",
|
||||
"up_blocks.2.resnets.1.conv_shortcut.bias": "output_blocks.7.0.skip_connection.bias",
|
||||
"up_blocks.2.resnets.2.norm1.weight": "output_blocks.8.0.in_layers.0.weight",
|
||||
"up_blocks.2.resnets.2.norm1.bias": "output_blocks.8.0.in_layers.0.bias",
|
||||
"up_blocks.2.resnets.2.conv1.weight": "output_blocks.8.0.in_layers.2.weight",
|
||||
"up_blocks.2.resnets.2.conv1.bias": "output_blocks.8.0.in_layers.2.bias",
|
||||
"up_blocks.2.resnets.2.time_emb_proj.weight": "output_blocks.8.0.emb_layers.1.weight",
|
||||
"up_blocks.2.resnets.2.time_emb_proj.bias": "output_blocks.8.0.emb_layers.1.bias",
|
||||
"up_blocks.2.resnets.2.norm2.weight": "output_blocks.8.0.out_layers.0.weight",
|
||||
"up_blocks.2.resnets.2.norm2.bias": "output_blocks.8.0.out_layers.0.bias",
|
||||
"up_blocks.2.resnets.2.conv2.weight": "output_blocks.8.0.out_layers.3.weight",
|
||||
"up_blocks.2.resnets.2.conv2.bias": "output_blocks.8.0.out_layers.3.bias",
|
||||
"up_blocks.2.resnets.2.conv_shortcut.weight": "output_blocks.8.0.skip_connection.weight",
|
||||
"up_blocks.2.resnets.2.conv_shortcut.bias": "output_blocks.8.0.skip_connection.bias",
|
||||
"up_blocks.2.upsamplers.0.conv.weight": "output_blocks.8.2.conv.weight",
|
||||
"up_blocks.2.upsamplers.0.conv.bias": "output_blocks.8.2.conv.bias",
|
||||
"up_blocks.3.attentions.0.norm.weight": "output_blocks.9.1.norm.weight",
|
||||
"up_blocks.3.attentions.0.norm.bias": "output_blocks.9.1.norm.bias",
|
||||
"up_blocks.3.attentions.0.proj_in.weight": "output_blocks.9.1.proj_in.weight",
|
||||
"up_blocks.3.attentions.0.proj_in.bias": "output_blocks.9.1.proj_in.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn1.to_q.weight": "output_blocks.9.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn1.to_k.weight": "output_blocks.9.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn1.to_v.weight": "output_blocks.9.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.9.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.9.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.9.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.9.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.ff.net.2.weight": "output_blocks.9.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.ff.net.2.bias": "output_blocks.9.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn2.to_q.weight": "output_blocks.9.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn2.to_k.weight": "output_blocks.9.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn2.to_v.weight": "output_blocks.9.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.9.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.9.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.norm1.weight": "output_blocks.9.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.norm1.bias": "output_blocks.9.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.norm2.weight": "output_blocks.9.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.norm2.bias": "output_blocks.9.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.norm3.weight": "output_blocks.9.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.3.attentions.0.transformer_blocks.0.norm3.bias": "output_blocks.9.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.3.attentions.0.proj_out.weight": "output_blocks.9.1.proj_out.weight",
|
||||
"up_blocks.3.attentions.0.proj_out.bias": "output_blocks.9.1.proj_out.bias",
|
||||
"up_blocks.3.attentions.1.norm.weight": "output_blocks.10.1.norm.weight",
|
||||
"up_blocks.3.attentions.1.norm.bias": "output_blocks.10.1.norm.bias",
|
||||
"up_blocks.3.attentions.1.proj_in.weight": "output_blocks.10.1.proj_in.weight",
|
||||
"up_blocks.3.attentions.1.proj_in.bias": "output_blocks.10.1.proj_in.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn1.to_q.weight": "output_blocks.10.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn1.to_k.weight": "output_blocks.10.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn1.to_v.weight": "output_blocks.10.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.10.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.10.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.10.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.10.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.ff.net.2.weight": "output_blocks.10.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.ff.net.2.bias": "output_blocks.10.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn2.to_q.weight": "output_blocks.10.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn2.to_k.weight": "output_blocks.10.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn2.to_v.weight": "output_blocks.10.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.10.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.10.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.norm1.weight": "output_blocks.10.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.norm1.bias": "output_blocks.10.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.norm2.weight": "output_blocks.10.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.norm2.bias": "output_blocks.10.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.norm3.weight": "output_blocks.10.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.3.attentions.1.transformer_blocks.0.norm3.bias": "output_blocks.10.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.3.attentions.1.proj_out.weight": "output_blocks.10.1.proj_out.weight",
|
||||
"up_blocks.3.attentions.1.proj_out.bias": "output_blocks.10.1.proj_out.bias",
|
||||
"up_blocks.3.attentions.2.norm.weight": "output_blocks.11.1.norm.weight",
|
||||
"up_blocks.3.attentions.2.norm.bias": "output_blocks.11.1.norm.bias",
|
||||
"up_blocks.3.attentions.2.proj_in.weight": "output_blocks.11.1.proj_in.weight",
|
||||
"up_blocks.3.attentions.2.proj_in.bias": "output_blocks.11.1.proj_in.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn1.to_q.weight": "output_blocks.11.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn1.to_k.weight": "output_blocks.11.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn1.to_v.weight": "output_blocks.11.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn1.to_out.0.weight": "output_blocks.11.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn1.to_out.0.bias": "output_blocks.11.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.ff.net.0.proj.weight": "output_blocks.11.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.ff.net.0.proj.bias": "output_blocks.11.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.ff.net.2.weight": "output_blocks.11.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.ff.net.2.bias": "output_blocks.11.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn2.to_q.weight": "output_blocks.11.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn2.to_k.weight": "output_blocks.11.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn2.to_v.weight": "output_blocks.11.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn2.to_out.0.weight": "output_blocks.11.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.attn2.to_out.0.bias": "output_blocks.11.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.norm1.weight": "output_blocks.11.1.transformer_blocks.0.norm1.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.norm1.bias": "output_blocks.11.1.transformer_blocks.0.norm1.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.norm2.weight": "output_blocks.11.1.transformer_blocks.0.norm2.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.norm2.bias": "output_blocks.11.1.transformer_blocks.0.norm2.bias",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.norm3.weight": "output_blocks.11.1.transformer_blocks.0.norm3.weight",
|
||||
"up_blocks.3.attentions.2.transformer_blocks.0.norm3.bias": "output_blocks.11.1.transformer_blocks.0.norm3.bias",
|
||||
"up_blocks.3.attentions.2.proj_out.weight": "output_blocks.11.1.proj_out.weight",
|
||||
"up_blocks.3.attentions.2.proj_out.bias": "output_blocks.11.1.proj_out.bias",
|
||||
"up_blocks.3.resnets.0.norm1.weight": "output_blocks.9.0.in_layers.0.weight",
|
||||
"up_blocks.3.resnets.0.norm1.bias": "output_blocks.9.0.in_layers.0.bias",
|
||||
"up_blocks.3.resnets.0.conv1.weight": "output_blocks.9.0.in_layers.2.weight",
|
||||
"up_blocks.3.resnets.0.conv1.bias": "output_blocks.9.0.in_layers.2.bias",
|
||||
"up_blocks.3.resnets.0.time_emb_proj.weight": "output_blocks.9.0.emb_layers.1.weight",
|
||||
"up_blocks.3.resnets.0.time_emb_proj.bias": "output_blocks.9.0.emb_layers.1.bias",
|
||||
"up_blocks.3.resnets.0.norm2.weight": "output_blocks.9.0.out_layers.0.weight",
|
||||
"up_blocks.3.resnets.0.norm2.bias": "output_blocks.9.0.out_layers.0.bias",
|
||||
"up_blocks.3.resnets.0.conv2.weight": "output_blocks.9.0.out_layers.3.weight",
|
||||
"up_blocks.3.resnets.0.conv2.bias": "output_blocks.9.0.out_layers.3.bias",
|
||||
"up_blocks.3.resnets.0.conv_shortcut.weight": "output_blocks.9.0.skip_connection.weight",
|
||||
"up_blocks.3.resnets.0.conv_shortcut.bias": "output_blocks.9.0.skip_connection.bias",
|
||||
"up_blocks.3.resnets.1.norm1.weight": "output_blocks.10.0.in_layers.0.weight",
|
||||
"up_blocks.3.resnets.1.norm1.bias": "output_blocks.10.0.in_layers.0.bias",
|
||||
"up_blocks.3.resnets.1.conv1.weight": "output_blocks.10.0.in_layers.2.weight",
|
||||
"up_blocks.3.resnets.1.conv1.bias": "output_blocks.10.0.in_layers.2.bias",
|
||||
"up_blocks.3.resnets.1.time_emb_proj.weight": "output_blocks.10.0.emb_layers.1.weight",
|
||||
"up_blocks.3.resnets.1.time_emb_proj.bias": "output_blocks.10.0.emb_layers.1.bias",
|
||||
"up_blocks.3.resnets.1.norm2.weight": "output_blocks.10.0.out_layers.0.weight",
|
||||
"up_blocks.3.resnets.1.norm2.bias": "output_blocks.10.0.out_layers.0.bias",
|
||||
"up_blocks.3.resnets.1.conv2.weight": "output_blocks.10.0.out_layers.3.weight",
|
||||
"up_blocks.3.resnets.1.conv2.bias": "output_blocks.10.0.out_layers.3.bias",
|
||||
"up_blocks.3.resnets.1.conv_shortcut.weight": "output_blocks.10.0.skip_connection.weight",
|
||||
"up_blocks.3.resnets.1.conv_shortcut.bias": "output_blocks.10.0.skip_connection.bias",
|
||||
"up_blocks.3.resnets.2.norm1.weight": "output_blocks.11.0.in_layers.0.weight",
|
||||
"up_blocks.3.resnets.2.norm1.bias": "output_blocks.11.0.in_layers.0.bias",
|
||||
"up_blocks.3.resnets.2.conv1.weight": "output_blocks.11.0.in_layers.2.weight",
|
||||
"up_blocks.3.resnets.2.conv1.bias": "output_blocks.11.0.in_layers.2.bias",
|
||||
"up_blocks.3.resnets.2.time_emb_proj.weight": "output_blocks.11.0.emb_layers.1.weight",
|
||||
"up_blocks.3.resnets.2.time_emb_proj.bias": "output_blocks.11.0.emb_layers.1.bias",
|
||||
"up_blocks.3.resnets.2.norm2.weight": "output_blocks.11.0.out_layers.0.weight",
|
||||
"up_blocks.3.resnets.2.norm2.bias": "output_blocks.11.0.out_layers.0.bias",
|
||||
"up_blocks.3.resnets.2.conv2.weight": "output_blocks.11.0.out_layers.3.weight",
|
||||
"up_blocks.3.resnets.2.conv2.bias": "output_blocks.11.0.out_layers.3.bias",
|
||||
"up_blocks.3.resnets.2.conv_shortcut.weight": "output_blocks.11.0.skip_connection.weight",
|
||||
"up_blocks.3.resnets.2.conv_shortcut.bias": "output_blocks.11.0.skip_connection.bias",
|
||||
"mid_block.attentions.0.norm.weight": "middle_block.1.norm.weight",
|
||||
"mid_block.attentions.0.norm.bias": "middle_block.1.norm.bias",
|
||||
"mid_block.attentions.0.proj_in.weight": "middle_block.1.proj_in.weight",
|
||||
"mid_block.attentions.0.proj_in.bias": "middle_block.1.proj_in.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn1.to_q.weight": "middle_block.1.transformer_blocks.0.attn1.to_q.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn1.to_k.weight": "middle_block.1.transformer_blocks.0.attn1.to_k.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn1.to_v.weight": "middle_block.1.transformer_blocks.0.attn1.to_v.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn1.to_out.0.weight": "middle_block.1.transformer_blocks.0.attn1.to_out.0.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn1.to_out.0.bias": "middle_block.1.transformer_blocks.0.attn1.to_out.0.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.ff.net.0.proj.weight": "middle_block.1.transformer_blocks.0.ff.net.0.proj.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.ff.net.0.proj.bias": "middle_block.1.transformer_blocks.0.ff.net.0.proj.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.ff.net.2.weight": "middle_block.1.transformer_blocks.0.ff.net.2.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.ff.net.2.bias": "middle_block.1.transformer_blocks.0.ff.net.2.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn2.to_q.weight": "middle_block.1.transformer_blocks.0.attn2.to_q.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn2.to_k.weight": "middle_block.1.transformer_blocks.0.attn2.to_k.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn2.to_v.weight": "middle_block.1.transformer_blocks.0.attn2.to_v.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn2.to_out.0.weight": "middle_block.1.transformer_blocks.0.attn2.to_out.0.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.attn2.to_out.0.bias": "middle_block.1.transformer_blocks.0.attn2.to_out.0.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.norm1.weight": "middle_block.1.transformer_blocks.0.norm1.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.norm1.bias": "middle_block.1.transformer_blocks.0.norm1.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.norm2.weight": "middle_block.1.transformer_blocks.0.norm2.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.norm2.bias": "middle_block.1.transformer_blocks.0.norm2.bias",
|
||||
"mid_block.attentions.0.transformer_blocks.0.norm3.weight": "middle_block.1.transformer_blocks.0.norm3.weight",
|
||||
"mid_block.attentions.0.transformer_blocks.0.norm3.bias": "middle_block.1.transformer_blocks.0.norm3.bias",
|
||||
"mid_block.attentions.0.proj_out.weight": "middle_block.1.proj_out.weight",
|
||||
"mid_block.attentions.0.proj_out.bias": "middle_block.1.proj_out.bias",
|
||||
"mid_block.resnets.0.norm1.weight": "middle_block.0.in_layers.0.weight",
|
||||
"mid_block.resnets.0.norm1.bias": "middle_block.0.in_layers.0.bias",
|
||||
"mid_block.resnets.0.conv1.weight": "middle_block.0.in_layers.2.weight",
|
||||
"mid_block.resnets.0.conv1.bias": "middle_block.0.in_layers.2.bias",
|
||||
"mid_block.resnets.0.time_emb_proj.weight": "middle_block.0.emb_layers.1.weight",
|
||||
"mid_block.resnets.0.time_emb_proj.bias": "middle_block.0.emb_layers.1.bias",
|
||||
"mid_block.resnets.0.norm2.weight": "middle_block.0.out_layers.0.weight",
|
||||
"mid_block.resnets.0.norm2.bias": "middle_block.0.out_layers.0.bias",
|
||||
"mid_block.resnets.0.conv2.weight": "middle_block.0.out_layers.3.weight",
|
||||
"mid_block.resnets.0.conv2.bias": "middle_block.0.out_layers.3.bias",
|
||||
"mid_block.resnets.1.norm1.weight": "middle_block.2.in_layers.0.weight",
|
||||
"mid_block.resnets.1.norm1.bias": "middle_block.2.in_layers.0.bias",
|
||||
"mid_block.resnets.1.conv1.weight": "middle_block.2.in_layers.2.weight",
|
||||
"mid_block.resnets.1.conv1.bias": "middle_block.2.in_layers.2.bias",
|
||||
"mid_block.resnets.1.time_emb_proj.weight": "middle_block.2.emb_layers.1.weight",
|
||||
"mid_block.resnets.1.time_emb_proj.bias": "middle_block.2.emb_layers.1.bias",
|
||||
"mid_block.resnets.1.norm2.weight": "middle_block.2.out_layers.0.weight",
|
||||
"mid_block.resnets.1.norm2.bias": "middle_block.2.out_layers.0.bias",
|
||||
"mid_block.resnets.1.conv2.weight": "middle_block.2.out_layers.3.weight",
|
||||
"mid_block.resnets.1.conv2.bias": "middle_block.2.out_layers.3.bias",
|
||||
"conv_norm_out.weight": "out.0.weight",
|
||||
"conv_norm_out.bias": "out.0.bias",
|
||||
"conv_out.weight": "out.2.weight",
|
||||
"conv_out.bias": "out.2.bias"
|
||||
}
|
||||
@@ -0,0 +1,318 @@
|
||||
import argparse
|
||||
import os
|
||||
import platform
|
||||
import struct
|
||||
import subprocess
|
||||
import time
|
||||
from typing import List
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch.multiprocessing as mp
|
||||
from numba import njit
|
||||
|
||||
import blender.histogram_blend as histogram_blend
|
||||
from blender.guide import (BaseGuide, ColorGuide, EdgeGuide, PositionalGuide,
|
||||
TemporalGuide)
|
||||
from blender.poisson_fusion import poisson_fusion
|
||||
from blender.video_sequence import VideoSequence
|
||||
from flow.flow_utils import flow_calc
|
||||
from src.video_util import frame_to_video
|
||||
|
||||
OPEN_EBSYNTH_LOG = False
|
||||
MAX_PROCESS = 8
|
||||
|
||||
os_str = platform.system()
|
||||
|
||||
if os_str == 'Windows':
|
||||
ebsynth_bin = '.\\deps\\ebsynth\\bin\\ebsynth.exe'
|
||||
elif os_str == 'Linux':
|
||||
ebsynth_bin = './deps/ebsynth/bin/ebsynth'
|
||||
elif os_str == 'Darwin':
|
||||
ebsynth_bin = './deps/ebsynth/bin/ebsynth.app'
|
||||
else:
|
||||
print('Cannot recognize OS. Run Ebsynth failed.')
|
||||
exit(0)
|
||||
|
||||
|
||||
@njit
|
||||
def g_error_mask_loop(H, W, dist1, dist2, output, weight1, weight2):
|
||||
for i in range(H):
|
||||
for j in range(W):
|
||||
if weight1 * dist1[i, j] < weight2 * dist2[i, j]:
|
||||
output[i, j] = 0
|
||||
else:
|
||||
output[i, j] = 1
|
||||
if weight1 == 0:
|
||||
output[i, j] = 0
|
||||
elif weight2 == 0:
|
||||
output[i, j] = 1
|
||||
|
||||
|
||||
def g_error_mask(dist1, dist2, weight1=1, weight2=1):
|
||||
H, W = dist1.shape
|
||||
output = np.empty_like(dist1, dtype=np.byte)
|
||||
g_error_mask_loop(H, W, dist1, dist2, output, weight1, weight2)
|
||||
return output
|
||||
|
||||
|
||||
def create_sequence(base_dir, beg, end, interval, key_dir):
|
||||
sequence = VideoSequence(base_dir, beg, end, interval, 'video', key_dir,
|
||||
'tmp', '%04d.png', '%04d.png')
|
||||
return sequence
|
||||
|
||||
|
||||
def process_one_sequence(i, video_sequence: VideoSequence):
|
||||
interval = video_sequence.interval
|
||||
for is_forward in [True, False]:
|
||||
input_seq = video_sequence.get_input_sequence(i, is_forward)
|
||||
output_seq = video_sequence.get_output_sequence(i, is_forward)
|
||||
flow_seq = video_sequence.get_flow_sequence(i, is_forward)
|
||||
key_img_id = i if is_forward else i + 1
|
||||
key_img = video_sequence.get_key_img(key_img_id)
|
||||
for j in range(interval - 1):
|
||||
i1 = cv2.imread(input_seq[j])
|
||||
i2 = cv2.imread(input_seq[j + 1])
|
||||
flow_calc.get_flow(i1, i2, flow_seq[j])
|
||||
|
||||
guides: List[BaseGuide] = [
|
||||
ColorGuide(input_seq),
|
||||
EdgeGuide(input_seq,
|
||||
video_sequence.get_edge_sequence(i, is_forward)),
|
||||
TemporalGuide(key_img, output_seq, flow_seq,
|
||||
video_sequence.get_temporal_sequence(i, is_forward)),
|
||||
PositionalGuide(flow_seq,
|
||||
video_sequence.get_pos_sequence(i, is_forward))
|
||||
]
|
||||
weights = [6, 0.5, 0.5, 2]
|
||||
for j in range(interval):
|
||||
# key frame
|
||||
if j == 0:
|
||||
img = cv2.imread(key_img)
|
||||
cv2.imwrite(output_seq[0], img)
|
||||
else:
|
||||
cmd = f'{ebsynth_bin} -style {os.path.abspath(key_img)}'
|
||||
for g, w in zip(guides, weights):
|
||||
cmd += ' ' + g.get_cmd(j, w)
|
||||
|
||||
cmd += (f' -output {os.path.abspath(output_seq[j])}'
|
||||
' -searchvoteiters 12 -patchmatchiters 6')
|
||||
if OPEN_EBSYNTH_LOG:
|
||||
print(cmd)
|
||||
subprocess.run(cmd,
|
||||
shell=True,
|
||||
capture_output=not OPEN_EBSYNTH_LOG)
|
||||
|
||||
|
||||
def process_sequences(i_arr, video_sequence: VideoSequence):
|
||||
for i in i_arr:
|
||||
process_one_sequence(i, video_sequence)
|
||||
|
||||
|
||||
def run_ebsynth(video_sequence: VideoSequence):
|
||||
|
||||
beg = time.time()
|
||||
|
||||
processes = []
|
||||
mp.set_start_method('spawn')
|
||||
|
||||
n_process = min(MAX_PROCESS, video_sequence.n_seq)
|
||||
cnt = video_sequence.n_seq // n_process
|
||||
remainder = video_sequence.n_seq % n_process
|
||||
|
||||
prev_idx = 0
|
||||
|
||||
for i in range(n_process):
|
||||
task_cnt = cnt + 1 if i < remainder else cnt
|
||||
i_arr = list(range(prev_idx, prev_idx + task_cnt))
|
||||
prev_idx += task_cnt
|
||||
p = mp.Process(target=process_sequences, args=(i_arr, video_sequence))
|
||||
p.start()
|
||||
processes.append(p)
|
||||
for p in processes:
|
||||
p.join()
|
||||
|
||||
end = time.time()
|
||||
|
||||
print(f'ebsynth: {end-beg}')
|
||||
|
||||
|
||||
@njit
|
||||
def assemble_min_error_img_loop(H, W, a, b, error_mask, out):
|
||||
for i in range(H):
|
||||
for j in range(W):
|
||||
if error_mask[i, j] == 0:
|
||||
out[i, j] = a[i, j]
|
||||
else:
|
||||
out[i, j] = b[i, j]
|
||||
|
||||
|
||||
def assemble_min_error_img(a, b, error_mask):
|
||||
H, W = a.shape[0:2]
|
||||
out = np.empty_like(a)
|
||||
assemble_min_error_img_loop(H, W, a, b, error_mask, out)
|
||||
return out
|
||||
|
||||
|
||||
def load_error(bin_path, img_shape):
|
||||
img_size = img_shape[0] * img_shape[1]
|
||||
with open(bin_path, 'rb') as fp:
|
||||
bytes = fp.read()
|
||||
|
||||
read_size = struct.unpack('q', bytes[:8])
|
||||
assert read_size[0] == img_size
|
||||
float_res = struct.unpack('f' * img_size, bytes[8:])
|
||||
res = np.array(float_res,
|
||||
dtype=np.float32).reshape(img_shape[0], img_shape[1])
|
||||
return res
|
||||
|
||||
|
||||
def process_seq(video_sequence: VideoSequence,
|
||||
i,
|
||||
blend_histogram=True,
|
||||
blend_gradient=True):
|
||||
|
||||
key1_img = cv2.imread(video_sequence.get_key_img(i))
|
||||
img_shape = key1_img.shape
|
||||
interval = video_sequence.interval
|
||||
beg_id = video_sequence.get_sequence_beg_id(i)
|
||||
|
||||
oas = video_sequence.get_output_sequence(i)
|
||||
obs = video_sequence.get_output_sequence(i, False)
|
||||
|
||||
binas = [x.replace('jpg', 'bin') for x in oas]
|
||||
binbs = [x.replace('jpg', 'bin') for x in obs]
|
||||
|
||||
obs = [obs[0]] + list(reversed(obs[1:]))
|
||||
inputs = video_sequence.get_input_sequence(i)
|
||||
oas = [cv2.imread(x) for x in oas]
|
||||
obs = [cv2.imread(x) for x in obs]
|
||||
inputs = [cv2.imread(x) for x in inputs]
|
||||
flow_seq = video_sequence.get_flow_sequence(i)
|
||||
|
||||
dist1s = []
|
||||
dist2s = []
|
||||
for i in range(interval - 1):
|
||||
bin_a = binas[i + 1]
|
||||
bin_b = binbs[i + 1]
|
||||
dist1s.append(load_error(bin_a, img_shape))
|
||||
dist2s.append(load_error(bin_b, img_shape))
|
||||
|
||||
lb = 0
|
||||
ub = 1
|
||||
beg = time.time()
|
||||
p_mask = None
|
||||
|
||||
# write key img
|
||||
blend_out_path = video_sequence.get_blending_img(beg_id)
|
||||
cv2.imwrite(blend_out_path, key1_img)
|
||||
|
||||
for i in range(interval - 1):
|
||||
c_id = beg_id + i + 1
|
||||
blend_out_path = video_sequence.get_blending_img(c_id)
|
||||
|
||||
dist1 = dist1s[i]
|
||||
dist2 = dist2s[i]
|
||||
oa = oas[i + 1]
|
||||
ob = obs[i + 1]
|
||||
weight1 = i / (interval - 1) * (ub - lb) + lb
|
||||
weight2 = 1 - weight1
|
||||
mask = g_error_mask(dist1, dist2, weight1, weight2)
|
||||
if p_mask is not None:
|
||||
flow_path = flow_seq[i]
|
||||
flow = flow_calc.get_flow(inputs[i], inputs[i + 1], flow_path)
|
||||
p_mask = flow_calc.warp(p_mask, flow, 'nearest')
|
||||
mask = p_mask | mask
|
||||
p_mask = mask
|
||||
|
||||
# Save tmp mask
|
||||
# out_mask = np.expand_dims(mask, 2)
|
||||
# cv2.imwrite(f'mask/mask_{c_id:04d}.jpg', out_mask * 255)
|
||||
|
||||
min_error_img = assemble_min_error_img(oa, ob, mask)
|
||||
if blend_histogram:
|
||||
hb_res = histogram_blend.blend(oa, ob, min_error_img,
|
||||
(1 - weight1), (1 - weight2))
|
||||
|
||||
else:
|
||||
# hb_res = min_error_img
|
||||
tmpa = oa.astype(np.float32)
|
||||
tmpb = ob.astype(np.float32)
|
||||
hb_res = (1 - weight1) * tmpa + (1 - weight2) * tmpb
|
||||
|
||||
# cv2.imwrite(blend_out_path, hb_res)
|
||||
|
||||
# gradient blend
|
||||
if blend_gradient:
|
||||
res = poisson_fusion(hb_res, oa, ob, mask)
|
||||
else:
|
||||
res = hb_res
|
||||
|
||||
cv2.imwrite(blend_out_path, res)
|
||||
end = time.time()
|
||||
print('others:', end - beg)
|
||||
|
||||
|
||||
def main(args):
|
||||
global MAX_PROCESS
|
||||
MAX_PROCESS = args.n_proc
|
||||
|
||||
video_sequence = create_sequence(f'{args.name}', args.beg, args.end,
|
||||
args.itv, args.key)
|
||||
if not args.ne:
|
||||
run_ebsynth(video_sequence)
|
||||
blend_histogram = True
|
||||
blend_gradient = args.ps
|
||||
for i in range(video_sequence.n_seq):
|
||||
process_seq(video_sequence, i, blend_histogram, blend_gradient)
|
||||
if args.output:
|
||||
frame_to_video(args.output, video_sequence.blending_dir, args.fps,
|
||||
False)
|
||||
if not args.tmp:
|
||||
video_sequence.remove_out_and_tmp()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('name', type=str, help='Path to input video')
|
||||
parser.add_argument('--output',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Path to output video')
|
||||
parser.add_argument('--fps',
|
||||
type=float,
|
||||
default=30,
|
||||
help='The FPS of output video')
|
||||
parser.add_argument('--beg',
|
||||
type=int,
|
||||
default=1,
|
||||
help='The index of the first frame to be stylized')
|
||||
parser.add_argument('--end',
|
||||
type=int,
|
||||
default=101,
|
||||
help='The index of the last frame to be stylized')
|
||||
parser.add_argument('--itv',
|
||||
type=int,
|
||||
default=10,
|
||||
help='The interval of key frame')
|
||||
parser.add_argument('--key',
|
||||
type=str,
|
||||
default='keys0',
|
||||
help='The subfolder name of stylized key frames')
|
||||
parser.add_argument('--n_proc',
|
||||
type=int,
|
||||
default=8,
|
||||
help='The max process count')
|
||||
parser.add_argument('-ps',
|
||||
action='store_true',
|
||||
help='Use poisson gradient blending')
|
||||
parser.add_argument(
|
||||
'-ne',
|
||||
action='store_true',
|
||||
help='Do not run ebsynth (use previous ebsynth output)')
|
||||
parser.add_argument('-tmp',
|
||||
action='store_true',
|
||||
help='Keep temporary output')
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,969 @@
|
||||
import os
|
||||
import shutil
|
||||
from enum import Enum
|
||||
|
||||
import cv2
|
||||
import einops
|
||||
import gradio as gr
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as T
|
||||
from blendmodes.blend import BlendType, blendLayers
|
||||
from PIL import Image
|
||||
from pytorch_lightning import seed_everything
|
||||
from safetensors.torch import load_file
|
||||
from skimage import exposure
|
||||
|
||||
import src.import_util # noqa: F401
|
||||
from deps.ControlNet.annotator.canny import CannyDetector
|
||||
from deps.ControlNet.annotator.hed import HEDdetector
|
||||
from deps.ControlNet.annotator.util import HWC3
|
||||
from deps.ControlNet.cldm.model import create_model, load_state_dict
|
||||
from deps.gmflow.gmflow.gmflow import GMFlow
|
||||
from flow.flow_utils import get_warped_and_mask
|
||||
from sd_model_cfg import model_dict
|
||||
from src.config import RerenderConfig
|
||||
from src.controller import AttentionControl
|
||||
from src.ddim_v_hacked import DDIMVSampler
|
||||
from src.freeu import freeu_forward
|
||||
from src.img_util import find_flat_region, numpy2tensor
|
||||
from src.video_util import (frame_to_video, get_fps, get_frame_count,
|
||||
prepare_frames)
|
||||
|
||||
inversed_model_dict = dict()
|
||||
for k, v in model_dict.items():
|
||||
inversed_model_dict[v] = k
|
||||
|
||||
to_tensor = T.PILToTensor()
|
||||
blur = T.GaussianBlur(kernel_size=(9, 9), sigma=(18, 18))
|
||||
|
||||
|
||||
class ProcessingState(Enum):
|
||||
NULL = 0
|
||||
FIRST_IMG = 1
|
||||
KEY_IMGS = 2
|
||||
|
||||
|
||||
class GlobalState:
|
||||
|
||||
def __init__(self):
|
||||
self.sd_model = None
|
||||
self.ddim_v_sampler = None
|
||||
self.detector_type = None
|
||||
self.detector = None
|
||||
self.controller = None
|
||||
self.processing_state = ProcessingState.NULL
|
||||
flow_model = GMFlow(
|
||||
feature_channels=128,
|
||||
num_scales=1,
|
||||
upsample_factor=8,
|
||||
num_head=1,
|
||||
attention_type='swin',
|
||||
ffn_dim_expansion=4,
|
||||
num_transformer_layers=6,
|
||||
).to('cuda')
|
||||
|
||||
checkpoint = torch.load('models/gmflow_sintel-0c07dcb3.pth',
|
||||
map_location=lambda storage, loc: storage)
|
||||
weights = checkpoint['model'] if 'model' in checkpoint else checkpoint
|
||||
flow_model.load_state_dict(weights, strict=False)
|
||||
flow_model.eval()
|
||||
self.flow_model = flow_model
|
||||
|
||||
def update_controller(self, inner_strength, mask_period, cross_period,
|
||||
ada_period, warp_period, loose_cfattn):
|
||||
self.controller = AttentionControl(inner_strength,
|
||||
mask_period,
|
||||
cross_period,
|
||||
ada_period,
|
||||
warp_period,
|
||||
loose_cfatnn=loose_cfattn)
|
||||
|
||||
def update_sd_model(self, sd_model, control_type, freeu_args):
|
||||
if sd_model == self.sd_model:
|
||||
return
|
||||
self.sd_model = sd_model
|
||||
model = create_model('./deps/ControlNet/models/cldm_v15.yaml').cpu()
|
||||
if control_type == 'HED':
|
||||
model.load_state_dict(
|
||||
load_state_dict('./models/control_sd15_hed.pth',
|
||||
location='cuda'))
|
||||
elif control_type == 'canny':
|
||||
model.load_state_dict(
|
||||
load_state_dict('./models/control_sd15_canny.pth',
|
||||
location='cuda'))
|
||||
model = model.cuda()
|
||||
sd_model_path = model_dict[sd_model]
|
||||
if len(sd_model_path) > 0:
|
||||
model_ext = os.path.splitext(sd_model_path)[1]
|
||||
if model_ext == '.safetensors':
|
||||
model.load_state_dict(load_file(sd_model_path), strict=False)
|
||||
elif model_ext == '.ckpt' or model_ext == '.pth':
|
||||
model.load_state_dict(torch.load(sd_model_path)['state_dict'],
|
||||
strict=False)
|
||||
|
||||
try:
|
||||
model.first_stage_model.load_state_dict(torch.load(
|
||||
'./models/vae-ft-mse-840000-ema-pruned.ckpt')['state_dict'],
|
||||
strict=False)
|
||||
except Exception:
|
||||
print('Warning: We suggest you download the fine-tuned VAE',
|
||||
'otherwise the generation quality will be degraded')
|
||||
|
||||
model.model.diffusion_model.forward = freeu_forward(
|
||||
model.model.diffusion_model, *freeu_args)
|
||||
self.ddim_v_sampler = DDIMVSampler(model)
|
||||
|
||||
def clear_sd_model(self):
|
||||
self.sd_model = None
|
||||
self.ddim_v_sampler = None
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def update_detector(self, control_type, canny_low=100, canny_high=200):
|
||||
if self.detector_type == control_type:
|
||||
return
|
||||
if control_type == 'HED':
|
||||
self.detector = HEDdetector()
|
||||
elif control_type == 'canny':
|
||||
canny_detector = CannyDetector()
|
||||
low_threshold = canny_low
|
||||
high_threshold = canny_high
|
||||
|
||||
def apply_canny(x):
|
||||
return canny_detector(x, low_threshold, high_threshold)
|
||||
|
||||
self.detector = apply_canny
|
||||
|
||||
|
||||
global_state = GlobalState()
|
||||
global_video_path = None
|
||||
video_frame_count = None
|
||||
|
||||
|
||||
def create_cfg(input_path, prompt, image_resolution, control_strength,
|
||||
color_preserve, left_crop, right_crop, top_crop, bottom_crop,
|
||||
control_type, low_threshold, high_threshold, ddim_steps, scale,
|
||||
seed, sd_model, a_prompt, n_prompt, interval, keyframe_count,
|
||||
x0_strength, use_constraints, cross_start, cross_end,
|
||||
style_update_freq, warp_start, warp_end, mask_start, mask_end,
|
||||
ada_start, ada_end, mask_strength, inner_strength,
|
||||
smooth_boundary, loose_cfattn, b1, b2, s1, s2):
|
||||
use_warp = 'shape-aware fusion' in use_constraints
|
||||
use_mask = 'pixel-aware fusion' in use_constraints
|
||||
use_ada = 'color-aware AdaIN' in use_constraints
|
||||
|
||||
if not use_warp:
|
||||
warp_start = 1
|
||||
warp_end = 0
|
||||
|
||||
if not use_mask:
|
||||
mask_start = 1
|
||||
mask_end = 0
|
||||
|
||||
if not use_ada:
|
||||
ada_start = 1
|
||||
ada_end = 0
|
||||
|
||||
input_name = os.path.split(input_path)[-1].split('.')[0]
|
||||
frame_count = 2 + keyframe_count * interval
|
||||
cfg = RerenderConfig()
|
||||
cfg.create_from_parameters(
|
||||
input_path,
|
||||
os.path.join('result', input_name, 'blend.mp4'),
|
||||
prompt,
|
||||
a_prompt=a_prompt,
|
||||
n_prompt=n_prompt,
|
||||
frame_count=frame_count,
|
||||
interval=interval,
|
||||
crop=[left_crop, right_crop, top_crop, bottom_crop],
|
||||
sd_model=sd_model,
|
||||
ddim_steps=ddim_steps,
|
||||
scale=scale,
|
||||
control_type=control_type,
|
||||
control_strength=control_strength,
|
||||
canny_low=low_threshold,
|
||||
canny_high=high_threshold,
|
||||
seed=seed,
|
||||
image_resolution=image_resolution,
|
||||
x0_strength=x0_strength,
|
||||
style_update_freq=style_update_freq,
|
||||
cross_period=(cross_start, cross_end),
|
||||
warp_period=(warp_start, warp_end),
|
||||
mask_period=(mask_start, mask_end),
|
||||
ada_period=(ada_start, ada_end),
|
||||
mask_strength=mask_strength,
|
||||
inner_strength=inner_strength,
|
||||
smooth_boundary=smooth_boundary,
|
||||
color_preserve=color_preserve,
|
||||
loose_cfattn=loose_cfattn,
|
||||
freeu_args=[b1, b2, s1, s2])
|
||||
return cfg
|
||||
|
||||
|
||||
def cfg_to_input(filename):
|
||||
|
||||
cfg = RerenderConfig()
|
||||
cfg.create_from_path(filename)
|
||||
keyframe_count = (cfg.frame_count - 2) // cfg.interval
|
||||
use_constraints = [
|
||||
'shape-aware fusion', 'pixel-aware fusion', 'color-aware AdaIN'
|
||||
]
|
||||
|
||||
sd_model = inversed_model_dict.get(cfg.sd_model, 'Stable Diffusion 1.5')
|
||||
|
||||
args = [
|
||||
cfg.input_path, cfg.prompt, cfg.image_resolution, cfg.control_strength,
|
||||
cfg.color_preserve, *cfg.crop, cfg.control_type, cfg.canny_low,
|
||||
cfg.canny_high, cfg.ddim_steps, cfg.scale, cfg.seed, sd_model,
|
||||
cfg.a_prompt, cfg.n_prompt, cfg.interval, keyframe_count,
|
||||
cfg.x0_strength, use_constraints, *cfg.cross_period,
|
||||
cfg.style_update_freq, *cfg.warp_period, *cfg.mask_period,
|
||||
*cfg.ada_period, cfg.mask_strength, cfg.inner_strength,
|
||||
cfg.smooth_boundary, cfg.loose_cfattn, *cfg.freeu_args
|
||||
]
|
||||
return args
|
||||
|
||||
|
||||
def setup_color_correction(image):
|
||||
correction_target = cv2.cvtColor(np.asarray(image.copy()),
|
||||
cv2.COLOR_RGB2LAB)
|
||||
return correction_target
|
||||
|
||||
|
||||
def apply_color_correction(correction, original_image):
|
||||
image = Image.fromarray(
|
||||
cv2.cvtColor(
|
||||
exposure.match_histograms(cv2.cvtColor(np.asarray(original_image),
|
||||
cv2.COLOR_RGB2LAB),
|
||||
correction,
|
||||
channel_axis=2),
|
||||
cv2.COLOR_LAB2RGB).astype('uint8'))
|
||||
|
||||
image = blendLayers(image, original_image, BlendType.LUMINOSITY)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def process(*args):
|
||||
args_wo_process3 = args[:-2]
|
||||
first_frame = process1(*args_wo_process3)
|
||||
|
||||
keypath = process2(*args_wo_process3)
|
||||
|
||||
fullpath = process3(*args)
|
||||
|
||||
return first_frame, keypath, fullpath
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def process1(*args):
|
||||
|
||||
global global_video_path
|
||||
cfg = create_cfg(global_video_path, *args)
|
||||
global global_state
|
||||
global_state.update_sd_model(cfg.sd_model, cfg.control_type,
|
||||
cfg.freeu_args)
|
||||
global_state.update_controller(cfg.inner_strength, cfg.mask_period,
|
||||
cfg.cross_period, cfg.ada_period,
|
||||
cfg.warp_period, cfg.loose_cfattn)
|
||||
global_state.update_detector(cfg.control_type, cfg.canny_low,
|
||||
cfg.canny_high)
|
||||
global_state.processing_state = ProcessingState.FIRST_IMG
|
||||
|
||||
prepare_frames(cfg.input_path, cfg.input_dir, cfg.image_resolution, cfg.crop, cfg.use_limit_device_resolution)
|
||||
|
||||
ddim_v_sampler = global_state.ddim_v_sampler
|
||||
model = ddim_v_sampler.model
|
||||
detector = global_state.detector
|
||||
controller = global_state.controller
|
||||
model.control_scales = [cfg.control_strength] * 13
|
||||
|
||||
num_samples = 1
|
||||
eta = 0.0
|
||||
imgs = sorted(os.listdir(cfg.input_dir))
|
||||
imgs = [os.path.join(cfg.input_dir, img) for img in imgs]
|
||||
|
||||
with torch.no_grad():
|
||||
frame = cv2.imread(imgs[0])
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
img = HWC3(frame)
|
||||
H, W, C = img.shape
|
||||
|
||||
img_ = numpy2tensor(img)
|
||||
|
||||
def generate_first_img(img_, strength):
|
||||
encoder_posterior = model.encode_first_stage(img_.cuda())
|
||||
x0 = model.get_first_stage_encoding(encoder_posterior).detach()
|
||||
|
||||
detected_map = detector(img)
|
||||
detected_map = HWC3(detected_map)
|
||||
|
||||
control = torch.from_numpy(
|
||||
detected_map.copy()).float().cuda() / 255.0
|
||||
control = torch.stack([control for _ in range(num_samples)], dim=0)
|
||||
control = einops.rearrange(control, 'b h w c -> b c h w').clone()
|
||||
cond = {
|
||||
'c_concat': [control],
|
||||
'c_crossattn': [
|
||||
model.get_learned_conditioning(
|
||||
[cfg.prompt + ', ' + cfg.a_prompt] * num_samples)
|
||||
]
|
||||
}
|
||||
un_cond = {
|
||||
'c_concat': [control],
|
||||
'c_crossattn':
|
||||
[model.get_learned_conditioning([cfg.n_prompt] * num_samples)]
|
||||
}
|
||||
shape = (4, H // 8, W // 8)
|
||||
|
||||
controller.set_task('initfirst')
|
||||
seed_everything(cfg.seed)
|
||||
|
||||
samples, _ = ddim_v_sampler.sample(
|
||||
cfg.ddim_steps,
|
||||
num_samples,
|
||||
shape,
|
||||
cond,
|
||||
verbose=False,
|
||||
eta=eta,
|
||||
unconditional_guidance_scale=cfg.scale,
|
||||
unconditional_conditioning=un_cond,
|
||||
controller=controller,
|
||||
x0=x0,
|
||||
strength=strength)
|
||||
x_samples = model.decode_first_stage(samples)
|
||||
x_samples_np = (
|
||||
einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +
|
||||
127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
return x_samples, x_samples_np
|
||||
|
||||
# When not preserve color, draw a different frame at first and use its
|
||||
# color to redraw the first frame.
|
||||
if not cfg.color_preserve:
|
||||
first_strength = -1
|
||||
else:
|
||||
first_strength = 1 - cfg.x0_strength
|
||||
|
||||
x_samples, x_samples_np = generate_first_img(img_, first_strength)
|
||||
|
||||
if not cfg.color_preserve:
|
||||
color_corrections = setup_color_correction(
|
||||
Image.fromarray(x_samples_np[0]))
|
||||
global_state.color_corrections = color_corrections
|
||||
img_ = apply_color_correction(color_corrections,
|
||||
Image.fromarray(img))
|
||||
img_ = to_tensor(img_).unsqueeze(0)[:, :3] / 127.5 - 1
|
||||
x_samples, x_samples_np = generate_first_img(
|
||||
img_, 1 - cfg.x0_strength)
|
||||
|
||||
global_state.first_result = x_samples
|
||||
global_state.first_img = img
|
||||
|
||||
Image.fromarray(x_samples_np[0]).save(
|
||||
os.path.join(cfg.first_dir, 'first.jpg'))
|
||||
|
||||
return x_samples_np[0]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def process2(*args):
|
||||
global global_state
|
||||
global global_video_path
|
||||
|
||||
if global_state.processing_state != ProcessingState.FIRST_IMG:
|
||||
raise gr.Error('Please generate the first key image before generating'
|
||||
' all key images')
|
||||
|
||||
cfg = create_cfg(global_video_path, *args)
|
||||
global_state.update_sd_model(cfg.sd_model, cfg.control_type,
|
||||
cfg.freeu_args)
|
||||
global_state.update_detector(cfg.control_type, cfg.canny_low,
|
||||
cfg.canny_high)
|
||||
global_state.processing_state = ProcessingState.KEY_IMGS
|
||||
|
||||
# reset key dir
|
||||
shutil.rmtree(cfg.key_dir)
|
||||
os.makedirs(cfg.key_dir, exist_ok=True)
|
||||
|
||||
ddim_v_sampler = global_state.ddim_v_sampler
|
||||
model = ddim_v_sampler.model
|
||||
detector = global_state.detector
|
||||
controller = global_state.controller
|
||||
flow_model = global_state.flow_model
|
||||
model.control_scales = [cfg.control_strength] * 13
|
||||
|
||||
num_samples = 1
|
||||
eta = 0.0
|
||||
firstx0 = True
|
||||
pixelfusion = cfg.use_mask
|
||||
imgs = sorted(os.listdir(cfg.input_dir))
|
||||
imgs = [os.path.join(cfg.input_dir, img) for img in imgs]
|
||||
|
||||
first_result = global_state.first_result
|
||||
first_img = global_state.first_img
|
||||
pre_result = first_result
|
||||
pre_img = first_img
|
||||
|
||||
for i in range(0, min(len(imgs), cfg.frame_count) - 1, cfg.interval):
|
||||
cid = i + 1
|
||||
print(cid)
|
||||
if cid <= (len(imgs) - 1):
|
||||
frame = cv2.imread(imgs[cid])
|
||||
else:
|
||||
frame = cv2.imread(imgs[len(imgs) - 1])
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
img = HWC3(frame)
|
||||
H, W, C = img.shape
|
||||
|
||||
if cfg.color_preserve or global_state.color_corrections is None:
|
||||
img_ = numpy2tensor(img)
|
||||
else:
|
||||
img_ = apply_color_correction(global_state.color_corrections,
|
||||
Image.fromarray(img))
|
||||
img_ = to_tensor(img_).unsqueeze(0)[:, :3] / 127.5 - 1
|
||||
encoder_posterior = model.encode_first_stage(img_.cuda())
|
||||
x0 = model.get_first_stage_encoding(encoder_posterior).detach()
|
||||
|
||||
detected_map = detector(img)
|
||||
detected_map = HWC3(detected_map)
|
||||
|
||||
control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0
|
||||
control = torch.stack([control for _ in range(num_samples)], dim=0)
|
||||
control = einops.rearrange(control, 'b h w c -> b c h w').clone()
|
||||
cond = {
|
||||
'c_concat': [control],
|
||||
'c_crossattn': [
|
||||
model.get_learned_conditioning(
|
||||
[cfg.prompt + ', ' + cfg.a_prompt] * num_samples)
|
||||
]
|
||||
}
|
||||
un_cond = {
|
||||
'c_concat': [control],
|
||||
'c_crossattn':
|
||||
[model.get_learned_conditioning([cfg.n_prompt] * num_samples)]
|
||||
}
|
||||
shape = (4, H // 8, W // 8)
|
||||
|
||||
cond['c_concat'] = [control]
|
||||
un_cond['c_concat'] = [control]
|
||||
|
||||
image1 = torch.from_numpy(pre_img).permute(2, 0, 1).float()
|
||||
image2 = torch.from_numpy(img).permute(2, 0, 1).float()
|
||||
warped_pre, bwd_occ_pre, bwd_flow_pre = get_warped_and_mask(
|
||||
flow_model, image1, image2, pre_result, False)
|
||||
blend_mask_pre = blur(
|
||||
F.max_pool2d(bwd_occ_pre, kernel_size=9, stride=1, padding=4))
|
||||
blend_mask_pre = torch.clamp(blend_mask_pre + bwd_occ_pre, 0, 1)
|
||||
|
||||
image1 = torch.from_numpy(first_img).permute(2, 0, 1).float()
|
||||
warped_0, bwd_occ_0, bwd_flow_0 = get_warped_and_mask(
|
||||
flow_model, image1, image2, first_result, False)
|
||||
blend_mask_0 = blur(
|
||||
F.max_pool2d(bwd_occ_0, kernel_size=9, stride=1, padding=4))
|
||||
blend_mask_0 = torch.clamp(blend_mask_0 + bwd_occ_0, 0, 1)
|
||||
|
||||
if firstx0:
|
||||
mask = 1 - F.max_pool2d(blend_mask_0, kernel_size=8)
|
||||
controller.set_warp(
|
||||
F.interpolate(bwd_flow_0 / 8.0,
|
||||
scale_factor=1. / 8,
|
||||
mode='bilinear'), mask)
|
||||
else:
|
||||
mask = 1 - F.max_pool2d(blend_mask_pre, kernel_size=8)
|
||||
controller.set_warp(
|
||||
F.interpolate(bwd_flow_pre / 8.0,
|
||||
scale_factor=1. / 8,
|
||||
mode='bilinear'), mask)
|
||||
|
||||
controller.set_task('keepx0, keepstyle')
|
||||
seed_everything(cfg.seed)
|
||||
samples, intermediates = ddim_v_sampler.sample(
|
||||
cfg.ddim_steps,
|
||||
num_samples,
|
||||
shape,
|
||||
cond,
|
||||
verbose=False,
|
||||
eta=eta,
|
||||
unconditional_guidance_scale=cfg.scale,
|
||||
unconditional_conditioning=un_cond,
|
||||
controller=controller,
|
||||
x0=x0,
|
||||
strength=1 - cfg.x0_strength)
|
||||
direct_result = model.decode_first_stage(samples)
|
||||
|
||||
if not pixelfusion:
|
||||
pre_result = direct_result
|
||||
pre_img = img
|
||||
viz = (
|
||||
einops.rearrange(direct_result, 'b c h w -> b h w c') * 127.5 +
|
||||
127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
|
||||
else:
|
||||
|
||||
blend_results = (1 - blend_mask_pre
|
||||
) * warped_pre + blend_mask_pre * direct_result
|
||||
blend_results = (
|
||||
1 - blend_mask_0) * warped_0 + blend_mask_0 * blend_results
|
||||
|
||||
bwd_occ = 1 - torch.clamp(1 - bwd_occ_pre + 1 - bwd_occ_0, 0, 1)
|
||||
blend_mask = blur(
|
||||
F.max_pool2d(bwd_occ, kernel_size=9, stride=1, padding=4))
|
||||
blend_mask = 1 - torch.clamp(blend_mask + bwd_occ, 0, 1)
|
||||
|
||||
encoder_posterior = model.encode_first_stage(blend_results)
|
||||
xtrg = model.get_first_stage_encoding(
|
||||
encoder_posterior).detach() # * mask
|
||||
blend_results_rec = model.decode_first_stage(xtrg)
|
||||
encoder_posterior = model.encode_first_stage(blend_results_rec)
|
||||
xtrg_rec = model.get_first_stage_encoding(
|
||||
encoder_posterior).detach()
|
||||
xtrg_ = (xtrg + 1 * (xtrg - xtrg_rec)) # * mask
|
||||
blend_results_rec_new = model.decode_first_stage(xtrg_)
|
||||
tmp = (abs(blend_results_rec_new - blend_results).mean(
|
||||
dim=1, keepdims=True) > 0.25).float()
|
||||
mask_x = F.max_pool2d((F.interpolate(
|
||||
tmp, scale_factor=1 / 8., mode='bilinear') > 0).float(),
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
|
||||
mask = (1 - F.max_pool2d(1 - blend_mask, kernel_size=8)
|
||||
) # * (1-mask_x)
|
||||
|
||||
if cfg.smooth_boundary:
|
||||
noise_rescale = find_flat_region(mask)
|
||||
else:
|
||||
noise_rescale = torch.ones_like(mask)
|
||||
masks = []
|
||||
for j in range(cfg.ddim_steps):
|
||||
if j <= cfg.ddim_steps * cfg.mask_period[
|
||||
0] or j >= cfg.ddim_steps * cfg.mask_period[1]:
|
||||
masks += [None]
|
||||
else:
|
||||
masks += [mask * cfg.mask_strength]
|
||||
|
||||
# mask 3
|
||||
# xtrg = ((1-mask_x) *
|
||||
# (xtrg + xtrg - xtrg_rec) + mask_x * samples) * mask
|
||||
# mask 2
|
||||
# xtrg = (xtrg + 1 * (xtrg - xtrg_rec)) * mask
|
||||
xtrg = (xtrg + (1 - mask_x) * (xtrg - xtrg_rec)) * mask # mask 1
|
||||
|
||||
tasks = 'keepstyle, keepx0'
|
||||
if not firstx0:
|
||||
tasks += ', updatex0'
|
||||
if i % cfg.style_update_freq == 0:
|
||||
tasks += ', updatestyle'
|
||||
controller.set_task(tasks, 1.0)
|
||||
|
||||
seed_everything(cfg.seed)
|
||||
samples, _ = ddim_v_sampler.sample(
|
||||
cfg.ddim_steps,
|
||||
num_samples,
|
||||
shape,
|
||||
cond,
|
||||
verbose=False,
|
||||
eta=eta,
|
||||
unconditional_guidance_scale=cfg.scale,
|
||||
unconditional_conditioning=un_cond,
|
||||
controller=controller,
|
||||
x0=x0,
|
||||
strength=1 - cfg.x0_strength,
|
||||
xtrg=xtrg,
|
||||
mask=masks,
|
||||
noise_rescale=noise_rescale)
|
||||
x_samples = model.decode_first_stage(samples)
|
||||
pre_result = x_samples
|
||||
pre_img = img
|
||||
|
||||
viz = (einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +
|
||||
127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
|
||||
Image.fromarray(viz[0]).save(
|
||||
os.path.join(cfg.key_dir, f'{cid:04d}.png'))
|
||||
|
||||
key_video_path = os.path.join(cfg.work_dir, 'key.mp4')
|
||||
fps = get_fps(cfg.input_path)
|
||||
fps //= cfg.interval
|
||||
frame_to_video(key_video_path, cfg.key_dir, fps, False)
|
||||
|
||||
return key_video_path
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def process3(*args):
|
||||
max_process = args[-2]
|
||||
use_poisson = args[-1]
|
||||
args = args[:-2]
|
||||
global global_video_path
|
||||
global global_state
|
||||
if global_state.processing_state != ProcessingState.KEY_IMGS:
|
||||
raise gr.Error('Please generate key images before propagation')
|
||||
|
||||
global_state.clear_sd_model()
|
||||
|
||||
cfg = create_cfg(global_video_path, *args)
|
||||
|
||||
# reset blend dir
|
||||
blend_dir = os.path.join(cfg.work_dir, 'blend')
|
||||
if os.path.exists(blend_dir):
|
||||
shutil.rmtree(blend_dir)
|
||||
os.makedirs(blend_dir, exist_ok=True)
|
||||
|
||||
video_base_dir = cfg.work_dir
|
||||
o_video = cfg.output_path
|
||||
fps = get_fps(cfg.input_path)
|
||||
|
||||
end_frame = cfg.frame_count - 1
|
||||
interval = cfg.interval
|
||||
key_dir = os.path.split(cfg.key_dir)[-1]
|
||||
o_video_cmd = f'--output {o_video}'
|
||||
ps = '-ps' if use_poisson else ''
|
||||
cmd = (f'python video_blend.py {video_base_dir} --beg 1 --end {end_frame} '
|
||||
f'--itv {interval} --key {key_dir} {o_video_cmd} --fps {fps} '
|
||||
f'--n_proc {max_process} {ps}')
|
||||
print(cmd)
|
||||
os.system(cmd)
|
||||
|
||||
return o_video
|
||||
|
||||
|
||||
block = gr.Blocks().queue()
|
||||
with block:
|
||||
with gr.Row():
|
||||
gr.Markdown('## Rerender A Video')
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
input_path = gr.Video(label='Input Video',
|
||||
source='upload',
|
||||
format='mp4',
|
||||
visible=True)
|
||||
prompt = gr.Textbox(label='Prompt')
|
||||
seed = gr.Slider(label='Seed',
|
||||
minimum=0,
|
||||
maximum=2147483647,
|
||||
step=1,
|
||||
value=0,
|
||||
randomize=True)
|
||||
run_button = gr.Button(value='Run All')
|
||||
with gr.Row():
|
||||
run_button1 = gr.Button(value='Run 1st Key Frame')
|
||||
run_button2 = gr.Button(value='Run Key Frames')
|
||||
run_button3 = gr.Button(value='Run Propagation')
|
||||
with gr.Accordion('Advanced options for the 1st frame translation',
|
||||
open=False):
|
||||
image_resolution = gr.Slider(label='Frame resolution',
|
||||
minimum=256,
|
||||
maximum=768,
|
||||
value=512,
|
||||
step=64)
|
||||
control_strength = gr.Slider(label='ControlNet strength',
|
||||
minimum=0.0,
|
||||
maximum=2.0,
|
||||
value=1.0,
|
||||
step=0.01)
|
||||
x0_strength = gr.Slider(
|
||||
label='Denoising strength',
|
||||
minimum=0.00,
|
||||
maximum=1.05,
|
||||
value=0.75,
|
||||
step=0.05,
|
||||
info=('0: fully recover the input.'
|
||||
'1.05: fully rerender the input.'))
|
||||
color_preserve = gr.Checkbox(
|
||||
label='Preserve color',
|
||||
value=True,
|
||||
info='Keep the color of the input video')
|
||||
with gr.Row():
|
||||
left_crop = gr.Slider(label='Left crop length',
|
||||
minimum=0,
|
||||
maximum=512,
|
||||
value=0,
|
||||
step=1)
|
||||
right_crop = gr.Slider(label='Right crop length',
|
||||
minimum=0,
|
||||
maximum=512,
|
||||
value=0,
|
||||
step=1)
|
||||
with gr.Row():
|
||||
top_crop = gr.Slider(label='Top crop length',
|
||||
minimum=0,
|
||||
maximum=512,
|
||||
value=0,
|
||||
step=1)
|
||||
bottom_crop = gr.Slider(label='Bottom crop length',
|
||||
minimum=0,
|
||||
maximum=512,
|
||||
value=0,
|
||||
step=1)
|
||||
with gr.Row():
|
||||
control_type = gr.Dropdown(['HED', 'canny'],
|
||||
label='Control type',
|
||||
value='HED')
|
||||
low_threshold = gr.Slider(label='Canny low threshold',
|
||||
minimum=1,
|
||||
maximum=255,
|
||||
value=100,
|
||||
step=1)
|
||||
high_threshold = gr.Slider(label='Canny high threshold',
|
||||
minimum=1,
|
||||
maximum=255,
|
||||
value=200,
|
||||
step=1)
|
||||
ddim_steps = gr.Slider(label='Steps',
|
||||
minimum=20,
|
||||
maximum=100,
|
||||
value=20,
|
||||
step=20)
|
||||
scale = gr.Slider(label='CFG scale',
|
||||
minimum=0.1,
|
||||
maximum=30.0,
|
||||
value=7.5,
|
||||
step=0.1)
|
||||
sd_model_list = list(model_dict.keys())
|
||||
sd_model = gr.Dropdown(sd_model_list,
|
||||
label='Base model',
|
||||
value='Stable Diffusion 1.5')
|
||||
a_prompt = gr.Textbox(label='Added prompt',
|
||||
value='best quality, extremely detailed')
|
||||
n_prompt = gr.Textbox(
|
||||
label='Negative prompt',
|
||||
value=('longbody, lowres, bad anatomy, bad hands, '
|
||||
'missing fingers, extra digit, fewer digits, '
|
||||
'cropped, worst quality, low quality'))
|
||||
with gr.Row():
|
||||
b1 = gr.Slider(label='FreeU first-stage backbone factor',
|
||||
minimum=1,
|
||||
maximum=1.6,
|
||||
value=1,
|
||||
step=0.01,
|
||||
info='FreeU to enhance texture and color')
|
||||
b2 = gr.Slider(label='FreeU second-stage backbone factor',
|
||||
minimum=1,
|
||||
maximum=1.6,
|
||||
value=1,
|
||||
step=0.01)
|
||||
with gr.Row():
|
||||
s1 = gr.Slider(label='FreeU first-stage skip factor',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=1,
|
||||
step=0.01)
|
||||
s2 = gr.Slider(label='FreeU second-stage skip factor',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=1,
|
||||
step=0.01)
|
||||
with gr.Accordion('Advanced options for the key fame translation',
|
||||
open=False):
|
||||
interval = gr.Slider(
|
||||
label='Key frame frequency (K)',
|
||||
minimum=1,
|
||||
maximum=1,
|
||||
value=1,
|
||||
step=1,
|
||||
info='Uniformly sample the key frames every K frames')
|
||||
keyframe_count = gr.Slider(label='Number of key frames',
|
||||
minimum=1,
|
||||
maximum=1,
|
||||
value=1,
|
||||
step=1)
|
||||
|
||||
use_constraints = gr.CheckboxGroup(
|
||||
[
|
||||
'shape-aware fusion', 'pixel-aware fusion',
|
||||
'color-aware AdaIN'
|
||||
],
|
||||
label='Select the cross-frame contraints to be used',
|
||||
value=[
|
||||
'shape-aware fusion', 'pixel-aware fusion',
|
||||
'color-aware AdaIN'
|
||||
]),
|
||||
with gr.Row():
|
||||
cross_start = gr.Slider(
|
||||
label='Cross-frame attention start',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0,
|
||||
step=0.05)
|
||||
cross_end = gr.Slider(label='Cross-frame attention end',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=1,
|
||||
step=0.05)
|
||||
style_update_freq = gr.Slider(
|
||||
label='Cross-frame attention update frequency',
|
||||
minimum=1,
|
||||
maximum=100,
|
||||
value=1,
|
||||
step=1,
|
||||
info=('Update the key and value for '
|
||||
'cross-frame attention every N key frames'))
|
||||
loose_cfattn = gr.Checkbox(
|
||||
label='Loose Cross-frame attention',
|
||||
value=True,
|
||||
info='Select to make output better match the input video')
|
||||
with gr.Row():
|
||||
warp_start = gr.Slider(label='Shape-aware fusion start',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0,
|
||||
step=0.05)
|
||||
warp_end = gr.Slider(label='Shape-aware fusion end',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0.1,
|
||||
step=0.05)
|
||||
with gr.Row():
|
||||
mask_start = gr.Slider(label='Pixel-aware fusion start',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0.5,
|
||||
step=0.05)
|
||||
mask_end = gr.Slider(label='Pixel-aware fusion end',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0.8,
|
||||
step=0.05)
|
||||
with gr.Row():
|
||||
ada_start = gr.Slider(label='Color-aware AdaIN start',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0.8,
|
||||
step=0.05)
|
||||
ada_end = gr.Slider(label='Color-aware AdaIN end',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=1,
|
||||
step=0.05)
|
||||
mask_strength = gr.Slider(label='Pixel-aware fusion strength',
|
||||
minimum=0,
|
||||
maximum=1,
|
||||
value=0.5,
|
||||
step=0.01)
|
||||
inner_strength = gr.Slider(
|
||||
label='Pixel-aware fusion detail level',
|
||||
minimum=0.5,
|
||||
maximum=1,
|
||||
value=0.9,
|
||||
step=0.01,
|
||||
info='Use a low value to prevent artifacts')
|
||||
smooth_boundary = gr.Checkbox(
|
||||
label='Smooth fusion boundary',
|
||||
value=True,
|
||||
info='Select to prevent artifacts at boundary')
|
||||
with gr.Accordion(
|
||||
'Advanced options for the full video translation',
|
||||
open=False):
|
||||
use_poisson = gr.Checkbox(
|
||||
label='Gradient blending',
|
||||
value=True,
|
||||
info=('Blend the output video in gradient, to reduce'
|
||||
' ghosting artifacts (but may increase flickers)'))
|
||||
max_process = gr.Slider(label='Number of parallel processes',
|
||||
minimum=1,
|
||||
maximum=16,
|
||||
value=4,
|
||||
step=1)
|
||||
|
||||
with gr.Accordion('Example configs', open=True):
|
||||
config_dir = 'config'
|
||||
config_list = [
|
||||
'real2sculpture.json', 'van_gogh_man.json', 'woman.json'
|
||||
]
|
||||
args_list = []
|
||||
for config in config_list:
|
||||
try:
|
||||
config_path = os.path.join(config_dir, config)
|
||||
args = cfg_to_input(config_path)
|
||||
args_list.append(args)
|
||||
except FileNotFoundError:
|
||||
# The video file does not exist, skipped
|
||||
pass
|
||||
|
||||
ips = [
|
||||
prompt, image_resolution, control_strength, color_preserve,
|
||||
left_crop, right_crop, top_crop, bottom_crop, control_type,
|
||||
low_threshold, high_threshold, ddim_steps, scale, seed,
|
||||
sd_model, a_prompt, n_prompt, interval, keyframe_count,
|
||||
x0_strength, use_constraints[0], cross_start, cross_end,
|
||||
style_update_freq, warp_start, warp_end, mask_start,
|
||||
mask_end, ada_start, ada_end, mask_strength,
|
||||
inner_strength, smooth_boundary, loose_cfattn, b1, b2, s1,
|
||||
s2
|
||||
]
|
||||
|
||||
gr.Examples(
|
||||
examples=args_list,
|
||||
inputs=[input_path, *ips],
|
||||
)
|
||||
|
||||
with gr.Column():
|
||||
result_image = gr.Image(label='Output first frame',
|
||||
type='numpy',
|
||||
interactive=False)
|
||||
result_keyframe = gr.Video(label='Output key frame video',
|
||||
format='mp4',
|
||||
interactive=False)
|
||||
result_video = gr.Video(label='Output full video',
|
||||
format='mp4',
|
||||
interactive=False)
|
||||
|
||||
def input_uploaded(path):
|
||||
frame_count = get_frame_count(path)
|
||||
if frame_count <= 2:
|
||||
raise gr.Error('The input video is too short!'
|
||||
'Please input another video.')
|
||||
|
||||
default_interval = min(10, frame_count - 2)
|
||||
max_keyframe = (frame_count - 2) // default_interval
|
||||
|
||||
global video_frame_count
|
||||
video_frame_count = frame_count
|
||||
global global_video_path
|
||||
global_video_path = path
|
||||
|
||||
return gr.Slider.update(value=default_interval,
|
||||
maximum=max_keyframe), gr.Slider.update(
|
||||
value=max_keyframe, maximum=max_keyframe)
|
||||
|
||||
def input_changed(path):
|
||||
frame_count = get_frame_count(path)
|
||||
if frame_count <= 2:
|
||||
return gr.Slider.update(maximum=1), gr.Slider.update(maximum=1)
|
||||
|
||||
default_interval = min(10, frame_count - 2)
|
||||
max_keyframe = (frame_count - 2) // default_interval
|
||||
|
||||
global video_frame_count
|
||||
video_frame_count = frame_count
|
||||
global global_video_path
|
||||
global_video_path = path
|
||||
|
||||
return gr.Slider.update(maximum=max_keyframe), \
|
||||
gr.Slider.update(maximum=max_keyframe)
|
||||
|
||||
def interval_changed(interval):
|
||||
global video_frame_count
|
||||
if video_frame_count is None:
|
||||
return gr.Slider.update()
|
||||
|
||||
max_keyframe = (video_frame_count - 2) // interval
|
||||
|
||||
return gr.Slider.update(value=max_keyframe, maximum=max_keyframe)
|
||||
|
||||
input_path.change(input_changed, input_path, [interval, keyframe_count])
|
||||
input_path.upload(input_uploaded, input_path, [interval, keyframe_count])
|
||||
interval.change(interval_changed, interval, keyframe_count)
|
||||
|
||||
ips_process3 = [*ips, max_process, use_poisson]
|
||||
run_button.click(fn=process,
|
||||
inputs=ips_process3,
|
||||
outputs=[result_image, result_keyframe, result_video])
|
||||
run_button1.click(fn=process1, inputs=ips, outputs=[result_image])
|
||||
run_button2.click(fn=process2, inputs=ips, outputs=[result_keyframe])
|
||||
run_button3.click(fn=process3, inputs=ips_process3, outputs=[result_video])
|
||||
|
||||
block.launch(server_name='localhost')
|
||||
Reference in New Issue
Block a user