Merge pull request #5 from aria1th/dghs-imgutils

Dghs imgutils taggers node
This commit is contained in:
AngelBottomless
2024-12-01 00:41:20 +09:00
committed by GitHub
7 changed files with 257 additions and 8 deletions
+28
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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"
+8 -3
View File
@@ -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)
+31
View File
@@ -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())