Add SD3
This commit is contained in:
@@ -12,6 +12,7 @@ test.py
|
|||||||
*.pth
|
*.pth
|
||||||
*.ckpt
|
*.ckpt
|
||||||
*.safetensors
|
*.safetensors
|
||||||
|
preprocess_*
|
||||||
|
|
||||||
# default github .gitignore follows
|
# default github .gitignore follows
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ config = {
|
|||||||
"xl-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12},
|
"xl-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12},
|
||||||
"ca-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12},
|
"ca-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12},
|
||||||
"ca-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12},
|
"ca-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12},
|
||||||
|
"v3-to-v1": {"ch_in":16, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12},
|
||||||
|
"v3-to-xl": {"ch_in":16, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12},
|
||||||
}
|
}
|
||||||
|
|
||||||
class ResBlock(nn.Module):
|
class ResBlock(nn.Module):
|
||||||
@@ -89,8 +91,8 @@ class ComfyLatentInterposer:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"samples": ("LATENT", ),
|
"samples": ("LATENT", ),
|
||||||
"latent_src": (["v1", "xl", "ca"],),
|
"latent_src": (["v1", "xl", "v3", "ca"],),
|
||||||
"latent_dst": (["v1", "xl", "ca"],),
|
"latent_dst": (["v1", "xl"],),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -144,7 +146,9 @@ class ComfyLatentInterposer:
|
|||||||
lt = samples["samples"]
|
lt = samples["samples"]
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
# force FP32, always run on CPU
|
# force FP32, always run on CPU
|
||||||
lt = self.model(lt.cpu().float()).to(lt.device).to(lt.dtype)
|
lt = self.model(
|
||||||
|
lt.cpu().float()
|
||||||
|
).to(lt.device).to(lt.dtype)
|
||||||
samples["samples"] = lt
|
samples["samples"] = lt
|
||||||
return (samples,)
|
return (samples,)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,40 @@
|
|||||||
|
steps: 20000
|
||||||
|
batch: 48
|
||||||
|
fconst: 0
|
||||||
|
cosine: False
|
||||||
|
resume: False
|
||||||
|
device: "cuda"
|
||||||
|
p_loss_weight: 1.0
|
||||||
|
r_loss_weight: 1.4
|
||||||
|
b_loss_weight: 1.0
|
||||||
|
h_loss_weight: 0.0
|
||||||
|
save_image: 100
|
||||||
|
eval_model: 10
|
||||||
|
|
||||||
|
model:
|
||||||
|
src: v3 # Stable Diffusion Version three point oh
|
||||||
|
dst: v1 # Stable Diffusion 1.x
|
||||||
|
rev: "v4.0-rc1"
|
||||||
|
args:
|
||||||
|
scale: 1.0
|
||||||
|
ch_in: 16
|
||||||
|
ch_out: 4
|
||||||
|
ch_mid: 64
|
||||||
|
blocks: 12
|
||||||
|
|
||||||
|
optim:
|
||||||
|
lr: 5.0e-4
|
||||||
|
beta1: 0.5
|
||||||
|
beta2: 0.95
|
||||||
|
|
||||||
|
dataset:
|
||||||
|
src: "./latents/v3_256px_combined.bin"
|
||||||
|
dst: "./latents/v1_256px_combined.bin"
|
||||||
|
preload: False
|
||||||
|
evals:
|
||||||
|
main:
|
||||||
|
src: "./latents/test_eru/test_v3_768px.npy"
|
||||||
|
dst: "./latents/test_eru/test_v1_768px.npy"
|
||||||
|
aux:
|
||||||
|
src: "./latents/test_bga/test_v3_768px.npy"
|
||||||
|
dst: "./latents/test_bga/test_v1_768px.npy"
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
steps: 20000
|
||||||
|
batch: 48
|
||||||
|
fconst: 0
|
||||||
|
cosine: False
|
||||||
|
resume: False
|
||||||
|
device: "cuda"
|
||||||
|
p_loss_weight: 1.0
|
||||||
|
r_loss_weight: 1.4
|
||||||
|
b_loss_weight: 1.0
|
||||||
|
h_loss_weight: 0.0
|
||||||
|
save_image: 100
|
||||||
|
eval_model: 10
|
||||||
|
|
||||||
|
model:
|
||||||
|
src: v3 # Stable Diffusion Version three point oh
|
||||||
|
dst: xl # Stable Diffusion Extra Large
|
||||||
|
rev: "v4.0-rc1"
|
||||||
|
args:
|
||||||
|
scale: 1.0
|
||||||
|
ch_in: 16
|
||||||
|
ch_out: 4
|
||||||
|
ch_mid: 64
|
||||||
|
blocks: 12
|
||||||
|
|
||||||
|
optim:
|
||||||
|
lr: 5.0e-4
|
||||||
|
beta1: 0.5
|
||||||
|
beta2: 0.95
|
||||||
|
|
||||||
|
dataset:
|
||||||
|
src: "./latents/v3_256px_combined.bin"
|
||||||
|
dst: "./latents/xl_256px_combined.bin"
|
||||||
|
preload: False
|
||||||
|
evals:
|
||||||
|
main:
|
||||||
|
src: "./latents/test_eru/test_v3_768px.npy"
|
||||||
|
dst: "./latents/test_eru/test_xl_768px.npy"
|
||||||
|
aux:
|
||||||
|
src: "./latents/test_bga/test_v3_768px.npy"
|
||||||
|
dst: "./latents/test_bga/test_xl_768px.npy"
|
||||||
@@ -1,123 +0,0 @@
|
|||||||
#
|
|
||||||
# This file just has all the random saving/logging/eval related code
|
|
||||||
#
|
|
||||||
import os
|
|
||||||
import torch
|
|
||||||
from tqdm import tqdm
|
|
||||||
from diffusers import AutoencoderKL
|
|
||||||
from safetensors.torch import save_file
|
|
||||||
from torchvision.utils import save_image
|
|
||||||
|
|
||||||
LOSS_MEMORY = 500
|
|
||||||
LOG_EVERY_N = 500
|
|
||||||
SAVE_FOLDER = "models"
|
|
||||||
|
|
||||||
class ModelWrapper:
|
|
||||||
def __init__(self, name, specs, model, optimizer, criterion, scheduler, device="cpu", evals=[None,None], stdout=True):
|
|
||||||
self.name = name
|
|
||||||
self.specs = specs
|
|
||||||
self.losses = []
|
|
||||||
|
|
||||||
self.model = model
|
|
||||||
self.optimizer = optimizer
|
|
||||||
self.criterion = criterion
|
|
||||||
self.scheduler = scheduler
|
|
||||||
|
|
||||||
self.device = device
|
|
||||||
self.vae = self.get_vae(self.specs[1], fp16=True)
|
|
||||||
self.eval_src = evals[0]
|
|
||||||
self.eval_dst = evals[1]
|
|
||||||
|
|
||||||
os.makedirs(SAVE_FOLDER, exist_ok=True)
|
|
||||||
self.csvlog = open(f"{SAVE_FOLDER}/{self.name}.csv", "w")
|
|
||||||
self.stdout = stdout
|
|
||||||
|
|
||||||
def log_step(self, loss, step=None):
|
|
||||||
self.losses.append(loss)
|
|
||||||
step = step if step else len(self.losses)
|
|
||||||
if step % LOG_EVERY_N == 0:
|
|
||||||
self.log_main(step)
|
|
||||||
|
|
||||||
def log_main(self, step=None):
|
|
||||||
lr = float(self.scheduler.get_last_lr()[0])
|
|
||||||
avg = sum(self.losses[-LOSS_MEMORY:])/LOSS_MEMORY
|
|
||||||
evl = self.eval_model()[0]
|
|
||||||
if self.stdout:
|
|
||||||
tqdm.write(f"{str(step):<10} {avg:.4e}|{evl:.4e} @ {lr:.4e}")
|
|
||||||
if self.csvlog:
|
|
||||||
self.csvlog.write(f"{step},{avg},{evl},{lr}\n")
|
|
||||||
self.csvlog.flush()
|
|
||||||
|
|
||||||
def eval_model(self):
|
|
||||||
with torch.no_grad():
|
|
||||||
pred = self.model(self.eval_src.to(self.device))
|
|
||||||
loss = self.criterion(pred, self.eval_dst.to(self.device))
|
|
||||||
return loss, pred
|
|
||||||
|
|
||||||
def save_model(self, step=None, epoch=None):
|
|
||||||
step = step if step else len(self.losses)
|
|
||||||
if epoch is None and step >= 10**6:
|
|
||||||
epoch = f"_e{round(step/10**6,2)}M"
|
|
||||||
elif epoch is None:
|
|
||||||
epoch = f"_e{round(step/10**3)}K"
|
|
||||||
output_name = f"./{SAVE_FOLDER}/{self.name}{epoch}"
|
|
||||||
if self.vae:
|
|
||||||
out = self.eval_model()[1]
|
|
||||||
img = self.vae_decode(out).detach()
|
|
||||||
save_image(img, f"{output_name}.png")
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
save_file(self.model.state_dict(), f"{output_name}.safetensors")
|
|
||||||
torch.save(self.optimizer.state_dict(), f"{output_name}.optim.pth")
|
|
||||||
|
|
||||||
def close(self):
|
|
||||||
del self.vae
|
|
||||||
self.csvlog.close()
|
|
||||||
|
|
||||||
def vae_decode(self, latent):
|
|
||||||
latent = latent.to(torch.float16).to("cuda")
|
|
||||||
out = self.vae.decode(latent).sample
|
|
||||||
out = out.float().to(latent.device)
|
|
||||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
|
||||||
return ((out + 1.0) / 2.0)
|
|
||||||
|
|
||||||
def get_vae(self, version, file_path=None, fp16=False):
|
|
||||||
"""Load VAE from file or default hf repo. fp16 only works from hf"""
|
|
||||||
vae = None
|
|
||||||
dtype = torch.float16 if fp16 else torch.float32
|
|
||||||
if version == "v1" and file_path:
|
|
||||||
vae = AutoencoderKL.from_single_file(
|
|
||||||
file_path,
|
|
||||||
image_size=512,
|
|
||||||
)
|
|
||||||
elif version == "v1":
|
|
||||||
vae = AutoencoderKL.from_pretrained(
|
|
||||||
"runwayml/stable-diffusion-v1-5",
|
|
||||||
subfolder="vae",
|
|
||||||
torch_dtype=dtype,
|
|
||||||
)
|
|
||||||
elif version == "xl" and file_path:
|
|
||||||
vae = AutoencoderKL.from_single_file(
|
|
||||||
file_path,
|
|
||||||
image_size=1024
|
|
||||||
)
|
|
||||||
elif version == "xl" and fp16:
|
|
||||||
vae = AutoencoderKL.from_pretrained(
|
|
||||||
"madebyollin/sdxl-vae-fp16-fix",
|
|
||||||
torch_dtype=torch.float16,
|
|
||||||
)
|
|
||||||
elif version == "xl":
|
|
||||||
vae = AutoencoderKL.from_pretrained(
|
|
||||||
"stabilityai/stable-diffusion-xl-base-1.0",
|
|
||||||
subfolder="vae"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
raise NotImplementedError(f"Unknown VAE version '{version}'")
|
|
||||||
|
|
||||||
# save VRAM
|
|
||||||
vae.to(dtype).to("cuda")
|
|
||||||
vae.decoder.eval()
|
|
||||||
vae.set_use_memory_efficient_attention_xformers(True)
|
|
||||||
vae.enable_xformers_memory_efficient_attention()
|
|
||||||
vae.enable_gradient_checkpointing()
|
|
||||||
del vae.encoder
|
|
||||||
return vae
|
|
||||||
@@ -48,6 +48,20 @@ class SDXL_VAE(SDv1_VAE):
|
|||||||
if dec_only:
|
if dec_only:
|
||||||
del self.model.encoder
|
del self.model.encoder
|
||||||
|
|
||||||
|
class SDv3_VAE(SDv1_VAE):
|
||||||
|
scale = 1/8
|
||||||
|
channels = 16
|
||||||
|
def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False):
|
||||||
|
self.device = device
|
||||||
|
self.dtype = dtype
|
||||||
|
self.model = AutoencoderKL.from_pretrained(
|
||||||
|
"stabilityai/stable-diffusion-3-medium-diffusers",
|
||||||
|
subfolder="vae"
|
||||||
|
)
|
||||||
|
self.model.eval().to(self.dtype).to(self.device)
|
||||||
|
if dec_only:
|
||||||
|
del self.model.encoder
|
||||||
|
|
||||||
class CascadeC_VAE(SDv1_VAE):
|
class CascadeC_VAE(SDv1_VAE):
|
||||||
scale = 1/32
|
scale = 1/32
|
||||||
channels = 16
|
channels = 16
|
||||||
@@ -75,7 +89,7 @@ class CascadeA_VAE():
|
|||||||
def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False):
|
def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False):
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
|
|
||||||
# not sure if this will change in the future?
|
# not sure if this will change in the future?
|
||||||
from diffusers.pipelines.wuerstchen.modeling_paella_vq_model import PaellaVQModel
|
from diffusers.pipelines.wuerstchen.modeling_paella_vq_model import PaellaVQModel
|
||||||
self.model = PaellaVQModel.from_pretrained(
|
self.model = PaellaVQModel.from_pretrained(
|
||||||
@@ -102,13 +116,28 @@ class CascadeA_VAE():
|
|||||||
out = torch.clamp(out, min=0.0, max=1.0)
|
out = torch.clamp(out, min=0.0, max=1.0)
|
||||||
return out.to(latent.dtype).to(latent.device)
|
return out.to(latent.dtype).to(latent.device)
|
||||||
|
|
||||||
|
class No_VAE():
|
||||||
|
scale = 1
|
||||||
|
channels = 3
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def encode(self, image):
|
||||||
|
return image
|
||||||
|
|
||||||
|
def decode(self, image):
|
||||||
|
return image
|
||||||
|
|
||||||
|
vae_vers = {
|
||||||
|
"no": No_VAE,
|
||||||
|
"v1": SDv1_VAE,
|
||||||
|
"xl": SDXL_VAE,
|
||||||
|
"v3": SDv3_VAE,
|
||||||
|
"cc": CascadeC_VAE,
|
||||||
|
"ca": CascadeA_VAE,
|
||||||
|
}
|
||||||
|
|
||||||
def load_vae(ver, *args, **kwargs):
|
def load_vae(ver, *args, **kwargs):
|
||||||
if ver == "v1":
|
assert ver in vae_vers.keys(), f"Unknown VAE '{ver}'"
|
||||||
VAE = SDv1_VAE
|
vae_class = vae_vers[ver]
|
||||||
elif ver == "xl":
|
return vae_class(*args, **kwargs)
|
||||||
VAE = SDXL_VAE
|
|
||||||
elif ver == "cc":
|
|
||||||
VAE = CascadeC_VAE
|
|
||||||
elif ver == "ca":
|
|
||||||
VAE = CascadeA_VAE
|
|
||||||
return VAE(*args, **kwargs)
|
|
||||||
|
|||||||
Reference in New Issue
Block a user