From b7a293cc27c1b2862a79f28c44e8604cb0daf08d Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Mon, 6 Nov 2023 23:19:42 +0100 Subject: [PATCH] Consistency decoder initial Quick and lazy, just ported the example code. Can't do much more without having the model arch. https://github.com/openai/consistencydecoder/issues/1 --- VAE/conf.py | 6 ++ VAE/loader.py | 44 ++++++---- VAE/models/LICENSE-Consistency-Decoder | 21 +++++ VAE/models/consistencydecoder.py | 116 +++++++++++++++++++++++++ 4 files changed, 168 insertions(+), 19 deletions(-) create mode 100644 VAE/models/LICENSE-Consistency-Decoder create mode 100644 VAE/models/consistencydecoder.py diff --git a/VAE/conf.py b/VAE/conf.py index c45e142..d81c9a5 100644 --- a/VAE/conf.py +++ b/VAE/conf.py @@ -105,4 +105,10 @@ vae_conf = { "num_res_blocks" : 2, "attn_resolutions" : [16], }, + # OpenAI Consistency Decoder + "Consistency-Decoder": { + "type" : "ConsistencyDecoder", + "embed_scale" : 8, + "embed_dim" : 4, + } } diff --git a/VAE/loader.py b/VAE/loader.py index f9e842a..b6031ff 100644 --- a/VAE/loader.py +++ b/VAE/loader.py @@ -13,29 +13,36 @@ vae_dtype_dict = { class EXVAE(comfy.sd.VAE): def __init__(self, model_path, model_conf, dtype=None): - sd = comfy.utils.load_torch_file(model_path) - if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format - sd = diffusers_convert.convert_vae_state_dict(sd) - self.latent_dim = model_conf["embed_dim"] self.latent_scale = model_conf["embed_scale"] - - if model_conf["type"] == "AutoencoderKL": - from .models.kl import AutoencoderKL - model = AutoencoderKL(config=model_conf) - if model_conf["type"] == "VQModel": - from .models.vq import VQModel - model = VQModel(config=model_conf) - - 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.device = model_management.vae_device() self.offload_device = model_management.vae_offload_device() self.vae_dtype = vae_dtype_dict.get(dtype, "auto") - self.first_stage_model.to(self.vae_dtype) + + sd = None + 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) + + self.first_stage_model.to(self.vae_dtype).to(self.offload_device) ### Encode/Decode functions below needed due to source repo having 4 VAE channels and a scale factor of 8 hardcoded def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap = 16): @@ -106,4 +113,3 @@ class EXVAE(comfy.sd.VAE): self.first_stage_model = self.first_stage_model.to(self.offload_device) return samples - diff --git a/VAE/models/LICENSE-Consistency-Decoder b/VAE/models/LICENSE-Consistency-Decoder new file mode 100644 index 0000000..b3841f6 --- /dev/null +++ b/VAE/models/LICENSE-Consistency-Decoder @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2023 OpenAI + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. 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. diff --git a/VAE/models/consistencydecoder.py b/VAE/models/consistencydecoder.py new file mode 100644 index 0000000..8b37afe --- /dev/null +++ b/VAE/models/consistencydecoder.py @@ -0,0 +1,116 @@ +import math +import torch + +def _extract_into_tensor(arr, timesteps, broadcast_shape): + # from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """ + res = arr[timesteps.to(torch.int).cpu()].float().to(timesteps.device) + dims_to_append = len(broadcast_shape) - len(res.shape) + return res[(...,) + (None,) * dims_to_append] + +def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999): + # from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L45 + betas = [] + for i in range(num_diffusion_timesteps): + t1 = i / num_diffusion_timesteps + t2 = (i + 1) / num_diffusion_timesteps + betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta)) + return torch.tensor(betas) + +class ConsistencyDecoder(torch.nn.Module): + # From https://github.com/openai/consistencydecoder + def __init__(self, model_path, device): + super().__init__() + self.ckpt = torch.jit.load(model_path, map_location=device) + self.n_distilled_steps = 64 + + sigma_data = 0.5 + betas = betas_for_alpha_bar( + 1024, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2 + ) + alphas = 1.0 - betas + alphas_cumprod = torch.cumprod(alphas, dim=0) + self.sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod) + self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod) + sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod) + sigmas = torch.sqrt(1.0 / alphas_cumprod - 1) + self.c_skip = ( + sqrt_recip_alphas_cumprod + * sigma_data**2 + / (sigmas**2 + sigma_data**2) + ) + self.c_out = sigmas * sigma_data / (sigmas**2 + sigma_data**2) ** 0.5 + self.c_in = sqrt_recip_alphas_cumprod / (sigmas**2 + sigma_data**2) ** 0.5 + + @staticmethod + def round_timesteps(timesteps, total_timesteps, n_distilled_steps, truncate_start=True): + with torch.no_grad(): + space = torch.div(total_timesteps, n_distilled_steps, rounding_mode="floor") + rounded_timesteps = ( + torch.div(timesteps, space, rounding_mode="floor") + 1 + ) * space + if truncate_start: + rounded_timesteps[rounded_timesteps == total_timesteps] -= space + else: + rounded_timesteps[rounded_timesteps == total_timesteps] -= space + rounded_timesteps[rounded_timesteps == 0] += space + return rounded_timesteps + + @staticmethod + def ldm_transform_latent(z, extra_scale_factor=1): + channel_means = [0.38862467, 0.02253063, 0.07381133, -0.0171294] + channel_stds = [0.9654121, 1.0440036, 0.76147926, 0.77022034] + + if len(z.shape) != 4: + raise ValueError() + + z = z * 0.18215 + channels = [z[:, i] for i in range(z.shape[1])] + + channels = [ + extra_scale_factor * (c - channel_means[i]) / channel_stds[i] + for i, c in enumerate(channels) + ] + return torch.stack(channels, dim=1) + + @torch.no_grad() + def decode(self, features: torch.Tensor, schedule=[1.0, 0.5]): + features = self.ldm_transform_latent(features) + ts = self.round_timesteps( + torch.arange(0, 1024), + 1024, + self.n_distilled_steps, + truncate_start=False, + ) + shape = ( + features.size(0), + 3, + 8 * features.size(2), + 8 * features.size(3), + ) + x_start = torch.zeros(shape, device=features.device, dtype=features.dtype) + 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) + x_start = ( + _extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape) + * x_start + + _extract_into_tensor( + self.sqrt_one_minus_alphas_cumprod, t_, x_start.shape + ) + * noise + ) + c_in = _extract_into_tensor(self.c_in, t_, x_start.shape) + model_output = self.ckpt(c_in * x_start, t_, features=features) + B, C = x_start.shape[:2] + model_output, _ = torch.split(model_output, C, dim=1) + pred_xstart = ( + _extract_into_tensor(self.c_out, t_, x_start.shape) * model_output + + _extract_into_tensor(self.c_skip, t_, x_start.shape) * x_start + ).clamp(-1, 1) + x_start = pred_xstart + return x_start + + def encode(self, *args, **kwargs): + raise NotImplementedError("ConsistencyDecoder can't be used for encoding!")