diff --git a/autonode.py b/autonode.py index ed2ec13..3ed9419 100644 --- a/autonode.py +++ b/autonode.py @@ -42,6 +42,8 @@ __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] # Then, you can use the registered nodes in the UI! """ +from inspect import signature + def get_node_names_mappings(classes): node_names = {} node_classes = {} @@ -64,6 +66,32 @@ def validate(container): for attr in ["FUNCTION", "INPUT_TYPES", "RETURN_TYPES", "CATEGORY"]: if not hasattr(cls, attr): raise Exception("Class {} doesn't have attribute {}".format(cls.__name__, attr)) + return_type = cls.RETURN_TYPES + if not isinstance(return_type, tuple): + raise Exception(f"RETURN_TYPES must be a tuple, got {type(return_type)} in {cls.__name__}") + if not all(isinstance(x, str) for x in return_type): + raise Exception(f"RETURN_TYPES must be a tuple of strings, got {return_type} in {cls.__name__}") + input_keys = ["self"] + for key in cls.INPUT_TYPES()["required"]: + input_keys.append(key) + for key in cls.INPUT_TYPES().get("optional", {}): + input_keys.append(key) + function_kwargs = signature(cls.__dict__[cls.FUNCTION]).parameters.keys() + function_kwargs = list(function_kwargs) + ["self"] + # if args/kwargs are in function kwargs, warn and skip + if "args" in function_kwargs: + print(f"Warning: args in function arguments in {cls.__name__}, skipping argument validation") + continue + if "kwargs" in function_kwargs: + print(f"Warning: kwargs in function arguments in {cls.__name__}, skipping argument validation") + continue + # input kwargs are subset of function kwargs + if not set(input_keys).issubset(function_kwargs): + raise Exception(f"INPUT_TYPES and function arguments must match in {cls.__name__}, input_types: {input_keys}, function arguments: {function_kwargs}") + # if not exact match, print warning + if len(set(input_keys)) != len(set(function_kwargs)): + print(f"Warning: INPUT_TYPES and function arguments don't match in {cls.__name__}, input_types: {input_keys}, function arguments: {function_kwargs}") + # AllTrue class hijacks the isinstance, issubclass, bool, str, jsonserializable, eq, ne methods to always return True class AllTrue(str): def __init__(self, representation=None) -> None: diff --git a/auxilary.py b/auxilary.py new file mode 100644 index 0000000..f9f47a7 --- /dev/null +++ b/auxilary.py @@ -0,0 +1,166 @@ +from .imgio.converter import PILHandlingHodes +from .autonode import node_wrapper, get_node_names_mappings, validate, anytype, PILImage +from .utils.tagger import get_tags + +auxilary_classes = [] +auxilary_node = node_wrapper(auxilary_classes) + +@auxilary_node +class GetRatingNode: + FUNCTION = "get_rating_class" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get Rating Class" + @staticmethod + def get_rating_class(image): + image = PILHandlingHodes.handle_input(image) + result_dict = get_tags(image) + return (result_dict['rating'], ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + }, + } + +@auxilary_node +class GetRatingFromTextNode: + FUNCTION = "get_rating_class" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get Rating Class From Text" + @staticmethod + def get_rating_class(image): + image = PILHandlingHodes.handle_input(image) + result_dict = get_tags(image) + return (result_dict['rating'], ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("STRING", {"default": "/path/to/image.jpg"}), + }, + } + +@auxilary_node +class GetTagsAboveThresholdNode: + FUNCTION = "get_tags_above_threshold" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get Tags Above Threshold" + @staticmethod + def get_tags_above_threshold(image, threshold, replace): + image = PILHandlingHodes.handle_input(image) + result_dict = get_tags(image, threshold=threshold, replace=replace) + return (", ".join(result_dict['tags']), ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + }, + "optional": { + "threshold": ("FLOAT", {"default": 0.4}), + "replace": ("BOOLEAN", {"default": False}), + } + } + +@auxilary_node +class GetTagsAboveThresholdFromTextNode: + FUNCTION = "get_tags_above_threshold" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get Tags Above Threshold From Text" + @staticmethod + def get_tags_above_threshold(image, threshold, replace): + image = PILHandlingHodes.handle_input(image) + result_dict = get_tags(image, threshold=threshold, replace=replace) + return (", ".join(result_dict['tags']), ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + }, + "optional": { + "threshold": ("FLOAT", {"default": 0.4}), + "replace": ("BOOLEAN", {"default": False}), + } + } + +@auxilary_node +class GetCharactersAboveThresholdNode: + FUNCTION = "get_tags_above_threshold" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get Chars Above Threshold" + @staticmethod + def get_tags_above_threshold(image, threshold, replace): + image = PILHandlingHodes.handle_input(image) + result_dict = get_tags(image, threshold=threshold, replace=replace) + return (", ".join(result_dict['chars']), ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + }, + "optional": { + "threshold": ("FLOAT", {"default": 0.4}), + "replace": ("BOOLEAN", {"default": False}), + } + } + +@auxilary_node +class GetCharactersAboveThresholdFromTextNode: + FUNCTION = "get_tags_above_threshold" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get Chars Above Threshold From Text" + @staticmethod + def get_tags_above_threshold(image, threshold, replace): + image = PILHandlingHodes.handle_input(image) + result_dict = get_tags(image, threshold=threshold, replace=replace) + return (", ".join(result_dict['chars']), ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + }, + "optional": { + "threshold": ("FLOAT", {"default": 0.4}), + "replace": ("BOOLEAN", {"default": False}), + } + } + +@auxilary_node +class GetAllTagsAboveThresholdNode: + FUNCTION = "get_tags" + RETURN_TYPES = ("STRING",) + CATEGORY = "tagger" + custom_name = "Get All Tags Above Threshold" + @staticmethod + def get_tags(image, threshold, replace): + image = PILHandlingHodes.handle_input(image) + result = get_tags(image, threshold=threshold, replace=replace) + result_list = [] + result_list.append(result['rating']) + result_list.extend(result['tags']) + result_list.extend(result['chars']) + return (", ".join(result_list), ) + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + }, + "optional": { + "threshold": ("FLOAT", {"default": 0.4}), + "replace": ("BOOLEAN", {"default": False}), + } + } + +CLASS_MAPPINGS, CLASS_NAMES = get_node_names_mappings(auxilary_classes) +validate(auxilary_classes) \ No newline at end of file diff --git a/imgio/converter.py b/imgio/converter.py index f5d6a73..76cc94f 100644 --- a/imgio/converter.py +++ b/imgio/converter.py @@ -63,7 +63,7 @@ class IOConverter: # if not first element is 1, then it is a batch of images so warning if input_data.shape[0] != 1: print("Warning: Batch of images detected, taking first image") - input_data = input_data[0] * 255.0 + input_data = input_data[0] * 255.0 if input_data.dtype == np.float32 else input_data[0] input_data = input_data.astype('uint8') return Image.fromarray(input_data) elif input_type == IOConverter.InputType.TORCH: diff --git a/install.py b/install.py index a723d96..64a74dd 100644 --- a/install.py +++ b/install.py @@ -41,9 +41,13 @@ def initialization(): except ImportError: run_installation("piexif") try: - import chadet + import chardet except ImportError: run_installation("chardet") + try: + from imgutils.tagging import get_wd14_tags + except ImportError: + run_installation("dghs-imgutils[gpu]") def run_installation(pkg_name: str): print(f"Installing {pkg_name}...") @@ -54,3 +58,4 @@ def run_installation(pkg_name: str): if __name__ == "__main__": initialization() + print("Installation completed.") diff --git a/io_node.py b/io_node.py index 0ec37d7..594990b 100644 --- a/io_node.py +++ b/io_node.py @@ -10,7 +10,7 @@ try: except ImportError: piexif_loaded = False -from .imgio.converter import IOConverter, PILHandlingHodes +from .imgio.converter import PILHandlingHodes from .autonode import node_wrapper, get_node_names_mappings, validate, anytype, PILImage import time import os @@ -166,7 +166,7 @@ class SaveImageCustomNode: FUNCTION = "save_images" OUTPUT_NODE = True - + RESULT_NODE = True CATEGORY = "image" custom_name = "Save Image Custom Node" @@ -220,7 +220,9 @@ class SaveTextCustomNode: FUNCTION = "save_text" custom_name = "Save Text Custom Node" CATEGORY = "text" - + RESULT_NODE = True + OUTPUT_NODE = True + def save_text(self, text, filename_prefix="ComfyUI",subfolder_dir="",filename=""): text = str(text) assert len(text) > 0 and len(filename) > 0, "Text and filename must be non-empty" @@ -263,6 +265,7 @@ class SaveImageWebpCustomNode: FUNCTION = "save_images" OUTPUT_NODE = True + RESULT_NODE = True CATEGORY = "image" custom_name = "Save Image Webp Node" diff --git a/nodes.py b/nodes.py index 851e237..abe1bad 100644 --- a/nodes.py +++ b/nodes.py @@ -1,13 +1,16 @@ +from .install import initialization + +initialization() + from .logic_gates import CLASS_MAPPINGS as LogicMapping, CLASS_NAMES as LogicNames from .randomness import CLASS_MAPPINGS as RandomMapping, CLASS_NAMES as RandomNames from .conversion import CLASS_MAPPINGS as ConversionMapping, CLASS_NAMES as ConversionNames from .math_nodes import CLASS_MAPPINGS as MathMapping, CLASS_NAMES as MathNames from .io_node import CLASS_MAPPINGS as IOMapping, CLASS_NAMES as IONames - +from .auxilary import CLASS_MAPPINGS as AuxilaryMapping, CLASS_NAMES as AuxilaryNames from .external import CLASS_MAPPINGS as ExternalMapping, CLASS_NAMES as ExternalNames - NODE_CLASS_MAPPINGS = { } NODE_CLASS_MAPPINGS.update(IOMapping) @@ -16,6 +19,7 @@ NODE_CLASS_MAPPINGS.update(RandomMapping) NODE_CLASS_MAPPINGS.update(ConversionMapping) NODE_CLASS_MAPPINGS.update(MathMapping) NODE_CLASS_MAPPINGS.update(ExternalMapping) +NODE_CLASS_MAPPINGS.update(AuxilaryMapping) NODE_DISPLAY_NAME_MAPPINGS = { @@ -25,4 +29,5 @@ NODE_DISPLAY_NAME_MAPPINGS.update(LogicNames) NODE_DISPLAY_NAME_MAPPINGS.update(RandomNames) NODE_DISPLAY_NAME_MAPPINGS.update(ConversionNames) NODE_DISPLAY_NAME_MAPPINGS.update(MathNames) -NODE_DISPLAY_NAME_MAPPINGS.update(ExternalNames) \ No newline at end of file +NODE_DISPLAY_NAME_MAPPINGS.update(ExternalNames) +NODE_DISPLAY_NAME_MAPPINGS.update(AuxilaryNames) diff --git a/utils/tagger.py b/utils/tagger.py new file mode 100644 index 0000000..f75e0eb --- /dev/null +++ b/utils/tagger.py @@ -0,0 +1,28 @@ +try: + from imgutils.tagging import get_wd14_tags +except ImportError: + def get_wd14_tags(image_path): + raise Exception("Tagger feature not available, please install dghs-imgutils") +from typing import Union +from PIL import Image + +def get_rating_class(rating): + # argmax + return max(rating, key=rating.get) + +def get_tags_above_threshold(tags, threshold=0.4): + return [tag for tag, score in tags.items() if score > threshold] + +def replace_underscore(tag): + return tag.replace('_', ' ') + +def get_tags(image_path:Union[str, Image.Image], threshold:float = 0.4, replace:bool = False) -> dict[str, list[str]]: + result = {} + rating, features, chars = get_wd14_tags(image_path) + result['rating'] = get_rating_class(rating) + result['tags'] = get_tags_above_threshold(features, threshold) + result['chars'] = get_tags_above_threshold(chars, threshold) + if replace: + result['tags'] = [replace_underscore(tag) for tag in result['tags']] + result['chars'] = [replace_underscore(tag) for tag in result['chars']] + return result