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
This commit is contained in:
City
2023-11-06 23:19:42 +01:00
parent 5a27b5b1b8
commit b7a293cc27
4 changed files with 168 additions and 19 deletions
+6
View File
@@ -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,
}
}
+15 -9
View File
@@ -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"]
self.device = model_management.vae_device()
self.offload_device = model_management.vae_offload_device()
self.vae_dtype = vae_dtype_dict.get(dtype, "auto")
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.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)
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
+21
View File
@@ -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.
+116
View File
@@ -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!")