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:
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
+25
-19
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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!")
|
||||
Reference in New Issue
Block a user