From 096b915a50cb0e72433307a4384a6692b05cbd28 Mon Sep 17 00:00:00 2001 From: Chris Date: Wed, 11 Mar 2026 13:16:24 +1100 Subject: [PATCH] #127and #130 and some other cleaning --- __init__.py | 9 +- image_filter.py | 208 ----------------------- image_filter_messaging.py | 11 +- image_filter_nodes.py | 234 ++++++++++++++++++++++++++ js/popup.js | 2 +- list_utility_nodes.py | 88 ---------- mask_utility_nodes.py | 39 ----- string_utility_nodes.py | 105 ------------ utility_nodes/list_utility_nodes.py | 81 +++++++++ utility_nodes/mask_utility_nodes.py | 44 +++++ utility_nodes/string_utility_nodes.py | 116 +++++++++++++ 11 files changed, 485 insertions(+), 452 deletions(-) delete mode 100644 image_filter.py create mode 100644 image_filter_nodes.py delete mode 100644 list_utility_nodes.py delete mode 100644 mask_utility_nodes.py delete mode 100644 string_utility_nodes.py create mode 100644 utility_nodes/list_utility_nodes.py create mode 100644 utility_nodes/mask_utility_nodes.py create mode 100644 utility_nodes/string_utility_nodes.py diff --git a/__init__.py b/__init__.py index 0a4b300..df1efbc 100644 --- a/__init__.py +++ b/__init__.py @@ -5,10 +5,10 @@ @description: A custom node that pauses the flow while you choose which image or images to pass on to the rest of the workflow. Simplified and improved version of cg-image-picker. """ -from .image_filter import ImageFilter, MaskImageFilter, TextImageFilterWithExtras -from .list_utility_nodes import PickFromList, BatchFromImageList, ImageListFromBatch, StringListFromStrings -from .string_utility_nodes import SplitByCommas, StringToFloat, StringToInt, AnyListToString, StringToStringList -from .mask_utility_nodes import MaskedSection +from .image_filter_nodes import ImageFilter, MaskImageFilter, TextImageFilterWithExtras +from .utility_nodes.list_utility_nodes import PickFromList, BatchFromImageList, ImageListFromBatch +from .utility_nodes.string_utility_nodes import SplitByCommas, StringToFloat, StringToInt, AnyListToString, StringToStringList +from .utility_nodes.mask_utility_nodes import MaskedSection VERSION = "1.7" WEB_DIRECTORY = "./js" @@ -24,7 +24,6 @@ NODE_CLASS_MAPPINGS= { "String to Float": StringToFloat, "Pick from List": PickFromList, "Any List to String": AnyListToString, - "String List from Strings": StringListFromStrings, "Batch from Image List": BatchFromImageList, "Image List From Batch": ImageListFromBatch, "Masked Section": MaskedSection, diff --git a/image_filter.py b/image_filter.py deleted file mode 100644 index 7898d10..0000000 --- a/image_filter.py +++ /dev/null @@ -1,208 +0,0 @@ -from nodes import PreviewImage, LoadImage -from comfy.model_management import InterruptProcessingException -import os, random -import torch - -import base64 -import io -from PIL import Image -import numpy as np - -from .image_filter_messaging import send_and_wait, Response, TimeoutResponse - -HIDDEN = { - "prompt": "PROMPT", - "extra_pnginfo": "EXTRA_PNGINFO", - "uid":"UNIQUE_ID", - - } - -class ImageFilter(PreviewImage): - RETURN_TYPES = ("IMAGE","LATENT","MASK","STRING","STRING","STRING","STRING") - RETURN_NAMES = ("images","latents","masks","extra1","extra2","extra3","indexes") - FUNCTION = "func" - CATEGORY = "image_filter" - OUTPUT_NODE = False - DESCRIPTION = "Allows you to preview images and choose which, if any to proceed with" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "images" : ("IMAGE", ), - "timeout": ("INT", {"default": 600, "min":1, "max":9999999, "tooltip": "Timeout in seconds."}), - "ontimeout": (["send none", "send all", "send first", "send last"], {}), - }, - "optional": { - "latents" : ("LATENT", {"tooltip": "Optional - if provided, will be output"}), - "masks" : ("MASK", {"tooltip": "Optional - if provided, will be output"}), - "tip" : ("STRING", {"default":"", "tooltip": "Optional - if provided, will be displayed in popup window"}), - "extra1" : ("STRING", {"default":""}), - "extra2" : ("STRING", {"default":""}), - "extra3" : ("STRING", {"default":""}), - "pick_list_start" : ("INT", {"default":0, "tooltip":"The number used in pick_list for the first image"}), - "pick_list" : ("STRING", {"default":"", "tooltip":"If a comma separated list of integers is provided, the images with these indices will be selected automatically."}), - "video_frames" : ("INT", {"default":1, "min":1, "tooltip": "treat each block of n images as a video"}), - "graph_id": ("STRING", {"default":""}), - }, - "hidden": HIDDEN, - } - - @classmethod - def IS_CHANGED(cls, pick_list, **kwargs): - return pick_list or float("NaN") - - def func(self, images, timeout, ontimeout, uid, graph_id, tip="", extra1="", extra2="", extra3="", latents=None, masks=None, pick_list_start:int=0, pick_list:str="", video_frames:int=1, **kwargs): - e1, e2, e3 = extra1, extra2, extra3 - B = images.shape[0] - - if video_frames>B: video_frames=1 - - try: - images_to_return:list[int] = [ int(x.strip())%B for x in pick_list.split(',') ] if pick_list else [] - except Exception as e: - print(f"{e} parsing pick_list - will manually select") - images_to_return = [] - - if len(images_to_return) == 0: - all_the_same = ( B and all( (images[i]==images[0]).all() for i in range(1,B) )) - urls:list[str] = self.save_images(images=images, **kwargs)['ui']['images'] - payload = {"uid": uid, "urls":urls, "allsame":all_the_same, "extras":[extra1, extra2, extra3], "tip":tip, "video_frames":video_frames} - - response:Response = send_and_wait(payload, timeout, uid, graph_id) - - if isinstance(response, TimeoutResponse): - if ontimeout=='send none': images_to_return = [] - if ontimeout=='send all': images_to_return = [*range(len(images)//video_frames)] - if ontimeout=='send first': images_to_return = [0,] - if ontimeout=='send last': images_to_return = [(len(images)//video_frames)-1,] - else: - e1, e2, e3 = response.get_extras([extra1, extra2, extra3]) - images_to_return = [ int(x) for x in response.selection ] if response.selection else [] - - if images_to_return is None or len(images_to_return) == 0: raise InterruptProcessingException() - - if video_frames>1: - images_to_return = [ key*video_frames + frm for key in images_to_return for frm in range(video_frames) ] - - images = torch.stack(list(images[int(i)] for i in images_to_return)) - latents = {"samples": torch.stack(list(latents['samples'][int(i)] for i in images_to_return))} if latents is not None else None - masks = torch.stack(list(masks[int(i)] for i in images_to_return)) if masks is not None else None - - try: int(pick_list_start) - except: pick_list_start = 0 - - return (images, latents, masks, e1, e2, e3, ",".join(str(int(x)+int(pick_list_start)) for x in images_to_return)) - -class TextImageFilterWithExtras(PreviewImage): - RETURN_TYPES = ("IMAGE","STRING","STRING","STRING","STRING") - RETURN_NAMES = ("image","text","extra1","extra2","extra3") - FUNCTION = "func" - CATEGORY = "image_filter" - OUTPUT_NODE = False - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image" : ("IMAGE", ), - "text" : ("STRING", {"default":""}), - "timeout": ("INT", {"default": 600, "min":1, "max":9999999, "tooltip": "Timeout in seconds."}), - }, - "optional": { - "mask" : ("MASK", {"tooltip": "Optional - if provided, will be overlaid on image"}), - "tip" : ("STRING", {"default":"", "tooltip": "Optional - if provided, will be displayed in popup window"}), - "extra1" : ("STRING", {"default":""}), - "extra2" : ("STRING", {"default":""}), - "extra3" : ("STRING", {"default":""}), - "textareaheight" : ("INT", {"default": 150, "min": 50, "max": 500, "tooltip": "Height of text area in pixels"}), - "graph_id": ("STRING", {"default":""}), - }, - "hidden": HIDDEN, - } - - @classmethod - def IS_CHANGED(cls, **kwargs): - return float("NaN") - - def func(self, image, text, timeout, uid, graph_id, extra1="", extra2="", extra3="", mask=None, tip="", textareaheight=None, **kwargs): - if image is None: image = torch.zeros((1,64,64,3)) - urls:list[str] = self.save_images(images=image, **kwargs)['ui']['images'] - payload = {"uid": uid, "urls":urls, "text":text, "extras":[extra1, extra2, extra3], "tip":tip} - if textareaheight is not None: payload['textareaheight'] = textareaheight - if mask is not None: payload['mask_urls'] = self.save_images(images=mask_to_image(mask), **kwargs)['ui']['images'] - - response = send_and_wait(payload, timeout, uid, graph_id) - if isinstance(response, TimeoutResponse): - return (image, text, extra1, extra2, extra3) - - return (image, response.text, *response.get_extras([extra1, extra2, extra3])) - -def mask_to_image(mask:torch.Tensor): - return torch.stack([mask, mask, mask, 1.0-mask], -1) - -class MaskImageFilter(PreviewImage, LoadImage): - RETURN_TYPES = ("IMAGE","MASK","STRING","STRING","STRING") - RETURN_NAMES = ("image","mask","extra1","extra2","extra3") - FUNCTION = "func" - CATEGORY = "image_filter" - OUTPUT_NODE = False - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image" : ("IMAGE", ), - "timeout": ("INT", {"default": 600, "min":1, "max":9999999, "tooltip": "Timeout in seconds."}), - "if_no_mask": (["cancel", "send blank"], {}), - }, - "optional": { - "mask" : ("MASK", {"tooltip":"optional initial mask"}), - "tip" : ("STRING", {"default":"", "tooltip": "Optional - if provided, will be displayed in popup window"}), - "extra1" : ("STRING", {"default":""}), - "extra2" : ("STRING", {"default":""}), - "extra3" : ("STRING", {"default":""}), - "graph_id": ("STRING", {"default":""}), - }, - "hidden": HIDDEN, - } - - @classmethod - def IS_CHANGED(cls, *args, **kwargs): - return f"{random.random()}" - - @classmethod - def VALIDATE_INPUTS(cls, *args, **kwargs): return True - - def func(self, image, timeout, uid, if_no_mask, graph_id, mask=None, extra1="", extra2="", extra3="", tip="", **kwargs): - if mask is not None and mask.shape[:3] == image.shape[:3] and not torch.all(mask==0): - saveable = torch.cat((image, mask.unsqueeze(-1)), dim=-1) - else: - saveable = image - - urls:list[dict[str,str]] = self.save_images(images=saveable, **kwargs)['ui']['images'] - payload = {"uid": uid, "urls":urls, "maskedit":True, "extras":[extra1, extra2, extra3], "tip":tip} - response = send_and_wait(payload, timeout, uid, graph_id) - - if (response.masked_image): - try: - return ( *(self.load_image(os.path.join('clipspace', response.masked_image)+" [input]")), *response.get_extras([extra1, extra2, extra3]) ) - except FileNotFoundError: - pass - elif (response.masked_data): - data = response.masked_data.split(',',1)[-1] - bytes_data = data.encode('utf-8') - image_data = base64.decodebytes(bytes_data) - data_io = io.BytesIO(image_data) - img = Image.open(data_io) - - mask = np.array(img.getchannel('A')).astype(np.float32) / 255.0 - mask = 1. - torch.from_numpy(mask) - mask = mask.unsqueeze(0) - - return ( image, mask, *response.get_extras([extra1, extra2, extra3]) ) - - - if if_no_mask == 'cancel': - raise InterruptProcessingException() - return ( *(self.load_image(urls[0]['filename']+" [temp]")), *response.get_extras([extra1, extra2, extra3]) ) diff --git a/image_filter_messaging.py b/image_filter_messaging.py index dd57760..af36122 100644 --- a/image_filter_messaging.py +++ b/image_filter_messaging.py @@ -101,26 +101,25 @@ async def cg_image_filter_message(request): return web.json_response({}) -def wait_for_response(secs, uid, graph_id) -> Response: +def wait_for_response(secs, graph_id) -> Response: MessageState.start_waiting(graph_id) try: end_time = time.monotonic() + secs while(time.monotonic() < end_time and MessageState.waiting()): throw_exception_if_processing_interrupted() - PromptServer.instance.send_sync("cg-image-filter-images", {"tick": int(end_time - time.monotonic()), "uid": uid, "graph_id":graph_id}) + PromptServer.instance.send_sync("cg-image-filter-images", {"tick": int(end_time - time.monotonic()), "graph_id":graph_id}) time.sleep(0.5) if MessageState.waiting(): - PromptServer.instance.send_sync("cg-image-filter-images", {"timeout": True, "uid": uid, "graph_id":graph_id}) + PromptServer.instance.send_sync("cg-image-filter-images", {"timeout": True, "graph_id":graph_id}) return MessageState.get_response() finally: MessageState.stop_waiting() -def send_and_wait(payload, timeout, uid, graph_id) -> Response: - payload['uid'] = uid +def send_and_wait(payload, timeout, graph_id) -> Response: payload['graph_id'] = graph_id while True: PromptServer.instance.send_sync("cg-image-filter-images", payload) - r = wait_for_response(timeout, uid, graph_id) + r = wait_for_response(timeout, graph_id) if isinstance(r,CancelledResponse): raise InterruptProcessingException() if (not isinstance(r, RequestResponse)): return r \ No newline at end of file diff --git a/image_filter_nodes.py b/image_filter_nodes.py new file mode 100644 index 0000000..f2d1c28 --- /dev/null +++ b/image_filter_nodes.py @@ -0,0 +1,234 @@ +from nodes import PreviewImage, LoadImage +from comfy.model_management import InterruptProcessingException +import os, random +import torch + +import base64 +from io import BytesIO +from PIL import Image +import numpy as np + +from .image_filter_messaging import send_and_wait, Response, TimeoutResponse +from comfy_api.latest import io + +class FilterNodeBase: + _preview_image = PreviewImage() + _load_image = LoadImage() + + @classmethod + def save_images_return_urls(cls, images:torch.Tensor, **kwargs) -> list[dict[str,str]]: + return cls._preview_image.save_images(images, **kwargs)['ui']['images'] + + @classmethod + def load_mask(cls, file:str, type:str="clipspace", append=" [input]") -> torch.Tensor: + return cls._load_image.load_image(os.path.join(type, file)+append)[1] + + @classmethod + def fingerprint_inputs(cls, **kwargs): # type: ignore + return random.random() + + @classmethod + def VALIDATE_INPUTS(cls, *args, **kwargs): return True + + +class ImageFilter(io.ComfyNode, FilterNodeBase): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Image Filter", + display_name = "Image Filter", + inputs = [ + io.Image.Input("images"), + io.Latent.Input("latents", optional=True, tooltip="optional"), + io.Mask.Input("masks", optional=True, tooltip="optional"), + io.Int.Input("timeout", default=600, min=1, max=1000000, tooltip="timeout in seconds"), + io.Combo.Input("ontimeout", options=["send none", "send all", "send first", "send last"]), + io.String.Input("tip", default="", optional=True), + io.String.Input("extra1", default="", optional=True), + io.String.Input("extra2", default="", optional=True), + io.String.Input("extra3", default="", optional=True), + io.Int.Input("pick_list_start", optional=True, default=0, tooltip="The index of the first image (normally 0 or 1)"), + io.String.Input("pick_list", optional=True, default="", tooltip="If a comma separated list of integers is provided, the images with these indices will be selected automatically."), + io.Int.Input("video_frames", optional=True, default=1, tooltip="Treat each block of n images as a video"), + io.String.Input("graph_id", default="") + ], + outputs = [ + io.Image.Output("images", display_name="images"), + io.Latent.Output("latents", display_name="latents"), + io.Mask.Output("masks", display_name="masks"), + io.String.Output("extra1", display_name="extra1"), + io.String.Output("extra2", display_name="extra2"), + io.String.Output("extra3", display_name="extra3"), + io.String.Output("indexes", display_name="indexes") + ], + category = "image_filter" + ) + + @classmethod + def parse_picklist(cls, pick_list:str, B:int=1) -> list[int]: + return [ int(x.strip())%B for x in pick_list.split(',') ] if pick_list else [] + + @classmethod + def fingerprint_inputs(cls, pick_list:str, **kwargs): # type: ignore + try: + if (pl:=cls.parse_picklist(pick_list)): return ",".join([str(p) for p in pl]) + except: + pass + return random.random() + + @classmethod + def execute( # type: ignore + cls, + images: torch.Tensor, latents=None, masks=None, + timeout:int=600, ontimeout:str="send none", + graph_id:str="", + tip:str="", extra1:str="", extra2:str="", extra3:str="", + pick_list_start:int=0, pick_list:str="", video_frames:int=1, + **kwargs + ) -> io.NodeOutput: + e1, e2, e3 = extra1, extra2, extra3 + B = images.shape[0] + + if video_frames>B: video_frames=1 + + try: + images_to_return:list[int] = cls.parse_picklist(pick_list, B) + except Exception as e: + print(f"{e} parsing pick_list - will manually select") + images_to_return = [] + + if len(images_to_return) == 0: + all_the_same = ( B and all( (images[i]==images[0]).all() for i in range(1,B) )) + urls:list[dict[str,str]] = cls.save_images_return_urls(images=images, **kwargs) + payload = { "urls":urls, "allsame":all_the_same, "extras":[extra1, extra2, extra3], "tip":tip, "video_frames":video_frames } + + response:Response = send_and_wait(payload, timeout, graph_id) + images_to_return:list[int] + + if isinstance(response, TimeoutResponse): + if ontimeout=='send none': images_to_return = [] + if ontimeout=='send all': images_to_return = [*range(len(images)//video_frames)] + if ontimeout=='send first': images_to_return = [0,] + if ontimeout=='send last': images_to_return = [(len(images)//video_frames)-1,] + else: + e1, e2, e3 = response.get_extras([extra1, extra2, extra3]) + images_to_return = response.selection or [] + + if not images_to_return: raise InterruptProcessingException() + + if video_frames>1: + images_to_return = [ key*video_frames + frm for key in images_to_return for frm in range(video_frames) ] + + images = torch.stack(list(images[i] for i in images_to_return)) + latents = {"samples": torch.stack(list(latents['samples'][int(i)] for i in images_to_return))} if latents is not None else None + masks = torch.stack(list(masks[i] for i in images_to_return)) if masks is not None else None + + return io.NodeOutput(images, latents, masks, e1, e2, e3, ",".join(str(x+pick_list_start) for x in images_to_return)) + +class TextImageFilterWithExtras(io.ComfyNode, FilterNodeBase): + + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Text Image Filter", + display_name = "Text Image Filter", + inputs = [ + io.Image.Input("image"), + io.String.Input("text", default=""), + io.Int.Input("timeout", default=600, min=1, max=1000000, tooltip="timeout in seconds"), + io.Mask.Input("mask", optional=True, tooltip="optional"), + io.String.Input("tip", default="", optional=True), + io.String.Input("extra1", default="", optional=True), + io.String.Input("extra2", default="", optional=True), + io.String.Input("extra3", default="", optional=True), + io.Int.Input("textareaheight", default=150, min=30, max=500), + io.String.Input("graph_id", default="") + ], + outputs = [ + io.Image.Output("images", display_name="images"), + io.String.Output("text", display_name="text"), + io.String.Output("extra1", display_name="extra1"), + io.String.Output("extra2", display_name="extra2"), + io.String.Output("extra3", display_name="extra3"), + ], + category = "image_filter" + ) + + @classmethod + def execute(cls, image, text, timeout, graph_id, extra1="", extra2="", extra3="", mask=None, tip="", textareaheight=None, **kwargs): # type: ignore + if image is None: image = torch.zeros((1,64,64,3)) + urls:list[dict[str,str]] = cls.save_images_return_urls(images=image, **kwargs) + payload = {"urls":urls, "text":text, "extras":[extra1, extra2, extra3], "tip":tip} + if textareaheight is not None: payload['textareaheight'] = textareaheight + if mask is not None: payload['mask_urls'] = cls.save_images_return_urls(images=mask_to_image(mask), **kwargs) + + response = send_and_wait(payload, timeout, graph_id) + if isinstance(response, TimeoutResponse): + return io.NodeOutput(image, text, extra1, extra2, extra3) + + return io.NodeOutput(image, response.text, *response.get_extras([extra1, extra2, extra3])) + + +def mask_to_image(mask:torch.Tensor): + return torch.stack([mask, mask, mask, 1.0-mask], -1) + +def mask_from_data(data) -> torch.Tensor: + bytes_data = data.encode('utf-8') + image_data = base64.decodebytes(bytes_data) + data_io = BytesIO(image_data) + img = Image.open(data_io) + + mask = np.array(img.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + return mask.unsqueeze(0) + +class MaskImageFilter(io.ComfyNode, FilterNodeBase): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Mask Image Filter", + display_name = "Mask Image Filter", + inputs = [ + io.Image.Input("image"), + io.Int.Input("timeout", default=600, min=1, max=1000000, tooltip="timeout in seconds"), + io.Combo.Input("if_no_mask", options=["cancel", "send blank"], default="send blank"), + io.Mask.Input("mask", optional=True, tooltip="optional"), + io.String.Input("tip", default="", optional=True), + io.String.Input("extra1", default="", optional=True), + io.String.Input("extra2", default="", optional=True), + io.String.Input("extra3", default="", optional=True), + io.String.Input("graph_id", default="") + ], + outputs = [ + io.Image.Output("image", display_name="image"), + io.Mask.Output("mask", display_name="mask"), + io.String.Output("extra1", display_name="extra1"), + io.String.Output("extra2", display_name="extra2"), + io.String.Output("extra3", display_name="extra3"), + ], + category = "image_filter" + ) + + @classmethod + def execute(cls, image, timeout, if_no_mask, graph_id, mask=None, extra1="", extra2="", extra3="", tip="", **kwargs): # type: ignore + if mask is not None and mask.shape[:3] == image.shape[:3] and not torch.all(mask==0): + saveable = torch.cat((image, mask.unsqueeze(-1)), dim=-1) + else: + saveable = image + + urls = cls.save_images_return_urls(images=saveable, **kwargs) + payload = { "urls":urls, "maskedit":True, "extras":[extra1, extra2, extra3], "tip":tip} + response = send_and_wait(payload, timeout, graph_id) + + if (response.masked_image): # old mask editor - uploads + try: + mask = cls.load_mask(response.masked_image) + except FileNotFoundError: # no mask was uploaded; reload the input mask, or the mask in the input image + mask = mask if mask is not None else cls.load_mask(urls[0]['filename']+" [temp]") + + elif (response.masked_data): # new mask editor - sends the blob + data = response.masked_data.split(',',1)[-1] + mask = mask_from_data(data) + + if if_no_mask == 'cancel' and torch.all(mask==0): raise InterruptProcessingException() + return io.NodeOutput( image, mask, *response.get_extras([extra1, extra2, extra3]) ) diff --git a/js/popup.js b/js/popup.js index ff4153a..fea149b 100644 --- a/js/popup.js +++ b/js/popup.js @@ -275,7 +275,7 @@ class Popup extends HTMLElement { _handle_message(message, using_saved) { const detail = message.detail - const uid = detail.uid + const uid = app.runningNodeId const the_node = this.find_node(uid) const graph_id = message.detail.graph_id diff --git a/list_utility_nodes.py b/list_utility_nodes.py deleted file mode 100644 index d7a1e53..0000000 --- a/list_utility_nodes.py +++ /dev/null @@ -1,88 +0,0 @@ -import torch -from comfy.comfy_types.node_typing import IO - -class BatchFromImageList: - @classmethod - def INPUT_TYPES(cls): - return {"required": { "images": ("IMAGE", ), } } - INPUT_IS_LIST = True - RETURN_TYPES = ("IMAGE", ) - FUNCTION = "func" - - CATEGORY = "image_filter/helpers" - - def func(self, images): - if len(images) <= 1: - return (images[0],) - else: - return (torch.cat(list(i for i in images), dim=0),) - -class ImageListFromBatch: - @classmethod - def INPUT_TYPES(cls): - return {"required": { "images": ("IMAGE", ), } } - INPUT_IS_LIST = False - OUTPUT_IS_LIST = [True,] - RETURN_TYPES = ("IMAGE", ) - FUNCTION = "func" - - CATEGORY = "image_filter/helpers" - - def func(self, images): - image_list = list( i.unsqueeze(0) for i in images ) - return (image_list,) - -class StringListFromStrings: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "s0": ("STRING", {"default":""}), - "s1": ("STRING", {"default":""}), - }, - "optional": { - "s2": ("STRING", {"default":""}), - "s3": ("STRING", {"default":""}), - } - - } - INPUT_IS_LIST = False - OUTPUT_IS_LIST = [True,] - RETURN_TYPES = ("STRING", ) - FUNCTION = "func" - - CATEGORY = "image_filter/helpers" - - def func(self, s0,s1,s2=None,s3=None): - lst = [s0,s1] - if s2: lst.append(s2) - if s3: lst.append(s3) - return (lst,) - - -class PickFromList: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "anything" : (IO.ANY, ), - "indexes": ("STRING", {"default": ""}) - }, - } - RETURN_TYPES = (IO.ANY,) - RETURN_NAMES = ("picks",) - - FUNCTION = "func" - CATEGORY = "image_filter/helpers" - INPUT_IS_LIST = True - OUTPUT_IS_LIST = [True,] - - def func(self, anything, indexes): - try: - if len(anything)==1 and isinstance(anything[0],list): anything = anything[0] - indexes = [int(x.strip()) for x in indexes[0].split(',') if x.strip()] - except Exception as e: - print(e) - indexes = [] - - return ([anything[i] for i in indexes], ) \ No newline at end of file diff --git a/mask_utility_nodes.py b/mask_utility_nodes.py deleted file mode 100644 index 051cd0d..0000000 --- a/mask_utility_nodes.py +++ /dev/null @@ -1,39 +0,0 @@ -import torch - -class MaskedSection: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "mask": ("MASK",), - "image": ("IMAGE",), - "minimum": ("INT", {"default":512, "min":16, "max":4096}) - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "func" - CATEGORY = "image_filter/helpers" - - def func(self, mask:torch.Tensor, image, minimum=512): - mbb = mask.squeeze() - H,W = mbb.shape - masked = mbb > 0.5 - - non_zero_positions = torch.nonzero(masked) - if len(non_zero_positions) < 2: return (image,) - - min_x = int(torch.min(non_zero_positions[:, 1])) - max_x = int(torch.max(non_zero_positions[:, 1])) - min_y = int(torch.min(non_zero_positions[:, 0])) - max_y = int(torch.max(non_zero_positions[:, 0])) - - if (x:=(minimum-(max_x-min_x))//2)>0: - min_x = max(min_x-x, 0) - max_x = min(max_x+x, W) - if (y:=(minimum-(max_y-min_y))//2)>0: - min_y = max(min_y-y, 0) - max_y = min(max_y+y, H) - - return (image[:,min_y:max_y,min_x:max_x,:],) - diff --git a/string_utility_nodes.py b/string_utility_nodes.py deleted file mode 100644 index e49b3b5..0000000 --- a/string_utility_nodes.py +++ /dev/null @@ -1,105 +0,0 @@ -from comfy.comfy_types.node_typing import IO -from comfy_api.latest import io - -class StringToStringList(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id = "StringToStringList", - display_name = "String to String List", - category = "quicknodes/prompting", - inputs = [ - io.String.Input("string"), - io.String.Input("split",default=",", tooltip="Split on this substring (or linebreak)"), - ], - outputs = [ - io.String.Output("string_list", is_output_list=True), - ], - ) - - @classmethod - def execute(cls, string, split): # type: ignore - if split == "linebreak": split = "\n" - bits:list[str] = [r.strip() for r in string.split(split)] - return io.NodeOutput(bits) - - -class SplitByCommas: - RETURN_TYPES = ("STRING","STRING","STRING","STRING","STRING","STRING") - FUNCTION = "func" - CATEGORY = "image_filter/helpers" - OUTPUT_NODE = False - OUTPUT_IS_LIST = [False, False, False, False, False, True] - - DESCRIPTION = "Split the input string into up to five pieces. Splits on commas (or | or ^) and then strips whitespace from front and end." - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { "string" : ("STRING", {"default":""}), }, - "optional": { "split": ([",", "|", "^", ":", "-", "_", "linebreak"], {}), }, - } - - def func(self, string:str, split:str=",") -> tuple[str,str,str,str,str,list[str]]: - if split == "linebreak": split = "\n" - bits:list[str] = [r.strip() for r in string.split(split)] - - while len(bits)<5: bits.append("") - if len(bits)>5: bits = bits[:4] + [",".join(bits[4:]),] - - return (bits[0], bits[1], bits[2], bits[3], bits[4], bits) - -class AnyListToString: - RETURN_TYPES = ("STRING",) - FUNCTION = "func" - CATEGORY = "image_filter/helpers" - INPUT_IS_LIST = True - OUTPUT_IS_LIST = (False,) - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "anything" : (IO.ANY, ), - "join" : ("STRING", {"default":""}), - } - } - - def func(self, anything, join:str): - return ( join[0].join( [f"{x}" for x in anything] ), ) - -class StringToInt: - RETURN_TYPES = ("INT",) - FUNCTION = "func" - CATEGORY = "image_filter/helpers" - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "string" : ("STRING", {"default":"", "forceInput":True, "tooltip":"whitespace will be stripped before parsing"}), - "default" : ("INT", {"default":0, "tooltip":"used if the string can't be parsed as an integer"}), - } - } - - def func(self, string:str, default:int): - try: return (int(string.strip()),) - except: return (default,) - -class StringToFloat: - RETURN_TYPES = ("FLOAT",) - FUNCTION = "func" - CATEGORY = "image_filter/helpers" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "string" : ("STRING", {"default":"", "forceInput":True, "tooltip":"whitespace will be stripped before parsing"}), - "default" : ("FLOAT", {"default":0, "tooltip":"used if the string can't be parsed as a float"}), - } - } - - def func(self, string:str, default:float): - try: return (float(string.strip()),) - except: return (default,) \ No newline at end of file diff --git a/utility_nodes/list_utility_nodes.py b/utility_nodes/list_utility_nodes.py new file mode 100644 index 0000000..d9c8252 --- /dev/null +++ b/utility_nodes/list_utility_nodes.py @@ -0,0 +1,81 @@ +import torch +from comfy_api.latest import io + +class BatchFromImageList(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Batch from Image List", + display_name = "Batch from Image List", + inputs = [ + io.Image.Input("images") + ], + outputs = [ + io.Image.Output("image") + ], + is_input_list = True, + category = "image_filter/helpers" + ) + + @classmethod + def execute(cls, images): # type: ignore + if len(images) <= 1: + return io.NodeOutput(images[0],) + else: + return io.NodeOutput(torch.cat(list(i for i in images), dim=0),) + +class ImageListFromBatch(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Image List From Batch", + display_name = "Image List From Batch", + inputs = [ + io.Image.Input("images") + ], + outputs = [ + io.Image.Output("image", is_output_list=True) + ], + category = "image_filter/helpers" + ) + + @classmethod + def execute(cls, images): # type: ignore + image_list = list( i.unsqueeze(0) for i in images ) + return io.NodeOutput(image_list,) + +class PickFromList(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Pick from List", + display_name = "Pick from List", + inputs = [ + io.AnyType.Input("anything"), + io.String.Input("indexes", display_name="indexes", tooltip="comma separated list of indexes. Whitespace stripped. Only these entries will be included. Zero indexed.") + ], + outputs = [ + io.String.Output("picks", display_name="picks", is_output_list=True) + ], + category = "image_filter/helpers", + is_input_list=True + ) + + + @classmethod + def execute(cls, anything:list, indexes:list[str]): # type: ignore + + if len(anything)==1 and isinstance(anything[0],list): + print("Warning: received list of lists. Processing just anything[0]") + anything = anything[0] + + index_str:str = indexes[0] + + result = [] + for x in [x.strip() for x in index_str.split(',')]: + try: + result.append(anything[int(x)]) + except Exception as e: + print(f"{e} when processing {x} from {index_str}") + + return io.NodeOutput(result, ) \ No newline at end of file diff --git a/utility_nodes/mask_utility_nodes.py b/utility_nodes/mask_utility_nodes.py new file mode 100644 index 0000000..fcd2c14 --- /dev/null +++ b/utility_nodes/mask_utility_nodes.py @@ -0,0 +1,44 @@ +import torch +from comfy_api.latest import io + +class MaskedSection(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Masked Section", + display_name = "Masked Section", + inputs = [ + io.Mask.Input("mask"), + io.Image.Input("image"), + io.Int.Input("minimum", default=512, min=16, max=16384, tooltip="Minimum image size to output") + ], + outputs = [ + io.Image.Output("image") + ], + category = "image_filter/helpers", + description = "return the image cropped to only include the masked section" + ) + + @classmethod + def execute(cls, mask:torch.Tensor, image, minimum=512): # type: ignore + mbb = mask.squeeze() + H,W = mbb.shape + masked = mbb > 0.5 + + non_zero_positions = torch.nonzero(masked) + if len(non_zero_positions) < 2: return (image,) + + min_x = int(torch.min(non_zero_positions[:, 1])) + max_x = int(torch.max(non_zero_positions[:, 1])) + min_y = int(torch.min(non_zero_positions[:, 0])) + max_y = int(torch.max(non_zero_positions[:, 0])) + + if (x:=(minimum-(max_x-min_x))//2)>0: + min_x = max(min_x-x, 0) + max_x = min(max_x+x, W) + if (y:=(minimum-(max_y-min_y))//2)>0: + min_y = max(min_y-y, 0) + max_y = min(max_y+y, H) + + return io.NodeOutput(image[:,min_y:max_y,min_x:max_x,:],) + diff --git a/utility_nodes/string_utility_nodes.py b/utility_nodes/string_utility_nodes.py new file mode 100644 index 0000000..b5c01f5 --- /dev/null +++ b/utility_nodes/string_utility_nodes.py @@ -0,0 +1,116 @@ +from comfy_api.latest import io +from typing import Any + +class StringToStringList(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "StringToStringList", + display_name = "String to String List", + category = "image_filter/helpers", + inputs = [ + io.String.Input("string"), + io.String.Input("split",default=",", tooltip="Split on this substring (or linebreak)"), + ], + outputs = [ + io.String.Output("string_list", is_output_list=True), + ], + ) + + @classmethod + def execute(cls, string, split): # type: ignore + if split == "linebreak": split = "\n" + bits:list[str] = [r.strip() for r in string.split(split)] + return io.NodeOutput(bits) + +class SplitByCommas(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Split String by Commas", + display_name = "Split String on character", + inputs = [ + io.String.Input("string"), + io.String.Input("split", default=",", tooltip="Split on this substring (or linebreak)"), + ], + outputs = [ + io.String.Output("string1", display_name="string", is_output_list=True), + io.String.Output("string2", display_name="string", is_output_list=True), + io.String.Output("string3", display_name="string", is_output_list=True), + io.String.Output("string4", display_name="string", is_output_list=True), + io.String.Output("string5", display_name="string", is_output_list=True), + io.String.Output("all_as_list", display_name="all", is_output_list=True), + ], + category = "image_filter/helpers", + description = "Split the input string and strips whitespace." + ) + + @classmethod + def execute(cls, string, split): # type: ignore + if split == "linebreak": split = "\n" + bits:list[str] = [r.strip() for r in string.split(split)] + five = (bits + [""*5])[:5] + return io.NodeOutput(*five, bits) + +class AnyListToString(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "Any List to String", + display_name = "Any List to String", + inputs = [ + io.AnyType.Input("anything"), + io.String.Input("join", default="") + ], + outputs = [ + io.String.Output("string") + ], + is_input_list = True, + category = "image_filter/helpers", + ) + + @classmethod + def execute(cls, anything:list[Any], join:list[str]): # type: ignore + return io.NodeOutput( join[0].join( [f"{x}" for x in anything] ), ) + +class StringToInt(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "String to Int", + display_name = "String to Int", + inputs = [ + io.String.Input("string"), + io.Int.Input("default") + ], + outputs = [ + io.Int.Output("int") + ], + category = "image_filter/helpers", + ) + + @classmethod + def execute(cls, string:str, default:int): # type: ignore + try: return io.NodeOutput(int(string.strip()),) + except: return io.NodeOutput(default,) + +class StringToFloat(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id = "String to Float", + display_name = "String to Float", + inputs = [ + io.String.Input("float"), + io.Float.Input("default") + ], + outputs = [ + io.Float.Output("float") + ], + category = "image_filter/helpers", + ) + + @classmethod + def execute(cls, string:str, default:float): # type: ignore + try: return io.NodeOutput(float(string.strip()),) + except: return io.NodeOutput(default,) \ No newline at end of file