From 58dcbeadd06af8cfe0c84981f13bc52f76938df6 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Tue, 15 Aug 2023 20:08:22 +0200 Subject: [PATCH] Add custom node --- __init__.py | 8 +++++ comfy_latent_upscaler.py | 73 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 81 insertions(+) create mode 100644 __init__.py create mode 100644 comfy_latent_upscaler.py diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..9229cd1 --- /dev/null +++ b/__init__.py @@ -0,0 +1,8 @@ +# only import if running as a custom node +try: + import comfy.utils +except ImportError: + pass +else: + from .comfy_latent_upscaler import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/comfy_latent_upscaler.py b/comfy_latent_upscaler.py new file mode 100644 index 0000000..478554c --- /dev/null +++ b/comfy_latent_upscaler.py @@ -0,0 +1,73 @@ +import torch +import torch.nn as nn +from safetensors.torch import load_file +from huggingface_hub import hf_hub_download + + +class Upscaler(nn.Module): + """ + Basic NN layout, ported from: + https://github.com/city96/SD-Latent-Upscaler/blob/main/upscaler.py + """ + version = 1.0 # network revision + def __init__(self, fac): + super().__init__() + + module_list = [ + nn.Conv2d(4, 64, kernel_size=5, padding=2), + nn.ReLU(), + nn.Upsample(scale_factor=fac, mode="nearest"), + nn.ReLU(), + nn.Conv2d(64, 64, kernel_size=7, padding=3), + nn.ReLU(), + nn.Conv2d(64, 64, kernel_size=7, padding=3), + nn.ReLU(), + nn.Conv2d(64, 32, kernel_size=7, padding=3), + nn.ReLU(), + nn.Conv2d(32, 4, kernel_size=5, padding=2), + ] + self.sequential = nn.Sequential(*module_list) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.sequential(x) + + +class LatentUpscaler: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "samples": ("LATENT", ), + "model_ver": (["v1", "xl"],), + "scale_factor": (["1.25", "1.5", "2.0"],), + } + } + + RETURN_TYPES = ("LATENT",) + FUNCTION = "upscale" + CATEGORY = "latent" + + def upscale(self, samples, model_ver, scale_factor): + model = Upscaler(scale_factor) + weights = str(hf_hub_download( + repo_id="city96/SD-Latent-Upscaler", + filename=f"latent-upscaler-v{model.version}_SD{model_ver}-x{scale_factor}.safetensors") + ) + # weights = f"./latent-upscaler-v{model.version}_SD{model_ver}-x{scale_factor}.safetensors" + + model.load_state_dict(load_file(weights)) + lt = samples["samples"] + lt = model(lt) + del model + return ({"samples": lt},) + +NODE_CLASS_MAPPINGS = { + "LatentUpscaler": LatentUpscaler, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LatentUpscaler": "Latent Upscaler" +}