From b2088d97a463d92401741d13175c99fa71624c57 Mon Sep 17 00:00:00 2001 From: yada Date: Sat, 1 Apr 2023 00:16:28 -0400 Subject: [PATCH] fix: model_router not working --- __init__.py | 4 +-- color_correction.py | 12 +++++--- model_router.py | 63 ++++++++++++++++++++++++++++++++++++++++++ shortcut_nodes.py | 67 +++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 140 insertions(+), 6 deletions(-) create mode 100644 shortcut_nodes.py diff --git a/__init__.py b/__init__.py index c153482..715422c 100644 --- a/__init__.py +++ b/__init__.py @@ -1,12 +1,12 @@ import custom_nodes.comfy_nodes_trojblue.image_layering as image_layering import custom_nodes.comfy_nodes_trojblue.color_correction as color_correction import custom_nodes.comfy_nodes_trojblue.model_router as model_router - +import custom_nodes.comfy_nodes_trojblue.shortcut_nodes as shortcut_nodes NODE_CLASS_MAPPINGS = { "trLayering": image_layering.Layering, # Layering "trColorCorrection": color_correction.ColorCorrectionNode, # ColorCorrectionNode "trRouter": model_router.ModelRouterPlugin, # ModelRouterPlugin - + "trRouterLonger": model_router.LongerModelRouterPlugin, # ModelRouterPlugin } diff --git a/color_correction.py b/color_correction.py index 5fd5657..d75ed08 100644 --- a/color_correction.py +++ b/color_correction.py @@ -10,8 +10,9 @@ class ColorCorrectionNode: def INPUT_TYPES(s): return { "required": { - "original_image": ("IMAGE",), - "target_image": ("IMAGE",), + "TARGET_IMAGE": ("IMAGE",), + "reference": ("IMAGE",), + "inverse selection": (["False", "True"],), }, } @@ -44,11 +45,14 @@ class ColorCorrectionNode: image = blendLayers(image, original_image, BlendType.LUMINOSITY) return image - def color_correct(self, original_image, target_image): + def color_correct(self, original_image, target_image, inverse_selection="False"): original_image = self.tensor_to_pil(original_image) target_image = self.tensor_to_pil(target_image) - corrected_image = self.apply_color_correction(target_image, original_image) + if inverse_selection == "False": + corrected_image = self.apply_color_correction(target_image, original_image) + else: + corrected_image = self.apply_color_correction(original_image, target_image) # convert to tensor corrected_image = corrected_image.convert('RGB') diff --git a/model_router.py b/model_router.py index 7f7468f..4e74379 100644 --- a/model_router.py +++ b/model_router.py @@ -31,6 +31,7 @@ class ModelRouterPlugin: @classmethod def INPUT_TYPES(s): return { + "required": {}, "optional": { "model": ("MODEL",), "clip": ("CLIP",), @@ -41,5 +42,67 @@ class ModelRouterPlugin: } def execute(self, model=None, clip=None, vae=None, conditioning1=None, conditioning2=None): + print(model, clip, vae) + return model, clip, vae, conditioning1, conditioning2 + + + +class LongerModelRouterPlugin: + """ + An example node + + Class methods + ------------- + INPUT_TYPES (dict): + Tell the main program input parameters of nodes. + + Attributes + ---------- + RETURN_TYPES (`tuple`): + The type of each element in the output tulple. + FUNCTION (`str`): + The name of the entry-point method. For example, if `FUNCTION = "execute"` then it will run Example().execute() + OUTPUT_NODE ([`bool`]): + If this node is an output node that outputs a result/image from the graph. The SaveImage node is an example. + The backend iterates on these output nodes and tries to execute all their parents if their parent graph is properly connected. + Assumed to be False if not present. + CATEGORY (`str`): + The category the node should appear in the UI. + execute(s) -> tuple || None: + The entry point method. The name of this method must be the same as the value of property `FUNCTION`. + For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. + """ + + FUNCTION = "execute" + CATEGORY = "trNodes" + RETURN_TYPES = ("MODEL", "CLIP", "VAE", "CONDITIONING", "IMAGE", "LATENT", "CONDITIONING") + + @classmethod + def INPUT_TYPES(s): + """ + Return a dictionary which contains config for all input fields. + Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT". + Input types "INT", "STRING" or "FLOAT" are special values for fields on the node. + The type can be a list for selection. + """ + return { + "required": {}, + "optional": { + "model": ("MODEL",), + "clip": ("CLIP",), + "vae": ("VAE",), + "conditioning1": ("CONDITIONING",), + "image": ("IMAGE",), + "latent": ("LATENT",), + "conditioning2": ("CONDITIONING",), + } + } + + + def execute(self, model=None, clip=None, vae=None, conditioning1=None, image=None, latent=None, conditioning2=None): + + print(model, clip, vae) + return (model, clip, vae, conditioning1, image, latent, conditioning2) + diff --git a/shortcut_nodes.py b/shortcut_nodes.py new file mode 100644 index 0000000..8d803cd --- /dev/null +++ b/shortcut_nodes.py @@ -0,0 +1,67 @@ +# import numpy as np +# from PIL import Image +# from PIL.PngImagePlugin import PngInfo +# import json +from nodes import VAEDecode, SaveImage +import os + + +""" +NOT WORKING +""" +class SaveVAEImageNode(): + """ + DECODE and show Image sample; Also export the image as an optional param to be saved + + Class methods + ------------- + INPUT_TYPES (dict): + Tell the main program input parameters of nodes. + + Attributes + ---------- + RETURN_TYPES (`tuple`): + The type of each element in the output tulple. + FUNCTION (`str`): + The name of the entry-point method. For example, if `FUNCTION = "execute"` then it will run Example().execute() + OUTPUT_NODE ([`bool`]): + If this node is an output node that outputs a result/image from the graph. The SaveImage node is an example. + The backend iterates on these output nodes and tries to execute all their parents if their parent graph is properly connected. + Assumed to be False if not present. + CATEGORY (`str`): + The category the node should appear in the UI. + execute(s) -> tuple || None: + The entry point method. The name of this method must be the same as the value of property `FUNCTION`. + For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. + """ + + def __init__(self, device="cpu"): + self.device = device + self.decode_handler = VAEDecode() + self.image_handler = SaveImage() + + RETURN_TYPES = () + FUNCTION = "do_decode_and_preview" + CATEGORY = "trNodes" + OUTPUT_NODE = True + + @classmethod + def INPUT_TYPES(s): + return {"required": { + "samples": ("LATENT",), + "vae": ("VAE",)} + } + + def _do_decode(self, vae, samples): + return (vae.decode(samples["samples"]),) + + def _do_save_image(self, image, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): + return self.image_handler.save_images(image, filename_prefix, prompt, extra_pnginfo) + + def do_decode_and_preview(self, vae, samples, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): + image = self._do_decode(vae, samples) + ui_out = self._do_save_image(image, filename_prefix, prompt, extra_pnginfo) + return (ui_out, image) + + +