fix: model_router not working
This commit is contained in:
+2
-2
@@ -1,12 +1,12 @@
|
|||||||
import custom_nodes.comfy_nodes_trojblue.image_layering as image_layering
|
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.color_correction as color_correction
|
||||||
import custom_nodes.comfy_nodes_trojblue.model_router as model_router
|
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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"trLayering": image_layering.Layering, # Layering
|
"trLayering": image_layering.Layering, # Layering
|
||||||
"trColorCorrection": color_correction.ColorCorrectionNode, # ColorCorrectionNode
|
"trColorCorrection": color_correction.ColorCorrectionNode, # ColorCorrectionNode
|
||||||
"trRouter": model_router.ModelRouterPlugin, # ModelRouterPlugin
|
"trRouter": model_router.ModelRouterPlugin, # ModelRouterPlugin
|
||||||
|
"trRouterLonger": model_router.LongerModelRouterPlugin, # ModelRouterPlugin
|
||||||
}
|
}
|
||||||
|
|||||||
+8
-4
@@ -10,8 +10,9 @@ class ColorCorrectionNode:
|
|||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"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)
|
image = blendLayers(image, original_image, BlendType.LUMINOSITY)
|
||||||
return image
|
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)
|
original_image = self.tensor_to_pil(original_image)
|
||||||
target_image = self.tensor_to_pil(target_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
|
# convert to tensor
|
||||||
corrected_image = corrected_image.convert('RGB')
|
corrected_image = corrected_image.convert('RGB')
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ class ModelRouterPlugin:
|
|||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
|
"required": {},
|
||||||
"optional": {
|
"optional": {
|
||||||
"model": ("MODEL",),
|
"model": ("MODEL",),
|
||||||
"clip": ("CLIP",),
|
"clip": ("CLIP",),
|
||||||
@@ -41,5 +42,67 @@ class ModelRouterPlugin:
|
|||||||
}
|
}
|
||||||
|
|
||||||
def execute(self, model=None, clip=None, vae=None, conditioning1=None, conditioning2=None):
|
def execute(self, model=None, clip=None, vae=None, conditioning1=None, conditioning2=None):
|
||||||
|
print(model, clip, vae)
|
||||||
|
|
||||||
return model, clip, vae, conditioning1, conditioning2
|
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)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
Reference in New Issue
Block a user