diff --git a/.gitignore b/.gitignore index 80720af..85011b7 100644 --- a/.gitignore +++ b/.gitignore @@ -12,6 +12,7 @@ test.py *.pth *.ckpt *.safetensors +preprocess_* # default github .gitignore follows diff --git a/comfy_latent_interposer.py b/comfy_latent_interposer.py index ecd7451..6c0d74a 100644 --- a/comfy_latent_interposer.py +++ b/comfy_latent_interposer.py @@ -13,6 +13,8 @@ config = { "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-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): @@ -89,8 +91,8 @@ class ComfyLatentInterposer: return { "required": { "samples": ("LATENT", ), - "latent_src": (["v1", "xl", "ca"],), - "latent_dst": (["v1", "xl", "ca"],), + "latent_src": (["v1", "xl", "v3", "ca"],), + "latent_dst": (["v1", "xl"],), } } @@ -144,7 +146,9 @@ class ComfyLatentInterposer: lt = samples["samples"] with torch.no_grad(): # 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 return (samples,) diff --git a/config/v3-to-v1.yaml b/config/v3-to-v1.yaml new file mode 100644 index 0000000..1c82812 --- /dev/null +++ b/config/v3-to-v1.yaml @@ -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" diff --git a/config/v3-to-xl.yaml b/config/v3-to-xl.yaml new file mode 100644 index 0000000..68ee5e7 --- /dev/null +++ b/config/v3-to-xl.yaml @@ -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" diff --git a/utils.py b/utils.py deleted file mode 100644 index ed550aa..0000000 --- a/utils.py +++ /dev/null @@ -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 diff --git a/vae.py b/vae.py index e4fa097..713ffa9 100644 --- a/vae.py +++ b/vae.py @@ -48,6 +48,20 @@ class SDXL_VAE(SDv1_VAE): if dec_only: 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): scale = 1/32 channels = 16 @@ -75,7 +89,7 @@ class CascadeA_VAE(): def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False): self.device = device self.dtype = dtype - + # not sure if this will change in the future? from diffusers.pipelines.wuerstchen.modeling_paella_vq_model import PaellaVQModel self.model = PaellaVQModel.from_pretrained( @@ -102,13 +116,28 @@ class CascadeA_VAE(): out = torch.clamp(out, min=0.0, max=1.0) 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): - if ver == "v1": - VAE = SDv1_VAE - elif ver == "xl": - VAE = SDXL_VAE - elif ver == "cc": - VAE = CascadeC_VAE - elif ver == "ca": - VAE = CascadeA_VAE - return VAE(*args, **kwargs) + assert ver in vae_vers.keys(), f"Unknown VAE '{ver}'" + vae_class = vae_vers[ver] + return vae_class(*args, **kwargs)