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,
|
"num_res_blocks" : 2,
|
||||||
"attn_resolutions" : [16],
|
"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):
|
class EXVAE(comfy.sd.VAE):
|
||||||
def __init__(self, model_path, model_conf, dtype=None):
|
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_dim = model_conf["embed_dim"]
|
||||||
self.latent_scale = model_conf["embed_scale"]
|
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.device = model_management.vae_device()
|
||||||
self.offload_device = model_management.vae_offload_device()
|
self.offload_device = model_management.vae_offload_device()
|
||||||
self.vae_dtype = vae_dtype_dict.get(dtype, "auto")
|
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
|
### 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):
|
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)
|
self.first_stage_model = self.first_stage_model.to(self.offload_device)
|
||||||
return samples
|
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