Initial push
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
__pycache__
|
||||
/venv
|
||||
.vscode
|
||||
*.ckpt
|
||||
*.safetensors
|
||||
*.pth
|
||||
types
|
||||
*.pyc
|
||||
@@ -0,0 +1,4 @@
|
||||
LLAVA_CLIP_PATH = None
|
||||
LLAVA_MODEL_PATH = None
|
||||
SDXL_CLIP1_PATH = None
|
||||
SDXL_CLIP2_CKPT_PTH = None
|
||||
@@ -0,0 +1,150 @@
|
||||
## (CVPR2024) Scaling Up to Excellence: Practicing Model Scaling for Photo-Realistic Image Restoration In the Wild
|
||||
|
||||
> [[Paper](https://arxiv.org/abs/2401.13627)]   [[Project Page](http://supir.xpixel.group/)]   [Online Demo (Coming soon)] <br>
|
||||
> Fanghua, Yu, [Jinjin Gu](https://www.jasongt.com/), Zheyuan Li, Jinfan Hu, Xiangtao Kong, [Xintao Wang](https://xinntao.github.io/), [Jingwen He](https://scholar.google.com.hk/citations?user=GUxrycUAAAAJ), [Yu Qiao](https://scholar.google.com.hk/citations?user=gFtI-8QAAAAJ), [Chao Dong](https://scholar.google.com.hk/citations?user=OSDCB0UAAAAJ) <br>
|
||||
> Shenzhen Institute of Advanced Technology; Shanghai AI Laboratory; University of Sydney; The Hong Kong Polytechnic University; ARC Lab, Tencent PCG; The Chinese University of Hong Kong <br>
|
||||
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/teaser.png">
|
||||
</p>
|
||||
|
||||
---
|
||||
#### ⚠ Due to the large RAM (60G) and VRAM (30G x2) costs of SUPIR, we are working on the online demo releasing.
|
||||
|
||||
---
|
||||
## 🔧 Dependencies and Installation
|
||||
|
||||
1. Clone repo
|
||||
```bash
|
||||
git clone https://github.com/Fanghua-Yu/SUPIR.git
|
||||
cd SUPIR
|
||||
```
|
||||
|
||||
2. Install dependent packages
|
||||
```bash
|
||||
conda create -n SUPIR python=3.8 -y
|
||||
conda activate SUPIR
|
||||
pip install --upgrade pip
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
3. Download Checkpoints
|
||||
|
||||
For users who can connect to huggingface, please setting `LLAVA_CLIP_PATH, SDXL_CLIP1_PATH, SDXL_CLIP2_CKPT_PTH` in `CKPT_PTH.py` as `None`. These CLIPs will be downloaded automatically.
|
||||
|
||||
#### Dependent Models
|
||||
* [SDXL CLIP Encoder-1](https://huggingface.co/openai/clip-vit-large-patch14)
|
||||
* [SDXL CLIP Encoder-2](https://huggingface.co/laion/CLIP-ViT-bigG-14-laion2B-39B-b160k)
|
||||
* [SDXL base 1.0_0.9vae](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/sd_xl_base_1.0_0.9vae.safetensors)
|
||||
* [LLaVA CLIP](https://huggingface.co/openai/clip-vit-large-patch14-336)
|
||||
* [LLaVA v1.5 13B](https://huggingface.co/liuhaotian/llava-v1.5-13b)
|
||||
|
||||
|
||||
#### Models we provided:
|
||||
* `SUPIR-v0Q`: [Baidu Netdisk](https://pan.baidu.com/s/1lnefCZhBTeDWijqbj1jIyw?pwd=pjq6), [Google Drive](https://drive.google.com/drive/folders/1yELzm5SvAi9e7kPcO_jPp2XkTs4vK6aR?usp=sharing)
|
||||
|
||||
Default training settings with paper. High generalization and high image quality in most cases.
|
||||
|
||||
* `SUPIR-v0F`: [Baidu Netdisk](https://pan.baidu.com/s/1AECN8NjiVuE3hvO8o-Ua6A?pwd=k2uz), [Google Drive](https://drive.google.com/drive/folders/1yELzm5SvAi9e7kPcO_jPp2XkTs4vK6aR?usp=sharing)
|
||||
|
||||
Training with light degradation settings. Stage1 encoder of `SUPIR-v0F` remains more details when facing light degradations.
|
||||
|
||||
4. Edit Custom Path for Checkpoints
|
||||
```
|
||||
* [CKPT_PTH.py] --> LLAVA_CLIP_PATH, LLAVA_MODEL_PATH, SDXL_CLIP1_PATH, SDXL_CLIP2_CACHE_DIR
|
||||
* [options/SUPIR_v0.yaml] --> SDXL_CKPT, SUPIR_CKPT_Q, SUPIR_CKPT_F
|
||||
```
|
||||
---
|
||||
|
||||
## ⚡ Quick Inference
|
||||
### Val Dataset
|
||||
RealPhoto60: [Baidu Netdisk](https://pan.baidu.com/s/1CJKsPGtyfs8QEVCQ97voBA?pwd=aocg), [Google Drive](https://drive.google.com/drive/folders/1yELzm5SvAi9e7kPcO_jPp2XkTs4vK6aR?usp=sharing)
|
||||
|
||||
### Usage of SUPIR
|
||||
```Shell
|
||||
Usage:
|
||||
-- python test.py [options]
|
||||
-- python gradio_demo.py [interactive options]
|
||||
|
||||
--img_dir Input folder.
|
||||
--save_dir Output folder.
|
||||
--upscale Upsampling ratio of given inputs. Default: 1
|
||||
--SUPIR_sign Model selection. Default: 'Q'; Options: ['F', 'Q']
|
||||
--seed Random seed. Default: 1234
|
||||
--min_size Minimum resolution of output images. Default: 1024
|
||||
--edm_steps Numb of steps for EDM Sampling Scheduler. Default: 50
|
||||
--s_stage1 Control Strength of Stage1. Default: -1 (negative means invalid)
|
||||
--s_churn Original hy-param of EDM. Default: 5
|
||||
--s_noise Original hy-param of EDM. Default: 1.003
|
||||
--s_cfg Classifier-free guidance scale for prompts. Default: 7.5
|
||||
--s_stage2 Control Strength of Stage2. Default: 1.0
|
||||
--num_samples Number of samples for each input. Default: 1
|
||||
--a_prompt Additive positive prompt for all inputs.
|
||||
Default: 'Cinematic, High Contrast, highly detailed, taken using a Canon EOS R camera,
|
||||
hyper detailed photo - realistic maximum detail, 32k, Color Grading, ultra HD, extreme
|
||||
meticulous detailing, skin pore detailing, hyper sharpness, perfect without deformations.'
|
||||
--n_prompt Fixed negative prompt for all inputs.
|
||||
Default: 'painting, oil painting, illustration, drawing, art, sketch, oil painting,
|
||||
cartoon, CG Style, 3D render, unreal engine, blurring, dirty, messy, worst quality,
|
||||
low quality, frames, watermark, signature, jpeg artifacts, deformed, lowres, over-smooth'
|
||||
--color_fix_type Color Fixing Type. Default: 'Wavelet'; Options: ['None', 'AdaIn', 'Wavelet']
|
||||
--linear_CFG Linearly (with sigma) increase CFG from 'spt_linear_CFG' to s_cfg. Default: False
|
||||
--linear_s_stage2 Linearly (with sigma) increase s_stage2 from 'spt_linear_s_stage2' to s_stage2. Default: False
|
||||
--spt_linear_CFG Start point of linearly increasing CFG. Default: 1.0
|
||||
--spt_linear_s_stage2 Start point of linearly increasing s_stage2. Default: 0.0
|
||||
--ae_dtype Inference data type of AutoEncoder. Default: 'bf16'; Options: ['fp32', 'bf16']
|
||||
--diff_dtype Inference data type of Diffusion. Default: 'fp16'; Options: ['fp32', 'fp16', 'bf16']
|
||||
```
|
||||
|
||||
### Python Script
|
||||
```Shell
|
||||
# Seek for best quality for most cases
|
||||
CUDA_VISIBLE_DEVICES=0,1 python test.py --img_dir '/opt/data/private/LV_Dataset/DiffGLV-Test-All/RealPhoto60/LQ' --save_dir ./results-Q --SUPIR_sign Q --upscale 2
|
||||
# for light degradation and high fidelity
|
||||
CUDA_VISIBLE_DEVICES=0,1 python test.py --img_dir '/opt/data/private/LV_Dataset/DiffGLV-Test-All/RealPhoto60/LQ' --save_dir ./results-F --SUPIR_sign F --upscale 2 --s_cfg 4.0 --linear_CFG
|
||||
```
|
||||
|
||||
### Gradio Demo
|
||||
```Shell
|
||||
CUDA_VISIBLE_DEVICES=0,1 python gradio_demo.py --ip 0.0.0.0 --port 6688 --use_image_slider --log_history
|
||||
|
||||
# less VRAM & slower (12G for Diffusion, 16G for LLaVA)
|
||||
CUDA_VISIBLE_DEVICES=0,1 python gradio_demo.py --ip 0.0.0.0 --port 6688 --use_image_slider --log_history --loading_half_params --use_tile_vae --load_8bit_llava
|
||||
```
|
||||
<p align="center">
|
||||
<img src="assets/DemoGuide.png">
|
||||
</p>
|
||||
|
||||
|
||||
### Online Demo (Coming Soon)
|
||||
|
||||
|
||||
---
|
||||
|
||||
## BibTeX
|
||||
@misc{yu2024scaling,
|
||||
title={Scaling Up to Excellence: Practicing Model Scaling for Photo-Realistic Image Restoration In the Wild},
|
||||
author={Fanghua Yu and Jinjin Gu and Zheyuan Li and Jinfan Hu and Xiangtao Kong and Xintao Wang and Jingwen He and Yu Qiao and Chao Dong},
|
||||
year={2024},
|
||||
eprint={2401.13627},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV}
|
||||
}
|
||||
|
||||
---
|
||||
|
||||
## 📧 Contact
|
||||
If you have any question, please email `fanghuayu96@gmail.com`.
|
||||
|
||||
---
|
||||
## Non-Commercial Use Only Declaration
|
||||
The SUPIR ("Software") is made available for use, reproduction, and distribution strictly for non-commercial purposes. For the purposes of this declaration, "non-commercial" is defined as not primarily intended for or directed towards commercial advantage or monetary compensation.
|
||||
|
||||
By using, reproducing, or distributing the Software, you agree to abide by this restriction and not to use the Software for any commercial purposes without obtaining prior written permission from Dr. Jinjin Gu.
|
||||
|
||||
This declaration does not in any way limit the rights under any open source license that may apply to the Software; it solely adds a condition that the Software shall not be used for commercial purposes.
|
||||
|
||||
IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||
|
||||
For inquiries or to obtain permission for commercial use, please contact Dr. Jinjin Gu (hellojasongt@gmail.com).
|
||||
@@ -0,0 +1,183 @@
|
||||
import torch
|
||||
from ...sgm.models.diffusion import DiffusionEngine
|
||||
from ...sgm.util import instantiate_from_config
|
||||
import copy
|
||||
from ...sgm.modules.distributions.distributions import DiagonalGaussianDistribution
|
||||
import random
|
||||
from ...SUPIR.utils.colorfix import wavelet_reconstruction, adaptive_instance_normalization
|
||||
from pytorch_lightning import seed_everything
|
||||
from torch.nn.functional import interpolate
|
||||
from ...SUPIR.utils.tilevae import VAEHook
|
||||
import importlib
|
||||
import os
|
||||
|
||||
class SUPIRModel(DiffusionEngine):
|
||||
def __init__(self, control_stage_config, ae_dtype='fp32', diffusion_dtype='fp32', p_p='', n_p='', *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
control_model = instantiate_from_config(control_stage_config)
|
||||
self.model.load_control_model(control_model)
|
||||
self.first_stage_model.denoise_encoder = copy.deepcopy(self.first_stage_model.encoder)
|
||||
self.sampler_config = kwargs['sampler_config']
|
||||
|
||||
assert (ae_dtype in ['fp32', 'fp16', 'bf16']) and (diffusion_dtype in ['fp32', 'fp16', 'bf16'])
|
||||
if ae_dtype == 'fp32':
|
||||
ae_dtype = torch.float32
|
||||
elif ae_dtype == 'fp16':
|
||||
raise RuntimeError('fp16 cause NaN in AE')
|
||||
elif ae_dtype == 'bf16':
|
||||
ae_dtype = torch.bfloat16
|
||||
|
||||
if diffusion_dtype == 'fp32':
|
||||
diffusion_dtype = torch.float32
|
||||
elif diffusion_dtype == 'fp16':
|
||||
diffusion_dtype = torch.float16
|
||||
elif diffusion_dtype == 'bf16':
|
||||
diffusion_dtype = torch.bfloat16
|
||||
|
||||
self.ae_dtype = ae_dtype
|
||||
self.model.dtype = diffusion_dtype
|
||||
|
||||
self.p_p = p_p
|
||||
self.n_p = n_p
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x):
|
||||
with torch.autocast("cuda", dtype=self.ae_dtype):
|
||||
z = self.first_stage_model.encode(x)
|
||||
z = self.scale_factor * z
|
||||
return z
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage_with_denoise(self, x, use_sample=True, is_stage1=False):
|
||||
with torch.autocast("cuda", dtype=self.ae_dtype):
|
||||
if is_stage1:
|
||||
h = self.first_stage_model.denoise_encoder_s1(x)
|
||||
else:
|
||||
h = self.first_stage_model.denoise_encoder(x)
|
||||
moments = self.first_stage_model.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
if use_sample:
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z = posterior.mode()
|
||||
z = self.scale_factor * z
|
||||
return z
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
z = 1.0 / self.scale_factor * z
|
||||
with torch.autocast("cuda", dtype=self.ae_dtype):
|
||||
out = self.first_stage_model.decode(z)
|
||||
return out.float()
|
||||
|
||||
@torch.no_grad()
|
||||
def batchify_denoise(self, x, is_stage1=False):
|
||||
'''
|
||||
[N, C, H, W], [-1, 1], RGB
|
||||
'''
|
||||
x = self.encode_first_stage_with_denoise(x, use_sample=False, is_stage1=is_stage1)
|
||||
return self.decode_first_stage(x)
|
||||
|
||||
@torch.no_grad()
|
||||
def batchify_sample(self, x, p, p_p='default', n_p='default', num_steps=100, restoration_scale=4.0, s_churn=0, s_noise=1.003, cfg_scale=4.0, seed=-1,
|
||||
num_samples=1, control_scale=1, color_fix_type='None', use_linear_CFG=False, use_linear_control_scale=False,
|
||||
cfg_scale_start=1.0, control_scale_start=0.0, **kwargs):
|
||||
'''
|
||||
[N, C], [-1, 1], RGB
|
||||
'''
|
||||
assert len(x) == len(p)
|
||||
assert color_fix_type in ['Wavelet', 'AdaIn', 'None']
|
||||
|
||||
N = len(x)
|
||||
if num_samples > 1:
|
||||
assert N == 1
|
||||
N = num_samples
|
||||
x = x.repeat(N, 1, 1, 1)
|
||||
p = p * N
|
||||
|
||||
if p_p == 'default':
|
||||
p_p = self.p_p
|
||||
if n_p == 'default':
|
||||
n_p = self.n_p
|
||||
|
||||
self.sampler_config.params.num_steps = num_steps
|
||||
if use_linear_CFG:
|
||||
self.sampler_config.params.guider_config.params.scale_min = cfg_scale
|
||||
self.sampler_config.params.guider_config.params.scale = cfg_scale_start
|
||||
else:
|
||||
self.sampler_config.params.guider_config.params.scale = cfg_scale
|
||||
self.sampler_config.params.restore_cfg = restoration_scale
|
||||
self.sampler_config.params.s_churn = s_churn
|
||||
self.sampler_config.params.s_noise = s_noise
|
||||
self.sampler = instantiate_from_config(self.sampler_config)
|
||||
|
||||
if seed == -1:
|
||||
seed = random.randint(0, 65535)
|
||||
seed_everything(seed)
|
||||
|
||||
_z = self.encode_first_stage_with_denoise(x, use_sample=False)
|
||||
|
||||
x_stage1 = self.decode_first_stage(_z)
|
||||
# x_stage1 = interpolate(x_stage1, scale_factor=scale_factor, mode='bilinear', antialias=True)
|
||||
# _z = self.encode_first_stage_with_denoise(x_stage1)
|
||||
|
||||
z_stage1 = self.encode_first_stage(x_stage1)
|
||||
|
||||
batch = {}
|
||||
batch['txt'] = [''.join([_p, p_p]) for _p in p]
|
||||
batch['original_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(x.device)
|
||||
batch['crop_coords_top_left'] = torch.tensor([0, 0]).repeat(N, 1).to(x.device)
|
||||
batch['target_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(x.device)
|
||||
batch['aesthetic_score'] = torch.tensor([9.0]).repeat(N, 1).to(x.device)
|
||||
batch['control'] = _z
|
||||
|
||||
batch_uc = copy.deepcopy(batch)
|
||||
batch_uc['txt'] = [n_p for _ in p]
|
||||
|
||||
with torch.cuda.amp.autocast(dtype=self.ae_dtype):
|
||||
c, uc = self.conditioner.get_unconditional_conditioning(batch, batch_uc)
|
||||
|
||||
denoiser = lambda input, sigma, c, control_scale: self.denoiser(
|
||||
self.model, input, sigma, c, control_scale, **kwargs
|
||||
)
|
||||
|
||||
noised_z = torch.randn_like(_z).to(_z.device)
|
||||
|
||||
_samples = self.sampler(denoiser, noised_z, cond=c, uc=uc, x_center=z_stage1, control_scale=control_scale,
|
||||
use_linear_control_scale=use_linear_control_scale, control_scale_start=control_scale_start)
|
||||
samples = self.decode_first_stage(_samples)
|
||||
if color_fix_type == 'Wavelet':
|
||||
samples = wavelet_reconstruction(samples, x_stage1)
|
||||
elif color_fix_type == 'AdaIn':
|
||||
samples = adaptive_instance_normalization(samples, x_stage1)
|
||||
return samples
|
||||
|
||||
def init_tile_vae(self, encoder_tile_size=512, decoder_tile_size=64):
|
||||
self.first_stage_model.denoise_encoder.original_forward = self.first_stage_model.denoise_encoder.forward
|
||||
self.first_stage_model.encoder.original_forward = self.first_stage_model.encoder.forward
|
||||
self.first_stage_model.decoder.original_forward = self.first_stage_model.decoder.forward
|
||||
self.first_stage_model.denoise_encoder.forward = VAEHook(
|
||||
self.first_stage_model.denoise_encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
|
||||
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||
self.first_stage_model.encoder.forward = VAEHook(
|
||||
self.first_stage_model.encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
|
||||
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||
self.first_stage_model.decoder.forward = VAEHook(
|
||||
self.first_stage_model.decoder, decoder_tile_size, is_decoder=True, fast_decoder=False,
|
||||
fast_encoder=False, color_fix=False, to_gpu=True)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
from SUPIR.util import create_model, load_state_dict
|
||||
|
||||
model = create_model('../../options/dev/SUPIR_paper_version.yaml')
|
||||
|
||||
SDXL_CKPT = '/opt/data/private/AIGC_pretrain/SDXL_cache/sd_xl_base_1.0_0.9vae.safetensors'
|
||||
SUPIR_CKPT = '/opt/data/private/AIGC_pretrain/SUPIR_cache/SUPIR-paper.ckpt'
|
||||
model.load_state_dict(load_state_dict(SDXL_CKPT), strict=False)
|
||||
model.load_state_dict(load_state_dict(SUPIR_CKPT), strict=False)
|
||||
model = model.cuda()
|
||||
|
||||
x = torch.randn(1, 3, 512, 512).cuda()
|
||||
p = ['a professional, detailed, high-quality photo']
|
||||
samples = model.batchify_sample(x, p, num_steps=50, restoration_scale=4.0, s_churn=0, cfg_scale=4.0, seed=-1, num_samples=1)
|
||||
@@ -0,0 +1,718 @@
|
||||
# from einops._torch_specific import allow_ops_in_compiled_graph
|
||||
# allow_ops_in_compiled_graph()
|
||||
import einops
|
||||
import torch
|
||||
import torch as th
|
||||
import torch.nn as nn
|
||||
from einops import rearrange, repeat
|
||||
|
||||
from ...sgm.modules.diffusionmodules.util import (
|
||||
avg_pool_nd,
|
||||
checkpoint,
|
||||
conv_nd,
|
||||
linear,
|
||||
normalization,
|
||||
timestep_embedding,
|
||||
zero_module,
|
||||
)
|
||||
|
||||
from ...sgm.modules.diffusionmodules.openaimodel import Downsample, Upsample, UNetModel, Timestep, \
|
||||
TimestepEmbedSequential, ResBlock, AttentionBlock, TimestepBlock
|
||||
from ...sgm.modules.attention import SpatialTransformer, MemoryEfficientCrossAttention, CrossAttention
|
||||
from ...sgm.util import default, log_txt_as_img, exists, instantiate_from_config
|
||||
import re
|
||||
import torch
|
||||
from functools import partial
|
||||
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
XFORMERS_IS_AVAILBLE = True
|
||||
except:
|
||||
XFORMERS_IS_AVAILBLE = False
|
||||
|
||||
|
||||
# dummy replace
|
||||
def convert_module_to_f16(x):
|
||||
pass
|
||||
|
||||
|
||||
def convert_module_to_f32(x):
|
||||
pass
|
||||
|
||||
|
||||
class ZeroConv(nn.Module):
|
||||
def __init__(self, label_nc, norm_nc, mask=False):
|
||||
super().__init__()
|
||||
self.zero_conv = zero_module(conv_nd(2, label_nc, norm_nc, 1, 1, 0))
|
||||
self.mask = mask
|
||||
|
||||
def forward(self, c, h, h_ori=None):
|
||||
# with torch.cuda.amp.autocast(enabled=False, dtype=torch.float32):
|
||||
if not self.mask:
|
||||
h = h + self.zero_conv(c)
|
||||
else:
|
||||
h = h + self.zero_conv(c) * torch.zeros_like(h)
|
||||
if h_ori is not None:
|
||||
h = th.cat([h_ori, h], dim=1)
|
||||
return h
|
||||
|
||||
|
||||
class ZeroSFT(nn.Module):
|
||||
def __init__(self, label_nc, norm_nc, concat_channels=0, norm=True, mask=False):
|
||||
super().__init__()
|
||||
|
||||
# param_free_norm_type = str(parsed.group(1))
|
||||
ks = 3
|
||||
pw = ks // 2
|
||||
|
||||
self.norm = norm
|
||||
if self.norm:
|
||||
self.param_free_norm = normalization(norm_nc + concat_channels)
|
||||
else:
|
||||
self.param_free_norm = nn.Identity()
|
||||
|
||||
nhidden = 128
|
||||
|
||||
self.mlp_shared = nn.Sequential(
|
||||
nn.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw),
|
||||
nn.SiLU()
|
||||
)
|
||||
self.zero_mul = zero_module(nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
|
||||
self.zero_add = zero_module(nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
|
||||
# self.zero_mul = nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
|
||||
# self.zero_add = nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
|
||||
|
||||
self.zero_conv = zero_module(conv_nd(2, label_nc, norm_nc, 1, 1, 0))
|
||||
self.pre_concat = bool(concat_channels != 0)
|
||||
self.mask = mask
|
||||
|
||||
def forward(self, c, h, h_ori=None, control_scale=1):
|
||||
assert self.mask is False
|
||||
if h_ori is not None and self.pre_concat:
|
||||
h_raw = th.cat([h_ori, h], dim=1)
|
||||
else:
|
||||
h_raw = h
|
||||
|
||||
if self.mask:
|
||||
h = h + self.zero_conv(c) * torch.zeros_like(h)
|
||||
else:
|
||||
h = h + self.zero_conv(c)
|
||||
if h_ori is not None and self.pre_concat:
|
||||
h = th.cat([h_ori, h], dim=1)
|
||||
actv = self.mlp_shared(c)
|
||||
gamma = self.zero_mul(actv)
|
||||
beta = self.zero_add(actv)
|
||||
if self.mask:
|
||||
gamma = gamma * torch.zeros_like(gamma)
|
||||
beta = beta * torch.zeros_like(beta)
|
||||
h = self.param_free_norm(h) * (gamma + 1) + beta
|
||||
if h_ori is not None and not self.pre_concat:
|
||||
h = th.cat([h_ori, h], dim=1)
|
||||
return h * control_scale + h_raw * (1 - control_scale)
|
||||
|
||||
|
||||
class ZeroCrossAttn(nn.Module):
|
||||
ATTENTION_MODES = {
|
||||
"softmax": CrossAttention, # vanilla attention
|
||||
"softmax-xformers": MemoryEfficientCrossAttention
|
||||
}
|
||||
|
||||
def __init__(self, context_dim, query_dim, zero_out=True, mask=False):
|
||||
super().__init__()
|
||||
attn_mode = "softmax-xformers" if XFORMERS_IS_AVAILBLE else "softmax"
|
||||
assert attn_mode in self.ATTENTION_MODES
|
||||
attn_cls = self.ATTENTION_MODES[attn_mode]
|
||||
self.attn = attn_cls(query_dim=query_dim, context_dim=context_dim, heads=query_dim//64, dim_head=64)
|
||||
self.norm1 = normalization(query_dim)
|
||||
self.norm2 = normalization(context_dim)
|
||||
|
||||
self.mask = mask
|
||||
|
||||
# if zero_out:
|
||||
# # for p in self.attn.to_out.parameters():
|
||||
# # p.detach().zero_()
|
||||
# self.attn.to_out = zero_module(self.attn.to_out)
|
||||
|
||||
def forward(self, context, x, control_scale=1):
|
||||
assert self.mask is False
|
||||
x_in = x
|
||||
x = self.norm1(x)
|
||||
context = self.norm2(context)
|
||||
b, c, h, w = x.shape
|
||||
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
|
||||
context = rearrange(context, 'b c h w -> b (h w) c').contiguous()
|
||||
x = self.attn(x, context)
|
||||
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
|
||||
if self.mask:
|
||||
x = x * torch.zeros_like(x)
|
||||
x = x_in + x * control_scale
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class GLVControl(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
model_channels,
|
||||
out_channels,
|
||||
num_res_blocks,
|
||||
attention_resolutions,
|
||||
dropout=0,
|
||||
channel_mult=(1, 2, 4, 8),
|
||||
conv_resample=True,
|
||||
dims=2,
|
||||
num_classes=None,
|
||||
use_checkpoint=False,
|
||||
use_fp16=False,
|
||||
num_heads=-1,
|
||||
num_head_channels=-1,
|
||||
num_heads_upsample=-1,
|
||||
use_scale_shift_norm=False,
|
||||
resblock_updown=False,
|
||||
use_new_attention_order=False,
|
||||
use_spatial_transformer=False, # custom transformer support
|
||||
transformer_depth=1, # custom transformer support
|
||||
context_dim=None, # custom transformer support
|
||||
n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model
|
||||
legacy=True,
|
||||
disable_self_attentions=None,
|
||||
num_attention_blocks=None,
|
||||
disable_middle_self_attn=False,
|
||||
use_linear_in_transformer=False,
|
||||
spatial_transformer_attn_type="softmax",
|
||||
adm_in_channels=None,
|
||||
use_fairscale_checkpoint=False,
|
||||
offload_to_cpu=False,
|
||||
transformer_depth_middle=None,
|
||||
input_upscale=1,
|
||||
):
|
||||
super().__init__()
|
||||
from omegaconf.listconfig import ListConfig
|
||||
|
||||
if use_spatial_transformer:
|
||||
assert (
|
||||
context_dim is not None
|
||||
), "Fool!! You forgot to include the dimension of your cross-attention conditioning..."
|
||||
|
||||
if context_dim is not None:
|
||||
assert (
|
||||
use_spatial_transformer
|
||||
), "Fool!! You forgot to use the spatial transformer for your cross-attention conditioning..."
|
||||
if type(context_dim) == ListConfig:
|
||||
context_dim = list(context_dim)
|
||||
|
||||
if num_heads_upsample == -1:
|
||||
num_heads_upsample = num_heads
|
||||
|
||||
if num_heads == -1:
|
||||
assert (
|
||||
num_head_channels != -1
|
||||
), "Either num_heads or num_head_channels has to be set"
|
||||
|
||||
if num_head_channels == -1:
|
||||
assert (
|
||||
num_heads != -1
|
||||
), "Either num_heads or num_head_channels has to be set"
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.model_channels = model_channels
|
||||
self.out_channels = out_channels
|
||||
if isinstance(transformer_depth, int):
|
||||
transformer_depth = len(channel_mult) * [transformer_depth]
|
||||
elif isinstance(transformer_depth, ListConfig):
|
||||
transformer_depth = list(transformer_depth)
|
||||
transformer_depth_middle = default(
|
||||
transformer_depth_middle, transformer_depth[-1]
|
||||
)
|
||||
|
||||
if isinstance(num_res_blocks, int):
|
||||
self.num_res_blocks = len(channel_mult) * [num_res_blocks]
|
||||
else:
|
||||
if len(num_res_blocks) != len(channel_mult):
|
||||
raise ValueError(
|
||||
"provide num_res_blocks either as an int (globally constant) or "
|
||||
"as a list/tuple (per-level) with the same length as channel_mult"
|
||||
)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
# self.num_res_blocks = num_res_blocks
|
||||
if disable_self_attentions is not None:
|
||||
# should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not
|
||||
assert len(disable_self_attentions) == len(channel_mult)
|
||||
if num_attention_blocks is not None:
|
||||
assert len(num_attention_blocks) == len(self.num_res_blocks)
|
||||
assert all(
|
||||
map(
|
||||
lambda i: self.num_res_blocks[i] >= num_attention_blocks[i],
|
||||
range(len(num_attention_blocks)),
|
||||
)
|
||||
)
|
||||
print(
|
||||
f"Constructor of UNetModel received num_attention_blocks={num_attention_blocks}. "
|
||||
f"This option has LESS priority than attention_resolutions {attention_resolutions}, "
|
||||
f"i.e., in cases where num_attention_blocks[i] > 0 but 2**i not in attention_resolutions, "
|
||||
f"attention will still not be set."
|
||||
) # todo: convert to warning
|
||||
|
||||
self.attention_resolutions = attention_resolutions
|
||||
self.dropout = dropout
|
||||
self.channel_mult = channel_mult
|
||||
self.conv_resample = conv_resample
|
||||
self.num_classes = num_classes
|
||||
self.use_checkpoint = use_checkpoint
|
||||
if use_fp16:
|
||||
print("WARNING: use_fp16 was dropped and has no effect anymore.")
|
||||
# self.dtype = th.float16 if use_fp16 else th.float32
|
||||
self.num_heads = num_heads
|
||||
self.num_head_channels = num_head_channels
|
||||
self.num_heads_upsample = num_heads_upsample
|
||||
self.predict_codebook_ids = n_embed is not None
|
||||
|
||||
assert use_fairscale_checkpoint != use_checkpoint or not (
|
||||
use_checkpoint or use_fairscale_checkpoint
|
||||
)
|
||||
|
||||
self.use_fairscale_checkpoint = False
|
||||
checkpoint_wrapper_fn = (
|
||||
partial(checkpoint_wrapper, offload_to_cpu=offload_to_cpu)
|
||||
if self.use_fairscale_checkpoint
|
||||
else lambda x: x
|
||||
)
|
||||
|
||||
time_embed_dim = model_channels * 4
|
||||
self.time_embed = checkpoint_wrapper_fn(
|
||||
nn.Sequential(
|
||||
linear(model_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
)
|
||||
)
|
||||
|
||||
if self.num_classes is not None:
|
||||
if isinstance(self.num_classes, int):
|
||||
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
|
||||
elif self.num_classes == "continuous":
|
||||
print("setting up linear c_adm embedding layer")
|
||||
self.label_emb = nn.Linear(1, time_embed_dim)
|
||||
elif self.num_classes == "timestep":
|
||||
self.label_emb = checkpoint_wrapper_fn(
|
||||
nn.Sequential(
|
||||
Timestep(model_channels),
|
||||
nn.Sequential(
|
||||
linear(model_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
),
|
||||
)
|
||||
)
|
||||
elif self.num_classes == "sequential":
|
||||
assert adm_in_channels is not None
|
||||
self.label_emb = nn.Sequential(
|
||||
nn.Sequential(
|
||||
linear(adm_in_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise ValueError()
|
||||
|
||||
self.input_blocks = nn.ModuleList(
|
||||
[
|
||||
TimestepEmbedSequential(
|
||||
conv_nd(dims, in_channels, model_channels, 3, padding=1)
|
||||
)
|
||||
]
|
||||
)
|
||||
self._feature_size = model_channels
|
||||
input_block_chans = [model_channels]
|
||||
ch = model_channels
|
||||
ds = 1
|
||||
for level, mult in enumerate(channel_mult):
|
||||
for nr in range(self.num_res_blocks[level]):
|
||||
layers = [
|
||||
checkpoint_wrapper_fn(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=mult * model_channels,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
)
|
||||
]
|
||||
ch = mult * model_channels
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
if legacy:
|
||||
# num_heads = 1
|
||||
dim_head = (
|
||||
ch // num_heads
|
||||
if use_spatial_transformer
|
||||
else num_head_channels
|
||||
)
|
||||
if exists(disable_self_attentions):
|
||||
disabled_sa = disable_self_attentions[level]
|
||||
else:
|
||||
disabled_sa = False
|
||||
|
||||
if (
|
||||
not exists(num_attention_blocks)
|
||||
or nr < num_attention_blocks[level]
|
||||
):
|
||||
layers.append(
|
||||
checkpoint_wrapper_fn(
|
||||
AttentionBlock(
|
||||
ch,
|
||||
use_checkpoint=use_checkpoint,
|
||||
num_heads=num_heads,
|
||||
num_head_channels=dim_head,
|
||||
use_new_attention_order=use_new_attention_order,
|
||||
)
|
||||
)
|
||||
if not use_spatial_transformer
|
||||
else checkpoint_wrapper_fn(
|
||||
SpatialTransformer(
|
||||
ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
depth=transformer_depth[level],
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa,
|
||||
use_linear=use_linear_in_transformer,
|
||||
attn_type=spatial_transformer_attn_type,
|
||||
use_checkpoint=use_checkpoint,
|
||||
)
|
||||
)
|
||||
)
|
||||
self.input_blocks.append(TimestepEmbedSequential(*layers))
|
||||
self._feature_size += ch
|
||||
input_block_chans.append(ch)
|
||||
if level != len(channel_mult) - 1:
|
||||
out_ch = ch
|
||||
self.input_blocks.append(
|
||||
TimestepEmbedSequential(
|
||||
checkpoint_wrapper_fn(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=out_ch,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
down=True,
|
||||
)
|
||||
)
|
||||
if resblock_updown
|
||||
else Downsample(
|
||||
ch, conv_resample, dims=dims, out_channels=out_ch
|
||||
)
|
||||
)
|
||||
)
|
||||
ch = out_ch
|
||||
input_block_chans.append(ch)
|
||||
ds *= 2
|
||||
self._feature_size += ch
|
||||
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
if legacy:
|
||||
# num_heads = 1
|
||||
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
|
||||
self.middle_block = TimestepEmbedSequential(
|
||||
checkpoint_wrapper_fn(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
),
|
||||
checkpoint_wrapper_fn(
|
||||
AttentionBlock(
|
||||
ch,
|
||||
use_checkpoint=use_checkpoint,
|
||||
num_heads=num_heads,
|
||||
num_head_channels=dim_head,
|
||||
use_new_attention_order=use_new_attention_order,
|
||||
)
|
||||
)
|
||||
if not use_spatial_transformer
|
||||
else checkpoint_wrapper_fn(
|
||||
SpatialTransformer( # always uses a self-attn
|
||||
ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
depth=transformer_depth_middle,
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disable_middle_self_attn,
|
||||
use_linear=use_linear_in_transformer,
|
||||
attn_type=spatial_transformer_attn_type,
|
||||
use_checkpoint=use_checkpoint,
|
||||
)
|
||||
),
|
||||
checkpoint_wrapper_fn(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
self.input_upscale = input_upscale
|
||||
self.input_hint_block = TimestepEmbedSequential(
|
||||
zero_module(conv_nd(dims, in_channels, model_channels, 3, padding=1))
|
||||
)
|
||||
|
||||
def convert_to_fp16(self):
|
||||
"""
|
||||
Convert the torso of the model to float16.
|
||||
"""
|
||||
self.input_blocks.apply(convert_module_to_f16)
|
||||
self.middle_block.apply(convert_module_to_f16)
|
||||
|
||||
def convert_to_fp32(self):
|
||||
"""
|
||||
Convert the torso of the model to float32.
|
||||
"""
|
||||
self.input_blocks.apply(convert_module_to_f32)
|
||||
self.middle_block.apply(convert_module_to_f32)
|
||||
|
||||
def forward(self, x, timesteps, xt, context=None, y=None, **kwargs):
|
||||
# with torch.cuda.amp.autocast(enabled=False, dtype=torch.float32):
|
||||
# x = x.to(torch.float32)
|
||||
# timesteps = timesteps.to(torch.float32)
|
||||
# xt = xt.to(torch.float32)
|
||||
# context = context.to(torch.float32)
|
||||
# y = y.to(torch.float32)
|
||||
# print(x.dtype)
|
||||
xt, context, y = xt.to(x.dtype), context.to(x.dtype), y.to(x.dtype)
|
||||
|
||||
if self.input_upscale != 1:
|
||||
x = nn.functional.interpolate(x, scale_factor=self.input_upscale, mode='bilinear', antialias=True)
|
||||
assert (y is not None) == (
|
||||
self.num_classes is not None
|
||||
), "must specify y if and only if the model is class-conditional"
|
||||
hs = []
|
||||
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
|
||||
# import pdb
|
||||
# pdb.set_trace()
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
if self.num_classes is not None:
|
||||
assert y.shape[0] == xt.shape[0]
|
||||
emb = emb + self.label_emb(y)
|
||||
|
||||
guided_hint = self.input_hint_block(x, emb, context)
|
||||
|
||||
# h = x.type(self.dtype)
|
||||
h = xt
|
||||
for module in self.input_blocks:
|
||||
if guided_hint is not None:
|
||||
h = module(h, emb, context)
|
||||
h += guided_hint
|
||||
guided_hint = None
|
||||
else:
|
||||
h = module(h, emb, context)
|
||||
hs.append(h)
|
||||
# print(module)
|
||||
# print(h.shape)
|
||||
h = self.middle_block(h, emb, context)
|
||||
hs.append(h)
|
||||
return hs
|
||||
|
||||
|
||||
class LightGLVUNet(UNetModel):
|
||||
def __init__(self, mode='', project_type='ZeroSFT', project_channel_scale=1,
|
||||
*args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if mode == 'XL-base':
|
||||
cond_output_channels = [320] * 4 + [640] * 3 + [1280] * 3
|
||||
project_channels = [160] * 4 + [320] * 3 + [640] * 3
|
||||
concat_channels = [320] * 2 + [640] * 3 + [1280] * 4 + [0]
|
||||
cross_attn_insert_idx = [6, 3]
|
||||
self.progressive_mask_nums = [0, 3, 7, 11]
|
||||
elif mode == 'XL-refine':
|
||||
cond_output_channels = [384] * 4 + [768] * 3 + [1536] * 6
|
||||
project_channels = [192] * 4 + [384] * 3 + [768] * 6
|
||||
concat_channels = [384] * 2 + [768] * 3 + [1536] * 7 + [0]
|
||||
cross_attn_insert_idx = [9, 6, 3]
|
||||
self.progressive_mask_nums = [0, 3, 6, 10, 14]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
project_channels = [int(c * project_channel_scale) for c in project_channels]
|
||||
|
||||
self.project_modules = nn.ModuleList()
|
||||
for i in range(len(cond_output_channels)):
|
||||
# if i == len(cond_output_channels) - 1:
|
||||
# _project_type = 'ZeroCrossAttn'
|
||||
# else:
|
||||
# _project_type = project_type
|
||||
_project_type = project_type
|
||||
if _project_type == 'ZeroSFT':
|
||||
self.project_modules.append(ZeroSFT(project_channels[i], cond_output_channels[i],
|
||||
concat_channels=concat_channels[i]))
|
||||
elif _project_type == 'ZeroCrossAttn':
|
||||
self.project_modules.append(ZeroCrossAttn(cond_output_channels[i], project_channels[i]))
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
for i in cross_attn_insert_idx:
|
||||
self.project_modules.insert(i, ZeroCrossAttn(cond_output_channels[i], concat_channels[i]))
|
||||
# print(self.project_modules[i])
|
||||
|
||||
def step_progressive_mask(self):
|
||||
if len(self.progressive_mask_nums) > 0:
|
||||
mask_num = self.progressive_mask_nums.pop()
|
||||
for i in range(len(self.project_modules)):
|
||||
if i < mask_num:
|
||||
self.project_modules[i].mask = True
|
||||
else:
|
||||
self.project_modules[i].mask = False
|
||||
return
|
||||
# print(f'step_progressive_mask, current masked layers: {mask_num}')
|
||||
else:
|
||||
return
|
||||
# print('step_progressive_mask, no more masked layers')
|
||||
# for i in range(len(self.project_modules)):
|
||||
# print(self.project_modules[i].mask)
|
||||
|
||||
|
||||
def forward(self, x, timesteps=None, context=None, y=None, control=None, control_scale=1, **kwargs):
|
||||
"""
|
||||
Apply the model to an input batch.
|
||||
:param x: an [N x C x ...] Tensor of inputs.
|
||||
:param timesteps: a 1-D batch of timesteps.
|
||||
:param context: conditioning plugged in via crossattn
|
||||
:param y: an [N] Tensor of labels, if class-conditional.
|
||||
:return: an [N x C x ...] Tensor of outputs.
|
||||
"""
|
||||
assert (y is not None) == (
|
||||
self.num_classes is not None
|
||||
), "must specify y if and only if the model is class-conditional"
|
||||
hs = []
|
||||
|
||||
_dtype = control[0].dtype
|
||||
x, context, y = x.to(_dtype), context.to(_dtype), y.to(_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
if self.num_classes is not None:
|
||||
assert y.shape[0] == x.shape[0]
|
||||
emb = emb + self.label_emb(y)
|
||||
|
||||
# h = x.type(self.dtype)
|
||||
h = x
|
||||
for module in self.input_blocks:
|
||||
h = module(h, emb, context)
|
||||
hs.append(h)
|
||||
|
||||
adapter_idx = len(self.project_modules) - 1
|
||||
control_idx = len(control) - 1
|
||||
h = self.middle_block(h, emb, context)
|
||||
h = self.project_modules[adapter_idx](control[control_idx], h, control_scale=control_scale)
|
||||
adapter_idx -= 1
|
||||
control_idx -= 1
|
||||
|
||||
for i, module in enumerate(self.output_blocks):
|
||||
_h = hs.pop()
|
||||
h = self.project_modules[adapter_idx](control[control_idx], _h, h, control_scale=control_scale)
|
||||
adapter_idx -= 1
|
||||
# h = th.cat([h, _h], dim=1)
|
||||
if len(module) == 3:
|
||||
assert isinstance(module[2], Upsample)
|
||||
for layer in module[:2]:
|
||||
if isinstance(layer, TimestepBlock):
|
||||
h = layer(h, emb)
|
||||
elif isinstance(layer, SpatialTransformer):
|
||||
h = layer(h, context)
|
||||
else:
|
||||
h = layer(h)
|
||||
# print('cross_attn_here')
|
||||
h = self.project_modules[adapter_idx](control[control_idx], h, control_scale=control_scale)
|
||||
adapter_idx -= 1
|
||||
h = module[2](h)
|
||||
else:
|
||||
h = module(h, emb, context)
|
||||
control_idx -= 1
|
||||
# print(module)
|
||||
# print(h.shape)
|
||||
|
||||
h = h.type(x.dtype)
|
||||
if self.predict_codebook_ids:
|
||||
assert False, "not supported anymore. what the f*** are you doing?"
|
||||
else:
|
||||
return self.out(h)
|
||||
|
||||
if __name__ == '__main__':
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
# refiner
|
||||
# opt = OmegaConf.load('../../options/train/debug_p2_xl.yaml')
|
||||
#
|
||||
# model = instantiate_from_config(opt.model.params.control_stage_config)
|
||||
# hint = model(torch.randn([1, 4, 64, 64]), torch.randn([1]), torch.randn([1, 4, 64, 64]))
|
||||
# hint = [h.cuda() for h in hint]
|
||||
# print(sum(map(lambda hint: hint.numel(), model.parameters())))
|
||||
#
|
||||
# unet = instantiate_from_config(opt.model.params.network_config)
|
||||
# unet = unet.cuda()
|
||||
#
|
||||
# _output = unet(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 77, 1280]).cuda(),
|
||||
# torch.randn([1, 2560]).cuda(), hint)
|
||||
# print(sum(map(lambda _output: _output.numel(), unet.parameters())))
|
||||
|
||||
# base
|
||||
with torch.no_grad():
|
||||
opt = OmegaConf.load('../../options/dev/SUPIR_tmp.yaml')
|
||||
|
||||
model = instantiate_from_config(opt.model.params.control_stage_config)
|
||||
model = model.cuda()
|
||||
|
||||
hint = model(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1, 77, 2048]).cuda(),
|
||||
torch.randn([1, 2816]).cuda())
|
||||
|
||||
for h in hint:
|
||||
print(h.shape)
|
||||
#
|
||||
unet = instantiate_from_config(opt.model.params.network_config)
|
||||
unet = unet.cuda()
|
||||
_output = unet(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 77, 2048]).cuda(),
|
||||
torch.randn([1, 2816]).cuda(), hint)
|
||||
|
||||
|
||||
# model = instantiate_from_config(opt.model.params.control_stage_config)
|
||||
# model = model.cuda()
|
||||
# # hint = model(torch.randn([1, 4, 64, 64]), torch.randn([1]), torch.randn([1, 4, 64, 64]))
|
||||
# hint = model(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1, 77, 1280]).cuda(),
|
||||
# torch.randn([1, 2560]).cuda())
|
||||
# # hint = [h.cuda() for h in hint]
|
||||
#
|
||||
# for h in hint:
|
||||
# print(h.shape)
|
||||
#
|
||||
# unet = instantiate_from_config(opt.model.params.network_config)
|
||||
# unet = unet.cuda()
|
||||
# _output = unet(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 77, 1280]).cuda(),
|
||||
# torch.randn([1, 2560]).cuda(), hint)
|
||||
@@ -0,0 +1,11 @@
|
||||
SDXL_BASE_CHANNEL_DICT = {
|
||||
'cond_output_channels': [320] * 4 + [640] * 3 + [1280] * 3,
|
||||
'project_channels': [160] * 4 + [320] * 3 + [640] * 3,
|
||||
'concat_channels': [320] * 2 + [640] * 3 + [1280] * 4 + [0]
|
||||
}
|
||||
|
||||
SDXL_REFINE_CHANNEL_DICT = {
|
||||
'cond_output_channels': [384] * 4 + [768] * 3 + [1536] * 6,
|
||||
'project_channels': [192] * 4 + [384] * 3 + [768] * 6,
|
||||
'concat_channels': [384] * 2 + [768] * 3 + [1536] * 7 + [0]
|
||||
}
|
||||
+173
@@ -0,0 +1,173 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
import cv2
|
||||
from PIL import Image
|
||||
from torch.nn.functional import interpolate
|
||||
from omegaconf import OmegaConf
|
||||
from ..sgm.util import instantiate_from_config
|
||||
|
||||
|
||||
def get_state_dict(d):
|
||||
return d.get('state_dict', d)
|
||||
|
||||
|
||||
def load_state_dict(ckpt_path, location='cpu'):
|
||||
_, extension = os.path.splitext(ckpt_path)
|
||||
if extension.lower() == ".safetensors":
|
||||
import safetensors.torch
|
||||
state_dict = safetensors.torch.load_file(ckpt_path, device=location)
|
||||
else:
|
||||
state_dict = get_state_dict(torch.load(ckpt_path, map_location=torch.device(location)))
|
||||
state_dict = get_state_dict(state_dict)
|
||||
print(f'Loaded state_dict from [{ckpt_path}]')
|
||||
return state_dict
|
||||
|
||||
|
||||
def create_model(config_path):
|
||||
config = OmegaConf.load(config_path)
|
||||
model = instantiate_from_config(config.model).cpu()
|
||||
print(f'Loaded model config from [{config_path}]')
|
||||
return model
|
||||
|
||||
|
||||
def create_SUPIR_model(config_path, SUPIR_sign=None):
|
||||
config = OmegaConf.load(config_path)
|
||||
model = instantiate_from_config(config.model).cpu()
|
||||
print(f'Loaded model config from [{config_path}]')
|
||||
if config.SDXL_CKPT is not None:
|
||||
model.load_state_dict(load_state_dict(config.SDXL_CKPT), strict=False)
|
||||
if config.SUPIR_CKPT is not None:
|
||||
model.load_state_dict(load_state_dict(config.SUPIR_CKPT), strict=False)
|
||||
if SUPIR_sign is not None:
|
||||
assert SUPIR_sign in ['F', 'Q']
|
||||
if SUPIR_sign == 'F':
|
||||
model.load_state_dict(load_state_dict(config.SUPIR_CKPT_F), strict=False)
|
||||
elif SUPIR_sign == 'Q':
|
||||
model.load_state_dict(load_state_dict(config.SUPIR_CKPT_Q), strict=False)
|
||||
return model
|
||||
|
||||
def load_QF_ckpt(config_path):
|
||||
config = OmegaConf.load(config_path)
|
||||
ckpt_F = torch.load(config.SUPIR_CKPT_F, map_location='cpu')
|
||||
ckpt_Q = torch.load(config.SUPIR_CKPT_Q, map_location='cpu')
|
||||
return ckpt_Q, ckpt_F
|
||||
|
||||
|
||||
def PIL2Tensor(img, upsacle=1, min_size=1024):
|
||||
'''
|
||||
PIL.Image -> Tensor[C, H, W], RGB, [-1, 1]
|
||||
'''
|
||||
# size
|
||||
w, h = img.size
|
||||
w *= upsacle
|
||||
h *= upsacle
|
||||
w0, h0 = round(w), round(h)
|
||||
if min(w, h) < min_size:
|
||||
_upsacle = min_size / min(w, h)
|
||||
w *= _upsacle
|
||||
h *= _upsacle
|
||||
else:
|
||||
_upsacle = 1
|
||||
w = int(np.round(w / 64.0)) * 64
|
||||
h = int(np.round(h / 64.0)) * 64
|
||||
x = img.resize((w, h), Image.BICUBIC)
|
||||
x = np.array(x).round().clip(0, 255).astype(np.uint8)
|
||||
x = x / 255 * 2 - 1
|
||||
x = torch.tensor(x, dtype=torch.float32).permute(2, 0, 1)
|
||||
return x, h0, w0
|
||||
|
||||
|
||||
def Tensor2PIL(x, h0, w0):
|
||||
'''
|
||||
Tensor[C, H, W], RGB, [-1, 1] -> PIL.Image
|
||||
'''
|
||||
x = x.unsqueeze(0)
|
||||
x = interpolate(x, size=(h0, w0), mode='bicubic')
|
||||
x = (x.squeeze(0).permute(1, 2, 0) * 127.5 + 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
return Image.fromarray(x)
|
||||
|
||||
|
||||
def HWC3(x):
|
||||
assert x.dtype == np.uint8
|
||||
if x.ndim == 2:
|
||||
x = x[:, :, None]
|
||||
assert x.ndim == 3
|
||||
H, W, C = x.shape
|
||||
assert C == 1 or C == 3 or C == 4
|
||||
if C == 3:
|
||||
return x
|
||||
if C == 1:
|
||||
return np.concatenate([x, x, x], axis=2)
|
||||
if C == 4:
|
||||
color = x[:, :, 0:3].astype(np.float32)
|
||||
alpha = x[:, :, 3:4].astype(np.float32) / 255.0
|
||||
y = color * alpha + 255.0 * (1.0 - alpha)
|
||||
y = y.clip(0, 255).astype(np.uint8)
|
||||
return y
|
||||
|
||||
|
||||
def upscale_image(input_image, upscale, min_size=None, unit_resolution=64):
|
||||
H, W, C = input_image.shape
|
||||
H = float(H)
|
||||
W = float(W)
|
||||
H *= upscale
|
||||
W *= upscale
|
||||
if min_size is not None:
|
||||
if min(H, W) < min_size:
|
||||
_upsacle = min_size / min(W, H)
|
||||
W *= _upsacle
|
||||
H *= _upsacle
|
||||
H = int(np.round(H / unit_resolution)) * unit_resolution
|
||||
W = int(np.round(W / unit_resolution)) * unit_resolution
|
||||
img = cv2.resize(input_image, (W, H), interpolation=cv2.INTER_LANCZOS4 if upscale > 1 else cv2.INTER_AREA)
|
||||
img = img.round().clip(0, 255).astype(np.uint8)
|
||||
return img
|
||||
|
||||
|
||||
def fix_resize(input_image, size=512, unit_resolution=64):
|
||||
H, W, C = input_image.shape
|
||||
H = float(H)
|
||||
W = float(W)
|
||||
upscale = size / min(H, W)
|
||||
H *= upscale
|
||||
W *= upscale
|
||||
H = int(np.round(H / unit_resolution)) * unit_resolution
|
||||
W = int(np.round(W / unit_resolution)) * unit_resolution
|
||||
img = cv2.resize(input_image, (W, H), interpolation=cv2.INTER_LANCZOS4 if upscale > 1 else cv2.INTER_AREA)
|
||||
img = img.round().clip(0, 255).astype(np.uint8)
|
||||
return img
|
||||
|
||||
|
||||
|
||||
def Numpy2Tensor(img):
|
||||
'''
|
||||
np.array[H, w, C] [0, 255] -> Tensor[C, H, W], RGB, [-1, 1]
|
||||
'''
|
||||
# size
|
||||
img = np.array(img) / 255 * 2 - 1
|
||||
img = torch.tensor(img, dtype=torch.float32).permute(2, 0, 1)
|
||||
return img
|
||||
|
||||
|
||||
def Tensor2Numpy(x, h0=None, w0=None):
|
||||
'''
|
||||
Tensor[C, H, W], RGB, [-1, 1] -> PIL.Image
|
||||
'''
|
||||
if h0 is not None and w0 is not None:
|
||||
x = x.unsqueeze(0)
|
||||
x = interpolate(x, size=(h0, w0), mode='bicubic')
|
||||
x = x.squeeze(0)
|
||||
x = (x.permute(1, 2, 0) * 127.5 + 127.5).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
return x
|
||||
|
||||
|
||||
def convert_dtype(dtype_str):
|
||||
if dtype_str == 'fp32':
|
||||
return torch.float32
|
||||
elif dtype_str == 'fp16':
|
||||
return torch.float16
|
||||
elif dtype_str == 'bf16':
|
||||
return torch.bfloat16
|
||||
else:
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,120 @@
|
||||
'''
|
||||
# --------------------------------------------------------------------------------
|
||||
# Color fixed script from Li Yi (https://github.com/pkuliyi2015/sd-webui-stablesr/blob/master/srmodule/colorfix.py)
|
||||
# --------------------------------------------------------------------------------
|
||||
'''
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch import Tensor
|
||||
from torch.nn import functional as F
|
||||
|
||||
from torchvision.transforms import ToTensor, ToPILImage
|
||||
|
||||
def adain_color_fix(target: Image, source: Image):
|
||||
# Convert images to tensors
|
||||
to_tensor = ToTensor()
|
||||
target_tensor = to_tensor(target).unsqueeze(0)
|
||||
source_tensor = to_tensor(source).unsqueeze(0)
|
||||
|
||||
# Apply adaptive instance normalization
|
||||
result_tensor = adaptive_instance_normalization(target_tensor, source_tensor)
|
||||
|
||||
# Convert tensor back to image
|
||||
to_image = ToPILImage()
|
||||
result_image = to_image(result_tensor.squeeze(0).clamp_(0.0, 1.0))
|
||||
|
||||
return result_image
|
||||
|
||||
def wavelet_color_fix(target: Image, source: Image):
|
||||
# Convert images to tensors
|
||||
to_tensor = ToTensor()
|
||||
target_tensor = to_tensor(target).unsqueeze(0)
|
||||
source_tensor = to_tensor(source).unsqueeze(0)
|
||||
|
||||
# Apply wavelet reconstruction
|
||||
result_tensor = wavelet_reconstruction(target_tensor, source_tensor)
|
||||
|
||||
# Convert tensor back to image
|
||||
to_image = ToPILImage()
|
||||
result_image = to_image(result_tensor.squeeze(0).clamp_(0.0, 1.0))
|
||||
|
||||
return result_image
|
||||
|
||||
def calc_mean_std(feat: Tensor, eps=1e-5):
|
||||
"""Calculate mean and std for adaptive_instance_normalization.
|
||||
Args:
|
||||
feat (Tensor): 4D tensor.
|
||||
eps (float): A small value added to the variance to avoid
|
||||
divide-by-zero. Default: 1e-5.
|
||||
"""
|
||||
size = feat.size()
|
||||
assert len(size) == 4, 'The input feature should be 4D tensor.'
|
||||
b, c = size[:2]
|
||||
feat_var = feat.reshape(b, c, -1).var(dim=2) + eps
|
||||
feat_std = feat_var.sqrt().reshape(b, c, 1, 1)
|
||||
feat_mean = feat.reshape(b, c, -1).mean(dim=2).reshape(b, c, 1, 1)
|
||||
return feat_mean, feat_std
|
||||
|
||||
def adaptive_instance_normalization(content_feat:Tensor, style_feat:Tensor):
|
||||
"""Adaptive instance normalization.
|
||||
Adjust the reference features to have the similar color and illuminations
|
||||
as those in the degradate features.
|
||||
Args:
|
||||
content_feat (Tensor): The reference feature.
|
||||
style_feat (Tensor): The degradate features.
|
||||
"""
|
||||
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)
|
||||
|
||||
def wavelet_blur(image: Tensor, radius: int):
|
||||
"""
|
||||
Apply wavelet blur to the input tensor.
|
||||
"""
|
||||
# input shape: (1, 3, H, W)
|
||||
# convolution kernel
|
||||
kernel_vals = [
|
||||
[0.0625, 0.125, 0.0625],
|
||||
[0.125, 0.25, 0.125],
|
||||
[0.0625, 0.125, 0.0625],
|
||||
]
|
||||
kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device)
|
||||
# add channel dimensions to the kernel to make it a 4D tensor
|
||||
kernel = kernel[None, None]
|
||||
# repeat the kernel across all input channels
|
||||
kernel = kernel.repeat(3, 1, 1, 1)
|
||||
image = F.pad(image, (radius, radius, radius, radius), mode='replicate')
|
||||
# apply convolution
|
||||
output = F.conv2d(image, kernel, groups=3, dilation=radius)
|
||||
return output
|
||||
|
||||
def wavelet_decomposition(image: Tensor, levels=5):
|
||||
"""
|
||||
Apply wavelet decomposition to the input tensor.
|
||||
This function only returns the low frequency & the high frequency.
|
||||
"""
|
||||
high_freq = torch.zeros_like(image)
|
||||
for i in range(levels):
|
||||
radius = 2 ** i
|
||||
low_freq = wavelet_blur(image, radius)
|
||||
high_freq += (image - low_freq)
|
||||
image = low_freq
|
||||
|
||||
return high_freq, low_freq
|
||||
|
||||
def wavelet_reconstruction(content_feat:Tensor, style_feat:Tensor):
|
||||
"""
|
||||
Apply wavelet decomposition, so that the content will have the same color as the style.
|
||||
"""
|
||||
# calculate the wavelet decomposition of the content feature
|
||||
content_high_freq, content_low_freq = wavelet_decomposition(content_feat)
|
||||
del content_low_freq
|
||||
# calculate the wavelet decomposition of the style feature
|
||||
style_high_freq, style_low_freq = wavelet_decomposition(style_feat)
|
||||
del style_high_freq
|
||||
# reconstruct the content feature with the style's high frequency
|
||||
return content_high_freq + style_low_freq
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
import sys
|
||||
import contextlib
|
||||
from functools import lru_cache
|
||||
|
||||
import torch
|
||||
#from modules import errors
|
||||
|
||||
if sys.platform == "darwin":
|
||||
from modules import mac_specific
|
||||
|
||||
|
||||
def has_mps() -> bool:
|
||||
if sys.platform != "darwin":
|
||||
return False
|
||||
else:
|
||||
return mac_specific.has_mps
|
||||
|
||||
|
||||
def get_cuda_device_string():
|
||||
return "cuda"
|
||||
|
||||
|
||||
def get_optimal_device_name():
|
||||
if torch.cuda.is_available():
|
||||
return get_cuda_device_string()
|
||||
|
||||
if has_mps():
|
||||
return "mps"
|
||||
|
||||
return "cpu"
|
||||
|
||||
|
||||
def get_optimal_device():
|
||||
return torch.device(get_optimal_device_name())
|
||||
|
||||
|
||||
def get_device_for(task):
|
||||
return get_optimal_device()
|
||||
|
||||
|
||||
def torch_gc():
|
||||
|
||||
if torch.cuda.is_available():
|
||||
with torch.cuda.device(get_cuda_device_string()):
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
if has_mps():
|
||||
mac_specific.torch_mps_gc()
|
||||
|
||||
|
||||
def enable_tf32():
|
||||
if torch.cuda.is_available():
|
||||
|
||||
# enabling benchmark option seems to enable a range of cards to do fp16 when they otherwise can't
|
||||
# see https://github.com/AUTOMATIC1111/stable-diffusion-webui/pull/4407
|
||||
if any(torch.cuda.get_device_capability(devid) == (7, 5) for devid in range(0, torch.cuda.device_count())):
|
||||
torch.backends.cudnn.benchmark = True
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
|
||||
enable_tf32()
|
||||
#errors.run(enable_tf32, "Enabling TF32")
|
||||
|
||||
cpu = torch.device("cpu")
|
||||
device = device_interrogate = device_gfpgan = device_esrgan = device_codeformer = torch.device("cuda")
|
||||
dtype = torch.float16
|
||||
dtype_vae = torch.float16
|
||||
dtype_unet = torch.float16
|
||||
unet_needs_upcast = False
|
||||
|
||||
|
||||
def cond_cast_unet(input):
|
||||
return input.to(dtype_unet) if unet_needs_upcast else input
|
||||
|
||||
|
||||
def cond_cast_float(input):
|
||||
return input.float() if unet_needs_upcast else input
|
||||
|
||||
|
||||
def randn(seed, shape):
|
||||
torch.manual_seed(seed)
|
||||
return torch.randn(shape, device=device)
|
||||
|
||||
|
||||
def randn_without_seed(shape):
|
||||
return torch.randn(shape, device=device)
|
||||
|
||||
|
||||
def autocast(disable=False):
|
||||
if disable:
|
||||
return contextlib.nullcontext()
|
||||
|
||||
return torch.autocast("cuda")
|
||||
|
||||
|
||||
def without_autocast(disable=False):
|
||||
return torch.autocast("cuda", enabled=False) if torch.is_autocast_enabled() and not disable else contextlib.nullcontext()
|
||||
|
||||
|
||||
class NansException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def test_for_nans(x, where):
|
||||
if not torch.all(torch.isnan(x)).item():
|
||||
return
|
||||
|
||||
if where == "unet":
|
||||
message = "A tensor with all NaNs was produced in Unet."
|
||||
|
||||
elif where == "vae":
|
||||
message = "A tensor with all NaNs was produced in VAE."
|
||||
|
||||
else:
|
||||
message = "A tensor with all NaNs was produced."
|
||||
|
||||
message += " Use --disable-nan-check commandline argument to disable this check."
|
||||
|
||||
raise NansException(message)
|
||||
|
||||
|
||||
@lru_cache
|
||||
def first_time_calculation():
|
||||
"""
|
||||
just do any calculation with pytorch layers - the first time this is done it allocaltes about 700MB of memory and
|
||||
spends about 2.7 seconds doing that, at least wih NVidia.
|
||||
"""
|
||||
|
||||
x = torch.zeros((1, 1)).to(device, dtype)
|
||||
linear = torch.nn.Linear(1, 1).to(device, dtype)
|
||||
linear(x)
|
||||
|
||||
x = torch.zeros((1, 1, 3, 3)).to(device, dtype)
|
||||
conv2d = torch.nn.Conv2d(1, 1, (3, 3)).to(device, dtype)
|
||||
conv2d(x)
|
||||
@@ -0,0 +1,974 @@
|
||||
# ------------------------------------------------------------------------
|
||||
#
|
||||
# Ultimate VAE Tile Optimization
|
||||
#
|
||||
# Introducing a revolutionary new optimization designed to make
|
||||
# the VAE work with giant images on limited VRAM!
|
||||
# Say goodbye to the frustration of OOM and hello to seamless output!
|
||||
#
|
||||
# ------------------------------------------------------------------------
|
||||
#
|
||||
# This script is a wild hack that splits the image into tiles,
|
||||
# encodes each tile separately, and merges the result back together.
|
||||
#
|
||||
# Advantages:
|
||||
# - The VAE can now work with giant images on limited VRAM
|
||||
# (~10 GB for 8K images!)
|
||||
# - The merged output is completely seamless without any post-processing.
|
||||
#
|
||||
# Drawbacks:
|
||||
# - Giant RAM needed. To store the intermediate results for a 4096x4096
|
||||
# images, you need 32 GB RAM it consumes ~20GB); for 8192x8192
|
||||
# you need 128 GB RAM machine (it consumes ~100 GB)
|
||||
# - NaNs always appear in for 8k images when you use fp16 (half) VAE
|
||||
# You must use --no-half-vae to disable half VAE for that giant image.
|
||||
# - Slow speed. With default tile size, it takes around 50/200 seconds
|
||||
# to encode/decode a 4096x4096 image; and 200/900 seconds to encode/decode
|
||||
# a 8192x8192 image. (The speed is limited by both the GPU and the CPU.)
|
||||
# - The gradient calculation is not compatible with this hack. It
|
||||
# will break any backward() or torch.autograd.grad() that passes VAE.
|
||||
# (But you can still use the VAE to generate training data.)
|
||||
#
|
||||
# How it works:
|
||||
# 1) The image is split into tiles.
|
||||
# - To ensure perfect results, each tile is padded with 32 pixels
|
||||
# on each side.
|
||||
# - Then the conv2d/silu/upsample/downsample can produce identical
|
||||
# results to the original image without splitting.
|
||||
# 2) The original forward is decomposed into a task queue and a task worker.
|
||||
# - The task queue is a list of functions that will be executed in order.
|
||||
# - The task worker is a loop that executes the tasks in the queue.
|
||||
# 3) The task queue is executed for each tile.
|
||||
# - Current tile is sent to GPU.
|
||||
# - local operations are directly executed.
|
||||
# - Group norm calculation is temporarily suspended until the mean
|
||||
# and var of all tiles are calculated.
|
||||
# - The residual is pre-calculated and stored and addded back later.
|
||||
# - When need to go to the next tile, the current tile is send to cpu.
|
||||
# 4) After all tiles are processed, tiles are merged on cpu and return.
|
||||
#
|
||||
# Enjoy!
|
||||
#
|
||||
# @author: LI YI @ Nanyang Technological University - Singapore
|
||||
# @date: 2023-03-02
|
||||
# @license: MIT License
|
||||
#
|
||||
# Please give me a star if you like this project!
|
||||
#
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
import gc
|
||||
from time import time
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
|
||||
import torch
|
||||
import torch.version
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from diffusers.utils.import_utils import is_xformers_available
|
||||
|
||||
#import SUPIR.utils.devices as devices
|
||||
|
||||
import comfy.model_management
|
||||
device = comfy.model_management.get_torch_device()
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
sd_flag = True
|
||||
|
||||
def get_recommend_encoder_tile_size():
|
||||
if torch.cuda.is_available():
|
||||
total_memory = torch.cuda.get_device_properties(
|
||||
device).total_memory // 2**20
|
||||
if total_memory > 16*1000:
|
||||
ENCODER_TILE_SIZE = 3072
|
||||
elif total_memory > 12*1000:
|
||||
ENCODER_TILE_SIZE = 2048
|
||||
elif total_memory > 8*1000:
|
||||
ENCODER_TILE_SIZE = 1536
|
||||
else:
|
||||
ENCODER_TILE_SIZE = 960
|
||||
else:
|
||||
ENCODER_TILE_SIZE = 512
|
||||
return ENCODER_TILE_SIZE
|
||||
|
||||
|
||||
def get_recommend_decoder_tile_size():
|
||||
if torch.cuda.is_available():
|
||||
total_memory = torch.cuda.get_device_properties(
|
||||
device).total_memory // 2**20
|
||||
if total_memory > 30*1000:
|
||||
DECODER_TILE_SIZE = 256
|
||||
elif total_memory > 16*1000:
|
||||
DECODER_TILE_SIZE = 192
|
||||
elif total_memory > 12*1000:
|
||||
DECODER_TILE_SIZE = 128
|
||||
elif total_memory > 8*1000:
|
||||
DECODER_TILE_SIZE = 96
|
||||
else:
|
||||
DECODER_TILE_SIZE = 64
|
||||
else:
|
||||
DECODER_TILE_SIZE = 64
|
||||
return DECODER_TILE_SIZE
|
||||
|
||||
|
||||
if 'global const':
|
||||
DEFAULT_ENABLED = False
|
||||
DEFAULT_MOVE_TO_GPU = False
|
||||
DEFAULT_FAST_ENCODER = True
|
||||
DEFAULT_FAST_DECODER = True
|
||||
DEFAULT_COLOR_FIX = 0
|
||||
DEFAULT_ENCODER_TILE_SIZE = get_recommend_encoder_tile_size()
|
||||
DEFAULT_DECODER_TILE_SIZE = get_recommend_decoder_tile_size()
|
||||
|
||||
|
||||
# inplace version of silu
|
||||
def inplace_nonlinearity(x):
|
||||
# Test: fix for Nans
|
||||
return F.silu(x, inplace=True)
|
||||
|
||||
# extracted from ldm.modules.diffusionmodules.model
|
||||
|
||||
# from diffusers lib
|
||||
def attn_forward_new(self, h_):
|
||||
batch_size, channel, height, width = h_.shape
|
||||
hidden_states = h_.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
attention_mask = None
|
||||
encoder_hidden_states = None
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
attention_mask = self.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif self.norm_cross:
|
||||
encoder_hidden_states = self.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
query = self.head_to_batch_dim(query)
|
||||
key = self.head_to_batch_dim(key)
|
||||
value = self.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = self.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = self.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
return hidden_states
|
||||
|
||||
def attn_forward_new_pt2_0(self, hidden_states,):
|
||||
scale = 1
|
||||
attention_mask = None
|
||||
encoder_hidden_states = None
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = self.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(batch_size, self.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
if self.group_norm is not None:
|
||||
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hidden_states, scale=scale)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif self.norm_cross:
|
||||
encoder_hidden_states = self.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = self.to_k(encoder_hidden_states, scale=scale)
|
||||
value = self.to_v(encoder_hidden_states, scale=scale)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // self.heads
|
||||
|
||||
query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
|
||||
key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states, scale=scale)
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
return hidden_states
|
||||
|
||||
def attn_forward_new_xformers(self, hidden_states):
|
||||
scale = 1
|
||||
attention_op = None
|
||||
attention_mask = None
|
||||
encoder_hidden_states = None
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, key_tokens, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
attention_mask = self.prepare_attention_mask(attention_mask, key_tokens, batch_size)
|
||||
if attention_mask is not None:
|
||||
# expand our mask's singleton query_tokens dimension:
|
||||
# [batch*heads, 1, key_tokens] ->
|
||||
# [batch*heads, query_tokens, key_tokens]
|
||||
# so that it can be added as a bias onto the attention scores that xformers computes:
|
||||
# [batch*heads, query_tokens, key_tokens]
|
||||
# we do this explicitly because xformers doesn't broadcast the singleton dimension for us.
|
||||
_, query_tokens, _ = hidden_states.shape
|
||||
attention_mask = attention_mask.expand(-1, query_tokens, -1)
|
||||
|
||||
if self.group_norm is not None:
|
||||
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hidden_states, scale=scale)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif self.norm_cross:
|
||||
encoder_hidden_states = self.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = self.to_k(encoder_hidden_states, scale=scale)
|
||||
value = self.to_v(encoder_hidden_states, scale=scale)
|
||||
|
||||
query = self.head_to_batch_dim(query).contiguous()
|
||||
key = self.head_to_batch_dim(key).contiguous()
|
||||
value = self.head_to_batch_dim(value).contiguous()
|
||||
|
||||
hidden_states = xformers.ops.memory_efficient_attention(
|
||||
query, key, value, attn_bias=attention_mask, op=attention_op#, scale=scale
|
||||
)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
hidden_states = self.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states, scale=scale)
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
return hidden_states
|
||||
|
||||
def attn_forward(self, h_):
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b, c, h, w = q.shape
|
||||
q = q.reshape(b, c, h*w)
|
||||
q = q.permute(0, 2, 1) # b,hw,c
|
||||
k = k.reshape(b, c, h*w) # b,c,hw
|
||||
w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
|
||||
w_ = w_ * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = v.reshape(b, c, h*w)
|
||||
w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q)
|
||||
# b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
|
||||
h_ = torch.bmm(v, w_)
|
||||
h_ = h_.reshape(b, c, h, w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return h_
|
||||
|
||||
|
||||
def xformer_attn_forward(self, h_):
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
B, C, H, W = q.shape
|
||||
q, k, v = map(lambda x: rearrange(x, 'b c h w -> b (h w) c'), (q, k, v))
|
||||
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3)
|
||||
.reshape(B, t.shape[1], 1, C)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(B * 1, t.shape[1], C)
|
||||
.contiguous(),
|
||||
(q, k, v),
|
||||
)
|
||||
out = xformers.ops.memory_efficient_attention(
|
||||
q, k, v, attn_bias=None, op=self.attention_op)
|
||||
|
||||
out = (
|
||||
out.unsqueeze(0)
|
||||
.reshape(B, 1, out.shape[1], C)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(B, out.shape[1], C)
|
||||
)
|
||||
out = rearrange(out, 'b (h w) c -> b c h w', b=B, h=H, w=W, c=C)
|
||||
out = self.proj_out(out)
|
||||
return out
|
||||
|
||||
|
||||
def attn2task(task_queue, net):
|
||||
if False: #isinstance(net, AttnBlock):
|
||||
task_queue.append(('store_res', lambda x: x))
|
||||
task_queue.append(('pre_norm', net.norm))
|
||||
task_queue.append(('attn', lambda x, net=net: attn_forward(net, x)))
|
||||
task_queue.append(['add_res', None])
|
||||
elif False: #isinstance(net, MemoryEfficientAttnBlock):
|
||||
task_queue.append(('store_res', lambda x: x))
|
||||
task_queue.append(('pre_norm', net.norm))
|
||||
task_queue.append(
|
||||
('attn', lambda x, net=net: xformer_attn_forward(net, x)))
|
||||
task_queue.append(['add_res', None])
|
||||
else:
|
||||
task_queue.append(('store_res', lambda x: x))
|
||||
task_queue.append(('pre_norm', net.norm))
|
||||
if is_xformers_available:
|
||||
# task_queue.append(('attn', lambda x, net=net: attn_forward_new_xformers(net, x)))
|
||||
task_queue.append(
|
||||
('attn', lambda x, net=net: xformer_attn_forward(net, x)))
|
||||
elif hasattr(F, "scaled_dot_product_attention"):
|
||||
task_queue.append(('attn', lambda x, net=net: attn_forward_new_pt2_0(net, x)))
|
||||
else:
|
||||
task_queue.append(('attn', lambda x, net=net: attn_forward_new(net, x)))
|
||||
task_queue.append(['add_res', None])
|
||||
|
||||
def resblock2task(queue, block):
|
||||
"""
|
||||
Turn a ResNetBlock into a sequence of tasks and append to the task queue
|
||||
|
||||
@param queue: the target task queue
|
||||
@param block: ResNetBlock
|
||||
|
||||
"""
|
||||
if block.in_channels != block.out_channels:
|
||||
if sd_flag:
|
||||
if block.use_conv_shortcut:
|
||||
queue.append(('store_res', block.conv_shortcut))
|
||||
else:
|
||||
queue.append(('store_res', block.nin_shortcut))
|
||||
else:
|
||||
if block.use_in_shortcut:
|
||||
queue.append(('store_res', block.conv_shortcut))
|
||||
else:
|
||||
queue.append(('store_res', block.nin_shortcut))
|
||||
|
||||
else:
|
||||
queue.append(('store_res', lambda x: x))
|
||||
queue.append(('pre_norm', block.norm1))
|
||||
queue.append(('silu', inplace_nonlinearity))
|
||||
queue.append(('conv1', block.conv1))
|
||||
queue.append(('pre_norm', block.norm2))
|
||||
queue.append(('silu', inplace_nonlinearity))
|
||||
queue.append(('conv2', block.conv2))
|
||||
queue.append(['add_res', None])
|
||||
|
||||
|
||||
def build_sampling(task_queue, net, is_decoder):
|
||||
"""
|
||||
Build the sampling part of a task queue
|
||||
@param task_queue: the target task queue
|
||||
@param net: the network
|
||||
@param is_decoder: currently building decoder or encoder
|
||||
"""
|
||||
if is_decoder:
|
||||
if sd_flag:
|
||||
resblock2task(task_queue, net.mid.block_1)
|
||||
attn2task(task_queue, net.mid.attn_1)
|
||||
print(task_queue)
|
||||
resblock2task(task_queue, net.mid.block_2)
|
||||
resolution_iter = reversed(range(net.num_resolutions))
|
||||
block_ids = net.num_res_blocks + 1
|
||||
condition = 0
|
||||
module = net.up
|
||||
func_name = 'upsample'
|
||||
else:
|
||||
resblock2task(task_queue, net.mid_block.resnets[0])
|
||||
attn2task(task_queue, net.mid_block.attentions[0])
|
||||
resblock2task(task_queue, net.mid_block.resnets[1])
|
||||
resolution_iter = (range(len(net.up_blocks))) # net.num_resolutions = 3
|
||||
block_ids = 2 + 1
|
||||
condition = len(net.up_blocks) - 1
|
||||
module = net.up_blocks
|
||||
func_name = 'upsamplers'
|
||||
else:
|
||||
if sd_flag:
|
||||
resolution_iter = range(net.num_resolutions)
|
||||
block_ids = net.num_res_blocks
|
||||
condition = net.num_resolutions - 1
|
||||
module = net.down
|
||||
func_name = 'downsample'
|
||||
else:
|
||||
resolution_iter = range(len(net.down_blocks))
|
||||
block_ids = 2
|
||||
condition = len(net.down_blocks) - 1
|
||||
module = net.down_blocks
|
||||
func_name = 'downsamplers'
|
||||
|
||||
for i_level in resolution_iter:
|
||||
for i_block in range(block_ids):
|
||||
if sd_flag:
|
||||
resblock2task(task_queue, module[i_level].block[i_block])
|
||||
else:
|
||||
resblock2task(task_queue, module[i_level].resnets[i_block])
|
||||
if i_level != condition:
|
||||
if sd_flag:
|
||||
task_queue.append((func_name, getattr(module[i_level], func_name)))
|
||||
else:
|
||||
if is_decoder:
|
||||
task_queue.append((func_name, module[i_level].upsamplers[0]))
|
||||
else:
|
||||
task_queue.append((func_name, module[i_level].downsamplers[0]))
|
||||
|
||||
if not is_decoder:
|
||||
if sd_flag:
|
||||
resblock2task(task_queue, net.mid.block_1)
|
||||
attn2task(task_queue, net.mid.attn_1)
|
||||
resblock2task(task_queue, net.mid.block_2)
|
||||
else:
|
||||
resblock2task(task_queue, net.mid_block.resnets[0])
|
||||
attn2task(task_queue, net.mid_block.attentions[0])
|
||||
resblock2task(task_queue, net.mid_block.resnets[1])
|
||||
|
||||
|
||||
def build_task_queue(net, is_decoder):
|
||||
"""
|
||||
Build a single task queue for the encoder or decoder
|
||||
@param net: the VAE decoder or encoder network
|
||||
@param is_decoder: currently building decoder or encoder
|
||||
@return: the task queue
|
||||
"""
|
||||
task_queue = []
|
||||
task_queue.append(('conv_in', net.conv_in))
|
||||
|
||||
# construct the sampling part of the task queue
|
||||
# because encoder and decoder share the same architecture, we extract the sampling part
|
||||
build_sampling(task_queue, net, is_decoder)
|
||||
if is_decoder and not sd_flag:
|
||||
net.give_pre_end = False
|
||||
net.tanh_out = False
|
||||
|
||||
if not is_decoder or not net.give_pre_end:
|
||||
if sd_flag:
|
||||
task_queue.append(('pre_norm', net.norm_out))
|
||||
else:
|
||||
task_queue.append(('pre_norm', net.conv_norm_out))
|
||||
task_queue.append(('silu', inplace_nonlinearity))
|
||||
task_queue.append(('conv_out', net.conv_out))
|
||||
if is_decoder and net.tanh_out:
|
||||
task_queue.append(('tanh', torch.tanh))
|
||||
|
||||
return task_queue
|
||||
|
||||
|
||||
def clone_task_queue(task_queue):
|
||||
"""
|
||||
Clone a task queue
|
||||
@param task_queue: the task queue to be cloned
|
||||
@return: the cloned task queue
|
||||
"""
|
||||
return [[item for item in task] for task in task_queue]
|
||||
|
||||
|
||||
def get_var_mean(input, num_groups, eps=1e-6):
|
||||
"""
|
||||
Get mean and var for group norm
|
||||
"""
|
||||
b, c = input.size(0), input.size(1)
|
||||
channel_in_group = int(c/num_groups)
|
||||
input_reshaped = input.contiguous().view(
|
||||
1, int(b * num_groups), channel_in_group, *input.size()[2:])
|
||||
var, mean = torch.var_mean(
|
||||
input_reshaped, dim=[0, 2, 3, 4], unbiased=False)
|
||||
return var, mean
|
||||
|
||||
|
||||
def custom_group_norm(input, num_groups, mean, var, weight=None, bias=None, eps=1e-6):
|
||||
"""
|
||||
Custom group norm with fixed mean and var
|
||||
|
||||
@param input: input tensor
|
||||
@param num_groups: number of groups. by default, num_groups = 32
|
||||
@param mean: mean, must be pre-calculated by get_var_mean
|
||||
@param var: var, must be pre-calculated by get_var_mean
|
||||
@param weight: weight, should be fetched from the original group norm
|
||||
@param bias: bias, should be fetched from the original group norm
|
||||
@param eps: epsilon, by default, eps = 1e-6 to match the original group norm
|
||||
|
||||
@return: normalized tensor
|
||||
"""
|
||||
b, c = input.size(0), input.size(1)
|
||||
channel_in_group = int(c/num_groups)
|
||||
input_reshaped = input.contiguous().view(
|
||||
1, int(b * num_groups), channel_in_group, *input.size()[2:])
|
||||
|
||||
out = F.batch_norm(input_reshaped, mean, var, weight=None, bias=None,
|
||||
training=False, momentum=0, eps=eps)
|
||||
|
||||
out = out.view(b, c, *input.size()[2:])
|
||||
|
||||
# post affine transform
|
||||
if weight is not None:
|
||||
out *= weight.view(1, -1, 1, 1)
|
||||
if bias is not None:
|
||||
out += bias.view(1, -1, 1, 1)
|
||||
return out
|
||||
|
||||
|
||||
def crop_valid_region(x, input_bbox, target_bbox, is_decoder):
|
||||
"""
|
||||
Crop the valid region from the tile
|
||||
@param x: input tile
|
||||
@param input_bbox: original input bounding box
|
||||
@param target_bbox: output bounding box
|
||||
@param scale: scale factor
|
||||
@return: cropped tile
|
||||
"""
|
||||
padded_bbox = [i * 8 if is_decoder else i//8 for i in input_bbox]
|
||||
margin = [target_bbox[i] - padded_bbox[i] for i in range(4)]
|
||||
return x[:, :, margin[2]:x.size(2)+margin[3], margin[0]:x.size(3)+margin[1]]
|
||||
|
||||
# ↓↓↓ https://github.com/Kahsolt/stable-diffusion-webui-vae-tile-infer ↓↓↓
|
||||
|
||||
|
||||
def perfcount(fn):
|
||||
def wrapper(*args, **kwargs):
|
||||
ts = time()
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
comfy.model_management.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
ret = fn(*args, **kwargs)
|
||||
|
||||
comfy.model_management.soft_empty_cache()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
vram = torch.cuda.max_memory_allocated(device) / 2**20
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
print(
|
||||
f'[Tiled VAE]: Done in {time() - ts:.3f}s, max VRAM alloc {vram:.3f} MB')
|
||||
else:
|
||||
print(f'[Tiled VAE]: Done in {time() - ts:.3f}s')
|
||||
|
||||
return ret
|
||||
return wrapper
|
||||
|
||||
# copy end :)
|
||||
|
||||
|
||||
class GroupNormParam:
|
||||
def __init__(self):
|
||||
self.var_list = []
|
||||
self.mean_list = []
|
||||
self.pixel_list = []
|
||||
self.weight = None
|
||||
self.bias = None
|
||||
|
||||
def add_tile(self, tile, layer):
|
||||
var, mean = get_var_mean(tile, 32)
|
||||
# For giant images, the variance can be larger than max float16
|
||||
# In this case we create a copy to float32
|
||||
if var.dtype == torch.float16 and var.isinf().any():
|
||||
fp32_tile = tile.float()
|
||||
var, mean = get_var_mean(fp32_tile, 32)
|
||||
# ============= DEBUG: test for infinite =============
|
||||
# if torch.isinf(var).any():
|
||||
# print('var: ', var)
|
||||
# ====================================================
|
||||
self.var_list.append(var)
|
||||
self.mean_list.append(mean)
|
||||
self.pixel_list.append(
|
||||
tile.shape[2]*tile.shape[3])
|
||||
if hasattr(layer, 'weight'):
|
||||
self.weight = layer.weight
|
||||
self.bias = layer.bias
|
||||
else:
|
||||
self.weight = None
|
||||
self.bias = None
|
||||
|
||||
def summary(self):
|
||||
"""
|
||||
summarize the mean and var and return a function
|
||||
that apply group norm on each tile
|
||||
"""
|
||||
if len(self.var_list) == 0:
|
||||
return None
|
||||
var = torch.vstack(self.var_list)
|
||||
mean = torch.vstack(self.mean_list)
|
||||
max_value = max(self.pixel_list)
|
||||
pixels = torch.tensor(
|
||||
self.pixel_list, dtype=torch.float32, device=device) / max_value
|
||||
sum_pixels = torch.sum(pixels)
|
||||
pixels = pixels.unsqueeze(
|
||||
1) / sum_pixels
|
||||
var = torch.sum(
|
||||
var * pixels, dim=0)
|
||||
mean = torch.sum(
|
||||
mean * pixels, dim=0)
|
||||
return lambda x: custom_group_norm(x, 32, mean, var, self.weight, self.bias)
|
||||
|
||||
@staticmethod
|
||||
def from_tile(tile, norm):
|
||||
"""
|
||||
create a function from a single tile without summary
|
||||
"""
|
||||
var, mean = get_var_mean(tile, 32)
|
||||
if var.dtype == torch.float16 and var.isinf().any():
|
||||
fp32_tile = tile.float()
|
||||
var, mean = get_var_mean(fp32_tile, 32)
|
||||
# if it is a macbook, we need to convert back to float16
|
||||
if var.device.type == 'mps':
|
||||
# clamp to avoid overflow
|
||||
var = torch.clamp(var, 0, 60000)
|
||||
var = var.half()
|
||||
mean = mean.half()
|
||||
if hasattr(norm, 'weight'):
|
||||
weight = norm.weight
|
||||
bias = norm.bias
|
||||
else:
|
||||
weight = None
|
||||
bias = None
|
||||
|
||||
def group_norm_func(x, mean=mean, var=var, weight=weight, bias=bias):
|
||||
return custom_group_norm(x, 32, mean, var, weight, bias, 1e-6)
|
||||
return group_norm_func
|
||||
|
||||
|
||||
class VAEHook:
|
||||
def __init__(self, net, tile_size, is_decoder, fast_decoder, fast_encoder, color_fix, to_gpu=False):
|
||||
self.net = net # encoder | decoder
|
||||
self.tile_size = tile_size
|
||||
self.is_decoder = is_decoder
|
||||
self.fast_mode = (fast_encoder and not is_decoder) or (
|
||||
fast_decoder and is_decoder)
|
||||
self.color_fix = color_fix and not is_decoder
|
||||
self.to_gpu = to_gpu
|
||||
self.pad = 11 if is_decoder else 32
|
||||
|
||||
def __call__(self, x):
|
||||
B, C, H, W = x.shape
|
||||
original_device = next(self.net.parameters()).device
|
||||
try:
|
||||
if self.to_gpu:
|
||||
self.net.to(device)
|
||||
if max(H, W) <= self.pad * 2 + self.tile_size:
|
||||
print("[Tiled VAE]: the input size is tiny and unnecessary to tile.")
|
||||
return self.net.original_forward(x)
|
||||
else:
|
||||
return self.vae_tile_forward(x)
|
||||
finally:
|
||||
self.net.to(original_device)
|
||||
|
||||
def get_best_tile_size(self, lowerbound, upperbound):
|
||||
"""
|
||||
Get the best tile size for GPU memory
|
||||
"""
|
||||
divider = 32
|
||||
while divider >= 2:
|
||||
remainer = lowerbound % divider
|
||||
if remainer == 0:
|
||||
return lowerbound
|
||||
candidate = lowerbound - remainer + divider
|
||||
if candidate <= upperbound:
|
||||
return candidate
|
||||
divider //= 2
|
||||
return lowerbound
|
||||
|
||||
def split_tiles(self, h, w):
|
||||
"""
|
||||
Tool function to split the image into tiles
|
||||
@param h: height of the image
|
||||
@param w: width of the image
|
||||
@return: tile_input_bboxes, tile_output_bboxes
|
||||
"""
|
||||
tile_input_bboxes, tile_output_bboxes = [], []
|
||||
tile_size = self.tile_size
|
||||
pad = self.pad
|
||||
num_height_tiles = math.ceil((h - 2 * pad) / tile_size)
|
||||
num_width_tiles = math.ceil((w - 2 * pad) / tile_size)
|
||||
# If any of the numbers are 0, we let it be 1
|
||||
# This is to deal with long and thin images
|
||||
num_height_tiles = max(num_height_tiles, 1)
|
||||
num_width_tiles = max(num_width_tiles, 1)
|
||||
|
||||
# Suggestions from https://github.com/Kahsolt: auto shrink the tile size
|
||||
real_tile_height = math.ceil((h - 2 * pad) / num_height_tiles)
|
||||
real_tile_width = math.ceil((w - 2 * pad) / num_width_tiles)
|
||||
real_tile_height = self.get_best_tile_size(real_tile_height, tile_size)
|
||||
real_tile_width = self.get_best_tile_size(real_tile_width, tile_size)
|
||||
|
||||
print(f'[Tiled VAE]: split to {num_height_tiles}x{num_width_tiles} = {num_height_tiles*num_width_tiles} tiles. ' +
|
||||
f'Optimal tile size {real_tile_width}x{real_tile_height}, original tile size {tile_size}x{tile_size}')
|
||||
|
||||
for i in range(num_height_tiles):
|
||||
for j in range(num_width_tiles):
|
||||
# bbox: [x1, x2, y1, y2]
|
||||
# the padding is is unnessary for image borders. So we directly start from (32, 32)
|
||||
input_bbox = [
|
||||
pad + j * real_tile_width,
|
||||
min(pad + (j + 1) * real_tile_width, w),
|
||||
pad + i * real_tile_height,
|
||||
min(pad + (i + 1) * real_tile_height, h),
|
||||
]
|
||||
|
||||
# if the output bbox is close to the image boundary, we extend it to the image boundary
|
||||
output_bbox = [
|
||||
input_bbox[0] if input_bbox[0] > pad else 0,
|
||||
input_bbox[1] if input_bbox[1] < w - pad else w,
|
||||
input_bbox[2] if input_bbox[2] > pad else 0,
|
||||
input_bbox[3] if input_bbox[3] < h - pad else h,
|
||||
]
|
||||
|
||||
# scale to get the final output bbox
|
||||
output_bbox = [x * 8 if self.is_decoder else x // 8 for x in output_bbox]
|
||||
tile_output_bboxes.append(output_bbox)
|
||||
|
||||
# indistinguishable expand the input bbox by pad pixels
|
||||
tile_input_bboxes.append([
|
||||
max(0, input_bbox[0] - pad),
|
||||
min(w, input_bbox[1] + pad),
|
||||
max(0, input_bbox[2] - pad),
|
||||
min(h, input_bbox[3] + pad),
|
||||
])
|
||||
|
||||
return tile_input_bboxes, tile_output_bboxes
|
||||
|
||||
@torch.no_grad()
|
||||
def estimate_group_norm(self, z, task_queue, color_fix):
|
||||
device = z.device
|
||||
tile = z
|
||||
last_id = len(task_queue) - 1
|
||||
while last_id >= 0 and task_queue[last_id][0] != 'pre_norm':
|
||||
last_id -= 1
|
||||
if last_id <= 0 or task_queue[last_id][0] != 'pre_norm':
|
||||
raise ValueError('No group norm found in the task queue')
|
||||
# estimate until the last group norm
|
||||
for i in range(last_id + 1):
|
||||
task = task_queue[i]
|
||||
if task[0] == 'pre_norm':
|
||||
group_norm_func = GroupNormParam.from_tile(tile, task[1])
|
||||
task_queue[i] = ('apply_norm', group_norm_func)
|
||||
if i == last_id:
|
||||
return True
|
||||
tile = group_norm_func(tile)
|
||||
elif task[0] == 'store_res':
|
||||
task_id = i + 1
|
||||
while task_id < last_id and task_queue[task_id][0] != 'add_res':
|
||||
task_id += 1
|
||||
if task_id >= last_id:
|
||||
continue
|
||||
task_queue[task_id][1] = task[1](tile)
|
||||
elif task[0] == 'add_res':
|
||||
tile += task[1].to(device)
|
||||
task[1] = None
|
||||
elif color_fix and task[0] == 'downsample':
|
||||
for j in range(i, last_id + 1):
|
||||
if task_queue[j][0] == 'store_res':
|
||||
task_queue[j] = ('store_res_cpu', task_queue[j][1])
|
||||
return True
|
||||
else:
|
||||
tile = task[1](tile)
|
||||
try:
|
||||
devices.test_for_nans(tile, "vae")
|
||||
except:
|
||||
print(f'Nan detected in fast mode estimation. Fast mode disabled.')
|
||||
return False
|
||||
|
||||
raise IndexError('Should not reach here')
|
||||
|
||||
@perfcount
|
||||
@torch.no_grad()
|
||||
def vae_tile_forward(self, z):
|
||||
"""
|
||||
Decode a latent vector z into an image in a tiled manner.
|
||||
@param z: latent vector
|
||||
@return: image
|
||||
"""
|
||||
device = next(self.net.parameters()).device
|
||||
dtype = z.dtype
|
||||
net = self.net
|
||||
tile_size = self.tile_size
|
||||
is_decoder = self.is_decoder
|
||||
|
||||
z = z.detach() # detach the input to avoid backprop
|
||||
|
||||
N, height, width = z.shape[0], z.shape[2], z.shape[3]
|
||||
net.last_z_shape = z.shape
|
||||
|
||||
# Split the input into tiles and build a task queue for each tile
|
||||
print(f'[Tiled VAE]: input_size: {z.shape}, tile_size: {tile_size}, padding: {self.pad}')
|
||||
|
||||
in_bboxes, out_bboxes = self.split_tiles(height, width)
|
||||
|
||||
# Prepare tiles by split the input latents
|
||||
tiles = []
|
||||
for input_bbox in in_bboxes:
|
||||
tile = z[:, :, input_bbox[2]:input_bbox[3], input_bbox[0]:input_bbox[1]].cpu()
|
||||
tiles.append(tile)
|
||||
|
||||
num_tiles = len(tiles)
|
||||
num_completed = 0
|
||||
|
||||
# Build task queues
|
||||
single_task_queue = build_task_queue(net, is_decoder)
|
||||
#print(single_task_queue)
|
||||
if self.fast_mode:
|
||||
# Fast mode: downsample the input image to the tile size,
|
||||
# then estimate the group norm parameters on the downsampled image
|
||||
scale_factor = tile_size / max(height, width)
|
||||
z = z.to(device)
|
||||
downsampled_z = F.interpolate(z, scale_factor=scale_factor, mode='nearest-exact')
|
||||
# use nearest-exact to keep statictics as close as possible
|
||||
print(f'[Tiled VAE]: Fast mode enabled, estimating group norm parameters on {downsampled_z.shape[3]} x {downsampled_z.shape[2]} image')
|
||||
|
||||
# ======= Special thanks to @Kahsolt for distribution shift issue ======= #
|
||||
# The downsampling will heavily distort its mean and std, so we need to recover it.
|
||||
std_old, mean_old = torch.std_mean(z, dim=[0, 2, 3], keepdim=True)
|
||||
std_new, mean_new = torch.std_mean(downsampled_z, dim=[0, 2, 3], keepdim=True)
|
||||
downsampled_z = (downsampled_z - mean_new) / std_new * std_old + mean_old
|
||||
del std_old, mean_old, std_new, mean_new
|
||||
# occasionally the std_new is too small or too large, which exceeds the range of float16
|
||||
# so we need to clamp it to max z's range.
|
||||
downsampled_z = torch.clamp_(downsampled_z, min=z.min(), max=z.max())
|
||||
estimate_task_queue = clone_task_queue(single_task_queue)
|
||||
if self.estimate_group_norm(downsampled_z, estimate_task_queue, color_fix=self.color_fix):
|
||||
single_task_queue = estimate_task_queue
|
||||
del downsampled_z
|
||||
|
||||
task_queues = [clone_task_queue(single_task_queue) for _ in range(num_tiles)]
|
||||
|
||||
# Dummy result
|
||||
result = None
|
||||
result_approx = None
|
||||
#try:
|
||||
# with devices.autocast():
|
||||
# result_approx = torch.cat([F.interpolate(cheap_approximation(x).unsqueeze(0), scale_factor=opt_f, mode='nearest-exact') for x in z], dim=0).cpu()
|
||||
#except: pass
|
||||
# Free memory of input latent tensor
|
||||
del z
|
||||
|
||||
# Task queue execution
|
||||
pbar = tqdm(total=num_tiles * len(task_queues[0]), desc=f"[Tiled VAE]: Executing {'Decoder' if is_decoder else 'Encoder'} Task Queue: ")
|
||||
|
||||
# execute the task back and forth when switch tiles so that we always
|
||||
# keep one tile on the GPU to reduce unnecessary data transfer
|
||||
forward = True
|
||||
interrupted = False
|
||||
#state.interrupted = interrupted
|
||||
while True:
|
||||
#if state.interrupted: interrupted = True ; break
|
||||
|
||||
group_norm_param = GroupNormParam()
|
||||
for i in range(num_tiles) if forward else reversed(range(num_tiles)):
|
||||
#if state.interrupted: interrupted = True ; break
|
||||
|
||||
tile = tiles[i].to(device)
|
||||
input_bbox = in_bboxes[i]
|
||||
task_queue = task_queues[i]
|
||||
|
||||
interrupted = False
|
||||
while len(task_queue) > 0:
|
||||
#if state.interrupted: interrupted = True ; break
|
||||
|
||||
# DEBUG: current task
|
||||
# print('Running task: ', task_queue[0][0], ' on tile ', i, '/', num_tiles, ' with shape ', tile.shape)
|
||||
task = task_queue.pop(0)
|
||||
if task[0] == 'pre_norm':
|
||||
group_norm_param.add_tile(tile, task[1])
|
||||
break
|
||||
elif task[0] == 'store_res' or task[0] == 'store_res_cpu':
|
||||
task_id = 0
|
||||
res = task[1](tile)
|
||||
if not self.fast_mode or task[0] == 'store_res_cpu':
|
||||
res = res.cpu()
|
||||
while task_queue[task_id][0] != 'add_res':
|
||||
task_id += 1
|
||||
task_queue[task_id][1] = res
|
||||
elif task[0] == 'add_res':
|
||||
tile += task[1].to(device)
|
||||
task[1] = None
|
||||
else:
|
||||
tile = task[1](tile)
|
||||
#print(tiles[i].shape, tile.shape, task)
|
||||
pbar.update(1)
|
||||
|
||||
if interrupted: break
|
||||
|
||||
# check for NaNs in the tile.
|
||||
# If there are NaNs, we abort the process to save user's time
|
||||
#devices.test_for_nans(tile, "vae")
|
||||
|
||||
#print(tiles[i].shape, tile.shape, i, num_tiles)
|
||||
if len(task_queue) == 0:
|
||||
tiles[i] = None
|
||||
num_completed += 1
|
||||
if result is None: # NOTE: dim C varies from different cases, can only be inited dynamically
|
||||
result = torch.zeros((N, tile.shape[1], height * 8 if is_decoder else height // 8, width * 8 if is_decoder else width // 8), device=device, requires_grad=False)
|
||||
result[:, :, out_bboxes[i][2]:out_bboxes[i][3], out_bboxes[i][0]:out_bboxes[i][1]] = crop_valid_region(tile, in_bboxes[i], out_bboxes[i], is_decoder)
|
||||
del tile
|
||||
elif i == num_tiles - 1 and forward:
|
||||
forward = False
|
||||
tiles[i] = tile
|
||||
elif i == 0 and not forward:
|
||||
forward = True
|
||||
tiles[i] = tile
|
||||
else:
|
||||
tiles[i] = tile.cpu()
|
||||
del tile
|
||||
|
||||
if interrupted: break
|
||||
if num_completed == num_tiles: break
|
||||
|
||||
# insert the group norm task to the head of each task queue
|
||||
group_norm_func = group_norm_param.summary()
|
||||
if group_norm_func is not None:
|
||||
for i in range(num_tiles):
|
||||
task_queue = task_queues[i]
|
||||
task_queue.insert(0, ('apply_norm', group_norm_func))
|
||||
|
||||
# Done!
|
||||
pbar.close()
|
||||
return result.to(dtype) if result is not None else result_approx.to(device)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -0,0 +1,117 @@
|
||||
import os
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
from contextlib import nullcontext
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
import comfy.model_management
|
||||
import folder_paths
|
||||
from nodes import ImageScaleBy
|
||||
from nodes import ImageScale
|
||||
import torch.cuda
|
||||
from .SUPIR.models.SUPIR_model import SUPIRModel
|
||||
from PIL import Image
|
||||
from .sgm.util import instantiate_from_config
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
class SUPIR_Upscale:
|
||||
upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"]
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"supir_model": (folder_paths.get_filename_list("checkpoints"), ),
|
||||
"sdxl_model": (folder_paths.get_filename_list("checkpoints"), ),
|
||||
"image": ("IMAGE", ),
|
||||
"resize_method": (s.upscale_methods, {"default": "lanczos"}),
|
||||
"scale_by": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 20.0, "step": 0.01}),
|
||||
"steps": ("INT", {"default": 45, "min": 3, "max": 4096, "step": 1}),
|
||||
"cfg_scale": ("FLOAT", {"default": 7.5,"min": 0, "max": 20, "step": 0.01}),
|
||||
"a_prompt": ("STRING", {"multiline": True, "default": "high quality",}),
|
||||
"n_prompt": ("STRING", {"multiline": True, "default": "illustration",}),
|
||||
|
||||
"min_size": ("INT", {"default": 1024, "min": 1, "max": 4096, "step": 1}),
|
||||
|
||||
"color_fix_type": (
|
||||
[
|
||||
'None',
|
||||
'AdaIn',
|
||||
'Wavelet',
|
||||
], {
|
||||
"default": 'adain'
|
||||
}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
"seed": ("INT", {"default": 123,"min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
},
|
||||
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES =("upscaled_image",)
|
||||
FUNCTION = "process"
|
||||
|
||||
CATEGORY = "SUPIR"
|
||||
|
||||
def process(self, steps, image, color_fix_type, seed, scale_by, min_size, cfg_scale, resize_method,
|
||||
a_prompt, n_prompt, sdxl_model, supir_model, keep_model_loaded):
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
comfy.model_management.unload_all_models()
|
||||
device = comfy.model_management.get_torch_device()
|
||||
image = image.to(device)
|
||||
SUPIR_MODEL_PATH = folder_paths.get_full_path("checkpoints", supir_model)
|
||||
SDXL_MODEL_PATH = folder_paths.get_full_path("checkpoints", sdxl_model)
|
||||
|
||||
config_path = os.path.join(script_directory, "options/SUPIR_v0.yaml")
|
||||
dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32
|
||||
if not hasattr(self, "model") or self.model is None:
|
||||
|
||||
config = OmegaConf.load(config_path)
|
||||
self.model = instantiate_from_config(config.model).cpu()
|
||||
from .SUPIR.util import load_state_dict
|
||||
supir_state_dict = load_state_dict(SUPIR_MODEL_PATH)
|
||||
sdxl_state_dict = load_state_dict(SDXL_MODEL_PATH)
|
||||
self.model.load_state_dict(supir_state_dict, strict=False)
|
||||
self.model.load_state_dict(sdxl_state_dict, strict=False)
|
||||
self.model.to(device).to(dtype)
|
||||
|
||||
autocast_condition = dtype == torch.float16 or torch.bfloat16 and not comfy.model_management.is_device_mps(device)
|
||||
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||
image, = ImageScaleBy.upscale(self, image, resize_method, scale_by)
|
||||
|
||||
# Assuming 'image' is a PyTorch tensor with shape [B, H, W, C] and you want to resize it.
|
||||
B, H, W, C = image.shape
|
||||
|
||||
# Calculate the new height and width, rounding down to the nearest multiple of 64.
|
||||
new_height = H // 64 * 64
|
||||
new_width = W // 64 * 64
|
||||
|
||||
# Reorder to [B, C, H, W] before using interpolate.
|
||||
image = image.permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
# Resize the image tensor.
|
||||
resized_image = F.interpolate(image, size=(new_height, new_width), mode='bicubic', align_corners=False)
|
||||
|
||||
captions = ['']
|
||||
print(captions)
|
||||
|
||||
# # step 3: Diffusion Process
|
||||
samples = self.model.batchify_sample(resized_image, captions, num_steps=steps, restoration_scale= -1, s_churn=5,
|
||||
s_noise=1.003, cfg_scale=cfg_scale, control_scale= 1, seed=seed,
|
||||
num_samples=1, p_p=a_prompt, n_p=n_prompt, color_fix_type=color_fix_type,
|
||||
use_linear_CFG=False, use_linear_control_scale=False,
|
||||
cfg_scale_start=1.0, control_scale_start=0)
|
||||
# save
|
||||
print(samples.shape)
|
||||
samples = samples.permute(0, 2, 3, 1).cpu()
|
||||
|
||||
return(samples,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SUPIR_Upscale": SUPIR_Upscale,
|
||||
"SUPIR_Upscale": SUPIR_Upscale
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SUPIR_Upscale": "SUPIR_Upscale",
|
||||
"SUPIR_Upscale": "SUPIR_Upscale"
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
model:
|
||||
target: .SUPIR.models.SUPIR_model.SUPIRModel
|
||||
params:
|
||||
ae_dtype: bf16
|
||||
diffusion_dtype: fp16
|
||||
scale_factor: 0.13025
|
||||
disable_first_stage_autocast: True
|
||||
network_wrapper: .sgm.modules.diffusionmodules.wrappers.ControlWrapper
|
||||
|
||||
denoiser_config:
|
||||
target: .sgm.modules.diffusionmodules.denoiser.DiscreteDenoiserWithControl
|
||||
params:
|
||||
num_idx: 1000
|
||||
weighting_config:
|
||||
target: .sgm.modules.diffusionmodules.denoiser_weighting.EpsWeighting
|
||||
scaling_config:
|
||||
target: .sgm.modules.diffusionmodules.denoiser_scaling.EpsScaling
|
||||
discretization_config:
|
||||
target: .sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization
|
||||
|
||||
control_stage_config:
|
||||
target: .SUPIR.modules.SUPIR_v0.GLVControl
|
||||
params:
|
||||
adm_in_channels: 2816
|
||||
num_classes: sequential
|
||||
use_checkpoint: True
|
||||
in_channels: 4
|
||||
out_channels: 4
|
||||
model_channels: 320
|
||||
attention_resolutions: [4, 2]
|
||||
num_res_blocks: 2
|
||||
channel_mult: [1, 2, 4]
|
||||
num_head_channels: 64
|
||||
use_spatial_transformer: True
|
||||
use_linear_in_transformer: True
|
||||
transformer_depth: [1, 2, 10] # note: the first is unused (due to attn_res starting at 2) 32, 16, 8 --> 64, 32, 16
|
||||
# transformer_depth: [1, 1, 4]
|
||||
context_dim: 2048
|
||||
spatial_transformer_attn_type: softmax-xformers
|
||||
legacy: False
|
||||
input_upscale: 1
|
||||
|
||||
network_config:
|
||||
target: .SUPIR.modules.SUPIR_v0.LightGLVUNet
|
||||
params:
|
||||
mode: XL-base
|
||||
project_type: ZeroSFT
|
||||
project_channel_scale: 2
|
||||
adm_in_channels: 2816
|
||||
num_classes: sequential
|
||||
use_checkpoint: True
|
||||
in_channels: 4
|
||||
out_channels: 4
|
||||
model_channels: 320
|
||||
attention_resolutions: [4, 2]
|
||||
num_res_blocks: 2
|
||||
channel_mult: [1, 2, 4]
|
||||
num_head_channels: 64
|
||||
use_spatial_transformer: True
|
||||
use_linear_in_transformer: True
|
||||
transformer_depth: [1, 2, 10] # note: the first is unused (due to attn_res starting at 2) 32, 16, 8 --> 64, 32, 16
|
||||
context_dim: 2048
|
||||
spatial_transformer_attn_type: softmax-xformers
|
||||
legacy: False
|
||||
|
||||
conditioner_config:
|
||||
target: .sgm.modules.GeneralConditionerWithControl
|
||||
params:
|
||||
emb_models:
|
||||
# crossattn cond
|
||||
- is_trainable: False
|
||||
input_key: txt
|
||||
target: .sgm.modules.encoders.modules.FrozenCLIPEmbedder
|
||||
params:
|
||||
layer: hidden
|
||||
layer_idx: 11
|
||||
# crossattn and vector cond
|
||||
- is_trainable: False
|
||||
input_key: txt
|
||||
target: .sgm.modules.encoders.modules.FrozenOpenCLIPEmbedder2
|
||||
params:
|
||||
arch: ViT-bigG-14
|
||||
version: laion2b_s39b_b160k
|
||||
freeze: True
|
||||
layer: penultimate
|
||||
always_return_pooled: True
|
||||
legacy: False
|
||||
# vector cond
|
||||
- is_trainable: False
|
||||
input_key: original_size_as_tuple
|
||||
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
|
||||
params:
|
||||
outdim: 256 # multiplied by two
|
||||
# vector cond
|
||||
- is_trainable: False
|
||||
input_key: crop_coords_top_left
|
||||
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
|
||||
params:
|
||||
outdim: 256 # multiplied by two
|
||||
# vector cond
|
||||
- is_trainable: False
|
||||
input_key: target_size_as_tuple
|
||||
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
|
||||
params:
|
||||
outdim: 256 # multiplied by two
|
||||
|
||||
first_stage_config:
|
||||
target: .sgm.models.autoencoder.AutoencoderKLInferenceWrapper
|
||||
params:
|
||||
ckpt_path: ~
|
||||
embed_dim: 4
|
||||
monitor: val/rec_loss
|
||||
ddconfig:
|
||||
attn_type: vanilla-xformers
|
||||
double_z: true
|
||||
z_channels: 4
|
||||
resolution: 256
|
||||
in_channels: 3
|
||||
out_ch: 3
|
||||
ch: 128
|
||||
ch_mult: [ 1, 2, 4, 4 ]
|
||||
num_res_blocks: 2
|
||||
attn_resolutions: [ ]
|
||||
dropout: 0.0
|
||||
lossconfig:
|
||||
target: torch.nn.Identity
|
||||
|
||||
sampler_config:
|
||||
target: .sgm.modules.diffusionmodules.sampling.RestoreEDMSampler
|
||||
params:
|
||||
num_steps: 100
|
||||
restore_cfg: 4.0
|
||||
s_churn: 0
|
||||
s_noise: 1.003
|
||||
discretization_config:
|
||||
target: .sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization
|
||||
guider_config:
|
||||
target: .sgm.modules.diffusionmodules.guiders.LinearCFG
|
||||
params:
|
||||
scale: 7.5
|
||||
scale_min: 4.0
|
||||
|
||||
p_p:
|
||||
'Cinematic, High Contrast, highly detailed, taken using a Canon EOS R camera,
|
||||
hyper detailed photo - realistic maximum detail, 32k, Color Grading, ultra HD, extreme meticulous detailing,
|
||||
skin pore detailing, hyper sharpness, perfect without deformations.'
|
||||
n_p:
|
||||
'painting, oil painting, illustration, drawing, art, sketch, oil painting, cartoon, CG Style, 3D render,
|
||||
unreal engine, blurring, dirty, messy, worst quality, low quality, frames, watermark, signature,
|
||||
jpeg artifacts, deformed, lowres, over-smooth'
|
||||
|
||||
SDXL_CKPT: /opt/data/private/AIGC_pretrain/SDXL_cache/sd_xl_base_1.0_0.9vae.safetensors
|
||||
SUPIR_CKPT_F: /opt/data/private/AIGC_pretrain/SUPIR_cache/SUPIR-v0F.ckpt
|
||||
SUPIR_CKPT_Q: /opt/data/private/AIGC_pretrain/SUPIR_cache/SUPIR-v0Q.ckpt
|
||||
SUPIR_CKPT: ~
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
fastapi==0.95.1
|
||||
gradio==4.16.0
|
||||
gradio_imageslider==0.0.17
|
||||
gradio_client==0.8.1
|
||||
Markdown==3.4.1
|
||||
numpy==1.24.2
|
||||
requests==2.28.2
|
||||
sentencepiece==0.1.98
|
||||
tokenizers==0.13.3
|
||||
torch>=2.1.0
|
||||
torchvision>=0.16.0
|
||||
uvicorn==0.21.1
|
||||
wandb==0.14.0
|
||||
httpx==0.24.0
|
||||
transformers==4.28.1
|
||||
accelerate==0.18.0
|
||||
scikit-learn==1.2.2
|
||||
sentencepiece==0.1.98
|
||||
einops==0.7.0
|
||||
einops-exts==0.0.4
|
||||
timm==0.9.8
|
||||
openai-clip==1.0.1
|
||||
fsspec==2023.4.0
|
||||
kornia==0.6.9
|
||||
matplotlib==3.7.1
|
||||
ninja==1.11.1
|
||||
omegaconf==2.3.0
|
||||
open-clip-torch==2.17.1
|
||||
opencv-python==4.7.0.72
|
||||
pandas==2.0.1
|
||||
Pillow==9.4.0
|
||||
pytorch-lightning==2.1.2
|
||||
PyYAML==6.0
|
||||
scipy==1.9.1
|
||||
tqdm==4.65.0
|
||||
triton==2.1.0
|
||||
urllib3==1.26.15
|
||||
webdataset==0.2.48
|
||||
xformers>=0.0.20
|
||||
@@ -0,0 +1,4 @@
|
||||
from .models import AutoencodingEngine, DiffusionEngine
|
||||
from .util import get_configs_path, instantiate_from_config
|
||||
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,135 @@
|
||||
import numpy as np
|
||||
|
||||
|
||||
class LambdaWarmUpCosineScheduler:
|
||||
"""
|
||||
note: use with a base_lr of 1.0
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
warm_up_steps,
|
||||
lr_min,
|
||||
lr_max,
|
||||
lr_start,
|
||||
max_decay_steps,
|
||||
verbosity_interval=0,
|
||||
):
|
||||
self.lr_warm_up_steps = warm_up_steps
|
||||
self.lr_start = lr_start
|
||||
self.lr_min = lr_min
|
||||
self.lr_max = lr_max
|
||||
self.lr_max_decay_steps = max_decay_steps
|
||||
self.last_lr = 0.0
|
||||
self.verbosity_interval = verbosity_interval
|
||||
|
||||
def schedule(self, n, **kwargs):
|
||||
if self.verbosity_interval > 0:
|
||||
if n % self.verbosity_interval == 0:
|
||||
print(f"current step: {n}, recent lr-multiplier: {self.last_lr}")
|
||||
if n < self.lr_warm_up_steps:
|
||||
lr = (
|
||||
self.lr_max - self.lr_start
|
||||
) / self.lr_warm_up_steps * n + self.lr_start
|
||||
self.last_lr = lr
|
||||
return lr
|
||||
else:
|
||||
t = (n - self.lr_warm_up_steps) / (
|
||||
self.lr_max_decay_steps - self.lr_warm_up_steps
|
||||
)
|
||||
t = min(t, 1.0)
|
||||
lr = self.lr_min + 0.5 * (self.lr_max - self.lr_min) * (
|
||||
1 + np.cos(t * np.pi)
|
||||
)
|
||||
self.last_lr = lr
|
||||
return lr
|
||||
|
||||
def __call__(self, n, **kwargs):
|
||||
return self.schedule(n, **kwargs)
|
||||
|
||||
|
||||
class LambdaWarmUpCosineScheduler2:
|
||||
"""
|
||||
supports repeated iterations, configurable via lists
|
||||
note: use with a base_lr of 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, warm_up_steps, f_min, f_max, f_start, cycle_lengths, verbosity_interval=0
|
||||
):
|
||||
assert (
|
||||
len(warm_up_steps)
|
||||
== len(f_min)
|
||||
== len(f_max)
|
||||
== len(f_start)
|
||||
== len(cycle_lengths)
|
||||
)
|
||||
self.lr_warm_up_steps = warm_up_steps
|
||||
self.f_start = f_start
|
||||
self.f_min = f_min
|
||||
self.f_max = f_max
|
||||
self.cycle_lengths = cycle_lengths
|
||||
self.cum_cycles = np.cumsum([0] + list(self.cycle_lengths))
|
||||
self.last_f = 0.0
|
||||
self.verbosity_interval = verbosity_interval
|
||||
|
||||
def find_in_interval(self, n):
|
||||
interval = 0
|
||||
for cl in self.cum_cycles[1:]:
|
||||
if n <= cl:
|
||||
return interval
|
||||
interval += 1
|
||||
|
||||
def schedule(self, n, **kwargs):
|
||||
cycle = self.find_in_interval(n)
|
||||
n = n - self.cum_cycles[cycle]
|
||||
if self.verbosity_interval > 0:
|
||||
if n % self.verbosity_interval == 0:
|
||||
print(
|
||||
f"current step: {n}, recent lr-multiplier: {self.last_f}, "
|
||||
f"current cycle {cycle}"
|
||||
)
|
||||
if n < self.lr_warm_up_steps[cycle]:
|
||||
f = (self.f_max[cycle] - self.f_start[cycle]) / self.lr_warm_up_steps[
|
||||
cycle
|
||||
] * n + self.f_start[cycle]
|
||||
self.last_f = f
|
||||
return f
|
||||
else:
|
||||
t = (n - self.lr_warm_up_steps[cycle]) / (
|
||||
self.cycle_lengths[cycle] - self.lr_warm_up_steps[cycle]
|
||||
)
|
||||
t = min(t, 1.0)
|
||||
f = self.f_min[cycle] + 0.5 * (self.f_max[cycle] - self.f_min[cycle]) * (
|
||||
1 + np.cos(t * np.pi)
|
||||
)
|
||||
self.last_f = f
|
||||
return f
|
||||
|
||||
def __call__(self, n, **kwargs):
|
||||
return self.schedule(n, **kwargs)
|
||||
|
||||
|
||||
class LambdaLinearScheduler(LambdaWarmUpCosineScheduler2):
|
||||
def schedule(self, n, **kwargs):
|
||||
cycle = self.find_in_interval(n)
|
||||
n = n - self.cum_cycles[cycle]
|
||||
if self.verbosity_interval > 0:
|
||||
if n % self.verbosity_interval == 0:
|
||||
print(
|
||||
f"current step: {n}, recent lr-multiplier: {self.last_f}, "
|
||||
f"current cycle {cycle}"
|
||||
)
|
||||
|
||||
if n < self.lr_warm_up_steps[cycle]:
|
||||
f = (self.f_max[cycle] - self.f_start[cycle]) / self.lr_warm_up_steps[
|
||||
cycle
|
||||
] * n + self.f_start[cycle]
|
||||
self.last_f = f
|
||||
return f
|
||||
else:
|
||||
f = self.f_min[cycle] + (self.f_max[cycle] - self.f_min[cycle]) * (
|
||||
self.cycle_lengths[cycle] - n
|
||||
) / (self.cycle_lengths[cycle])
|
||||
self.last_f = f
|
||||
return f
|
||||
@@ -0,0 +1,2 @@
|
||||
from .autoencoder import AutoencodingEngine
|
||||
from .diffusion import DiffusionEngine
|
||||
@@ -0,0 +1,335 @@
|
||||
import re
|
||||
from abc import abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Dict, Tuple, Union
|
||||
|
||||
import pytorch_lightning as pl
|
||||
import torch
|
||||
from omegaconf import ListConfig
|
||||
from packaging import version
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
|
||||
from ..modules.diffusionmodules.model import Decoder, Encoder
|
||||
from ..modules.distributions.distributions import DiagonalGaussianDistribution
|
||||
from ..modules.ema import LitEma
|
||||
from ..util import default, get_obj_from_str, instantiate_from_config
|
||||
|
||||
|
||||
class AbstractAutoencoder(pl.LightningModule):
|
||||
"""
|
||||
This is the base class for all autoencoders, including image autoencoders, image autoencoders with discriminators,
|
||||
unCLIP models, etc. Hence, it is fairly general, and specific features
|
||||
(e.g. discriminator training, encoding, decoding) must be implemented in subclasses.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ema_decay: Union[None, float] = None,
|
||||
monitor: Union[None, str] = None,
|
||||
input_key: str = "jpg",
|
||||
ckpt_path: Union[None, str] = None,
|
||||
ignore_keys: Union[Tuple, list, ListConfig] = (),
|
||||
):
|
||||
super().__init__()
|
||||
self.input_key = input_key
|
||||
self.use_ema = ema_decay is not None
|
||||
if monitor is not None:
|
||||
self.monitor = monitor
|
||||
|
||||
if self.use_ema:
|
||||
self.model_ema = LitEma(self, decay=ema_decay)
|
||||
print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
|
||||
|
||||
if ckpt_path is not None:
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
|
||||
|
||||
if version.parse(torch.__version__) >= version.parse("2.0.0"):
|
||||
self.automatic_optimization = False
|
||||
|
||||
def init_from_ckpt(
|
||||
self, path: str, ignore_keys: Union[Tuple, list, ListConfig] = tuple()
|
||||
) -> None:
|
||||
if path.endswith("ckpt"):
|
||||
sd = torch.load(path, map_location="cpu")["state_dict"]
|
||||
elif path.endswith("safetensors"):
|
||||
sd = load_safetensors(path)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
keys = list(sd.keys())
|
||||
for k in keys:
|
||||
for ik in ignore_keys:
|
||||
if re.match(ik, k):
|
||||
print("Deleting key {} from state_dict.".format(k))
|
||||
del sd[k]
|
||||
missing, unexpected = self.load_state_dict(sd, strict=False)
|
||||
print(
|
||||
f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys"
|
||||
)
|
||||
if len(missing) > 0:
|
||||
print(f"Missing Keys: {missing}")
|
||||
if len(unexpected) > 0:
|
||||
print(f"Unexpected Keys: {unexpected}")
|
||||
|
||||
@abstractmethod
|
||||
def get_input(self, batch) -> Any:
|
||||
raise NotImplementedError()
|
||||
|
||||
def on_train_batch_end(self, *args, **kwargs):
|
||||
# for EMA computation
|
||||
if self.use_ema:
|
||||
self.model_ema(self)
|
||||
|
||||
@contextmanager
|
||||
def ema_scope(self, context=None):
|
||||
if self.use_ema:
|
||||
self.model_ema.store(self.parameters())
|
||||
self.model_ema.copy_to(self)
|
||||
if context is not None:
|
||||
print(f"{context}: Switched to EMA weights")
|
||||
try:
|
||||
yield None
|
||||
finally:
|
||||
if self.use_ema:
|
||||
self.model_ema.restore(self.parameters())
|
||||
if context is not None:
|
||||
print(f"{context}: Restored training weights")
|
||||
|
||||
@abstractmethod
|
||||
def encode(self, *args, **kwargs) -> torch.Tensor:
|
||||
raise NotImplementedError("encode()-method of abstract base class called")
|
||||
|
||||
@abstractmethod
|
||||
def decode(self, *args, **kwargs) -> torch.Tensor:
|
||||
raise NotImplementedError("decode()-method of abstract base class called")
|
||||
|
||||
def instantiate_optimizer_from_config(self, params, lr, cfg):
|
||||
print(f"loading >>> {cfg['target']} <<< optimizer from config")
|
||||
return get_obj_from_str(cfg["target"])(
|
||||
params, lr=lr, **cfg.get("params", dict())
|
||||
)
|
||||
|
||||
def configure_optimizers(self) -> Any:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class AutoencodingEngine(AbstractAutoencoder):
|
||||
"""
|
||||
Base class for all image autoencoders that we train, like VQGAN or AutoencoderKL
|
||||
(we also restore them explicitly as special cases for legacy reasons).
|
||||
Regularizations such as KL or VQ are moved to the regularizer class.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
encoder_config: Dict,
|
||||
decoder_config: Dict,
|
||||
loss_config: Dict,
|
||||
regularizer_config: Dict,
|
||||
optimizer_config: Union[Dict, None] = None,
|
||||
lr_g_factor: float = 1.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
# todo: add options to freeze encoder/decoder
|
||||
self.encoder = instantiate_from_config(encoder_config)
|
||||
self.decoder = instantiate_from_config(decoder_config)
|
||||
self.loss = instantiate_from_config(loss_config)
|
||||
self.regularization = instantiate_from_config(regularizer_config)
|
||||
self.optimizer_config = default(
|
||||
optimizer_config, {"target": "torch.optim.Adam"}
|
||||
)
|
||||
self.lr_g_factor = lr_g_factor
|
||||
|
||||
def get_input(self, batch: Dict) -> torch.Tensor:
|
||||
# assuming unified data format, dataloader returns a dict.
|
||||
# image tensors should be scaled to -1 ... 1 and in channels-first format (e.g., bchw instead if bhwc)
|
||||
return batch[self.input_key]
|
||||
|
||||
def get_autoencoder_params(self) -> list:
|
||||
params = (
|
||||
list(self.encoder.parameters())
|
||||
+ list(self.decoder.parameters())
|
||||
+ list(self.regularization.get_trainable_parameters())
|
||||
+ list(self.loss.get_trainable_autoencoder_parameters())
|
||||
)
|
||||
return params
|
||||
|
||||
def get_discriminator_params(self) -> list:
|
||||
params = list(self.loss.get_trainable_parameters()) # e.g., discriminator
|
||||
return params
|
||||
|
||||
def get_last_layer(self):
|
||||
return self.decoder.get_last_layer()
|
||||
|
||||
def encode(self, x: Any, return_reg_log: bool = False) -> Any:
|
||||
z = self.encoder(x)
|
||||
z, reg_log = self.regularization(z)
|
||||
if return_reg_log:
|
||||
return z, reg_log
|
||||
return z
|
||||
|
||||
def decode(self, z: Any) -> torch.Tensor:
|
||||
x = self.decoder(z)
|
||||
return x
|
||||
|
||||
def forward(self, x: Any) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
z, reg_log = self.encode(x, return_reg_log=True)
|
||||
dec = self.decode(z)
|
||||
return z, dec, reg_log
|
||||
|
||||
def training_step(self, batch, batch_idx, optimizer_idx) -> Any:
|
||||
x = self.get_input(batch)
|
||||
z, xrec, regularization_log = self(x)
|
||||
|
||||
if optimizer_idx == 0:
|
||||
# autoencode
|
||||
aeloss, log_dict_ae = self.loss(
|
||||
regularization_log,
|
||||
x,
|
||||
xrec,
|
||||
optimizer_idx,
|
||||
self.global_step,
|
||||
last_layer=self.get_last_layer(),
|
||||
split="train",
|
||||
)
|
||||
|
||||
self.log_dict(
|
||||
log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True
|
||||
)
|
||||
return aeloss
|
||||
|
||||
if optimizer_idx == 1:
|
||||
# discriminator
|
||||
discloss, log_dict_disc = self.loss(
|
||||
regularization_log,
|
||||
x,
|
||||
xrec,
|
||||
optimizer_idx,
|
||||
self.global_step,
|
||||
last_layer=self.get_last_layer(),
|
||||
split="train",
|
||||
)
|
||||
self.log_dict(
|
||||
log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=True
|
||||
)
|
||||
return discloss
|
||||
|
||||
def validation_step(self, batch, batch_idx) -> Dict:
|
||||
log_dict = self._validation_step(batch, batch_idx)
|
||||
with self.ema_scope():
|
||||
log_dict_ema = self._validation_step(batch, batch_idx, postfix="_ema")
|
||||
log_dict.update(log_dict_ema)
|
||||
return log_dict
|
||||
|
||||
def _validation_step(self, batch, batch_idx, postfix="") -> Dict:
|
||||
x = self.get_input(batch)
|
||||
|
||||
z, xrec, regularization_log = self(x)
|
||||
aeloss, log_dict_ae = self.loss(
|
||||
regularization_log,
|
||||
x,
|
||||
xrec,
|
||||
0,
|
||||
self.global_step,
|
||||
last_layer=self.get_last_layer(),
|
||||
split="val" + postfix,
|
||||
)
|
||||
|
||||
discloss, log_dict_disc = self.loss(
|
||||
regularization_log,
|
||||
x,
|
||||
xrec,
|
||||
1,
|
||||
self.global_step,
|
||||
last_layer=self.get_last_layer(),
|
||||
split="val" + postfix,
|
||||
)
|
||||
self.log(f"val{postfix}/rec_loss", log_dict_ae[f"val{postfix}/rec_loss"])
|
||||
log_dict_ae.update(log_dict_disc)
|
||||
self.log_dict(log_dict_ae)
|
||||
return log_dict_ae
|
||||
|
||||
def configure_optimizers(self) -> Any:
|
||||
ae_params = self.get_autoencoder_params()
|
||||
disc_params = self.get_discriminator_params()
|
||||
|
||||
opt_ae = self.instantiate_optimizer_from_config(
|
||||
ae_params,
|
||||
default(self.lr_g_factor, 1.0) * self.learning_rate,
|
||||
self.optimizer_config,
|
||||
)
|
||||
opt_disc = self.instantiate_optimizer_from_config(
|
||||
disc_params, self.learning_rate, self.optimizer_config
|
||||
)
|
||||
|
||||
return [opt_ae, opt_disc], []
|
||||
|
||||
@torch.no_grad()
|
||||
def log_images(self, batch: Dict, **kwargs) -> Dict:
|
||||
log = dict()
|
||||
x = self.get_input(batch)
|
||||
_, xrec, _ = self(x)
|
||||
log["inputs"] = x
|
||||
log["reconstructions"] = xrec
|
||||
with self.ema_scope():
|
||||
_, xrec_ema, _ = self(x)
|
||||
log["reconstructions_ema"] = xrec_ema
|
||||
return log
|
||||
|
||||
|
||||
class AutoencoderKL(AutoencodingEngine):
|
||||
def __init__(self, embed_dim: int, **kwargs):
|
||||
ddconfig = kwargs.pop("ddconfig")
|
||||
ckpt_path = kwargs.pop("ckpt_path", None)
|
||||
ignore_keys = kwargs.pop("ignore_keys", ())
|
||||
super().__init__(
|
||||
encoder_config={"target": "torch.nn.Identity"},
|
||||
decoder_config={"target": "torch.nn.Identity"},
|
||||
regularizer_config={"target": "torch.nn.Identity"},
|
||||
loss_config=kwargs.pop("lossconfig"),
|
||||
**kwargs,
|
||||
)
|
||||
assert ddconfig["double_z"]
|
||||
self.encoder = Encoder(**ddconfig)
|
||||
self.decoder = Decoder(**ddconfig)
|
||||
self.quant_conv = torch.nn.Conv2d(2 * ddconfig["z_channels"], 2 * embed_dim, 1)
|
||||
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
if ckpt_path is not None:
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
|
||||
|
||||
def encode(self, x):
|
||||
assert (
|
||||
not self.training
|
||||
), f"{self.__class__.__name__} only supports inference currently"
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def decode(self, z, **decoder_kwargs):
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z, **decoder_kwargs)
|
||||
return dec
|
||||
|
||||
|
||||
class AutoencoderKLInferenceWrapper(AutoencoderKL):
|
||||
def encode(self, x):
|
||||
return super().encode(x).sample()
|
||||
|
||||
|
||||
class IdentityFirstStage(AbstractAutoencoder):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def get_input(self, x: Any) -> Any:
|
||||
return x
|
||||
|
||||
def encode(self, x: Any, *args, **kwargs) -> Any:
|
||||
return x
|
||||
|
||||
def decode(self, x: Any, *args, **kwargs) -> Any:
|
||||
return x
|
||||
@@ -0,0 +1,320 @@
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
import pytorch_lightning as pl
|
||||
import torch
|
||||
from omegaconf import ListConfig, OmegaConf
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
from torch.optim.lr_scheduler import LambdaLR
|
||||
|
||||
from ..modules import UNCONDITIONAL_CONFIG
|
||||
from ..modules.diffusionmodules.wrappers import OPENAIUNETWRAPPER
|
||||
from ..modules.ema import LitEma
|
||||
from ..util import (
|
||||
default,
|
||||
disabled_train,
|
||||
get_obj_from_str,
|
||||
instantiate_from_config,
|
||||
log_txt_as_img,
|
||||
)
|
||||
|
||||
|
||||
class DiffusionEngine(pl.LightningModule):
|
||||
def __init__(
|
||||
self,
|
||||
network_config,
|
||||
denoiser_config,
|
||||
first_stage_config,
|
||||
conditioner_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||
sampler_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||
optimizer_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||
scheduler_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||
loss_fn_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||
network_wrapper: Union[None, str] = None,
|
||||
ckpt_path: Union[None, str] = None,
|
||||
use_ema: bool = False,
|
||||
ema_decay_rate: float = 0.9999,
|
||||
scale_factor: float = 1.0,
|
||||
disable_first_stage_autocast=False,
|
||||
input_key: str = "jpg",
|
||||
log_keys: Union[List, None] = None,
|
||||
no_cond_log: bool = False,
|
||||
compile_model: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.log_keys = log_keys
|
||||
self.input_key = input_key
|
||||
self.optimizer_config = default(
|
||||
optimizer_config, {"target": "torch.optim.AdamW"}
|
||||
)
|
||||
model = instantiate_from_config(network_config)
|
||||
self.model = get_obj_from_str(default(network_wrapper, OPENAIUNETWRAPPER))(
|
||||
model, compile_model=compile_model
|
||||
)
|
||||
|
||||
self.denoiser = instantiate_from_config(denoiser_config)
|
||||
self.sampler = (
|
||||
instantiate_from_config(sampler_config)
|
||||
if sampler_config is not None
|
||||
else None
|
||||
)
|
||||
self.conditioner = instantiate_from_config(
|
||||
default(conditioner_config, UNCONDITIONAL_CONFIG)
|
||||
)
|
||||
self.scheduler_config = scheduler_config
|
||||
self._init_first_stage(first_stage_config)
|
||||
|
||||
self.loss_fn = (
|
||||
instantiate_from_config(loss_fn_config)
|
||||
if loss_fn_config is not None
|
||||
else None
|
||||
)
|
||||
|
||||
self.use_ema = use_ema
|
||||
if self.use_ema:
|
||||
self.model_ema = LitEma(self.model, decay=ema_decay_rate)
|
||||
print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
|
||||
|
||||
self.scale_factor = scale_factor
|
||||
self.disable_first_stage_autocast = disable_first_stage_autocast
|
||||
self.no_cond_log = no_cond_log
|
||||
|
||||
if ckpt_path is not None:
|
||||
self.init_from_ckpt(ckpt_path)
|
||||
|
||||
def init_from_ckpt(
|
||||
self,
|
||||
path: str,
|
||||
) -> None:
|
||||
if path.endswith("ckpt"):
|
||||
sd = torch.load(path, map_location="cpu")["state_dict"]
|
||||
elif path.endswith("safetensors"):
|
||||
sd = load_safetensors(path)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
missing, unexpected = self.load_state_dict(sd, strict=False)
|
||||
print(
|
||||
f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys"
|
||||
)
|
||||
if len(missing) > 0:
|
||||
print(f"Missing Keys: {missing}")
|
||||
if len(unexpected) > 0:
|
||||
print(f"Unexpected Keys: {unexpected}")
|
||||
|
||||
def _init_first_stage(self, config):
|
||||
model = instantiate_from_config(config).eval()
|
||||
model.train = disabled_train
|
||||
for param in model.parameters():
|
||||
param.requires_grad = False
|
||||
self.first_stage_model = model
|
||||
|
||||
def get_input(self, batch):
|
||||
# assuming unified data format, dataloader returns a dict.
|
||||
# image tensors should be scaled to -1 ... 1 and in bchw format
|
||||
return batch[self.input_key]
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
z = 1.0 / self.scale_factor * z
|
||||
with torch.autocast("cuda", enabled=not self.disable_first_stage_autocast):
|
||||
out = self.first_stage_model.decode(z)
|
||||
return out
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x):
|
||||
with torch.autocast("cuda", enabled=not self.disable_first_stage_autocast):
|
||||
z = self.first_stage_model.encode(x)
|
||||
z = self.scale_factor * z
|
||||
return z
|
||||
|
||||
def forward(self, x, batch):
|
||||
loss = self.loss_fn(self.model, self.denoiser, self.conditioner, x, batch)
|
||||
loss_mean = loss.mean()
|
||||
loss_dict = {"loss": loss_mean}
|
||||
return loss_mean, loss_dict
|
||||
|
||||
def shared_step(self, batch: Dict) -> Any:
|
||||
x = self.get_input(batch)
|
||||
x = self.encode_first_stage(x)
|
||||
batch["global_step"] = self.global_step
|
||||
loss, loss_dict = self(x, batch)
|
||||
return loss, loss_dict
|
||||
|
||||
def training_step(self, batch, batch_idx):
|
||||
loss, loss_dict = self.shared_step(batch)
|
||||
|
||||
self.log_dict(
|
||||
loss_dict, prog_bar=True, logger=True, on_step=True, on_epoch=False
|
||||
)
|
||||
|
||||
self.log(
|
||||
"global_step",
|
||||
self.global_step,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
on_step=True,
|
||||
on_epoch=False,
|
||||
)
|
||||
|
||||
# if self.scheduler_config is not None:
|
||||
lr = self.optimizers().param_groups[0]["lr"]
|
||||
self.log(
|
||||
"lr_abs", lr, prog_bar=True, logger=True, on_step=True, on_epoch=False
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
def on_train_start(self, *args, **kwargs):
|
||||
if self.sampler is None or self.loss_fn is None:
|
||||
raise ValueError("Sampler and loss function need to be set for training.")
|
||||
|
||||
def on_train_batch_end(self, *args, **kwargs):
|
||||
if self.use_ema:
|
||||
self.model_ema(self.model)
|
||||
|
||||
@contextmanager
|
||||
def ema_scope(self, context=None):
|
||||
if self.use_ema:
|
||||
self.model_ema.store(self.model.parameters())
|
||||
self.model_ema.copy_to(self.model)
|
||||
if context is not None:
|
||||
print(f"{context}: Switched to EMA weights")
|
||||
try:
|
||||
yield None
|
||||
finally:
|
||||
if self.use_ema:
|
||||
self.model_ema.restore(self.model.parameters())
|
||||
if context is not None:
|
||||
print(f"{context}: Restored training weights")
|
||||
|
||||
def instantiate_optimizer_from_config(self, params, lr, cfg):
|
||||
return get_obj_from_str(cfg["target"])(
|
||||
params, lr=lr, **cfg.get("params", dict())
|
||||
)
|
||||
|
||||
def configure_optimizers(self):
|
||||
lr = self.learning_rate
|
||||
params = list(self.model.parameters())
|
||||
for embedder in self.conditioner.embedders:
|
||||
if embedder.is_trainable:
|
||||
params = params + list(embedder.parameters())
|
||||
opt = self.instantiate_optimizer_from_config(params, lr, self.optimizer_config)
|
||||
if self.scheduler_config is not None:
|
||||
scheduler = instantiate_from_config(self.scheduler_config)
|
||||
print("Setting up LambdaLR scheduler...")
|
||||
scheduler = [
|
||||
{
|
||||
"scheduler": LambdaLR(opt, lr_lambda=scheduler.schedule),
|
||||
"interval": "step",
|
||||
"frequency": 1,
|
||||
}
|
||||
]
|
||||
return [opt], scheduler
|
||||
return opt
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(
|
||||
self,
|
||||
cond: Dict,
|
||||
uc: Union[Dict, None] = None,
|
||||
batch_size: int = 16,
|
||||
shape: Union[None, Tuple, List] = None,
|
||||
**kwargs,
|
||||
):
|
||||
randn = torch.randn(batch_size, *shape).to(self.device)
|
||||
|
||||
denoiser = lambda input, sigma, c: self.denoiser(
|
||||
self.model, input, sigma, c, **kwargs
|
||||
)
|
||||
samples = self.sampler(denoiser, randn, cond, uc=uc)
|
||||
return samples
|
||||
|
||||
@torch.no_grad()
|
||||
def log_conditionings(self, batch: Dict, n: int) -> Dict:
|
||||
"""
|
||||
Defines heuristics to log different conditionings.
|
||||
These can be lists of strings (text-to-image), tensors, ints, ...
|
||||
"""
|
||||
image_h, image_w = batch[self.input_key].shape[2:]
|
||||
log = dict()
|
||||
|
||||
for embedder in self.conditioner.embedders:
|
||||
if (
|
||||
(self.log_keys is None) or (embedder.input_key in self.log_keys)
|
||||
) and not self.no_cond_log:
|
||||
x = batch[embedder.input_key][:n]
|
||||
if isinstance(x, torch.Tensor):
|
||||
if x.dim() == 1:
|
||||
# class-conditional, convert integer to string
|
||||
x = [str(x[i].item()) for i in range(x.shape[0])]
|
||||
xc = log_txt_as_img((image_h, image_w), x, size=image_h // 4)
|
||||
elif x.dim() == 2:
|
||||
# size and crop cond and the like
|
||||
x = [
|
||||
"x".join([str(xx) for xx in x[i].tolist()])
|
||||
for i in range(x.shape[0])
|
||||
]
|
||||
xc = log_txt_as_img((image_h, image_w), x, size=image_h // 20)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
elif isinstance(x, (List, ListConfig)):
|
||||
if isinstance(x[0], str):
|
||||
# strings
|
||||
xc = log_txt_as_img((image_h, image_w), x, size=image_h // 20)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
log[embedder.input_key] = xc
|
||||
return log
|
||||
|
||||
@torch.no_grad()
|
||||
def log_images(
|
||||
self,
|
||||
batch: Dict,
|
||||
N: int = 8,
|
||||
sample: bool = True,
|
||||
ucg_keys: List[str] = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
conditioner_input_keys = [e.input_key for e in self.conditioner.embedders]
|
||||
if ucg_keys:
|
||||
assert all(map(lambda x: x in conditioner_input_keys, ucg_keys)), (
|
||||
"Each defined ucg key for sampling must be in the provided conditioner input keys,"
|
||||
f"but we have {ucg_keys} vs. {conditioner_input_keys}"
|
||||
)
|
||||
else:
|
||||
ucg_keys = conditioner_input_keys
|
||||
log = dict()
|
||||
|
||||
x = self.get_input(batch)
|
||||
|
||||
c, uc = self.conditioner.get_unconditional_conditioning(
|
||||
batch,
|
||||
force_uc_zero_embeddings=ucg_keys
|
||||
if len(self.conditioner.embedders) > 0
|
||||
else [],
|
||||
)
|
||||
|
||||
sampling_kwargs = {}
|
||||
|
||||
N = min(x.shape[0], N)
|
||||
x = x.to(self.device)[:N]
|
||||
log["inputs"] = x
|
||||
z = self.encode_first_stage(x)
|
||||
log["reconstructions"] = self.decode_first_stage(z)
|
||||
log.update(self.log_conditionings(batch, N))
|
||||
|
||||
for k in c:
|
||||
if isinstance(c[k], torch.Tensor):
|
||||
c[k], uc[k] = map(lambda y: y[k][:N].to(self.device), (c, uc))
|
||||
|
||||
if sample:
|
||||
with self.ema_scope("Plotting"):
|
||||
samples = self.sample(
|
||||
c, shape=z.shape[1:], uc=uc, batch_size=N, **sampling_kwargs
|
||||
)
|
||||
samples = self.decode_first_stage(samples)
|
||||
log["samples"] = samples
|
||||
return log
|
||||
@@ -0,0 +1,8 @@
|
||||
from .encoders.modules import GeneralConditioner
|
||||
from .encoders.modules import GeneralConditionerWithControl
|
||||
from .encoders.modules import PreparedConditioner
|
||||
|
||||
UNCONDITIONAL_CONFIG = {
|
||||
"target": ".sgm.modules.GeneralConditioner",
|
||||
"params": {"emb_models": []},
|
||||
}
|
||||
@@ -0,0 +1,635 @@
|
||||
import math
|
||||
from inspect import isfunction
|
||||
from typing import Any, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
# from einops._torch_specific import allow_ops_in_compiled_graph
|
||||
# allow_ops_in_compiled_graph()
|
||||
from einops import rearrange, repeat
|
||||
from packaging import version
|
||||
from torch import nn
|
||||
|
||||
if version.parse(torch.__version__) >= version.parse("2.0.0"):
|
||||
SDP_IS_AVAILABLE = True
|
||||
from torch.backends.cuda import SDPBackend, sdp_kernel
|
||||
|
||||
BACKEND_MAP = {
|
||||
SDPBackend.MATH: {
|
||||
"enable_math": True,
|
||||
"enable_flash": False,
|
||||
"enable_mem_efficient": False,
|
||||
},
|
||||
SDPBackend.FLASH_ATTENTION: {
|
||||
"enable_math": False,
|
||||
"enable_flash": True,
|
||||
"enable_mem_efficient": False,
|
||||
},
|
||||
SDPBackend.EFFICIENT_ATTENTION: {
|
||||
"enable_math": False,
|
||||
"enable_flash": False,
|
||||
"enable_mem_efficient": True,
|
||||
},
|
||||
None: {"enable_math": True, "enable_flash": True, "enable_mem_efficient": True},
|
||||
}
|
||||
else:
|
||||
from contextlib import nullcontext
|
||||
|
||||
SDP_IS_AVAILABLE = False
|
||||
sdp_kernel = nullcontext
|
||||
BACKEND_MAP = {}
|
||||
print(
|
||||
f"No SDP backend available, likely because you are running in pytorch versions < 2.0. In fact, "
|
||||
f"you are using PyTorch {torch.__version__}. You might want to consider upgrading."
|
||||
)
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
|
||||
XFORMERS_IS_AVAILABLE = True
|
||||
except:
|
||||
XFORMERS_IS_AVAILABLE = False
|
||||
print("no module 'xformers'. Processing without...")
|
||||
|
||||
from .diffusionmodules.util import checkpoint
|
||||
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
|
||||
def uniq(arr):
|
||||
return {el: True for el in arr}.keys()
|
||||
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if isfunction(d) else d
|
||||
|
||||
|
||||
def max_neg_value(t):
|
||||
return -torch.finfo(t.dtype).max
|
||||
|
||||
|
||||
def init_(tensor):
|
||||
dim = tensor.shape[-1]
|
||||
std = 1 / math.sqrt(dim)
|
||||
tensor.uniform_(-std, std)
|
||||
return tensor
|
||||
|
||||
|
||||
# feedforward
|
||||
class GEGLU(nn.Module):
|
||||
def __init__(self, dim_in, dim_out):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out * 2)
|
||||
|
||||
def forward(self, x):
|
||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
||||
return x * F.gelu(gate)
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.0):
|
||||
super().__init__()
|
||||
inner_dim = int(dim * mult)
|
||||
dim_out = default(dim_out, dim)
|
||||
project_in = (
|
||||
nn.Sequential(nn.Linear(dim, inner_dim), nn.GELU())
|
||||
if not glu
|
||||
else GEGLU(dim, inner_dim)
|
||||
)
|
||||
|
||||
self.net = nn.Sequential(
|
||||
project_in, nn.Dropout(dropout), nn.Linear(inner_dim, dim_out)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
"""
|
||||
Zero out the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
def Normalize(in_channels):
|
||||
return torch.nn.GroupNorm(
|
||||
num_groups=32, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
|
||||
|
||||
class LinearAttention(nn.Module):
|
||||
def __init__(self, dim, heads=4, dim_head=32):
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
hidden_dim = dim_head * heads
|
||||
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias=False)
|
||||
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
|
||||
|
||||
def forward(self, x):
|
||||
b, c, h, w = x.shape
|
||||
qkv = self.to_qkv(x)
|
||||
q, k, v = rearrange(
|
||||
qkv, "b (qkv heads c) h w -> qkv b heads c (h w)", heads=self.heads, qkv=3
|
||||
)
|
||||
k = k.softmax(dim=-1)
|
||||
context = torch.einsum("bhdn,bhen->bhde", k, v)
|
||||
out = torch.einsum("bhde,bhdn->bhen", context, q)
|
||||
out = rearrange(
|
||||
out, "b heads c (h w) -> b (heads c) h w", heads=self.heads, h=h, w=w
|
||||
)
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class SpatialSelfAttention(nn.Module):
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.k = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.v = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.proj_out = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b, c, h, w = q.shape
|
||||
q = rearrange(q, "b c h w -> b (h w) c")
|
||||
k = rearrange(k, "b c h w -> b c (h w)")
|
||||
w_ = torch.einsum("bij,bjk->bik", q, k)
|
||||
|
||||
w_ = w_ * (int(c) ** (-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = rearrange(v, "b c h w -> b c (h w)")
|
||||
w_ = rearrange(w_, "b i j -> b j i")
|
||||
h_ = torch.einsum("bij,bjk->bik", v, w_)
|
||||
h_ = rearrange(h_, "b c (h w) -> b c h w", h=h)
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x + h_
|
||||
|
||||
|
||||
class CrossAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
query_dim,
|
||||
context_dim=None,
|
||||
heads=8,
|
||||
dim_head=64,
|
||||
dropout=0.0,
|
||||
backend=None,
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = default(context_dim, query_dim)
|
||||
|
||||
self.scale = dim_head**-0.5
|
||||
self.heads = heads
|
||||
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||
)
|
||||
self.backend = backend
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
context=None,
|
||||
mask=None,
|
||||
additional_tokens=None,
|
||||
n_times_crossframe_attn_in_self=0,
|
||||
):
|
||||
h = self.heads
|
||||
|
||||
if additional_tokens is not None:
|
||||
# get the number of masked tokens at the beginning of the output sequence
|
||||
n_tokens_to_mask = additional_tokens.shape[1]
|
||||
# add additional token
|
||||
x = torch.cat([additional_tokens, x], dim=1)
|
||||
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
if n_times_crossframe_attn_in_self:
|
||||
# reprogramming cross-frame attention as in https://arxiv.org/abs/2303.13439
|
||||
assert x.shape[0] % n_times_crossframe_attn_in_self == 0
|
||||
n_cp = x.shape[0] // n_times_crossframe_attn_in_self
|
||||
k = repeat(
|
||||
k[::n_times_crossframe_attn_in_self], "b ... -> (b n) ...", n=n_cp
|
||||
)
|
||||
v = repeat(
|
||||
v[::n_times_crossframe_attn_in_self], "b ... -> (b n) ...", n=n_cp
|
||||
)
|
||||
|
||||
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
||||
|
||||
## old
|
||||
"""
|
||||
sim = einsum('b i d, b j d -> b i j', q, k) * self.scale
|
||||
del q, k
|
||||
|
||||
if exists(mask):
|
||||
mask = rearrange(mask, 'b ... -> b (...)')
|
||||
max_neg_value = -torch.finfo(sim.dtype).max
|
||||
mask = 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 = einsum('b i j, b j d -> b i d', sim, v)
|
||||
"""
|
||||
## new
|
||||
with sdp_kernel(**BACKEND_MAP[self.backend]):
|
||||
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=mask
|
||||
) # scale is dim_head ** -0.5 per default
|
||||
|
||||
del q, k, v
|
||||
out = rearrange(out, "b h n d -> b n (h d)", h=h)
|
||||
|
||||
if additional_tokens is not None:
|
||||
# remove additional token
|
||||
out = out[:, n_tokens_to_mask:]
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class MemoryEfficientCrossAttention(nn.Module):
|
||||
# https://github.com/MatthieuTPHR/diffusers/blob/d80b531ff8060ec1ea982b65a1b8df70f73aa67c/src/diffusers/models/attention.py#L223
|
||||
def __init__(
|
||||
self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.0, **kwargs
|
||||
):
|
||||
super().__init__()
|
||||
print(
|
||||
f"Setting up {self.__class__.__name__}. Query dim is {query_dim}, context_dim is {context_dim} and using "
|
||||
f"{heads} heads with a dimension of {dim_head}."
|
||||
)
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = default(context_dim, query_dim)
|
||||
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||
)
|
||||
self.attention_op: Optional[Any] = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
context=None,
|
||||
mask=None,
|
||||
additional_tokens=None,
|
||||
n_times_crossframe_attn_in_self=0,
|
||||
):
|
||||
if additional_tokens is not None:
|
||||
# get the number of masked tokens at the beginning of the output sequence
|
||||
n_tokens_to_mask = additional_tokens.shape[1]
|
||||
# add additional token
|
||||
x = torch.cat([additional_tokens, x], dim=1)
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
if n_times_crossframe_attn_in_self:
|
||||
# reprogramming cross-frame attention as in https://arxiv.org/abs/2303.13439
|
||||
assert x.shape[0] % n_times_crossframe_attn_in_self == 0
|
||||
# n_cp = x.shape[0]//n_times_crossframe_attn_in_self
|
||||
k = repeat(
|
||||
k[::n_times_crossframe_attn_in_self],
|
||||
"b ... -> (b n) ...",
|
||||
n=n_times_crossframe_attn_in_self,
|
||||
)
|
||||
v = repeat(
|
||||
v[::n_times_crossframe_attn_in_self],
|
||||
"b ... -> (b n) ...",
|
||||
n=n_times_crossframe_attn_in_self,
|
||||
)
|
||||
|
||||
b, _, _ = q.shape
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3)
|
||||
.reshape(b, t.shape[1], self.heads, self.dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b * self.heads, t.shape[1], self.dim_head)
|
||||
.contiguous(),
|
||||
(q, k, v),
|
||||
)
|
||||
|
||||
# actually compute the attention, what we cannot get enough of
|
||||
out = xformers.ops.memory_efficient_attention(
|
||||
q, k, v, attn_bias=None, op=self.attention_op
|
||||
)
|
||||
|
||||
# TODO: Use this directly in the attention operation, as a bias
|
||||
if exists(mask):
|
||||
raise NotImplementedError
|
||||
out = (
|
||||
out.unsqueeze(0)
|
||||
.reshape(b, self.heads, out.shape[1], self.dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b, out.shape[1], self.heads * self.dim_head)
|
||||
)
|
||||
if additional_tokens is not None:
|
||||
# remove additional token
|
||||
out = out[:, n_tokens_to_mask:]
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
ATTENTION_MODES = {
|
||||
"softmax": CrossAttention, # vanilla attention
|
||||
"softmax-xformers": MemoryEfficientCrossAttention, # ampere
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=0.0,
|
||||
context_dim=None,
|
||||
gated_ff=True,
|
||||
checkpoint=True,
|
||||
disable_self_attn=False,
|
||||
attn_mode="softmax",
|
||||
sdp_backend=None,
|
||||
):
|
||||
super().__init__()
|
||||
assert attn_mode in self.ATTENTION_MODES
|
||||
if attn_mode != "softmax" and not XFORMERS_IS_AVAILABLE:
|
||||
print(
|
||||
f"Attention mode '{attn_mode}' is not available. Falling back to native attention. "
|
||||
f"This is not a problem in Pytorch >= 2.0. FYI, you are running with PyTorch version {torch.__version__}"
|
||||
)
|
||||
attn_mode = "softmax"
|
||||
elif attn_mode == "softmax" and not SDP_IS_AVAILABLE:
|
||||
print(
|
||||
"We do not support vanilla attention anymore, as it is too expensive. Sorry."
|
||||
)
|
||||
if not XFORMERS_IS_AVAILABLE:
|
||||
assert (
|
||||
False
|
||||
), "Please install xformers via e.g. 'pip install xformers==0.0.16'"
|
||||
else:
|
||||
print("Falling back to xformers efficient attention.")
|
||||
attn_mode = "softmax-xformers"
|
||||
attn_cls = self.ATTENTION_MODES[attn_mode]
|
||||
if version.parse(torch.__version__) >= version.parse("2.0.0"):
|
||||
assert sdp_backend is None or isinstance(sdp_backend, SDPBackend)
|
||||
else:
|
||||
assert sdp_backend is None
|
||||
self.disable_self_attn = disable_self_attn
|
||||
self.attn1 = attn_cls(
|
||||
query_dim=dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim if self.disable_self_attn else None,
|
||||
backend=sdp_backend,
|
||||
) # is a self-attention if not self.disable_self_attn
|
||||
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)
|
||||
self.attn2 = attn_cls(
|
||||
query_dim=dim,
|
||||
context_dim=context_dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
dropout=dropout,
|
||||
backend=sdp_backend,
|
||||
) # is self-attn if context is none
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
self.norm3 = nn.LayerNorm(dim)
|
||||
self.checkpoint = checkpoint
|
||||
if self.checkpoint:
|
||||
print(f"{self.__class__.__name__} is using checkpointing")
|
||||
|
||||
def forward(
|
||||
self, x, context=None, additional_tokens=None, n_times_crossframe_attn_in_self=0
|
||||
):
|
||||
kwargs = {"x": x}
|
||||
|
||||
if context is not None:
|
||||
kwargs.update({"context": context})
|
||||
|
||||
if additional_tokens is not None:
|
||||
kwargs.update({"additional_tokens": additional_tokens})
|
||||
|
||||
if n_times_crossframe_attn_in_self:
|
||||
kwargs.update(
|
||||
{"n_times_crossframe_attn_in_self": n_times_crossframe_attn_in_self}
|
||||
)
|
||||
|
||||
# return mixed_checkpoint(self._forward, kwargs, self.parameters(), self.checkpoint)
|
||||
return checkpoint(
|
||||
self._forward, (x, context), self.parameters(), self.checkpoint
|
||||
)
|
||||
|
||||
def _forward(
|
||||
self, x, context=None, additional_tokens=None, n_times_crossframe_attn_in_self=0
|
||||
):
|
||||
x = (
|
||||
self.attn1(
|
||||
self.norm1(x),
|
||||
context=context if self.disable_self_attn else None,
|
||||
additional_tokens=additional_tokens,
|
||||
n_times_crossframe_attn_in_self=n_times_crossframe_attn_in_self
|
||||
if not self.disable_self_attn
|
||||
else 0,
|
||||
)
|
||||
+ x
|
||||
)
|
||||
x = (
|
||||
self.attn2(
|
||||
self.norm2(x), context=context, additional_tokens=additional_tokens
|
||||
)
|
||||
+ x
|
||||
)
|
||||
x = self.ff(self.norm3(x)) + x
|
||||
return x
|
||||
|
||||
|
||||
class BasicTransformerSingleLayerBlock(nn.Module):
|
||||
ATTENTION_MODES = {
|
||||
"softmax": CrossAttention, # vanilla attention
|
||||
"softmax-xformers": MemoryEfficientCrossAttention # on the A100s not quite as fast as the above version
|
||||
# (todo might depend on head_dim, check, falls back to semi-optimized kernels for dim!=[16,32,64,128])
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=0.0,
|
||||
context_dim=None,
|
||||
gated_ff=True,
|
||||
checkpoint=True,
|
||||
attn_mode="softmax",
|
||||
):
|
||||
super().__init__()
|
||||
assert attn_mode in self.ATTENTION_MODES
|
||||
attn_cls = self.ATTENTION_MODES[attn_mode]
|
||||
self.attn1 = attn_cls(
|
||||
query_dim=dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim,
|
||||
)
|
||||
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
self.checkpoint = checkpoint
|
||||
|
||||
def forward(self, x, context=None):
|
||||
return checkpoint(
|
||||
self._forward, (x, context), self.parameters(), self.checkpoint
|
||||
)
|
||||
|
||||
def _forward(self, x, context=None):
|
||||
x = self.attn1(self.norm1(x), context=context) + x
|
||||
x = self.ff(self.norm2(x)) + x
|
||||
return x
|
||||
|
||||
|
||||
class SpatialTransformer(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data.
|
||||
First, project the input (aka embedding)
|
||||
and reshape to b, t, d.
|
||||
Then apply standard transformer action.
|
||||
Finally, reshape to image
|
||||
NEW: use_linear for more efficiency instead of the 1x1 convs
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
n_heads,
|
||||
d_head,
|
||||
depth=1,
|
||||
dropout=0.0,
|
||||
context_dim=None,
|
||||
disable_self_attn=False,
|
||||
use_linear=False,
|
||||
attn_type="softmax",
|
||||
use_checkpoint=True,
|
||||
# sdp_backend=SDPBackend.FLASH_ATTENTION
|
||||
sdp_backend=None,
|
||||
):
|
||||
super().__init__()
|
||||
print(
|
||||
f"constructing {self.__class__.__name__} of depth {depth} w/ {in_channels} channels and {n_heads} heads"
|
||||
)
|
||||
from omegaconf import ListConfig
|
||||
|
||||
if exists(context_dim) and not isinstance(context_dim, (list, ListConfig)):
|
||||
context_dim = [context_dim]
|
||||
if exists(context_dim) and isinstance(context_dim, list):
|
||||
if depth != len(context_dim):
|
||||
print(
|
||||
f"WARNING: {self.__class__.__name__}: Found context dims {context_dim} of depth {len(context_dim)}, "
|
||||
f"which does not match the specified 'depth' of {depth}. Setting context_dim to {depth * [context_dim[0]]} now."
|
||||
)
|
||||
# depth does not match context dims.
|
||||
assert all(
|
||||
map(lambda x: x == context_dim[0], context_dim)
|
||||
), "need homogenous context_dim to match depth automatically"
|
||||
context_dim = depth * [context_dim[0]]
|
||||
elif context_dim is None:
|
||||
context_dim = [None] * depth
|
||||
self.in_channels = in_channels
|
||||
inner_dim = n_heads * d_head
|
||||
self.norm = Normalize(in_channels)
|
||||
if not use_linear:
|
||||
self.proj_in = nn.Conv2d(
|
||||
in_channels, inner_dim, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
else:
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim[d],
|
||||
disable_self_attn=disable_self_attn,
|
||||
attn_mode=attn_type,
|
||||
checkpoint=use_checkpoint,
|
||||
sdp_backend=sdp_backend,
|
||||
)
|
||||
for d in range(depth)
|
||||
]
|
||||
)
|
||||
if not use_linear:
|
||||
self.proj_out = zero_module(
|
||||
nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
)
|
||||
else:
|
||||
# self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))
|
||||
self.proj_out = zero_module(nn.Linear(inner_dim, in_channels))
|
||||
self.use_linear = use_linear
|
||||
|
||||
def forward(self, x, context=None):
|
||||
# note: if no context is given, cross-attention defaults to self-attention
|
||||
if not isinstance(context, list):
|
||||
context = [context]
|
||||
b, c, h, w = x.shape
|
||||
x_in = x
|
||||
x = self.norm(x)
|
||||
if not self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
x = rearrange(x, "b c h w -> b (h w) c").contiguous()
|
||||
if self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
if i > 0 and len(context) == 1:
|
||||
i = 0 # use same context for each block
|
||||
x = block(x, context=context[i])
|
||||
if self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
x = rearrange(x, "b (h w) c -> b c h w", h=h, w=w).contiguous()
|
||||
if not self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
return x + x_in
|
||||
@@ -0,0 +1,246 @@
|
||||
from typing import Any, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from ....util import default, instantiate_from_config
|
||||
from ..lpips.loss.lpips import LPIPS
|
||||
from ..lpips.model.model import NLayerDiscriminator, weights_init
|
||||
from ..lpips.vqperceptual import hinge_d_loss, vanilla_d_loss
|
||||
|
||||
|
||||
def adopt_weight(weight, global_step, threshold=0, value=0.0):
|
||||
if global_step < threshold:
|
||||
weight = value
|
||||
return weight
|
||||
|
||||
|
||||
class LatentLPIPS(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
decoder_config,
|
||||
perceptual_weight=1.0,
|
||||
latent_weight=1.0,
|
||||
scale_input_to_tgt_size=False,
|
||||
scale_tgt_to_input_size=False,
|
||||
perceptual_weight_on_inputs=0.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.scale_input_to_tgt_size = scale_input_to_tgt_size
|
||||
self.scale_tgt_to_input_size = scale_tgt_to_input_size
|
||||
self.init_decoder(decoder_config)
|
||||
self.perceptual_loss = LPIPS().eval()
|
||||
self.perceptual_weight = perceptual_weight
|
||||
self.latent_weight = latent_weight
|
||||
self.perceptual_weight_on_inputs = perceptual_weight_on_inputs
|
||||
|
||||
def init_decoder(self, config):
|
||||
self.decoder = instantiate_from_config(config)
|
||||
if hasattr(self.decoder, "encoder"):
|
||||
del self.decoder.encoder
|
||||
|
||||
def forward(self, latent_inputs, latent_predictions, image_inputs, split="train"):
|
||||
log = dict()
|
||||
loss = (latent_inputs - latent_predictions) ** 2
|
||||
log[f"{split}/latent_l2_loss"] = loss.mean().detach()
|
||||
image_reconstructions = None
|
||||
if self.perceptual_weight > 0.0:
|
||||
image_reconstructions = self.decoder.decode(latent_predictions)
|
||||
image_targets = self.decoder.decode(latent_inputs)
|
||||
perceptual_loss = self.perceptual_loss(
|
||||
image_targets.contiguous(), image_reconstructions.contiguous()
|
||||
)
|
||||
loss = (
|
||||
self.latent_weight * loss.mean()
|
||||
+ self.perceptual_weight * perceptual_loss.mean()
|
||||
)
|
||||
log[f"{split}/perceptual_loss"] = perceptual_loss.mean().detach()
|
||||
|
||||
if self.perceptual_weight_on_inputs > 0.0:
|
||||
image_reconstructions = default(
|
||||
image_reconstructions, self.decoder.decode(latent_predictions)
|
||||
)
|
||||
if self.scale_input_to_tgt_size:
|
||||
image_inputs = torch.nn.functional.interpolate(
|
||||
image_inputs,
|
||||
image_reconstructions.shape[2:],
|
||||
mode="bicubic",
|
||||
antialias=True,
|
||||
)
|
||||
elif self.scale_tgt_to_input_size:
|
||||
image_reconstructions = torch.nn.functional.interpolate(
|
||||
image_reconstructions,
|
||||
image_inputs.shape[2:],
|
||||
mode="bicubic",
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
perceptual_loss2 = self.perceptual_loss(
|
||||
image_inputs.contiguous(), image_reconstructions.contiguous()
|
||||
)
|
||||
loss = loss + self.perceptual_weight_on_inputs * perceptual_loss2.mean()
|
||||
log[f"{split}/perceptual_loss_on_inputs"] = perceptual_loss2.mean().detach()
|
||||
return loss, log
|
||||
|
||||
|
||||
class GeneralLPIPSWithDiscriminator(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
disc_start: int,
|
||||
logvar_init: float = 0.0,
|
||||
pixelloss_weight=1.0,
|
||||
disc_num_layers: int = 3,
|
||||
disc_in_channels: int = 3,
|
||||
disc_factor: float = 1.0,
|
||||
disc_weight: float = 1.0,
|
||||
perceptual_weight: float = 1.0,
|
||||
disc_loss: str = "hinge",
|
||||
scale_input_to_tgt_size: bool = False,
|
||||
dims: int = 2,
|
||||
learn_logvar: bool = False,
|
||||
regularization_weights: Union[None, dict] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dims = dims
|
||||
if self.dims > 2:
|
||||
print(
|
||||
f"running with dims={dims}. This means that for perceptual loss calculation, "
|
||||
f"the LPIPS loss will be applied to each frame independently. "
|
||||
)
|
||||
self.scale_input_to_tgt_size = scale_input_to_tgt_size
|
||||
assert disc_loss in ["hinge", "vanilla"]
|
||||
self.pixel_weight = pixelloss_weight
|
||||
self.perceptual_loss = LPIPS().eval()
|
||||
self.perceptual_weight = perceptual_weight
|
||||
# output log variance
|
||||
self.logvar = nn.Parameter(torch.ones(size=()) * logvar_init)
|
||||
self.learn_logvar = learn_logvar
|
||||
|
||||
self.discriminator = NLayerDiscriminator(
|
||||
input_nc=disc_in_channels, n_layers=disc_num_layers, use_actnorm=False
|
||||
).apply(weights_init)
|
||||
self.discriminator_iter_start = disc_start
|
||||
self.disc_loss = hinge_d_loss if disc_loss == "hinge" else vanilla_d_loss
|
||||
self.disc_factor = disc_factor
|
||||
self.discriminator_weight = disc_weight
|
||||
self.regularization_weights = default(regularization_weights, {})
|
||||
|
||||
def get_trainable_parameters(self) -> Any:
|
||||
return self.discriminator.parameters()
|
||||
|
||||
def get_trainable_autoencoder_parameters(self) -> Any:
|
||||
if self.learn_logvar:
|
||||
yield self.logvar
|
||||
yield from ()
|
||||
|
||||
def calculate_adaptive_weight(self, nll_loss, g_loss, last_layer=None):
|
||||
if last_layer is not None:
|
||||
nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0]
|
||||
g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
|
||||
else:
|
||||
nll_grads = torch.autograd.grad(
|
||||
nll_loss, self.last_layer[0], retain_graph=True
|
||||
)[0]
|
||||
g_grads = torch.autograd.grad(
|
||||
g_loss, self.last_layer[0], retain_graph=True
|
||||
)[0]
|
||||
|
||||
d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
|
||||
d_weight = torch.clamp(d_weight, 0.0, 1e4).detach()
|
||||
d_weight = d_weight * self.discriminator_weight
|
||||
return d_weight
|
||||
|
||||
def forward(
|
||||
self,
|
||||
regularization_log,
|
||||
inputs,
|
||||
reconstructions,
|
||||
optimizer_idx,
|
||||
global_step,
|
||||
last_layer=None,
|
||||
split="train",
|
||||
weights=None,
|
||||
):
|
||||
if self.scale_input_to_tgt_size:
|
||||
inputs = torch.nn.functional.interpolate(
|
||||
inputs, reconstructions.shape[2:], mode="bicubic", antialias=True
|
||||
)
|
||||
|
||||
if self.dims > 2:
|
||||
inputs, reconstructions = map(
|
||||
lambda x: rearrange(x, "b c t h w -> (b t) c h w"),
|
||||
(inputs, reconstructions),
|
||||
)
|
||||
|
||||
rec_loss = torch.abs(inputs.contiguous() - reconstructions.contiguous())
|
||||
if self.perceptual_weight > 0:
|
||||
p_loss = self.perceptual_loss(
|
||||
inputs.contiguous(), reconstructions.contiguous()
|
||||
)
|
||||
rec_loss = rec_loss + self.perceptual_weight * p_loss
|
||||
|
||||
nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
|
||||
weighted_nll_loss = nll_loss
|
||||
if weights is not None:
|
||||
weighted_nll_loss = weights * nll_loss
|
||||
weighted_nll_loss = torch.sum(weighted_nll_loss) / weighted_nll_loss.shape[0]
|
||||
nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
|
||||
|
||||
# now the GAN part
|
||||
if optimizer_idx == 0:
|
||||
# generator update
|
||||
logits_fake = self.discriminator(reconstructions.contiguous())
|
||||
g_loss = -torch.mean(logits_fake)
|
||||
|
||||
if self.disc_factor > 0.0:
|
||||
try:
|
||||
d_weight = self.calculate_adaptive_weight(
|
||||
nll_loss, g_loss, last_layer=last_layer
|
||||
)
|
||||
except RuntimeError:
|
||||
assert not self.training
|
||||
d_weight = torch.tensor(0.0)
|
||||
else:
|
||||
d_weight = torch.tensor(0.0)
|
||||
|
||||
disc_factor = adopt_weight(
|
||||
self.disc_factor, global_step, threshold=self.discriminator_iter_start
|
||||
)
|
||||
loss = weighted_nll_loss + d_weight * disc_factor * g_loss
|
||||
log = dict()
|
||||
for k in regularization_log:
|
||||
if k in self.regularization_weights:
|
||||
loss = loss + self.regularization_weights[k] * regularization_log[k]
|
||||
log[f"{split}/{k}"] = regularization_log[k].detach().mean()
|
||||
|
||||
log.update(
|
||||
{
|
||||
"{}/total_loss".format(split): loss.clone().detach().mean(),
|
||||
"{}/logvar".format(split): self.logvar.detach(),
|
||||
"{}/nll_loss".format(split): nll_loss.detach().mean(),
|
||||
"{}/rec_loss".format(split): rec_loss.detach().mean(),
|
||||
"{}/d_weight".format(split): d_weight.detach(),
|
||||
"{}/disc_factor".format(split): torch.tensor(disc_factor),
|
||||
"{}/g_loss".format(split): g_loss.detach().mean(),
|
||||
}
|
||||
)
|
||||
|
||||
return loss, log
|
||||
|
||||
if optimizer_idx == 1:
|
||||
# second pass for discriminator update
|
||||
logits_real = self.discriminator(inputs.contiguous().detach())
|
||||
logits_fake = self.discriminator(reconstructions.contiguous().detach())
|
||||
|
||||
disc_factor = adopt_weight(
|
||||
self.disc_factor, global_step, threshold=self.discriminator_iter_start
|
||||
)
|
||||
d_loss = disc_factor * self.disc_loss(logits_real, logits_fake)
|
||||
|
||||
log = {
|
||||
"{}/disc_loss".format(split): d_loss.clone().detach().mean(),
|
||||
"{}/logits_real".format(split): logits_real.detach().mean(),
|
||||
"{}/logits_fake".format(split): logits_fake.detach().mean(),
|
||||
}
|
||||
return d_loss, log
|
||||
@@ -0,0 +1 @@
|
||||
vgg.pth
|
||||
@@ -0,0 +1,23 @@
|
||||
Copyright (c) 2018, Richard Zhang, Phillip Isola, Alexei A. Efros, Eli Shechtman, Oliver Wang
|
||||
All rights reserved.
|
||||
|
||||
Redistribution and use in source and binary forms, with or without
|
||||
modification, are permitted provided that the following conditions are met:
|
||||
|
||||
* Redistributions of source code must retain the above copyright notice, this
|
||||
list of conditions and the following disclaimer.
|
||||
|
||||
* 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.
|
||||
|
||||
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.
|
||||
@@ -0,0 +1,147 @@
|
||||
"""Stripped version of https://github.com/richzhang/PerceptualSimilarity/tree/master/models"""
|
||||
|
||||
from collections import namedtuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision import models
|
||||
|
||||
from ..util import get_ckpt_path
|
||||
|
||||
|
||||
class LPIPS(nn.Module):
|
||||
# Learned perceptual metric
|
||||
def __init__(self, use_dropout=True):
|
||||
super().__init__()
|
||||
self.scaling_layer = ScalingLayer()
|
||||
self.chns = [64, 128, 256, 512, 512] # vg16 features
|
||||
self.net = vgg16(pretrained=True, requires_grad=False)
|
||||
self.lin0 = NetLinLayer(self.chns[0], use_dropout=use_dropout)
|
||||
self.lin1 = NetLinLayer(self.chns[1], use_dropout=use_dropout)
|
||||
self.lin2 = NetLinLayer(self.chns[2], use_dropout=use_dropout)
|
||||
self.lin3 = NetLinLayer(self.chns[3], use_dropout=use_dropout)
|
||||
self.lin4 = NetLinLayer(self.chns[4], use_dropout=use_dropout)
|
||||
self.load_from_pretrained()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def load_from_pretrained(self, name="vgg_lpips"):
|
||||
ckpt = get_ckpt_path(name, "sgm/modules/autoencoding/lpips/loss")
|
||||
self.load_state_dict(
|
||||
torch.load(ckpt, map_location=torch.device("cpu")), strict=False
|
||||
)
|
||||
print("loaded pretrained LPIPS loss from {}".format(ckpt))
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, name="vgg_lpips"):
|
||||
if name != "vgg_lpips":
|
||||
raise NotImplementedError
|
||||
model = cls()
|
||||
ckpt = get_ckpt_path(name)
|
||||
model.load_state_dict(
|
||||
torch.load(ckpt, map_location=torch.device("cpu")), strict=False
|
||||
)
|
||||
return model
|
||||
|
||||
def forward(self, input, target):
|
||||
in0_input, in1_input = (self.scaling_layer(input), self.scaling_layer(target))
|
||||
outs0, outs1 = self.net(in0_input), self.net(in1_input)
|
||||
feats0, feats1, diffs = {}, {}, {}
|
||||
lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4]
|
||||
for kk in range(len(self.chns)):
|
||||
feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(
|
||||
outs1[kk]
|
||||
)
|
||||
diffs[kk] = (feats0[kk] - feats1[kk]) ** 2
|
||||
|
||||
res = [
|
||||
spatial_average(lins[kk].model(diffs[kk]), keepdim=True)
|
||||
for kk in range(len(self.chns))
|
||||
]
|
||||
val = res[0]
|
||||
for l in range(1, len(self.chns)):
|
||||
val += res[l]
|
||||
return val
|
||||
|
||||
|
||||
class ScalingLayer(nn.Module):
|
||||
def __init__(self):
|
||||
super(ScalingLayer, self).__init__()
|
||||
self.register_buffer(
|
||||
"shift", torch.Tensor([-0.030, -0.088, -0.188])[None, :, None, None]
|
||||
)
|
||||
self.register_buffer(
|
||||
"scale", torch.Tensor([0.458, 0.448, 0.450])[None, :, None, None]
|
||||
)
|
||||
|
||||
def forward(self, inp):
|
||||
return (inp - self.shift) / self.scale
|
||||
|
||||
|
||||
class NetLinLayer(nn.Module):
|
||||
"""A single linear layer which does a 1x1 conv"""
|
||||
|
||||
def __init__(self, chn_in, chn_out=1, use_dropout=False):
|
||||
super(NetLinLayer, self).__init__()
|
||||
layers = (
|
||||
[
|
||||
nn.Dropout(),
|
||||
]
|
||||
if (use_dropout)
|
||||
else []
|
||||
)
|
||||
layers += [
|
||||
nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
|
||||
]
|
||||
self.model = nn.Sequential(*layers)
|
||||
|
||||
|
||||
class vgg16(torch.nn.Module):
|
||||
def __init__(self, requires_grad=False, pretrained=True):
|
||||
super(vgg16, self).__init__()
|
||||
vgg_pretrained_features = models.vgg16(pretrained=pretrained).features
|
||||
self.slice1 = torch.nn.Sequential()
|
||||
self.slice2 = torch.nn.Sequential()
|
||||
self.slice3 = torch.nn.Sequential()
|
||||
self.slice4 = torch.nn.Sequential()
|
||||
self.slice5 = torch.nn.Sequential()
|
||||
self.N_slices = 5
|
||||
for x in range(4):
|
||||
self.slice1.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(4, 9):
|
||||
self.slice2.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(9, 16):
|
||||
self.slice3.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(16, 23):
|
||||
self.slice4.add_module(str(x), vgg_pretrained_features[x])
|
||||
for x in range(23, 30):
|
||||
self.slice5.add_module(str(x), vgg_pretrained_features[x])
|
||||
if not requires_grad:
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, X):
|
||||
h = self.slice1(X)
|
||||
h_relu1_2 = h
|
||||
h = self.slice2(h)
|
||||
h_relu2_2 = h
|
||||
h = self.slice3(h)
|
||||
h_relu3_3 = h
|
||||
h = self.slice4(h)
|
||||
h_relu4_3 = h
|
||||
h = self.slice5(h)
|
||||
h_relu5_3 = h
|
||||
vgg_outputs = namedtuple(
|
||||
"VggOutputs", ["relu1_2", "relu2_2", "relu3_3", "relu4_3", "relu5_3"]
|
||||
)
|
||||
out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3, h_relu5_3)
|
||||
return out
|
||||
|
||||
|
||||
def normalize_tensor(x, eps=1e-10):
|
||||
norm_factor = torch.sqrt(torch.sum(x**2, dim=1, keepdim=True))
|
||||
return x / (norm_factor + eps)
|
||||
|
||||
|
||||
def spatial_average(x, keepdim=True):
|
||||
return x.mean([2, 3], keepdim=keepdim)
|
||||
@@ -0,0 +1,58 @@
|
||||
Copyright (c) 2017, Jun-Yan Zhu and Taesung Park
|
||||
All rights reserved.
|
||||
|
||||
Redistribution and use in source and binary forms, with or without
|
||||
modification, are permitted provided that the following conditions are met:
|
||||
|
||||
* Redistributions of source code must retain the above copyright notice, this
|
||||
list of conditions and the following disclaimer.
|
||||
|
||||
* 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.
|
||||
|
||||
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.
|
||||
|
||||
|
||||
--------------------------- LICENSE FOR pix2pix --------------------------------
|
||||
BSD License
|
||||
|
||||
For pix2pix software
|
||||
Copyright (c) 2016, Phillip Isola and Jun-Yan Zhu
|
||||
All rights reserved.
|
||||
|
||||
Redistribution and use in source and binary forms, with or without
|
||||
modification, are permitted provided that the following conditions are met:
|
||||
|
||||
* Redistributions of source code must retain the above copyright notice, this
|
||||
list of conditions and the following disclaimer.
|
||||
|
||||
* 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.
|
||||
|
||||
----------------------------- LICENSE FOR DCGAN --------------------------------
|
||||
BSD License
|
||||
|
||||
For dcgan.torch software
|
||||
|
||||
Copyright (c) 2015, Facebook, Inc. All rights reserved.
|
||||
|
||||
Redistribution and use in source and binary forms, with or without modification, are permitted provided that the following conditions are met:
|
||||
|
||||
Redistributions of source code must retain the above copyright notice, this list of conditions and the following disclaimer.
|
||||
|
||||
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.
|
||||
|
||||
Neither the name Facebook 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.
|
||||
@@ -0,0 +1,88 @@
|
||||
import functools
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
from ..util import ActNorm
|
||||
|
||||
|
||||
def weights_init(m):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
nn.init.normal_(m.weight.data, 0.0, 0.02)
|
||||
elif classname.find("BatchNorm") != -1:
|
||||
nn.init.normal_(m.weight.data, 1.0, 0.02)
|
||||
nn.init.constant_(m.bias.data, 0)
|
||||
|
||||
|
||||
class NLayerDiscriminator(nn.Module):
|
||||
"""Defines a PatchGAN discriminator as in Pix2Pix
|
||||
--> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py
|
||||
"""
|
||||
|
||||
def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False):
|
||||
"""Construct a PatchGAN discriminator
|
||||
Parameters:
|
||||
input_nc (int) -- the number of channels in input images
|
||||
ndf (int) -- the number of filters in the last conv layer
|
||||
n_layers (int) -- the number of conv layers in the discriminator
|
||||
norm_layer -- normalization layer
|
||||
"""
|
||||
super(NLayerDiscriminator, self).__init__()
|
||||
if not use_actnorm:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
else:
|
||||
norm_layer = ActNorm
|
||||
if (
|
||||
type(norm_layer) == functools.partial
|
||||
): # no need to use bias as BatchNorm2d has affine parameters
|
||||
use_bias = norm_layer.func != nn.BatchNorm2d
|
||||
else:
|
||||
use_bias = norm_layer != nn.BatchNorm2d
|
||||
|
||||
kw = 4
|
||||
padw = 1
|
||||
sequence = [
|
||||
nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
nf_mult = 1
|
||||
nf_mult_prev = 1
|
||||
for n in range(1, n_layers): # gradually increase the number of filters
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n, 8)
|
||||
sequence += [
|
||||
nn.Conv2d(
|
||||
ndf * nf_mult_prev,
|
||||
ndf * nf_mult,
|
||||
kernel_size=kw,
|
||||
stride=2,
|
||||
padding=padw,
|
||||
bias=use_bias,
|
||||
),
|
||||
norm_layer(ndf * nf_mult),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n_layers, 8)
|
||||
sequence += [
|
||||
nn.Conv2d(
|
||||
ndf * nf_mult_prev,
|
||||
ndf * nf_mult,
|
||||
kernel_size=kw,
|
||||
stride=1,
|
||||
padding=padw,
|
||||
bias=use_bias,
|
||||
),
|
||||
norm_layer(ndf * nf_mult),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
|
||||
sequence += [
|
||||
nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
|
||||
] # output 1 channel prediction map
|
||||
self.main = nn.Sequential(*sequence)
|
||||
|
||||
def forward(self, input):
|
||||
"""Standard forward."""
|
||||
return self.main(input)
|
||||
@@ -0,0 +1,128 @@
|
||||
import hashlib
|
||||
import os
|
||||
|
||||
import requests
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from tqdm import tqdm
|
||||
|
||||
URL_MAP = {"vgg_lpips": "https://heibox.uni-heidelberg.de/f/607503859c864bc1b30b/?dl=1"}
|
||||
|
||||
CKPT_MAP = {"vgg_lpips": "vgg.pth"}
|
||||
|
||||
MD5_MAP = {"vgg_lpips": "d507d7349b931f0638a25a48a722f98a"}
|
||||
|
||||
|
||||
def download(url, local_path, chunk_size=1024):
|
||||
os.makedirs(os.path.split(local_path)[0], exist_ok=True)
|
||||
with requests.get(url, stream=True) as r:
|
||||
total_size = int(r.headers.get("content-length", 0))
|
||||
with tqdm(total=total_size, unit="B", unit_scale=True) as pbar:
|
||||
with open(local_path, "wb") as f:
|
||||
for data in r.iter_content(chunk_size=chunk_size):
|
||||
if data:
|
||||
f.write(data)
|
||||
pbar.update(chunk_size)
|
||||
|
||||
|
||||
def md5_hash(path):
|
||||
with open(path, "rb") as f:
|
||||
content = f.read()
|
||||
return hashlib.md5(content).hexdigest()
|
||||
|
||||
|
||||
def get_ckpt_path(name, root, check=False):
|
||||
assert name in URL_MAP
|
||||
path = os.path.join(root, CKPT_MAP[name])
|
||||
if not os.path.exists(path) or (check and not md5_hash(path) == MD5_MAP[name]):
|
||||
print("Downloading {} model from {} to {}".format(name, URL_MAP[name], path))
|
||||
download(URL_MAP[name], path)
|
||||
md5 = md5_hash(path)
|
||||
assert md5 == MD5_MAP[name], md5
|
||||
return path
|
||||
|
||||
|
||||
class ActNorm(nn.Module):
|
||||
def __init__(
|
||||
self, num_features, logdet=False, affine=True, allow_reverse_init=False
|
||||
):
|
||||
assert affine
|
||||
super().__init__()
|
||||
self.logdet = logdet
|
||||
self.loc = nn.Parameter(torch.zeros(1, num_features, 1, 1))
|
||||
self.scale = nn.Parameter(torch.ones(1, num_features, 1, 1))
|
||||
self.allow_reverse_init = allow_reverse_init
|
||||
|
||||
self.register_buffer("initialized", torch.tensor(0, dtype=torch.uint8))
|
||||
|
||||
def initialize(self, input):
|
||||
with torch.no_grad():
|
||||
flatten = input.permute(1, 0, 2, 3).contiguous().view(input.shape[1], -1)
|
||||
mean = (
|
||||
flatten.mean(1)
|
||||
.unsqueeze(1)
|
||||
.unsqueeze(2)
|
||||
.unsqueeze(3)
|
||||
.permute(1, 0, 2, 3)
|
||||
)
|
||||
std = (
|
||||
flatten.std(1)
|
||||
.unsqueeze(1)
|
||||
.unsqueeze(2)
|
||||
.unsqueeze(3)
|
||||
.permute(1, 0, 2, 3)
|
||||
)
|
||||
|
||||
self.loc.data.copy_(-mean)
|
||||
self.scale.data.copy_(1 / (std + 1e-6))
|
||||
|
||||
def forward(self, input, reverse=False):
|
||||
if reverse:
|
||||
return self.reverse(input)
|
||||
if len(input.shape) == 2:
|
||||
input = input[:, :, None, None]
|
||||
squeeze = True
|
||||
else:
|
||||
squeeze = False
|
||||
|
||||
_, _, height, width = input.shape
|
||||
|
||||
if self.training and self.initialized.item() == 0:
|
||||
self.initialize(input)
|
||||
self.initialized.fill_(1)
|
||||
|
||||
h = self.scale * (input + self.loc)
|
||||
|
||||
if squeeze:
|
||||
h = h.squeeze(-1).squeeze(-1)
|
||||
|
||||
if self.logdet:
|
||||
log_abs = torch.log(torch.abs(self.scale))
|
||||
logdet = height * width * torch.sum(log_abs)
|
||||
logdet = logdet * torch.ones(input.shape[0]).to(input)
|
||||
return h, logdet
|
||||
|
||||
return h
|
||||
|
||||
def reverse(self, output):
|
||||
if self.training and self.initialized.item() == 0:
|
||||
if not self.allow_reverse_init:
|
||||
raise RuntimeError(
|
||||
"Initializing ActNorm in reverse direction is "
|
||||
"disabled by default. Use allow_reverse_init=True to enable."
|
||||
)
|
||||
else:
|
||||
self.initialize(output)
|
||||
self.initialized.fill_(1)
|
||||
|
||||
if len(output.shape) == 2:
|
||||
output = output[:, :, None, None]
|
||||
squeeze = True
|
||||
else:
|
||||
squeeze = False
|
||||
|
||||
h = output / self.scale - self.loc
|
||||
|
||||
if squeeze:
|
||||
h = h.squeeze(-1).squeeze(-1)
|
||||
return h
|
||||
@@ -0,0 +1,17 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def hinge_d_loss(logits_real, logits_fake):
|
||||
loss_real = torch.mean(F.relu(1.0 - logits_real))
|
||||
loss_fake = torch.mean(F.relu(1.0 + logits_fake))
|
||||
d_loss = 0.5 * (loss_real + loss_fake)
|
||||
return d_loss
|
||||
|
||||
|
||||
def vanilla_d_loss(logits_real, logits_fake):
|
||||
d_loss = 0.5 * (
|
||||
torch.mean(torch.nn.functional.softplus(-logits_real))
|
||||
+ torch.mean(torch.nn.functional.softplus(logits_fake))
|
||||
)
|
||||
return d_loss
|
||||
@@ -0,0 +1,53 @@
|
||||
from abc import abstractmethod
|
||||
from typing import Any, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ....modules.distributions.distributions import DiagonalGaussianDistribution
|
||||
|
||||
|
||||
class AbstractRegularizer(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, z: torch.Tensor) -> Tuple[torch.Tensor, dict]:
|
||||
raise NotImplementedError()
|
||||
|
||||
@abstractmethod
|
||||
def get_trainable_parameters(self) -> Any:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class DiagonalGaussianRegularizer(AbstractRegularizer):
|
||||
def __init__(self, sample: bool = True):
|
||||
super().__init__()
|
||||
self.sample = sample
|
||||
|
||||
def get_trainable_parameters(self) -> Any:
|
||||
yield from ()
|
||||
|
||||
def forward(self, z: torch.Tensor) -> Tuple[torch.Tensor, dict]:
|
||||
log = dict()
|
||||
posterior = DiagonalGaussianDistribution(z)
|
||||
if self.sample:
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z = posterior.mode()
|
||||
kl_loss = posterior.kl()
|
||||
kl_loss = torch.sum(kl_loss) / kl_loss.shape[0]
|
||||
log["kl_loss"] = kl_loss
|
||||
return z, log
|
||||
|
||||
|
||||
def measure_perplexity(predicted_indices, num_centroids):
|
||||
# src: https://github.com/karpathy/deep-vector-quantization/blob/main/model.py
|
||||
# eval cluster perplexity. when perplexity == num_embeddings then all clusters are used exactly equally
|
||||
encodings = (
|
||||
F.one_hot(predicted_indices, num_centroids).float().reshape(-1, num_centroids)
|
||||
)
|
||||
avg_probs = encodings.mean(0)
|
||||
perplexity = (-(avg_probs * torch.log(avg_probs + 1e-10)).sum()).exp()
|
||||
cluster_use = torch.sum(avg_probs > 0)
|
||||
return perplexity, cluster_use
|
||||
@@ -0,0 +1,7 @@
|
||||
from .denoiser import Denoiser
|
||||
from .discretizer import Discretization
|
||||
from .loss import StandardDiffusionLoss
|
||||
from .model import Decoder, Encoder, Model
|
||||
from .openaimodel import UNetModel
|
||||
from .sampling import BaseDiffusionSampler
|
||||
from .wrappers import OpenAIWrapper
|
||||
@@ -0,0 +1,73 @@
|
||||
import torch.nn as nn
|
||||
|
||||
from ...util import append_dims, instantiate_from_config
|
||||
|
||||
|
||||
class Denoiser(nn.Module):
|
||||
def __init__(self, weighting_config, scaling_config):
|
||||
super().__init__()
|
||||
|
||||
self.weighting = instantiate_from_config(weighting_config)
|
||||
self.scaling = instantiate_from_config(scaling_config)
|
||||
|
||||
def possibly_quantize_sigma(self, sigma):
|
||||
return sigma
|
||||
|
||||
def possibly_quantize_c_noise(self, c_noise):
|
||||
return c_noise
|
||||
|
||||
def w(self, sigma):
|
||||
return self.weighting(sigma)
|
||||
|
||||
def __call__(self, network, input, sigma, cond):
|
||||
sigma = self.possibly_quantize_sigma(sigma)
|
||||
sigma_shape = sigma.shape
|
||||
sigma = append_dims(sigma, input.ndim)
|
||||
c_skip, c_out, c_in, c_noise = self.scaling(sigma)
|
||||
c_noise = self.possibly_quantize_c_noise(c_noise.reshape(sigma_shape))
|
||||
return network(input * c_in, c_noise, cond) * c_out + input * c_skip
|
||||
|
||||
|
||||
class DiscreteDenoiser(Denoiser):
|
||||
def __init__(
|
||||
self,
|
||||
weighting_config,
|
||||
scaling_config,
|
||||
num_idx,
|
||||
discretization_config,
|
||||
do_append_zero=False,
|
||||
quantize_c_noise=True,
|
||||
flip=True,
|
||||
):
|
||||
super().__init__(weighting_config, scaling_config)
|
||||
sigmas = instantiate_from_config(discretization_config)(
|
||||
num_idx, do_append_zero=do_append_zero, flip=flip
|
||||
)
|
||||
self.register_buffer("sigmas", sigmas)
|
||||
self.quantize_c_noise = quantize_c_noise
|
||||
|
||||
def sigma_to_idx(self, sigma):
|
||||
dists = sigma - self.sigmas[:, None]
|
||||
return dists.abs().argmin(dim=0).view(sigma.shape)
|
||||
|
||||
def idx_to_sigma(self, idx):
|
||||
return self.sigmas[idx]
|
||||
|
||||
def possibly_quantize_sigma(self, sigma):
|
||||
return self.idx_to_sigma(self.sigma_to_idx(sigma))
|
||||
|
||||
def possibly_quantize_c_noise(self, c_noise):
|
||||
if self.quantize_c_noise:
|
||||
return self.sigma_to_idx(c_noise)
|
||||
else:
|
||||
return c_noise
|
||||
|
||||
|
||||
class DiscreteDenoiserWithControl(DiscreteDenoiser):
|
||||
def __call__(self, network, input, sigma, cond, control_scale):
|
||||
sigma = self.possibly_quantize_sigma(sigma)
|
||||
sigma_shape = sigma.shape
|
||||
sigma = append_dims(sigma, input.ndim)
|
||||
c_skip, c_out, c_in, c_noise = self.scaling(sigma)
|
||||
c_noise = self.possibly_quantize_c_noise(c_noise.reshape(sigma_shape))
|
||||
return network(input * c_in, c_noise, cond, control_scale) * c_out + input * c_skip
|
||||
@@ -0,0 +1,31 @@
|
||||
import torch
|
||||
|
||||
|
||||
class EDMScaling:
|
||||
def __init__(self, sigma_data=0.5):
|
||||
self.sigma_data = sigma_data
|
||||
|
||||
def __call__(self, sigma):
|
||||
c_skip = self.sigma_data**2 / (sigma**2 + self.sigma_data**2)
|
||||
c_out = sigma * self.sigma_data / (sigma**2 + self.sigma_data**2) ** 0.5
|
||||
c_in = 1 / (sigma**2 + self.sigma_data**2) ** 0.5
|
||||
c_noise = 0.25 * sigma.log()
|
||||
return c_skip, c_out, c_in, c_noise
|
||||
|
||||
|
||||
class EpsScaling:
|
||||
def __call__(self, sigma):
|
||||
c_skip = torch.ones_like(sigma, device=sigma.device)
|
||||
c_out = -sigma
|
||||
c_in = 1 / (sigma**2 + 1.0) ** 0.5
|
||||
c_noise = sigma.clone()
|
||||
return c_skip, c_out, c_in, c_noise
|
||||
|
||||
|
||||
class VScaling:
|
||||
def __call__(self, sigma):
|
||||
c_skip = 1.0 / (sigma**2 + 1.0)
|
||||
c_out = -sigma / (sigma**2 + 1.0) ** 0.5
|
||||
c_in = 1.0 / (sigma**2 + 1.0) ** 0.5
|
||||
c_noise = sigma.clone()
|
||||
return c_skip, c_out, c_in, c_noise
|
||||
@@ -0,0 +1,24 @@
|
||||
import torch
|
||||
|
||||
class UnitWeighting:
|
||||
def __call__(self, sigma):
|
||||
return torch.ones_like(sigma, device=sigma.device)
|
||||
|
||||
|
||||
class EDMWeighting:
|
||||
def __init__(self, sigma_data=0.5):
|
||||
self.sigma_data = sigma_data
|
||||
|
||||
def __call__(self, sigma):
|
||||
return (sigma**2 + self.sigma_data**2) / (sigma * self.sigma_data) ** 2
|
||||
|
||||
|
||||
class VWeighting(EDMWeighting):
|
||||
def __init__(self):
|
||||
super().__init__(sigma_data=1.0)
|
||||
|
||||
|
||||
class EpsWeighting:
|
||||
def __call__(self, sigma):
|
||||
return sigma**-2
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
from abc import abstractmethod
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from ...modules.diffusionmodules.util import make_beta_schedule
|
||||
from ...util import append_zero
|
||||
|
||||
|
||||
def generate_roughly_equally_spaced_steps(
|
||||
num_substeps: int, max_step: int
|
||||
) -> np.ndarray:
|
||||
return np.linspace(max_step - 1, 0, num_substeps, endpoint=False).astype(int)[::-1]
|
||||
|
||||
|
||||
class Discretization:
|
||||
def __call__(self, n, do_append_zero=True, device="cpu", flip=False):
|
||||
sigmas = self.get_sigmas(n, device=device)
|
||||
sigmas = append_zero(sigmas) if do_append_zero else sigmas
|
||||
return sigmas if not flip else torch.flip(sigmas, (0,))
|
||||
|
||||
@abstractmethod
|
||||
def get_sigmas(self, n, device):
|
||||
pass
|
||||
|
||||
|
||||
class EDMDiscretization(Discretization):
|
||||
def __init__(self, sigma_min=0.02, sigma_max=80.0, rho=7.0):
|
||||
self.sigma_min = sigma_min
|
||||
self.sigma_max = sigma_max
|
||||
self.rho = rho
|
||||
|
||||
def get_sigmas(self, n, device="cpu"):
|
||||
ramp = torch.linspace(0, 1, n, device=device)
|
||||
min_inv_rho = self.sigma_min ** (1 / self.rho)
|
||||
max_inv_rho = self.sigma_max ** (1 / self.rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** self.rho
|
||||
return sigmas
|
||||
|
||||
|
||||
class LegacyDDPMDiscretization(Discretization):
|
||||
def __init__(
|
||||
self,
|
||||
linear_start=0.00085,
|
||||
linear_end=0.0120,
|
||||
num_timesteps=1000,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_timesteps = num_timesteps
|
||||
betas = make_beta_schedule(
|
||||
"linear", num_timesteps, linear_start=linear_start, linear_end=linear_end
|
||||
)
|
||||
alphas = 1.0 - betas
|
||||
self.alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
self.to_torch = partial(torch.tensor, dtype=torch.float32)
|
||||
|
||||
def get_sigmas(self, n, device="cpu"):
|
||||
if n < self.num_timesteps:
|
||||
timesteps = generate_roughly_equally_spaced_steps(n, self.num_timesteps)
|
||||
alphas_cumprod = self.alphas_cumprod[timesteps]
|
||||
elif n == self.num_timesteps:
|
||||
alphas_cumprod = self.alphas_cumprod
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
to_torch = partial(torch.tensor, dtype=torch.float32, device=device)
|
||||
sigmas = to_torch((1 - alphas_cumprod) / alphas_cumprod) ** 0.5
|
||||
return torch.flip(sigmas, (0,))
|
||||
@@ -0,0 +1,88 @@
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
|
||||
from ...util import default, instantiate_from_config
|
||||
|
||||
|
||||
class VanillaCFG:
|
||||
"""
|
||||
implements parallelized CFG
|
||||
"""
|
||||
|
||||
def __init__(self, scale, dyn_thresh_config=None):
|
||||
scale_schedule = lambda scale, sigma: scale # independent of step
|
||||
self.scale_schedule = partial(scale_schedule, scale)
|
||||
self.dyn_thresh = instantiate_from_config(
|
||||
default(
|
||||
dyn_thresh_config,
|
||||
{
|
||||
"target": ".sgm.modules.diffusionmodules.sampling_utils.NoDynamicThresholding"
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
def __call__(self, x, sigma):
|
||||
x_u, x_c = x.chunk(2)
|
||||
scale_value = self.scale_schedule(sigma)
|
||||
x_pred = self.dyn_thresh(x_u, x_c, scale_value)
|
||||
return x_pred
|
||||
|
||||
def prepare_inputs(self, x, s, c, uc):
|
||||
c_out = dict()
|
||||
|
||||
for k in c:
|
||||
if k in ["vector", "crossattn", "concat", "control", 'control_vector', 'mask_x']:
|
||||
c_out[k] = torch.cat((uc[k], c[k]), 0)
|
||||
else:
|
||||
assert c[k] == uc[k]
|
||||
c_out[k] = c[k]
|
||||
return torch.cat([x] * 2), torch.cat([s] * 2), c_out
|
||||
|
||||
|
||||
|
||||
class LinearCFG:
|
||||
def __init__(self, scale, scale_min=None, dyn_thresh_config=None):
|
||||
if scale_min is None:
|
||||
scale_min = scale
|
||||
scale_schedule = lambda scale, scale_min, sigma: (scale - scale_min) * sigma / 14.6146 + scale_min
|
||||
self.scale_schedule = partial(scale_schedule, scale, scale_min)
|
||||
self.dyn_thresh = instantiate_from_config(
|
||||
default(
|
||||
dyn_thresh_config,
|
||||
{
|
||||
"target": ".sgm.modules.diffusionmodules.sampling_utils.NoDynamicThresholding"
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
def __call__(self, x, sigma):
|
||||
x_u, x_c = x.chunk(2)
|
||||
scale_value = self.scale_schedule(sigma)
|
||||
x_pred = self.dyn_thresh(x_u, x_c, scale_value)
|
||||
return x_pred
|
||||
|
||||
def prepare_inputs(self, x, s, c, uc):
|
||||
c_out = dict()
|
||||
|
||||
for k in c:
|
||||
if k in ["vector", "crossattn", "concat", "control", 'control_vector', 'mask_x']:
|
||||
c_out[k] = torch.cat((uc[k], c[k]), 0)
|
||||
else:
|
||||
assert c[k] == uc[k]
|
||||
c_out[k] = c[k]
|
||||
return torch.cat([x] * 2), torch.cat([s] * 2), c_out
|
||||
|
||||
|
||||
|
||||
class IdentityGuider:
|
||||
def __call__(self, x, sigma):
|
||||
return x
|
||||
|
||||
def prepare_inputs(self, x, s, c, uc):
|
||||
c_out = dict()
|
||||
|
||||
for k in c:
|
||||
c_out[k] = c[k]
|
||||
|
||||
return x, s, c_out
|
||||
@@ -0,0 +1,69 @@
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from omegaconf import ListConfig
|
||||
|
||||
from ...util import append_dims, instantiate_from_config
|
||||
from ...modules.autoencoding.lpips.loss.lpips import LPIPS
|
||||
|
||||
|
||||
class StandardDiffusionLoss(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
sigma_sampler_config,
|
||||
type="l2",
|
||||
offset_noise_level=0.0,
|
||||
batch2model_keys: Optional[Union[str, List[str], ListConfig]] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert type in ["l2", "l1", "lpips"]
|
||||
|
||||
self.sigma_sampler = instantiate_from_config(sigma_sampler_config)
|
||||
|
||||
self.type = type
|
||||
self.offset_noise_level = offset_noise_level
|
||||
|
||||
if type == "lpips":
|
||||
self.lpips = LPIPS().eval()
|
||||
|
||||
if not batch2model_keys:
|
||||
batch2model_keys = []
|
||||
|
||||
if isinstance(batch2model_keys, str):
|
||||
batch2model_keys = [batch2model_keys]
|
||||
|
||||
self.batch2model_keys = set(batch2model_keys)
|
||||
|
||||
def __call__(self, network, denoiser, conditioner, input, batch):
|
||||
cond = conditioner(batch)
|
||||
additional_model_inputs = {
|
||||
key: batch[key] for key in self.batch2model_keys.intersection(batch)
|
||||
}
|
||||
|
||||
sigmas = self.sigma_sampler(input.shape[0]).to(input.device)
|
||||
noise = torch.randn_like(input)
|
||||
if self.offset_noise_level > 0.0:
|
||||
noise = noise + self.offset_noise_level * append_dims(
|
||||
torch.randn(input.shape[0], device=input.device), input.ndim
|
||||
)
|
||||
noised_input = input + noise * append_dims(sigmas, input.ndim)
|
||||
model_output = denoiser(
|
||||
network, noised_input, sigmas, cond, **additional_model_inputs
|
||||
)
|
||||
w = append_dims(denoiser.w(sigmas), input.ndim)
|
||||
return self.get_loss(model_output, input, w)
|
||||
|
||||
def get_loss(self, model_output, target, w):
|
||||
if self.type == "l2":
|
||||
return torch.mean(
|
||||
(w * (model_output - target) ** 2).reshape(target.shape[0], -1), 1
|
||||
)
|
||||
elif self.type == "l1":
|
||||
return torch.mean(
|
||||
(w * (model_output - target).abs()).reshape(target.shape[0], -1), 1
|
||||
)
|
||||
elif self.type == "lpips":
|
||||
loss = self.lpips(model_output, target).reshape(-1)
|
||||
return loss
|
||||
@@ -0,0 +1,743 @@
|
||||
# pytorch_diffusion + derived encoder decoder
|
||||
import math
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from packaging import version
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
|
||||
XFORMERS_IS_AVAILABLE = True
|
||||
except:
|
||||
XFORMERS_IS_AVAILABLE = False
|
||||
print("no module 'xformers'. Processing without...")
|
||||
|
||||
from ...modules.attention import LinearAttention, MemoryEfficientCrossAttention
|
||||
|
||||
|
||||
def get_timestep_embedding(timesteps, embedding_dim):
|
||||
"""
|
||||
This matches the implementation in Denoising Diffusion Probabilistic Models:
|
||||
From Fairseq.
|
||||
Build sinusoidal embeddings.
|
||||
This matches the implementation in tensor2tensor, but differs slightly
|
||||
from the description in Section 3.5 of "Attention Is All You Need".
|
||||
"""
|
||||
assert len(timesteps.shape) == 1
|
||||
|
||||
half_dim = embedding_dim // 2
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb)
|
||||
emb = emb.to(device=timesteps.device)
|
||||
emb = timesteps.float()[:, None] * emb[None, :]
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
|
||||
if embedding_dim % 2 == 1: # zero pad
|
||||
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
|
||||
return emb
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
# swish
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
def Normalize(in_channels, num_groups=32):
|
||||
return torch.nn.GroupNorm(
|
||||
num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
|
||||
|
||||
class Upsample(nn.Module):
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
self.conv = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
|
||||
if self.with_conv:
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class Downsample(nn.Module):
|
||||
def __init__(self, in_channels, with_conv):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
# no asymmetric padding in torch conv, must do it ourselves
|
||||
self.conv = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=2, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
if self.with_conv:
|
||||
pad = (0, 1, 0, 1)
|
||||
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
|
||||
x = self.conv(x)
|
||||
else:
|
||||
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
|
||||
return x
|
||||
|
||||
|
||||
class ResnetBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels,
|
||||
out_channels=None,
|
||||
conv_shortcut=False,
|
||||
dropout,
|
||||
temb_channels=512,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
|
||||
self.norm1 = Normalize(in_channels)
|
||||
self.conv1 = torch.nn.Conv2d(
|
||||
in_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
if temb_channels > 0:
|
||||
self.temb_proj = torch.nn.Linear(temb_channels, out_channels)
|
||||
self.norm2 = Normalize(out_channels)
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = torch.nn.Conv2d(
|
||||
out_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
if self.in_channels != self.out_channels:
|
||||
if self.use_conv_shortcut:
|
||||
self.conv_shortcut = torch.nn.Conv2d(
|
||||
in_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
else:
|
||||
self.nin_shortcut = torch.nn.Conv2d(
|
||||
in_channels, out_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x, temb):
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
|
||||
if temb is not None:
|
||||
h = h + self.temb_proj(nonlinearity(temb))[:, :, None, None]
|
||||
|
||||
h = self.norm2(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.dropout(h)
|
||||
h = self.conv2(h)
|
||||
|
||||
if self.in_channels != self.out_channels:
|
||||
if self.use_conv_shortcut:
|
||||
x = self.conv_shortcut(x)
|
||||
else:
|
||||
x = self.nin_shortcut(x)
|
||||
|
||||
return x + h
|
||||
|
||||
|
||||
class LinAttnBlock(LinearAttention):
|
||||
"""to match AttnBlock usage"""
|
||||
|
||||
def __init__(self, in_channels):
|
||||
super().__init__(dim=in_channels, heads=1, dim_head=in_channels)
|
||||
|
||||
|
||||
class AttnBlock(nn.Module):
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.k = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.v = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.proj_out = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
def attention(self, h_: torch.Tensor) -> torch.Tensor:
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
b, c, h, w = q.shape
|
||||
q, k, v = map(
|
||||
lambda x: rearrange(x, "b c h w -> b 1 (h w) c").contiguous(), (q, k, v)
|
||||
)
|
||||
h_ = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v
|
||||
) # scale is dim ** -0.5 per default
|
||||
# compute attention
|
||||
|
||||
return rearrange(h_, "b 1 (h w) c -> b c h w", h=h, w=w, c=c, b=b)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
h_ = x
|
||||
h_ = self.attention(h_)
|
||||
h_ = self.proj_out(h_)
|
||||
return x + h_
|
||||
|
||||
|
||||
class MemoryEfficientAttnBlock(nn.Module):
|
||||
"""
|
||||
Uses xformers efficient implementation,
|
||||
see https://github.com/MatthieuTPHR/diffusers/blob/d80b531ff8060ec1ea982b65a1b8df70f73aa67c/src/diffusers/models/attention.py#L223
|
||||
Note: this is a single-head self-attention operation
|
||||
"""
|
||||
|
||||
#
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.k = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.v = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.proj_out = torch.nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.attention_op: Optional[Any] = None
|
||||
|
||||
def attention(self, h_: torch.Tensor) -> torch.Tensor:
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
B, C, H, W = q.shape
|
||||
q, k, v = map(lambda x: rearrange(x, "b c h w -> b (h w) c"), (q, k, v))
|
||||
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3)
|
||||
.reshape(B, t.shape[1], 1, C)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(B * 1, t.shape[1], C)
|
||||
.contiguous(),
|
||||
(q, k, v),
|
||||
)
|
||||
out = xformers.ops.memory_efficient_attention(
|
||||
q, k, v, attn_bias=None, op=self.attention_op
|
||||
)
|
||||
|
||||
out = (
|
||||
out.unsqueeze(0)
|
||||
.reshape(B, 1, out.shape[1], C)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(B, out.shape[1], C)
|
||||
)
|
||||
return rearrange(out, "b (h w) c -> b c h w", b=B, h=H, w=W, c=C)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
h_ = x
|
||||
h_ = self.attention(h_)
|
||||
h_ = self.proj_out(h_)
|
||||
return x + h_
|
||||
|
||||
|
||||
class MemoryEfficientCrossAttentionWrapper(MemoryEfficientCrossAttention):
|
||||
def forward(self, x, context=None, mask=None, **unused_kwargs):
|
||||
b, c, h, w = x.shape
|
||||
x = rearrange(x, "b c h w -> b (h w) c")
|
||||
out = super().forward(x, context=context, mask=mask)
|
||||
out = rearrange(out, "b (h w) c -> b c h w", h=h, w=w, c=c)
|
||||
return x + out
|
||||
|
||||
|
||||
def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None):
|
||||
assert attn_type in [
|
||||
"vanilla",
|
||||
"vanilla-xformers",
|
||||
"memory-efficient-cross-attn",
|
||||
"linear",
|
||||
"none",
|
||||
], f"attn_type {attn_type} unknown"
|
||||
if (
|
||||
version.parse(torch.__version__) < version.parse("2.0.0")
|
||||
and attn_type != "none"
|
||||
):
|
||||
assert XFORMERS_IS_AVAILABLE, (
|
||||
f"We do not support vanilla attention in {torch.__version__} anymore, "
|
||||
f"as it is too expensive. Please install xformers via e.g. 'pip install xformers==0.0.16'"
|
||||
)
|
||||
attn_type = "vanilla-xformers"
|
||||
print(f"making attention of type '{attn_type}' with {in_channels} in_channels")
|
||||
if attn_type == "vanilla":
|
||||
assert attn_kwargs is None
|
||||
return AttnBlock(in_channels)
|
||||
elif attn_type == "vanilla-xformers":
|
||||
print(f"building MemoryEfficientAttnBlock with {in_channels} in_channels...")
|
||||
return MemoryEfficientAttnBlock(in_channels)
|
||||
elif type == "memory-efficient-cross-attn":
|
||||
attn_kwargs["query_dim"] = in_channels
|
||||
return MemoryEfficientCrossAttentionWrapper(**attn_kwargs)
|
||||
elif attn_type == "none":
|
||||
return nn.Identity(in_channels)
|
||||
else:
|
||||
return LinAttnBlock(in_channels)
|
||||
|
||||
|
||||
class Model(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ch,
|
||||
out_ch,
|
||||
ch_mult=(1, 2, 4, 8),
|
||||
num_res_blocks,
|
||||
attn_resolutions,
|
||||
dropout=0.0,
|
||||
resamp_with_conv=True,
|
||||
in_channels,
|
||||
resolution,
|
||||
use_timestep=True,
|
||||
use_linear_attn=False,
|
||||
attn_type="vanilla",
|
||||
):
|
||||
super().__init__()
|
||||
if use_linear_attn:
|
||||
attn_type = "linear"
|
||||
self.ch = ch
|
||||
self.temb_ch = self.ch * 4
|
||||
self.num_resolutions = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.resolution = resolution
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.use_timestep = use_timestep
|
||||
if self.use_timestep:
|
||||
# timestep embedding
|
||||
self.temb = nn.Module()
|
||||
self.temb.dense = nn.ModuleList(
|
||||
[
|
||||
torch.nn.Linear(self.ch, self.temb_ch),
|
||||
torch.nn.Linear(self.temb_ch, self.temb_ch),
|
||||
]
|
||||
)
|
||||
|
||||
# downsampling
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
in_channels, self.ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
curr_res = resolution
|
||||
in_ch_mult = (1,) + tuple(ch_mult)
|
||||
self.down = nn.ModuleList()
|
||||
for i_level in range(self.num_resolutions):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = ch * in_ch_mult[i_level]
|
||||
block_out = ch * ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks):
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(make_attn(block_in, attn_type=attn_type))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if i_level != self.num_resolutions - 1:
|
||||
down.downsample = Downsample(block_in, resamp_with_conv)
|
||||
curr_res = curr_res // 2
|
||||
self.down.append(down)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.mid.attn_1 = make_attn(block_in, attn_type=attn_type)
|
||||
self.mid.block_2 = ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
|
||||
# upsampling
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = ch * ch_mult[i_level]
|
||||
skip_in = ch * ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
if i_block == self.num_res_blocks:
|
||||
skip_in = ch * in_ch_mult[i_level]
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in + skip_in,
|
||||
out_channels=block_out,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(make_attn(block_in, attn_type=attn_type))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if i_level != 0:
|
||||
up.upsample = Upsample(block_in, resamp_with_conv)
|
||||
curr_res = curr_res * 2
|
||||
self.up.insert(0, up) # prepend to get consistent order
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
block_in, out_ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, x, t=None, context=None):
|
||||
# assert x.shape[2] == x.shape[3] == self.resolution
|
||||
if context is not None:
|
||||
# assume aligned context, cat along channel axis
|
||||
x = torch.cat((x, context), dim=1)
|
||||
if self.use_timestep:
|
||||
# timestep embedding
|
||||
assert t is not None
|
||||
temb = get_timestep_embedding(t, self.ch)
|
||||
temb = self.temb.dense[0](temb)
|
||||
temb = nonlinearity(temb)
|
||||
temb = self.temb.dense[1](temb)
|
||||
else:
|
||||
temb = None
|
||||
|
||||
# downsampling
|
||||
hs = [self.conv_in(x)]
|
||||
for i_level in range(self.num_resolutions):
|
||||
for i_block in range(self.num_res_blocks):
|
||||
h = self.down[i_level].block[i_block](hs[-1], temb)
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
hs.append(h)
|
||||
if i_level != self.num_resolutions - 1:
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
|
||||
# middle
|
||||
h = hs[-1]
|
||||
h = self.mid.block_1(h, temb)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h, temb)
|
||||
|
||||
# upsampling
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
h = self.up[i_level].block[i_block](
|
||||
torch.cat([h, hs.pop()], dim=1), temb
|
||||
)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h)
|
||||
if i_level != 0:
|
||||
h = self.up[i_level].upsample(h)
|
||||
|
||||
# end
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
return h
|
||||
|
||||
def get_last_layer(self):
|
||||
return self.conv_out.weight
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ch,
|
||||
out_ch,
|
||||
ch_mult=(1, 2, 4, 8),
|
||||
num_res_blocks,
|
||||
attn_resolutions,
|
||||
dropout=0.0,
|
||||
resamp_with_conv=True,
|
||||
in_channels,
|
||||
resolution,
|
||||
z_channels,
|
||||
double_z=True,
|
||||
use_linear_attn=False,
|
||||
attn_type="vanilla",
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
if use_linear_attn:
|
||||
attn_type = "linear"
|
||||
self.ch = ch
|
||||
self.temb_ch = 0
|
||||
self.num_resolutions = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.resolution = resolution
|
||||
self.in_channels = in_channels
|
||||
|
||||
# downsampling
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
in_channels, self.ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
curr_res = resolution
|
||||
in_ch_mult = (1,) + tuple(ch_mult)
|
||||
self.in_ch_mult = in_ch_mult
|
||||
self.down = nn.ModuleList()
|
||||
for i_level in range(self.num_resolutions):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = ch * in_ch_mult[i_level]
|
||||
block_out = ch * ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks):
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(make_attn(block_in, attn_type=attn_type))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if i_level != self.num_resolutions - 1:
|
||||
down.downsample = Downsample(block_in, resamp_with_conv)
|
||||
curr_res = curr_res // 2
|
||||
self.down.append(down)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.mid.attn_1 = make_attn(block_in, attn_type=attn_type)
|
||||
self.mid.block_2 = ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
block_in,
|
||||
2 * z_channels if double_z else z_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
# timestep embedding
|
||||
temb = None
|
||||
|
||||
# downsampling
|
||||
hs = [self.conv_in(x)]
|
||||
for i_level in range(self.num_resolutions):
|
||||
for i_block in range(self.num_res_blocks):
|
||||
h = self.down[i_level].block[i_block](hs[-1], temb)
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
hs.append(h)
|
||||
if i_level != self.num_resolutions - 1:
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
|
||||
# middle
|
||||
h = hs[-1]
|
||||
h = self.mid.block_1(h, temb)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h, temb)
|
||||
|
||||
# end
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
return h
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ch,
|
||||
out_ch,
|
||||
ch_mult=(1, 2, 4, 8),
|
||||
num_res_blocks,
|
||||
attn_resolutions,
|
||||
dropout=0.0,
|
||||
resamp_with_conv=True,
|
||||
in_channels,
|
||||
resolution,
|
||||
z_channels,
|
||||
give_pre_end=False,
|
||||
tanh_out=False,
|
||||
use_linear_attn=False,
|
||||
attn_type="vanilla",
|
||||
**ignorekwargs,
|
||||
):
|
||||
super().__init__()
|
||||
if use_linear_attn:
|
||||
attn_type = "linear"
|
||||
self.ch = ch
|
||||
self.temb_ch = 0
|
||||
self.num_resolutions = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.resolution = resolution
|
||||
self.in_channels = in_channels
|
||||
self.give_pre_end = give_pre_end
|
||||
self.tanh_out = tanh_out
|
||||
|
||||
# compute in_ch_mult, block_in and curr_res at lowest res
|
||||
in_ch_mult = (1,) + tuple(ch_mult)
|
||||
block_in = ch * ch_mult[self.num_resolutions - 1]
|
||||
curr_res = resolution // 2 ** (self.num_resolutions - 1)
|
||||
self.z_shape = (1, z_channels, curr_res, curr_res)
|
||||
print(
|
||||
"Working with z of shape {} = {} dimensions.".format(
|
||||
self.z_shape, np.prod(self.z_shape)
|
||||
)
|
||||
)
|
||||
|
||||
make_attn_cls = self._make_attn()
|
||||
make_resblock_cls = self._make_resblock()
|
||||
make_conv_cls = self._make_conv()
|
||||
# z to block_in
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
z_channels, block_in, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = make_resblock_cls(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.mid.attn_1 = make_attn_cls(block_in, attn_type=attn_type)
|
||||
self.mid.block_2 = make_resblock_cls(
|
||||
in_channels=block_in,
|
||||
out_channels=block_in,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
|
||||
# upsampling
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = ch * ch_mult[i_level]
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
block.append(
|
||||
make_resblock_cls(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
temb_channels=self.temb_ch,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(make_attn_cls(block_in, attn_type=attn_type))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if i_level != 0:
|
||||
up.upsample = Upsample(block_in, resamp_with_conv)
|
||||
curr_res = curr_res * 2
|
||||
self.up.insert(0, up) # prepend to get consistent order
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = make_conv_cls(
|
||||
block_in, out_ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def _make_attn(self) -> Callable:
|
||||
return make_attn
|
||||
|
||||
def _make_resblock(self) -> Callable:
|
||||
return ResnetBlock
|
||||
|
||||
def _make_conv(self) -> Callable:
|
||||
return torch.nn.Conv2d
|
||||
|
||||
def get_last_layer(self, **kwargs):
|
||||
return self.conv_out.weight
|
||||
|
||||
def forward(self, z, **kwargs):
|
||||
# assert z.shape[1:] == self.z_shape[1:]
|
||||
self.last_z_shape = z.shape
|
||||
|
||||
# timestep embedding
|
||||
temb = None
|
||||
|
||||
# z to block_in
|
||||
h = self.conv_in(z)
|
||||
|
||||
# middle
|
||||
h = self.mid.block_1(h, temb, **kwargs)
|
||||
h = self.mid.attn_1(h, **kwargs)
|
||||
h = self.mid.block_2(h, temb, **kwargs)
|
||||
|
||||
# upsampling
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
h = self.up[i_level].block[i_block](h, temb, **kwargs)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h, **kwargs)
|
||||
if i_level != 0:
|
||||
h = self.up[i_level].upsample(h)
|
||||
|
||||
# end
|
||||
if self.give_pre_end:
|
||||
return h
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h, **kwargs)
|
||||
if self.tanh_out:
|
||||
h = torch.tanh(h)
|
||||
return h
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,449 @@
|
||||
"""
|
||||
Partially ported from https://github.com/crowsonkb/k-diffusion/blob/master/k_diffusion/sampling.py
|
||||
"""
|
||||
|
||||
|
||||
from typing import Dict, Union
|
||||
|
||||
import torch
|
||||
from omegaconf import ListConfig, OmegaConf
|
||||
from tqdm import tqdm
|
||||
|
||||
from ...modules.diffusionmodules.sampling_utils import (
|
||||
get_ancestral_step,
|
||||
linear_multistep_coeff,
|
||||
to_d,
|
||||
to_neg_log_sigma,
|
||||
to_sigma,
|
||||
)
|
||||
from ...util import append_dims, default, instantiate_from_config
|
||||
|
||||
DEFAULT_GUIDER = {"target": ".sgm.modules.diffusionmodules.guiders.IdentityGuider"}
|
||||
|
||||
|
||||
class BaseDiffusionSampler:
|
||||
def __init__(
|
||||
self,
|
||||
discretization_config: Union[Dict, ListConfig, OmegaConf],
|
||||
num_steps: Union[int, None] = None,
|
||||
guider_config: Union[Dict, ListConfig, OmegaConf, None] = None,
|
||||
verbose: bool = False,
|
||||
device: str = "cuda",
|
||||
):
|
||||
self.num_steps = num_steps
|
||||
self.discretization = instantiate_from_config(discretization_config)
|
||||
self.guider = instantiate_from_config(
|
||||
default(
|
||||
guider_config,
|
||||
DEFAULT_GUIDER,
|
||||
)
|
||||
)
|
||||
self.verbose = verbose
|
||||
self.device = device
|
||||
|
||||
def prepare_sampling_loop(self, x, cond, uc=None, num_steps=None):
|
||||
sigmas = self.discretization(
|
||||
self.num_steps if num_steps is None else num_steps, device=self.device
|
||||
)
|
||||
uc = default(uc, cond)
|
||||
|
||||
x *= torch.sqrt(1.0 + sigmas[0] ** 2.0)
|
||||
num_sigmas = len(sigmas)
|
||||
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
return x, s_in, sigmas, num_sigmas, cond, uc
|
||||
|
||||
def denoise(self, x, denoiser, sigma, cond, uc):
|
||||
denoised = denoiser(*self.guider.prepare_inputs(x, sigma, cond, uc))
|
||||
denoised = self.guider(denoised, sigma)
|
||||
return denoised
|
||||
|
||||
def get_sigma_gen(self, num_sigmas):
|
||||
sigma_generator = range(num_sigmas - 1)
|
||||
if self.verbose:
|
||||
print("#" * 30, " Sampling setting ", "#" * 30)
|
||||
print(f"Sampler: {self.__class__.__name__}")
|
||||
print(f"Discretization: {self.discretization.__class__.__name__}")
|
||||
print(f"Guider: {self.guider.__class__.__name__}")
|
||||
sigma_generator = tqdm(
|
||||
sigma_generator,
|
||||
total=num_sigmas,
|
||||
desc=f"Sampling with {self.__class__.__name__} for {num_sigmas} steps",
|
||||
)
|
||||
return sigma_generator
|
||||
|
||||
|
||||
class SingleStepDiffusionSampler(BaseDiffusionSampler):
|
||||
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def euler_step(self, x, d, dt):
|
||||
return x + dt * d
|
||||
|
||||
|
||||
class EDMSampler(SingleStepDiffusionSampler):
|
||||
def __init__(
|
||||
self, s_churn=0.0, s_tmin=0.0, s_tmax=float("inf"), s_noise=1.0, *args, **kwargs
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.s_churn = s_churn
|
||||
self.s_tmin = s_tmin
|
||||
self.s_tmax = s_tmax
|
||||
self.s_noise = s_noise
|
||||
|
||||
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc=None, gamma=0.0):
|
||||
sigma_hat = sigma * (gamma + 1.0)
|
||||
if gamma > 0:
|
||||
eps = torch.randn_like(x) * self.s_noise
|
||||
x = x + eps * append_dims(sigma_hat**2 - sigma**2, x.ndim) ** 0.5
|
||||
|
||||
denoised = self.denoise(x, denoiser, sigma_hat, cond, uc)
|
||||
# print('denoised', denoised.mean(axis=[0, 2, 3]))
|
||||
d = to_d(x, sigma_hat, denoised)
|
||||
dt = append_dims(next_sigma - sigma_hat, x.ndim)
|
||||
|
||||
euler_step = self.euler_step(x, d, dt)
|
||||
x = self.possible_correction_step(
|
||||
euler_step, x, d, dt, next_sigma, denoiser, cond, uc
|
||||
)
|
||||
return x
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
|
||||
for i in self.get_sigma_gen(num_sigmas):
|
||||
gamma = (
|
||||
min(self.s_churn / (num_sigmas - 1), 2**0.5 - 1)
|
||||
if self.s_tmin <= sigmas[i] <= self.s_tmax
|
||||
else 0.0
|
||||
)
|
||||
x = self.sampler_step(
|
||||
s_in * sigmas[i],
|
||||
s_in * sigmas[i + 1],
|
||||
denoiser,
|
||||
x,
|
||||
cond,
|
||||
uc,
|
||||
gamma,
|
||||
)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class AncestralSampler(SingleStepDiffusionSampler):
|
||||
def __init__(self, eta=1.0, s_noise=1.0, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.eta = eta
|
||||
self.s_noise = s_noise
|
||||
self.noise_sampler = lambda x: torch.randn_like(x)
|
||||
|
||||
def ancestral_euler_step(self, x, denoised, sigma, sigma_down):
|
||||
d = to_d(x, sigma, denoised)
|
||||
dt = append_dims(sigma_down - sigma, x.ndim)
|
||||
|
||||
return self.euler_step(x, d, dt)
|
||||
|
||||
def ancestral_step(self, x, sigma, next_sigma, sigma_up):
|
||||
x = torch.where(
|
||||
append_dims(next_sigma, x.ndim) > 0.0,
|
||||
x + self.noise_sampler(x) * self.s_noise * append_dims(sigma_up, x.ndim),
|
||||
x,
|
||||
)
|
||||
return x
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
|
||||
for i in self.get_sigma_gen(num_sigmas):
|
||||
x = self.sampler_step(
|
||||
s_in * sigmas[i],
|
||||
s_in * sigmas[i + 1],
|
||||
denoiser,
|
||||
x,
|
||||
cond,
|
||||
uc,
|
||||
)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class LinearMultistepSampler(BaseDiffusionSampler):
|
||||
def __init__(
|
||||
self,
|
||||
order=4,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.order = order
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, **kwargs):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
|
||||
ds = []
|
||||
sigmas_cpu = sigmas.detach().cpu().numpy()
|
||||
for i in self.get_sigma_gen(num_sigmas):
|
||||
sigma = s_in * sigmas[i]
|
||||
denoised = denoiser(
|
||||
*self.guider.prepare_inputs(x, sigma, cond, uc), **kwargs
|
||||
)
|
||||
denoised = self.guider(denoised, sigma)
|
||||
d = to_d(x, sigma, denoised)
|
||||
ds.append(d)
|
||||
if len(ds) > self.order:
|
||||
ds.pop(0)
|
||||
cur_order = min(i + 1, self.order)
|
||||
coeffs = [
|
||||
linear_multistep_coeff(cur_order, sigmas_cpu, i, j)
|
||||
for j in range(cur_order)
|
||||
]
|
||||
x = x + sum(coeff * d for coeff, d in zip(coeffs, reversed(ds)))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class EulerEDMSampler(EDMSampler):
|
||||
def possible_correction_step(
|
||||
self, euler_step, x, d, dt, next_sigma, denoiser, cond, uc
|
||||
):
|
||||
# print("euler_step: ", euler_step.mean(axis=[0, 2, 3]))
|
||||
return euler_step
|
||||
|
||||
|
||||
class HeunEDMSampler(EDMSampler):
|
||||
def possible_correction_step(
|
||||
self, euler_step, x, d, dt, next_sigma, denoiser, cond, uc
|
||||
):
|
||||
if torch.sum(next_sigma) < 1e-14:
|
||||
# Save a network evaluation if all noise levels are 0
|
||||
return euler_step
|
||||
else:
|
||||
denoised = self.denoise(euler_step, denoiser, next_sigma, cond, uc)
|
||||
d_new = to_d(euler_step, next_sigma, denoised)
|
||||
d_prime = (d + d_new) / 2.0
|
||||
|
||||
# apply correction if noise level is not 0
|
||||
x = torch.where(
|
||||
append_dims(next_sigma, x.ndim) > 0.0, x + d_prime * dt, euler_step
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
class EulerAncestralSampler(AncestralSampler):
|
||||
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc):
|
||||
sigma_down, sigma_up = get_ancestral_step(sigma, next_sigma, eta=self.eta)
|
||||
denoised = self.denoise(x, denoiser, sigma, cond, uc)
|
||||
x = self.ancestral_euler_step(x, denoised, sigma, sigma_down)
|
||||
x = self.ancestral_step(x, sigma, next_sigma, sigma_up)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class DPMPP2SAncestralSampler(AncestralSampler):
|
||||
def get_variables(self, sigma, sigma_down):
|
||||
t, t_next = [to_neg_log_sigma(s) for s in (sigma, sigma_down)]
|
||||
h = t_next - t
|
||||
s = t + 0.5 * h
|
||||
return h, s, t, t_next
|
||||
|
||||
def get_mult(self, h, s, t, t_next):
|
||||
mult1 = to_sigma(s) / to_sigma(t)
|
||||
mult2 = (-0.5 * h).expm1()
|
||||
mult3 = to_sigma(t_next) / to_sigma(t)
|
||||
mult4 = (-h).expm1()
|
||||
|
||||
return mult1, mult2, mult3, mult4
|
||||
|
||||
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc=None, **kwargs):
|
||||
sigma_down, sigma_up = get_ancestral_step(sigma, next_sigma, eta=self.eta)
|
||||
denoised = self.denoise(x, denoiser, sigma, cond, uc)
|
||||
x_euler = self.ancestral_euler_step(x, denoised, sigma, sigma_down)
|
||||
|
||||
if torch.sum(sigma_down) < 1e-14:
|
||||
# Save a network evaluation if all noise levels are 0
|
||||
x = x_euler
|
||||
else:
|
||||
h, s, t, t_next = self.get_variables(sigma, sigma_down)
|
||||
mult = [
|
||||
append_dims(mult, x.ndim) for mult in self.get_mult(h, s, t, t_next)
|
||||
]
|
||||
|
||||
x2 = mult[0] * x - mult[1] * denoised
|
||||
denoised2 = self.denoise(x2, denoiser, to_sigma(s), cond, uc)
|
||||
x_dpmpp2s = mult[2] * x - mult[3] * denoised2
|
||||
|
||||
# apply correction if noise level is not 0
|
||||
x = torch.where(append_dims(sigma_down, x.ndim) > 0.0, x_dpmpp2s, x_euler)
|
||||
|
||||
x = self.ancestral_step(x, sigma, next_sigma, sigma_up)
|
||||
return x
|
||||
|
||||
|
||||
class DPMPP2MSampler(BaseDiffusionSampler):
|
||||
def get_variables(self, sigma, next_sigma, previous_sigma=None):
|
||||
t, t_next = [to_neg_log_sigma(s) for s in (sigma, next_sigma)]
|
||||
h = t_next - t
|
||||
|
||||
if previous_sigma is not None:
|
||||
h_last = t - to_neg_log_sigma(previous_sigma)
|
||||
r = h_last / h
|
||||
return h, r, t, t_next
|
||||
else:
|
||||
return h, None, t, t_next
|
||||
|
||||
def get_mult(self, h, r, t, t_next, previous_sigma):
|
||||
mult1 = to_sigma(t_next) / to_sigma(t)
|
||||
mult2 = (-h).expm1()
|
||||
|
||||
if previous_sigma is not None:
|
||||
mult3 = 1 + 1 / (2 * r)
|
||||
mult4 = 1 / (2 * r)
|
||||
return mult1, mult2, mult3, mult4
|
||||
else:
|
||||
return mult1, mult2
|
||||
|
||||
def sampler_step(
|
||||
self,
|
||||
old_denoised,
|
||||
previous_sigma,
|
||||
sigma,
|
||||
next_sigma,
|
||||
denoiser,
|
||||
x,
|
||||
cond,
|
||||
uc=None,
|
||||
):
|
||||
denoised = self.denoise(x, denoiser, sigma, cond, uc)
|
||||
|
||||
h, r, t, t_next = self.get_variables(sigma, next_sigma, previous_sigma)
|
||||
mult = [
|
||||
append_dims(mult, x.ndim)
|
||||
for mult in self.get_mult(h, r, t, t_next, previous_sigma)
|
||||
]
|
||||
|
||||
x_standard = mult[0] * x - mult[1] * denoised
|
||||
if old_denoised is None or torch.sum(next_sigma) < 1e-14:
|
||||
# Save a network evaluation if all noise levels are 0 or on the first step
|
||||
return x_standard, denoised
|
||||
else:
|
||||
denoised_d = mult[2] * denoised - mult[3] * old_denoised
|
||||
x_advanced = mult[0] * x - mult[1] * denoised_d
|
||||
|
||||
# apply correction if noise level is not 0 and not first step
|
||||
x = torch.where(
|
||||
append_dims(next_sigma, x.ndim) > 0.0, x_advanced, x_standard
|
||||
)
|
||||
|
||||
return x, denoised
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, **kwargs):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
|
||||
old_denoised = None
|
||||
for i in self.get_sigma_gen(num_sigmas):
|
||||
x, old_denoised = self.sampler_step(
|
||||
old_denoised,
|
||||
None if i == 0 else s_in * sigmas[i - 1],
|
||||
s_in * sigmas[i],
|
||||
s_in * sigmas[i + 1],
|
||||
denoiser,
|
||||
x,
|
||||
cond,
|
||||
uc=uc,
|
||||
)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class RestoreEDMSampler(SingleStepDiffusionSampler):
|
||||
def __init__(
|
||||
self, s_churn=0.0, s_tmin=0.0, s_tmax=float("inf"), s_noise=1.0, restore_cfg=4.0,
|
||||
restore_cfg_s_tmin=0.05, *args, **kwargs
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.s_churn = s_churn
|
||||
self.s_tmin = s_tmin
|
||||
self.s_tmax = s_tmax
|
||||
self.s_noise = s_noise
|
||||
self.restore_cfg = restore_cfg
|
||||
self.restore_cfg_s_tmin = restore_cfg_s_tmin
|
||||
self.sigma_max = 14.6146
|
||||
|
||||
def denoise(self, x, denoiser, sigma, cond, uc, control_scale=1.0):
|
||||
denoised = denoiser(*self.guider.prepare_inputs(x, sigma, cond, uc), control_scale)
|
||||
denoised = self.guider(denoised, sigma)
|
||||
return denoised
|
||||
|
||||
|
||||
def sampler_step(self, sigma, next_sigma, denoiser, x, cond, uc=None, gamma=0.0, x_center=None, eps_noise=None,
|
||||
control_scale=1.0, use_linear_control_scale=False, control_scale_start=0.0):
|
||||
sigma_hat = sigma * (gamma + 1.0)
|
||||
if gamma > 0:
|
||||
if eps_noise is not None:
|
||||
eps = eps_noise * self.s_noise
|
||||
else:
|
||||
eps = torch.randn_like(x) * self.s_noise
|
||||
x = x + eps * append_dims(sigma_hat**2 - sigma**2, x.ndim) ** 0.5
|
||||
|
||||
if use_linear_control_scale:
|
||||
control_scale = (sigma[0].item() / self.sigma_max) * (control_scale_start - control_scale) + control_scale
|
||||
|
||||
denoised = self.denoise(x, denoiser, sigma_hat, cond, uc, control_scale=control_scale)
|
||||
|
||||
if (next_sigma[0] > self.restore_cfg_s_tmin) and (self.restore_cfg > 0):
|
||||
d_center = (denoised - x_center)
|
||||
denoised = denoised - d_center * ((sigma.view(-1, 1, 1, 1) / self.sigma_max) ** self.restore_cfg)
|
||||
|
||||
d = to_d(x, sigma_hat, denoised)
|
||||
dt = append_dims(next_sigma - sigma_hat, x.ndim)
|
||||
x = self.euler_step(x, d, dt)
|
||||
return x
|
||||
|
||||
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, x_center=None, control_scale=1.0,
|
||||
use_linear_control_scale=False, control_scale_start=0.0):
|
||||
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
|
||||
x, cond, uc, num_steps
|
||||
)
|
||||
|
||||
for _idx, i in enumerate(self.get_sigma_gen(num_sigmas)):
|
||||
gamma = (
|
||||
min(self.s_churn / (num_sigmas - 1), 2**0.5 - 1)
|
||||
if self.s_tmin <= sigmas[i] <= self.s_tmax
|
||||
else 0.0
|
||||
)
|
||||
x = self.sampler_step(
|
||||
s_in * sigmas[i],
|
||||
s_in * sigmas[i + 1],
|
||||
denoiser,
|
||||
x,
|
||||
cond,
|
||||
uc,
|
||||
gamma,
|
||||
x_center,
|
||||
control_scale=control_scale,
|
||||
use_linear_control_scale=use_linear_control_scale,
|
||||
control_scale_start=control_scale_start,
|
||||
)
|
||||
return x
|
||||
|
||||
def to_d_center(denoised, x_center, x):
|
||||
b = denoised.shape[0]
|
||||
v_center = (denoised - x_center).view(b, -1)
|
||||
v_denoise = (x - denoised).view(b, -1)
|
||||
d_center = v_center - v_denoise * (v_center * v_denoise).sum(dim=1).view(b, 1) / \
|
||||
(v_denoise * v_denoise).sum(dim=1).view(b, 1)
|
||||
d_center = d_center / d_center.view(x.shape[0], -1).norm(dim=1).view(-1, 1)
|
||||
return d_center.view(denoised.shape)
|
||||
@@ -0,0 +1,48 @@
|
||||
import torch
|
||||
from scipy import integrate
|
||||
|
||||
from ...util import append_dims
|
||||
|
||||
|
||||
class NoDynamicThresholding:
|
||||
def __call__(self, uncond, cond, scale):
|
||||
return uncond + scale.view(-1, 1, 1, 1) * (cond - uncond)
|
||||
|
||||
|
||||
def linear_multistep_coeff(order, t, i, j, epsrel=1e-4):
|
||||
if order - 1 > i:
|
||||
raise ValueError(f"Order {order} too high for step {i}")
|
||||
|
||||
def fn(tau):
|
||||
prod = 1.0
|
||||
for k in range(order):
|
||||
if j == k:
|
||||
continue
|
||||
prod *= (tau - t[i - k]) / (t[i - j] - t[i - k])
|
||||
return prod
|
||||
|
||||
return integrate.quad(fn, t[i], t[i + 1], epsrel=epsrel)[0]
|
||||
|
||||
|
||||
def get_ancestral_step(sigma_from, sigma_to, eta=1.0):
|
||||
if not eta:
|
||||
return sigma_to, 0.0
|
||||
sigma_up = torch.minimum(
|
||||
sigma_to,
|
||||
eta
|
||||
* (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5,
|
||||
)
|
||||
sigma_down = (sigma_to**2 - sigma_up**2) ** 0.5
|
||||
return sigma_down, sigma_up
|
||||
|
||||
|
||||
def to_d(x, sigma, denoised):
|
||||
return (x - denoised) / append_dims(sigma, x.ndim)
|
||||
|
||||
|
||||
def to_neg_log_sigma(sigma):
|
||||
return sigma.log().neg()
|
||||
|
||||
|
||||
def to_sigma(neg_log_sigma):
|
||||
return neg_log_sigma.neg().exp()
|
||||
@@ -0,0 +1,40 @@
|
||||
import torch
|
||||
|
||||
from ...util import default, instantiate_from_config
|
||||
|
||||
|
||||
class EDMSampling:
|
||||
def __init__(self, p_mean=-1.2, p_std=1.2):
|
||||
self.p_mean = p_mean
|
||||
self.p_std = p_std
|
||||
|
||||
def __call__(self, n_samples, rand=None):
|
||||
log_sigma = self.p_mean + self.p_std * default(rand, torch.randn((n_samples,)))
|
||||
return log_sigma.exp()
|
||||
|
||||
|
||||
class DiscreteSampling:
|
||||
def __init__(self, discretization_config, num_idx, do_append_zero=False, flip=True, idx_range=None):
|
||||
self.num_idx = num_idx
|
||||
self.sigmas = instantiate_from_config(discretization_config)(
|
||||
num_idx, do_append_zero=do_append_zero, flip=flip
|
||||
)
|
||||
self.idx_range = idx_range
|
||||
|
||||
def idx_to_sigma(self, idx):
|
||||
# print(self.sigmas[idx])
|
||||
return self.sigmas[idx]
|
||||
|
||||
def __call__(self, n_samples, rand=None):
|
||||
if self.idx_range is None:
|
||||
idx = default(
|
||||
rand,
|
||||
torch.randint(0, self.num_idx, (n_samples,)),
|
||||
)
|
||||
else:
|
||||
idx = default(
|
||||
rand,
|
||||
torch.randint(self.idx_range[0], self.idx_range[1], (n_samples,)),
|
||||
)
|
||||
return self.idx_to_sigma(idx)
|
||||
|
||||
@@ -0,0 +1,309 @@
|
||||
"""
|
||||
adopted from
|
||||
https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
||||
and
|
||||
https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
|
||||
and
|
||||
https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py
|
||||
|
||||
thanks!
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import repeat
|
||||
|
||||
|
||||
def make_beta_schedule(
|
||||
schedule,
|
||||
n_timestep,
|
||||
linear_start=1e-4,
|
||||
linear_end=2e-2,
|
||||
):
|
||||
if schedule == "linear":
|
||||
betas = (
|
||||
torch.linspace(
|
||||
linear_start**0.5, linear_end**0.5, n_timestep, dtype=torch.float64
|
||||
)
|
||||
** 2
|
||||
)
|
||||
return betas.numpy()
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
|
||||
def mixed_checkpoint(func, inputs: dict, params, flag):
|
||||
"""
|
||||
Evaluate a function without caching intermediate activations, allowing for
|
||||
reduced memory at the expense of extra compute in the backward pass. This differs from the original checkpoint function
|
||||
borrowed from https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py in that
|
||||
it also works with non-tensor inputs
|
||||
:param func: the function to evaluate.
|
||||
:param inputs: the argument dictionary to pass to `func`.
|
||||
:param params: a sequence of parameters `func` depends on but does not
|
||||
explicitly take as arguments.
|
||||
:param flag: if False, disable gradient checkpointing.
|
||||
"""
|
||||
if flag:
|
||||
tensor_keys = [key for key in inputs if isinstance(inputs[key], torch.Tensor)]
|
||||
tensor_inputs = [
|
||||
inputs[key] for key in inputs if isinstance(inputs[key], torch.Tensor)
|
||||
]
|
||||
non_tensor_keys = [
|
||||
key for key in inputs if not isinstance(inputs[key], torch.Tensor)
|
||||
]
|
||||
non_tensor_inputs = [
|
||||
inputs[key] for key in inputs if not isinstance(inputs[key], torch.Tensor)
|
||||
]
|
||||
args = tuple(tensor_inputs) + tuple(non_tensor_inputs) + tuple(params)
|
||||
return MixedCheckpointFunction.apply(
|
||||
func,
|
||||
len(tensor_inputs),
|
||||
len(non_tensor_inputs),
|
||||
tensor_keys,
|
||||
non_tensor_keys,
|
||||
*args,
|
||||
)
|
||||
else:
|
||||
return func(**inputs)
|
||||
|
||||
|
||||
class MixedCheckpointFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
run_function,
|
||||
length_tensors,
|
||||
length_non_tensors,
|
||||
tensor_keys,
|
||||
non_tensor_keys,
|
||||
*args,
|
||||
):
|
||||
ctx.end_tensors = length_tensors
|
||||
ctx.end_non_tensors = length_tensors + length_non_tensors
|
||||
ctx.gpu_autocast_kwargs = {
|
||||
"enabled": torch.is_autocast_enabled(),
|
||||
"dtype": torch.get_autocast_gpu_dtype(),
|
||||
"cache_enabled": torch.is_autocast_cache_enabled(),
|
||||
}
|
||||
assert (
|
||||
len(tensor_keys) == length_tensors
|
||||
and len(non_tensor_keys) == length_non_tensors
|
||||
)
|
||||
|
||||
ctx.input_tensors = {
|
||||
key: val for (key, val) in zip(tensor_keys, list(args[: ctx.end_tensors]))
|
||||
}
|
||||
ctx.input_non_tensors = {
|
||||
key: val
|
||||
for (key, val) in zip(
|
||||
non_tensor_keys, list(args[ctx.end_tensors : ctx.end_non_tensors])
|
||||
)
|
||||
}
|
||||
ctx.run_function = run_function
|
||||
ctx.input_params = list(args[ctx.end_non_tensors :])
|
||||
|
||||
with torch.no_grad():
|
||||
output_tensors = ctx.run_function(
|
||||
**ctx.input_tensors, **ctx.input_non_tensors
|
||||
)
|
||||
return output_tensors
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, *output_grads):
|
||||
# additional_args = {key: ctx.input_tensors[key] for key in ctx.input_tensors if not isinstance(ctx.input_tensors[key],torch.Tensor)}
|
||||
ctx.input_tensors = {
|
||||
key: ctx.input_tensors[key].detach().requires_grad_(True)
|
||||
for key in ctx.input_tensors
|
||||
}
|
||||
|
||||
with torch.enable_grad(), torch.cuda.amp.autocast(**ctx.gpu_autocast_kwargs):
|
||||
# Fixes a bug where the first op in run_function modifies the
|
||||
# Tensor storage in place, which is not allowed for detach()'d
|
||||
# Tensors.
|
||||
shallow_copies = {
|
||||
key: ctx.input_tensors[key].view_as(ctx.input_tensors[key])
|
||||
for key in ctx.input_tensors
|
||||
}
|
||||
# shallow_copies.update(additional_args)
|
||||
output_tensors = ctx.run_function(**shallow_copies, **ctx.input_non_tensors)
|
||||
input_grads = torch.autograd.grad(
|
||||
output_tensors,
|
||||
list(ctx.input_tensors.values()) + ctx.input_params,
|
||||
output_grads,
|
||||
allow_unused=True,
|
||||
)
|
||||
del ctx.input_tensors
|
||||
del ctx.input_params
|
||||
del output_tensors
|
||||
return (
|
||||
(None, None, None, None, None)
|
||||
+ input_grads[: ctx.end_tensors]
|
||||
+ (None,) * (ctx.end_non_tensors - ctx.end_tensors)
|
||||
+ input_grads[ctx.end_tensors :]
|
||||
)
|
||||
|
||||
|
||||
def checkpoint(func, inputs, params, flag):
|
||||
"""
|
||||
Evaluate a function without caching intermediate activations, allowing for
|
||||
reduced memory at the expense of extra compute in the backward pass.
|
||||
:param func: the function to evaluate.
|
||||
:param inputs: the argument sequence to pass to `func`.
|
||||
:param params: a sequence of parameters `func` depends on but does not
|
||||
explicitly take as arguments.
|
||||
:param flag: if False, disable gradient checkpointing.
|
||||
"""
|
||||
if flag:
|
||||
args = tuple(inputs) + tuple(params)
|
||||
return CheckpointFunction.apply(func, len(inputs), *args)
|
||||
else:
|
||||
return func(*inputs)
|
||||
|
||||
|
||||
class CheckpointFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, run_function, length, *args):
|
||||
ctx.run_function = run_function
|
||||
ctx.input_tensors = list(args[:length])
|
||||
ctx.input_params = list(args[length:])
|
||||
ctx.gpu_autocast_kwargs = {
|
||||
"enabled": torch.is_autocast_enabled(),
|
||||
"dtype": torch.get_autocast_gpu_dtype(),
|
||||
"cache_enabled": torch.is_autocast_cache_enabled(),
|
||||
}
|
||||
with torch.no_grad():
|
||||
output_tensors = ctx.run_function(*ctx.input_tensors)
|
||||
return output_tensors
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, *output_grads):
|
||||
ctx.input_tensors = [x.detach().requires_grad_(True) for x in ctx.input_tensors]
|
||||
with torch.enable_grad(), torch.cuda.amp.autocast(**ctx.gpu_autocast_kwargs):
|
||||
# Fixes a bug where the first op in run_function modifies the
|
||||
# Tensor storage in place, which is not allowed for detach()'d
|
||||
# Tensors.
|
||||
shallow_copies = [x.view_as(x) for x in ctx.input_tensors]
|
||||
output_tensors = ctx.run_function(*shallow_copies)
|
||||
input_grads = torch.autograd.grad(
|
||||
output_tensors,
|
||||
ctx.input_tensors + ctx.input_params,
|
||||
output_grads,
|
||||
allow_unused=True,
|
||||
)
|
||||
del ctx.input_tensors
|
||||
del ctx.input_params
|
||||
del output_tensors
|
||||
return (None, None) + input_grads
|
||||
|
||||
|
||||
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param timesteps: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an [N x dim] Tensor of positional embeddings.
|
||||
"""
|
||||
if not repeat_only:
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period)
|
||||
* torch.arange(start=0, end=half, dtype=torch.float32)
|
||||
/ half
|
||||
).to(device=timesteps.device)
|
||||
args = timesteps[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1
|
||||
)
|
||||
else:
|
||||
embedding = repeat(timesteps, "b -> b d", d=dim)
|
||||
return embedding
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
"""
|
||||
Zero out the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
def scale_module(module, scale):
|
||||
"""
|
||||
Scale the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().mul_(scale)
|
||||
return module
|
||||
|
||||
|
||||
def mean_flat(tensor):
|
||||
"""
|
||||
Take the mean over all non-batch dimensions.
|
||||
"""
|
||||
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
||||
|
||||
|
||||
def normalization(channels):
|
||||
"""
|
||||
Make a standard normalization layer.
|
||||
:param channels: number of input channels.
|
||||
:return: an nn.Module for normalization.
|
||||
"""
|
||||
return GroupNorm32(32, channels)
|
||||
|
||||
|
||||
# PyTorch 1.7 has SiLU, but we support PyTorch 1.5.
|
||||
class SiLU(nn.Module):
|
||||
def forward(self, x):
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
class GroupNorm32(nn.GroupNorm):
|
||||
def forward(self, x):
|
||||
# return super().forward(x.float()).type(x.dtype)
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
def conv_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D convolution module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.Conv1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return nn.Conv2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return nn.Conv3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
|
||||
|
||||
def linear(*args, **kwargs):
|
||||
"""
|
||||
Create a linear module.
|
||||
"""
|
||||
return nn.Linear(*args, **kwargs)
|
||||
|
||||
|
||||
def avg_pool_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D average pooling module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.AvgPool1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return nn.AvgPool2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return nn.AvgPool3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
@@ -0,0 +1,103 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from packaging import version
|
||||
# import torch._dynamo
|
||||
# torch._dynamo.config.suppress_errors = True
|
||||
# torch._dynamo.config.cache_size_limit = 512
|
||||
|
||||
OPENAIUNETWRAPPER = ".sgm.modules.diffusionmodules.wrappers.OpenAIWrapper"
|
||||
|
||||
|
||||
class IdentityWrapper(nn.Module):
|
||||
def __init__(self, diffusion_model, compile_model: bool = False):
|
||||
super().__init__()
|
||||
compile = (
|
||||
torch.compile
|
||||
if (version.parse(torch.__version__) >= version.parse("2.0.0"))
|
||||
and compile_model
|
||||
else lambda x: x
|
||||
)
|
||||
self.diffusion_model = compile(diffusion_model)
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
return self.diffusion_model(*args, **kwargs)
|
||||
|
||||
|
||||
class OpenAIWrapper(IdentityWrapper):
|
||||
def forward(
|
||||
self, x: torch.Tensor, t: torch.Tensor, c: dict, **kwargs
|
||||
) -> torch.Tensor:
|
||||
x = torch.cat((x, c.get("concat", torch.Tensor([]).type_as(x))), dim=1)
|
||||
return self.diffusion_model(
|
||||
x,
|
||||
timesteps=t,
|
||||
context=c.get("crossattn", None),
|
||||
y=c.get("vector", None),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class OpenAIHalfWrapper(IdentityWrapper):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.diffusion_model = self.diffusion_model.half()
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, t: torch.Tensor, c: dict, **kwargs
|
||||
) -> torch.Tensor:
|
||||
x = torch.cat((x, c.get("concat", torch.Tensor([]).type_as(x))), dim=1)
|
||||
_context = c.get("crossattn", None)
|
||||
_y = c.get("vector", None)
|
||||
if _context is not None:
|
||||
_context = _context.half()
|
||||
if _y is not None:
|
||||
_y = _y.half()
|
||||
x = x.half()
|
||||
t = t.half()
|
||||
|
||||
out = self.diffusion_model(
|
||||
x,
|
||||
timesteps=t,
|
||||
context=_context,
|
||||
y=_y,
|
||||
**kwargs,
|
||||
)
|
||||
return out.float()
|
||||
|
||||
|
||||
class ControlWrapper(nn.Module):
|
||||
def __init__(self, diffusion_model, compile_model: bool = False, dtype=torch.float32):
|
||||
super().__init__()
|
||||
self.compile = (
|
||||
torch.compile
|
||||
if (version.parse(torch.__version__) >= version.parse("2.0.0"))
|
||||
and compile_model
|
||||
else lambda x: x
|
||||
)
|
||||
self.diffusion_model = self.compile(diffusion_model)
|
||||
self.control_model = None
|
||||
self.dtype = dtype
|
||||
|
||||
def load_control_model(self, control_model):
|
||||
self.control_model = self.compile(control_model)
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, t: torch.Tensor, c: dict, control_scale=1, **kwargs
|
||||
) -> torch.Tensor:
|
||||
with torch.autocast("cuda", dtype=self.dtype):
|
||||
control = self.control_model(x=c.get("control", None), timesteps=t, xt=x,
|
||||
control_vector=c.get("control_vector", None),
|
||||
mask_x=c.get("mask_x", None),
|
||||
context=c.get("crossattn", None),
|
||||
y=c.get("vector", None))
|
||||
out = self.diffusion_model(
|
||||
x,
|
||||
timesteps=t,
|
||||
context=c.get("crossattn", None),
|
||||
y=c.get("vector", None),
|
||||
control=control,
|
||||
control_scale=control_scale,
|
||||
**kwargs,
|
||||
)
|
||||
return out.float()
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
class AbstractDistribution:
|
||||
def sample(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
def mode(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class DiracDistribution(AbstractDistribution):
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
|
||||
def sample(self):
|
||||
return self.value
|
||||
|
||||
def mode(self):
|
||||
return self.value
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution(object):
|
||||
def __init__(self, parameters, deterministic=False):
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if self.deterministic:
|
||||
self.var = self.std = torch.zeros_like(self.mean).to(
|
||||
device=self.parameters.device
|
||||
)
|
||||
|
||||
def sample(self):
|
||||
x = self.mean + self.std * torch.randn(self.mean.shape).to(
|
||||
device=self.parameters.device
|
||||
)
|
||||
return x
|
||||
|
||||
def kl(self, other=None):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
else:
|
||||
if other is None:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
|
||||
dim=[1, 2, 3],
|
||||
)
|
||||
else:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean - other.mean, 2) / other.var
|
||||
+ self.var / other.var
|
||||
- 1.0
|
||||
- self.logvar
|
||||
+ other.logvar,
|
||||
dim=[1, 2, 3],
|
||||
)
|
||||
|
||||
def nll(self, sample, dims=[1, 2, 3]):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
return 0.5 * torch.sum(
|
||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims,
|
||||
)
|
||||
|
||||
def mode(self):
|
||||
return self.mean
|
||||
|
||||
|
||||
def normal_kl(mean1, logvar1, mean2, logvar2):
|
||||
"""
|
||||
source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
|
||||
Compute the KL divergence between two gaussians.
|
||||
Shapes are automatically broadcasted, so batches can be compared to
|
||||
scalars, among other use cases.
|
||||
"""
|
||||
tensor = None
|
||||
for obj in (mean1, logvar1, mean2, logvar2):
|
||||
if isinstance(obj, torch.Tensor):
|
||||
tensor = obj
|
||||
break
|
||||
assert tensor is not None, "at least one argument must be a Tensor"
|
||||
|
||||
# Force variances to be Tensors. Broadcasting helps convert scalars to
|
||||
# Tensors, but it does not work for torch.exp().
|
||||
logvar1, logvar2 = [
|
||||
x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor)
|
||||
for x in (logvar1, logvar2)
|
||||
]
|
||||
|
||||
return 0.5 * (
|
||||
-1.0
|
||||
+ logvar2
|
||||
- logvar1
|
||||
+ torch.exp(logvar1 - logvar2)
|
||||
+ ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
|
||||
)
|
||||
@@ -0,0 +1,86 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class LitEma(nn.Module):
|
||||
def __init__(self, model, decay=0.9999, use_num_upates=True):
|
||||
super().__init__()
|
||||
if decay < 0.0 or decay > 1.0:
|
||||
raise ValueError("Decay must be between 0 and 1")
|
||||
|
||||
self.m_name2s_name = {}
|
||||
self.register_buffer("decay", torch.tensor(decay, dtype=torch.float32))
|
||||
self.register_buffer(
|
||||
"num_updates",
|
||||
torch.tensor(0, dtype=torch.int)
|
||||
if use_num_upates
|
||||
else torch.tensor(-1, dtype=torch.int),
|
||||
)
|
||||
|
||||
for name, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
# remove as '.'-character is not allowed in buffers
|
||||
s_name = name.replace(".", "")
|
||||
self.m_name2s_name.update({name: s_name})
|
||||
self.register_buffer(s_name, p.clone().detach().data)
|
||||
|
||||
self.collected_params = []
|
||||
|
||||
def reset_num_updates(self):
|
||||
del self.num_updates
|
||||
self.register_buffer("num_updates", torch.tensor(0, dtype=torch.int))
|
||||
|
||||
def forward(self, model):
|
||||
decay = self.decay
|
||||
|
||||
if self.num_updates >= 0:
|
||||
self.num_updates += 1
|
||||
decay = min(self.decay, (1 + self.num_updates) / (10 + self.num_updates))
|
||||
|
||||
one_minus_decay = 1.0 - decay
|
||||
|
||||
with torch.no_grad():
|
||||
m_param = dict(model.named_parameters())
|
||||
shadow_params = dict(self.named_buffers())
|
||||
|
||||
for key in m_param:
|
||||
if m_param[key].requires_grad:
|
||||
sname = self.m_name2s_name[key]
|
||||
shadow_params[sname] = shadow_params[sname].type_as(m_param[key])
|
||||
shadow_params[sname].sub_(
|
||||
one_minus_decay * (shadow_params[sname] - m_param[key])
|
||||
)
|
||||
else:
|
||||
assert not key in self.m_name2s_name
|
||||
|
||||
def copy_to(self, model):
|
||||
m_param = dict(model.named_parameters())
|
||||
shadow_params = dict(self.named_buffers())
|
||||
for key in m_param:
|
||||
if m_param[key].requires_grad:
|
||||
m_param[key].data.copy_(shadow_params[self.m_name2s_name[key]].data)
|
||||
else:
|
||||
assert not key in self.m_name2s_name
|
||||
|
||||
def store(self, parameters):
|
||||
"""
|
||||
Save the current parameters for restoring later.
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
temporarily stored.
|
||||
"""
|
||||
self.collected_params = [param.clone() for param in parameters]
|
||||
|
||||
def restore(self, parameters):
|
||||
"""
|
||||
Restore the parameters stored with the `store` method.
|
||||
Useful to validate the model with EMA parameters without affecting the
|
||||
original optimization process. Store the parameters before the
|
||||
`copy_to` method. After validation (or model saving), use this to
|
||||
restore the former parameters.
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters.
|
||||
"""
|
||||
for c_param, param in zip(self.collected_params, parameters):
|
||||
param.data.copy_(c_param.data)
|
||||
File diff suppressed because it is too large
Load Diff
+248
@@ -0,0 +1,248 @@
|
||||
import functools
|
||||
import importlib
|
||||
import os
|
||||
from functools import partial
|
||||
from inspect import isfunction
|
||||
|
||||
import fsspec
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
|
||||
|
||||
def disabled_train(self, mode=True):
|
||||
"""Overwrite model.train with this function to make sure train/eval mode
|
||||
does not change anymore."""
|
||||
return self
|
||||
|
||||
|
||||
def get_string_from_tuple(s):
|
||||
try:
|
||||
# Check if the string starts and ends with parentheses
|
||||
if s[0] == "(" and s[-1] == ")":
|
||||
# Convert the string to a tuple
|
||||
t = eval(s)
|
||||
# Check if the type of t is tuple
|
||||
if type(t) == tuple:
|
||||
return t[0]
|
||||
else:
|
||||
pass
|
||||
except:
|
||||
pass
|
||||
return s
|
||||
|
||||
|
||||
def is_power_of_two(n):
|
||||
"""
|
||||
chat.openai.com/chat
|
||||
Return True if n is a power of 2, otherwise return False.
|
||||
|
||||
The function is_power_of_two takes an integer n as input and returns True if n is a power of 2, otherwise it returns False.
|
||||
The function works by first checking if n is less than or equal to 0. If n is less than or equal to 0, it can't be a power of 2, so the function returns False.
|
||||
If n is greater than 0, the function checks whether n is a power of 2 by using a bitwise AND operation between n and n-1. If n is a power of 2, then it will have only one bit set to 1 in its binary representation. When we subtract 1 from a power of 2, all the bits to the right of that bit become 1, and the bit itself becomes 0. So, when we perform a bitwise AND between n and n-1, we get 0 if n is a power of 2, and a non-zero value otherwise.
|
||||
Thus, if the result of the bitwise AND operation is 0, then n is a power of 2 and the function returns True. Otherwise, the function returns False.
|
||||
|
||||
"""
|
||||
if n <= 0:
|
||||
return False
|
||||
return (n & (n - 1)) == 0
|
||||
|
||||
|
||||
def autocast(f, enabled=True):
|
||||
def do_autocast(*args, **kwargs):
|
||||
with torch.cuda.amp.autocast(
|
||||
enabled=enabled,
|
||||
dtype=torch.get_autocast_gpu_dtype(),
|
||||
cache_enabled=torch.is_autocast_cache_enabled(),
|
||||
):
|
||||
return f(*args, **kwargs)
|
||||
|
||||
return do_autocast
|
||||
|
||||
|
||||
def load_partial_from_config(config):
|
||||
return partial(get_obj_from_str(config["target"]), **config.get("params", dict()))
|
||||
|
||||
|
||||
def log_txt_as_img(wh, xc, size=10):
|
||||
# wh a tuple of (width, height)
|
||||
# xc a list of captions to plot
|
||||
b = len(xc)
|
||||
txts = list()
|
||||
for bi in range(b):
|
||||
txt = Image.new("RGB", wh, color="white")
|
||||
draw = ImageDraw.Draw(txt)
|
||||
font = ImageFont.truetype("data/DejaVuSans.ttf", size=size)
|
||||
nc = int(40 * (wh[0] / 256))
|
||||
if isinstance(xc[bi], list):
|
||||
text_seq = xc[bi][0]
|
||||
else:
|
||||
text_seq = xc[bi]
|
||||
lines = "\n".join(
|
||||
text_seq[start : start + nc] for start in range(0, len(text_seq), nc)
|
||||
)
|
||||
|
||||
try:
|
||||
draw.text((0, 0), lines, fill="black", font=font)
|
||||
except UnicodeEncodeError:
|
||||
print("Cant encode string for logging. Skipping.")
|
||||
|
||||
txt = np.array(txt).transpose(2, 0, 1) / 127.5 - 1.0
|
||||
txts.append(txt)
|
||||
txts = np.stack(txts)
|
||||
txts = torch.tensor(txts)
|
||||
return txts
|
||||
|
||||
|
||||
def partialclass(cls, *args, **kwargs):
|
||||
class NewCls(cls):
|
||||
__init__ = functools.partialmethod(cls.__init__, *args, **kwargs)
|
||||
|
||||
return NewCls
|
||||
|
||||
|
||||
def make_path_absolute(path):
|
||||
fs, p = fsspec.core.url_to_fs(path)
|
||||
if fs.protocol == "file":
|
||||
return os.path.abspath(p)
|
||||
return path
|
||||
|
||||
|
||||
def ismap(x):
|
||||
if not isinstance(x, torch.Tensor):
|
||||
return False
|
||||
return (len(x.shape) == 4) and (x.shape[1] > 3)
|
||||
|
||||
|
||||
def isimage(x):
|
||||
if not isinstance(x, torch.Tensor):
|
||||
return False
|
||||
return (len(x.shape) == 4) and (x.shape[1] == 3 or x.shape[1] == 1)
|
||||
|
||||
|
||||
def isheatmap(x):
|
||||
if not isinstance(x, torch.Tensor):
|
||||
return False
|
||||
|
||||
return x.ndim == 2
|
||||
|
||||
|
||||
def isneighbors(x):
|
||||
if not isinstance(x, torch.Tensor):
|
||||
return False
|
||||
return x.ndim == 5 and (x.shape[2] == 3 or x.shape[2] == 1)
|
||||
|
||||
|
||||
def exists(x):
|
||||
return x is not None
|
||||
|
||||
|
||||
def expand_dims_like(x, y):
|
||||
while x.dim() != y.dim():
|
||||
x = x.unsqueeze(-1)
|
||||
return x
|
||||
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if isfunction(d) else d
|
||||
|
||||
|
||||
def mean_flat(tensor):
|
||||
"""
|
||||
https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/nn.py#L86
|
||||
Take the mean over all non-batch dimensions.
|
||||
"""
|
||||
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
||||
|
||||
|
||||
def count_params(model, verbose=False):
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
if verbose:
|
||||
print(f"{model.__class__.__name__} has {total_params * 1.e-6:.2f} M params.")
|
||||
return total_params
|
||||
|
||||
|
||||
def instantiate_from_config(config):
|
||||
if not "target" in config:
|
||||
if config == "__is_first_stage__":
|
||||
return None
|
||||
elif config == "__is_unconditional__":
|
||||
return None
|
||||
raise KeyError("Expected key `target` to instantiate.")
|
||||
return get_obj_from_str(config["target"])(**config.get("params", dict()))
|
||||
|
||||
|
||||
def get_obj_from_str(string, reload=False, invalidate_cache=True):
|
||||
module, cls = string.rsplit(".", 1)
|
||||
if invalidate_cache:
|
||||
importlib.invalidate_caches()
|
||||
if reload:
|
||||
module_imp = importlib.import_module(module)
|
||||
importlib.reload(module_imp)
|
||||
return getattr(importlib.import_module(module, package='ComfyUI-SUPIR'), cls)
|
||||
|
||||
|
||||
def append_zero(x):
|
||||
return torch.cat([x, x.new_zeros([1])])
|
||||
|
||||
|
||||
def append_dims(x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(
|
||||
f"input has {x.ndim} dims but target_dims is {target_dims}, which is less"
|
||||
)
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
def load_model_from_config(config, ckpt, verbose=True, freeze=True):
|
||||
print(f"Loading model from {ckpt}")
|
||||
if ckpt.endswith("ckpt"):
|
||||
pl_sd = torch.load(ckpt, map_location="cpu")
|
||||
if "global_step" in pl_sd:
|
||||
print(f"Global Step: {pl_sd['global_step']}")
|
||||
sd = pl_sd["state_dict"]
|
||||
elif ckpt.endswith("safetensors"):
|
||||
sd = load_safetensors(ckpt)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
model = instantiate_from_config(config.model)
|
||||
|
||||
m, u = model.load_state_dict(sd, strict=False)
|
||||
|
||||
if len(m) > 0 and verbose:
|
||||
print("missing keys:")
|
||||
print(m)
|
||||
if len(u) > 0 and verbose:
|
||||
print("unexpected keys:")
|
||||
print(u)
|
||||
|
||||
if freeze:
|
||||
for param in model.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
def get_configs_path() -> str:
|
||||
"""
|
||||
Get the `configs` directory.
|
||||
For a working copy, this is the one in the root of the repository,
|
||||
but for an installed copy, it's in the `sgm` package (see pyproject.toml).
|
||||
"""
|
||||
this_dir = os.path.dirname(__file__)
|
||||
candidates = (
|
||||
os.path.join(this_dir, "configs"),
|
||||
os.path.join(this_dir, "..", "configs"),
|
||||
)
|
||||
for candidate in candidates:
|
||||
candidate = os.path.abspath(candidate)
|
||||
if os.path.isdir(candidate):
|
||||
return candidate
|
||||
raise FileNotFoundError(f"Could not find SGM configs in {candidates}")
|
||||
Reference in New Issue
Block a user