From a3e1804defe2faf91a7050cd2083475dea687270 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Tue, 7 Nov 2023 21:26:54 +0100 Subject: [PATCH] Consistency decoder full With model definitions from https://gist.github.com/madebyollin/865fa6a18d9099351ddbdfbe7299ccbf --- README.md | 7 +- VAE/loader.py | 31 ++-- VAE/models/consistencydecoder.py | 270 ++++++++++++++++++++++++++++++- 3 files changed, 284 insertions(+), 24 deletions(-) diff --git a/README.md b/README.md index 76fe833..944b208 100644 --- a/README.md +++ b/README.md @@ -121,12 +121,13 @@ A few custom VAE models are supported. The option to select a different dtype wh [Original Repo](https://github.com/openai/consistencydecoder) -Proof of concept until [the model definitions are released](https://github.com/openai/consistencydecoder/issues/1) +This now works thanks to the work of @mrsteyk and @madebyollin - [Gist with more info](https://gist.github.com/madebyollin/865fa6a18d9099351ddbdfbe7299ccbf). -- Download the VAE from [the link in the OpenAI code](https://github.com/openai/consistencydecoder/blob/main/consistencydecoder/__init__.py#L79) / [Direct link](https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt) +- Download the converted safetensor VAE from [this HF repository](https://huggingface.co/mrsteyk/consistency-decoder-sd15/blob/main/stk_consistency_decoder_amalgamated.safetensors). If you downloaded the OpenAI model before, it won't work, as it is a TorchScript file. Feel free to delete it. - Put the file in your VAE folder - Load it with the ExtraVAELoader -- Run out of VRAM +- Set it to fp16 or bf16 to not run out of VRAM +- Use tiled VAE decode if required ### AutoencoderKL / VQModel diff --git a/VAE/loader.py b/VAE/loader.py index b6031ff..ebb87a1 100644 --- a/VAE/loader.py +++ b/VAE/loader.py @@ -19,28 +19,27 @@ class EXVAE(comfy.sd.VAE): self.offload_device = model_management.vae_offload_device() self.vae_dtype = vae_dtype_dict.get(dtype, "auto") - sd = None + sd = comfy.utils.load_torch_file(model_path) model = None if model_conf["type"] == "AutoencoderKL": from .models.kl import AutoencoderKL model = AutoencoderKL(config=model_conf) - sd = comfy.utils.load_torch_file(model_path) - if model_conf["type"] == "VQModel": - from .models.vq import VQModel - model = VQModel(config=model_conf) - sd = comfy.utils.load_torch_file(model_path) - if model_conf["type"] == "ConsistencyDecoder": - from .models.consistencydecoder import ConsistencyDecoder - model = ConsistencyDecoder(model_path, self.device) - sd = model.ckpt.state_dict() - - if sd: if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): sd = diffusers_convert.convert_vae_state_dict(sd) - self.first_stage_model = model.eval() - m, u = self.first_stage_model.load_state_dict(sd, strict=False) - if len(m) > 0: print("Missing VAE keys", m) - if len(u) > 0: print("Leftover VAE keys", u) + elif model_conf["type"] == "VQModel": + from .models.vq import VQModel + model = VQModel(config=model_conf) + elif model_conf["type"] == "ConsistencyDecoder": + from .models.consistencydecoder import ConsistencyDecoder + model = ConsistencyDecoder() + sd = {f"model.{k}":v for k,v in sd.items()} + else: + raise NotImplementedError(f"Unknown VAE type '{model_conf['type']}'") + + self.first_stage_model = model.eval() + m, u = self.first_stage_model.load_state_dict(sd, strict=False) + if len(m) > 0: print("Missing VAE keys", m) + if len(u) > 0: print("Leftover VAE keys", u) self.first_stage_model.to(self.vae_dtype).to(self.offload_device) diff --git a/VAE/models/consistencydecoder.py b/VAE/models/consistencydecoder.py index 8b37afe..4dc1902 100644 --- a/VAE/models/consistencydecoder.py +++ b/VAE/models/consistencydecoder.py @@ -1,5 +1,11 @@ import math import torch +import torch.nn.functional as F +import torch.nn as nn + +""" +Code below ported from https://github.com/openai/consistencydecoder +""" def _extract_into_tensor(arr, timesteps, broadcast_shape): # from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """ @@ -18,9 +24,9 @@ def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999): class ConsistencyDecoder(torch.nn.Module): # From https://github.com/openai/consistencydecoder - def __init__(self, model_path, device): + def __init__(self): super().__init__() - self.ckpt = torch.jit.load(model_path, map_location=device) + self.model = ConvUNetVAE() self.n_distilled_steps = 64 sigma_data = 0.5 @@ -91,8 +97,8 @@ class ConsistencyDecoder(torch.nn.Module): schedule_timesteps = [int((1024 - 1) * s) for s in schedule] for i in schedule_timesteps: t = ts[i].item() - t_ = torch.tensor([t] * features.shape[0]).to(features.device) - noise = torch.randn_like(x_start) + t_ = torch.tensor([t] * features.shape[0], device=features.device) + noise = torch.randn_like(x_start, device=features.device) x_start = ( _extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape) * x_start @@ -102,7 +108,7 @@ class ConsistencyDecoder(torch.nn.Module): * noise ) c_in = _extract_into_tensor(self.c_in, t_, x_start.shape) - model_output = self.ckpt(c_in * x_start, t_, features=features) + model_output = self.model((c_in * x_start).to(features.dtype), t_, features=features) B, C = x_start.shape[:2] model_output, _ = torch.split(model_output, C, dim=1) pred_xstart = ( @@ -114,3 +120,257 @@ class ConsistencyDecoder(torch.nn.Module): def encode(self, *args, **kwargs): raise NotImplementedError("ConsistencyDecoder can't be used for encoding!") + +""" +Model definitions ported from: +https://gist.github.com/madebyollin/865fa6a18d9099351ddbdfbe7299ccbf +https://gist.github.com/mrsteyk/74ad3ec2f6f823111ae4c90e168505ac. +""" + +class TimestepEmbedding(nn.Module): + def __init__(self, n_time=1024, n_emb=320, n_out=1280) -> None: + super().__init__() + self.emb = nn.Embedding(n_time, n_emb) + self.f_1 = nn.Linear(n_emb, n_out) + self.f_2 = nn.Linear(n_out, n_out) + + def forward(self, x) -> torch.Tensor: + x = self.emb(x) + x = self.f_1(x) + x = F.silu(x) + return self.f_2(x) + + +class ImageEmbedding(nn.Module): + def __init__(self, in_channels=7, out_channels=320) -> None: + super().__init__() + self.f = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1) + + def forward(self, x) -> torch.Tensor: + return self.f(x) + + +class ImageUnembedding(nn.Module): + def __init__(self, in_channels=320, out_channels=6) -> None: + super().__init__() + self.gn = nn.GroupNorm(32, in_channels) + self.f = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1) + + def forward(self, x) -> torch.Tensor: + return self.f(F.silu(self.gn(x))) + + +class ConvResblock(nn.Module): + def __init__(self, in_features=320, out_features=320) -> None: + super().__init__() + self.f_t = nn.Linear(1280, out_features * 2) + + self.gn_1 = nn.GroupNorm(32, in_features) + self.f_1 = nn.Conv2d(in_features, out_features, kernel_size=3, padding=1) + + self.gn_2 = nn.GroupNorm(32, out_features) + self.f_2 = nn.Conv2d(out_features, out_features, kernel_size=3, padding=1) + + skip_conv = in_features != out_features + self.f_s = ( + nn.Conv2d(in_features, out_features, kernel_size=1, padding=0) + if skip_conv + else nn.Identity() + ) + + def forward(self, x, t): + x_skip = x + t = self.f_t(F.silu(t)) + t = t.chunk(2, dim=1) + t_1 = t[0].unsqueeze(dim=2).unsqueeze(dim=3) + 1 + t_2 = t[1].unsqueeze(dim=2).unsqueeze(dim=3) + + gn_1 = F.silu(self.gn_1(x)) + f_1 = self.f_1(gn_1) + + gn_2 = self.gn_2(f_1) + + return self.f_s(x_skip) + self.f_2(F.silu(gn_2 * t_1 + t_2)) + + +# Also ConvResblock +class Downsample(nn.Module): + def __init__(self, in_channels=320) -> None: + super().__init__() + self.f_t = nn.Linear(1280, in_channels * 2) + + self.gn_1 = nn.GroupNorm(32, in_channels) + self.f_1 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) + self.gn_2 = nn.GroupNorm(32, in_channels) + + self.f_2 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) + + def forward(self, x, t) -> torch.Tensor: + x_skip = x + + t = self.f_t(F.silu(t)) + t_1, t_2 = t.chunk(2, dim=1) + t_1 = t_1.unsqueeze(2).unsqueeze(3) + 1 + t_2 = t_2.unsqueeze(2).unsqueeze(3) + + gn_1 = F.silu(self.gn_1(x)) + avg_pool2d = F.avg_pool2d(gn_1, kernel_size=(2, 2), stride=None) + f_1 = self.f_1(avg_pool2d) + gn_2 = self.gn_2(f_1) + + f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) + + return f_2 + F.avg_pool2d(x_skip, kernel_size=(2, 2), stride=None) + + +# Also ConvResblock +class Upsample(nn.Module): + def __init__(self, in_channels=1024) -> None: + super().__init__() + self.f_t = nn.Linear(1280, in_channels * 2) + + self.gn_1 = nn.GroupNorm(32, in_channels) + self.f_1 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) + self.gn_2 = nn.GroupNorm(32, in_channels) + + self.f_2 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) + + def forward(self, x, t) -> torch.Tensor: + x_skip = x + + t = self.f_t(F.silu(t)) + t_1, t_2 = t.chunk(2, dim=1) + t_1 = t_1.unsqueeze(2).unsqueeze(3) + 1 + t_2 = t_2.unsqueeze(2).unsqueeze(3) + + gn_1 = F.silu(self.gn_1(x)) + upsample = F.interpolate(gn_1.float(), scale_factor=2, mode="nearest").to(gn_1.dtype) + + f_1 = self.f_1(upsample) + gn_2 = self.gn_2(f_1) + + f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) + + return f_2 + F.interpolate(x_skip.float(), scale_factor=2, mode="nearest").to(x_skip.dtype) + + +class ConvUNetVAE(nn.Module): + def __init__(self) -> None: + super().__init__() + self.embed_image = ImageEmbedding() + self.embed_time = TimestepEmbedding() + + down_0 = nn.ModuleList( + [ + ConvResblock(320, 320), + ConvResblock(320, 320), + ConvResblock(320, 320), + Downsample(320), + ] + ) + down_1 = nn.ModuleList( + [ + ConvResblock(320, 640), + ConvResblock(640, 640), + ConvResblock(640, 640), + Downsample(640), + ] + ) + down_2 = nn.ModuleList( + [ + ConvResblock(640, 1024), + ConvResblock(1024, 1024), + ConvResblock(1024, 1024), + Downsample(1024), + ] + ) + down_3 = nn.ModuleList( + [ + ConvResblock(1024, 1024), + ConvResblock(1024, 1024), + ConvResblock(1024, 1024), + ] + ) + self.down = nn.ModuleList( + [ + down_0, + down_1, + down_2, + down_3, + ] + ) + + self.mid = nn.ModuleList( + [ + ConvResblock(1024, 1024), + ConvResblock(1024, 1024), + ] + ) + + up_3 = nn.ModuleList( + [ + ConvResblock(1024 * 2, 1024), + ConvResblock(1024 * 2, 1024), + ConvResblock(1024 * 2, 1024), + ConvResblock(1024 * 2, 1024), + Upsample(1024), + ] + ) + up_2 = nn.ModuleList( + [ + ConvResblock(1024 * 2, 1024), + ConvResblock(1024 * 2, 1024), + ConvResblock(1024 * 2, 1024), + ConvResblock(1024 + 640, 1024), + Upsample(1024), + ] + ) + up_1 = nn.ModuleList( + [ + ConvResblock(1024 + 640, 640), + ConvResblock(640 * 2, 640), + ConvResblock(640 * 2, 640), + ConvResblock(320 + 640, 640), + Upsample(640), + ] + ) + up_0 = nn.ModuleList( + [ + ConvResblock(320 + 640, 320), + ConvResblock(320 * 2, 320), + ConvResblock(320 * 2, 320), + ConvResblock(320 * 2, 320), + ] + ) + self.up = nn.ModuleList( + [ + up_0, + up_1, + up_2, + up_3, + ] + ) + + self.output = ImageUnembedding() + + def forward(self, x, t, features) -> torch.Tensor: + x = torch.cat([x, F.interpolate(features.float(),scale_factor=8,mode="nearest").to(features.dtype)], dim=1) + t = self.embed_time(t) + x = self.embed_image(x) + + skips = [x] + for down in self.down: + for block in down: + x = block(x, t) + skips.append(x) + + for i in range(2): + x = self.mid[i](x, t) + + for up in self.up[::-1]: + for block in up: + if isinstance(block, ConvResblock): + x = torch.concat([x, skips.pop()], dim=1) + x = block(x, t) + + return self.output(x)