Consistency decoder full
With model definitions from https://gist.github.com/madebyollin/865fa6a18d9099351ddbdfbe7299ccbf
This commit is contained in:
@@ -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)
|
[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
|
- Put the file in your VAE folder
|
||||||
- Load it with the ExtraVAELoader
|
- 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
|
### AutoencoderKL / VQModel
|
||||||
|
|
||||||
|
|||||||
+15
-16
@@ -19,28 +19,27 @@ class EXVAE(comfy.sd.VAE):
|
|||||||
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")
|
||||||
|
|
||||||
sd = None
|
sd = comfy.utils.load_torch_file(model_path)
|
||||||
model = None
|
model = None
|
||||||
if model_conf["type"] == "AutoencoderKL":
|
if model_conf["type"] == "AutoencoderKL":
|
||||||
from .models.kl import AutoencoderKL
|
from .models.kl import AutoencoderKL
|
||||||
model = AutoencoderKL(config=model_conf)
|
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():
|
if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys():
|
||||||
sd = diffusers_convert.convert_vae_state_dict(sd)
|
sd = diffusers_convert.convert_vae_state_dict(sd)
|
||||||
self.first_stage_model = model.eval()
|
elif model_conf["type"] == "VQModel":
|
||||||
m, u = self.first_stage_model.load_state_dict(sd, strict=False)
|
from .models.vq import VQModel
|
||||||
if len(m) > 0: print("Missing VAE keys", m)
|
model = VQModel(config=model_conf)
|
||||||
if len(u) > 0: print("Leftover VAE keys", u)
|
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)
|
self.first_stage_model.to(self.vae_dtype).to(self.offload_device)
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,11 @@
|
|||||||
import math
|
import math
|
||||||
import torch
|
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):
|
def _extract_into_tensor(arr, timesteps, broadcast_shape):
|
||||||
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """
|
# 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):
|
class ConsistencyDecoder(torch.nn.Module):
|
||||||
# From https://github.com/openai/consistencydecoder
|
# From https://github.com/openai/consistencydecoder
|
||||||
def __init__(self, model_path, device):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.ckpt = torch.jit.load(model_path, map_location=device)
|
self.model = ConvUNetVAE()
|
||||||
self.n_distilled_steps = 64
|
self.n_distilled_steps = 64
|
||||||
|
|
||||||
sigma_data = 0.5
|
sigma_data = 0.5
|
||||||
@@ -91,8 +97,8 @@ class ConsistencyDecoder(torch.nn.Module):
|
|||||||
schedule_timesteps = [int((1024 - 1) * s) for s in schedule]
|
schedule_timesteps = [int((1024 - 1) * s) for s in schedule]
|
||||||
for i in schedule_timesteps:
|
for i in schedule_timesteps:
|
||||||
t = ts[i].item()
|
t = ts[i].item()
|
||||||
t_ = torch.tensor([t] * features.shape[0]).to(features.device)
|
t_ = torch.tensor([t] * features.shape[0], device=features.device)
|
||||||
noise = torch.randn_like(x_start)
|
noise = torch.randn_like(x_start, device=features.device)
|
||||||
x_start = (
|
x_start = (
|
||||||
_extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape)
|
_extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape)
|
||||||
* x_start
|
* x_start
|
||||||
@@ -102,7 +108,7 @@ class ConsistencyDecoder(torch.nn.Module):
|
|||||||
* noise
|
* noise
|
||||||
)
|
)
|
||||||
c_in = _extract_into_tensor(self.c_in, t_, x_start.shape)
|
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]
|
B, C = x_start.shape[:2]
|
||||||
model_output, _ = torch.split(model_output, C, dim=1)
|
model_output, _ = torch.split(model_output, C, dim=1)
|
||||||
pred_xstart = (
|
pred_xstart = (
|
||||||
@@ -114,3 +120,257 @@ class ConsistencyDecoder(torch.nn.Module):
|
|||||||
|
|
||||||
def encode(self, *args, **kwargs):
|
def encode(self, *args, **kwargs):
|
||||||
raise NotImplementedError("ConsistencyDecoder can't be used for encoding!")
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user