fix: model_router not working

This commit is contained in:
yada
2023-04-01 00:16:28 -04:00
parent ddd4854207
commit b2088d97a4
4 changed files with 140 additions and 6 deletions
+2 -2
View File
@@ -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
}
+8 -4
View File
@@ -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')
+63
View File
@@ -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)
+67
View File
@@ -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)