add model name option

This commit is contained in:
aria1th
2024-12-01 00:40:46 +09:00
parent d3ad23d586
commit eece3b48a4
2 changed files with 31 additions and 17 deletions
+26 -15
View File
@@ -1,6 +1,6 @@
from .imgio.converter import PILHandlingHodes
from .autonode import node_wrapper, get_node_names_mappings, validate, anytype, PILImage
from .utils.tagger import get_tags
from .utils.tagger import get_tags, tagger_keys
auxilary_classes = []
auxilary_node = node_wrapper(auxilary_classes)
@@ -12,9 +12,9 @@ class GetRatingNode:
CATEGORY = "tagger"
custom_name = "Get Rating Class"
@staticmethod
def get_rating_class(image):
def get_rating_class(image, model_name):
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image)
result_dict = get_tags(image, model_name=model_name)
return (result_dict['rating'], )
@classmethod
def INPUT_TYPES(cls):
@@ -22,6 +22,9 @@ class GetRatingNode:
"required": {
"image": ("IMAGE",),
},
"optional": {
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@auxilary_node
@@ -31,9 +34,9 @@ class GetRatingFromTextNode:
CATEGORY = "tagger"
custom_name = "Get Rating Class From Text"
@staticmethod
def get_rating_class(image):
def get_rating_class(image, model_name):
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image)
result_dict = get_tags(image, model_name=model_name)
return (result_dict['rating'], )
@classmethod
def INPUT_TYPES(cls):
@@ -41,6 +44,9 @@ class GetRatingFromTextNode:
"required": {
"image": ("STRING", {"default": "/path/to/image.jpg"}),
},
"optional": {
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@auxilary_node
@@ -50,9 +56,9 @@ class GetTagsAboveThresholdNode:
CATEGORY = "tagger"
custom_name = "Get Tags Above Threshold"
@staticmethod
def get_tags_above_threshold(image, threshold, replace):
def get_tags_above_threshold(image, threshold, replace, model_name):
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image, threshold=threshold, replace=replace)
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
return (", ".join(result_dict['tags']), )
@classmethod
def INPUT_TYPES(cls):
@@ -63,6 +69,7 @@ class GetTagsAboveThresholdNode:
"optional": {
"threshold": ("FLOAT", {"default": 0.4}),
"replace": ("BOOLEAN", {"default": False}),
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@@ -73,9 +80,9 @@ class GetTagsAboveThresholdFromTextNode:
CATEGORY = "tagger"
custom_name = "Get Tags Above Threshold From Text"
@staticmethod
def get_tags_above_threshold(image, threshold, replace):
def get_tags_above_threshold(image, threshold, replace, model_name):
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image, threshold=threshold, replace=replace)
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
return (", ".join(result_dict['tags']), )
@classmethod
def INPUT_TYPES(cls):
@@ -86,6 +93,7 @@ class GetTagsAboveThresholdFromTextNode:
"optional": {
"threshold": ("FLOAT", {"default": 0.4}),
"replace": ("BOOLEAN", {"default": False}),
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@@ -96,9 +104,9 @@ class GetCharactersAboveThresholdNode:
CATEGORY = "tagger"
custom_name = "Get Chars Above Threshold"
@staticmethod
def get_tags_above_threshold(image, threshold, replace):
def get_tags_above_threshold(image, threshold, replace, model_name):
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image, threshold=threshold, replace=replace)
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
return (", ".join(result_dict['chars']), )
@classmethod
def INPUT_TYPES(cls):
@@ -109,6 +117,7 @@ class GetCharactersAboveThresholdNode:
"optional": {
"threshold": ("FLOAT", {"default": 0.4}),
"replace": ("BOOLEAN", {"default": False}),
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@@ -119,9 +128,9 @@ class GetCharactersAboveThresholdFromTextNode:
CATEGORY = "tagger"
custom_name = "Get Chars Above Threshold From Text"
@staticmethod
def get_tags_above_threshold(image, threshold, replace):
def get_tags_above_threshold(image, threshold, replace, model_name):
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image, threshold=threshold, replace=replace)
result_dict = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
return (", ".join(result_dict['chars']), )
@classmethod
def INPUT_TYPES(cls):
@@ -132,6 +141,7 @@ class GetCharactersAboveThresholdFromTextNode:
"optional": {
"threshold": ("FLOAT", {"default": 0.4}),
"replace": ("BOOLEAN", {"default": False}),
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@@ -142,9 +152,9 @@ class GetAllTagsAboveThresholdNode:
CATEGORY = "tagger"
custom_name = "Get All Tags Above Threshold"
@staticmethod
def get_tags(image, threshold, replace):
def get_tags(image, threshold, replace, model_name):
image = PILHandlingHodes.handle_input(image)
result = get_tags(image, threshold=threshold, replace=replace)
result = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
result_list = []
result_list.append(result['rating'])
result_list.extend(result['tags'])
@@ -159,6 +169,7 @@ class GetAllTagsAboveThresholdNode:
"optional": {
"threshold": ("FLOAT", {"default": 0.4}),
"replace": ("BOOLEAN", {"default": False}),
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
+5 -2
View File
@@ -1,5 +1,6 @@
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")
@@ -16,9 +17,9 @@ def get_tags_above_threshold(tags, threshold=0.4):
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]]:
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)
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)
@@ -26,3 +27,5 @@ def get_tags(image_path:Union[str, Image.Image], threshold:float = 0.4, 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())