Merge pull request #5 from aria1th/dghs-imgutils
Dghs imgutils taggers node
This commit is contained in:
+28
@@ -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:
|
||||
|
||||
+177
@@ -0,0 +1,177 @@
|
||||
from .imgio.converter import PILHandlingHodes
|
||||
from .autonode import node_wrapper, get_node_names_mappings, validate, anytype, PILImage
|
||||
from .utils.tagger import get_tags, tagger_keys
|
||||
|
||||
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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result_dict = get_tags(image, model_name=model_name)
|
||||
return (result_dict['rating'], )
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
"optional": {
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
@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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result_dict = get_tags(image, model_name=model_name)
|
||||
return (result_dict['rating'], )
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("STRING", {"default": "/path/to/image.jpg"}),
|
||||
},
|
||||
"optional": {
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
@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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
|
||||
return (", ".join(result_dict['tags']), )
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
"optional": {
|
||||
"threshold": ("FLOAT", {"default": 0.4}),
|
||||
"replace": ("BOOLEAN", {"default": False}),
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
@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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
|
||||
return (", ".join(result_dict['tags']), )
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
"optional": {
|
||||
"threshold": ("FLOAT", {"default": 0.4}),
|
||||
"replace": ("BOOLEAN", {"default": False}),
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
@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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
|
||||
return (", ".join(result_dict['chars']), )
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
"optional": {
|
||||
"threshold": ("FLOAT", {"default": 0.4}),
|
||||
"replace": ("BOOLEAN", {"default": False}),
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
@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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
|
||||
return (", ".join(result_dict['chars']), )
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
"optional": {
|
||||
"threshold": ("FLOAT", {"default": 0.4}),
|
||||
"replace": ("BOOLEAN", {"default": False}),
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
@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, model_name):
|
||||
image = PILHandlingHodes.handle_input(image)
|
||||
result = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
|
||||
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}),
|
||||
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
|
||||
}
|
||||
}
|
||||
|
||||
CLASS_MAPPINGS, CLASS_NAMES = get_node_names_mappings(auxilary_classes)
|
||||
validate(auxilary_classes)
|
||||
+1
-1
@@ -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:
|
||||
|
||||
+6
-1
@@ -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.")
|
||||
|
||||
+6
-3
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(ExternalNames)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(AuxilaryNames)
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
try:
|
||||
from imgutils.tagging import get_wd14_tags
|
||||
from imgutils.tagging.wd14 import MODEL_NAMES as tagger_model_names
|
||||
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, model_name:str = "SwinV2") -> dict[str, list[str]]:
|
||||
result = {}
|
||||
rating, features, chars = get_wd14_tags(image_path, model_name)
|
||||
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
|
||||
|
||||
tagger_keys = list(tagger_model_names.keys())
|
||||
Reference in New Issue
Block a user