diff --git a/wd14-tagger/README.md b/wd14-tagger/README.md new file mode 100644 index 0000000..bff0f18 --- /dev/null +++ b/wd14-tagger/README.md @@ -0,0 +1,17 @@ +Based on https://huggingface.co/spaces/SmilingWolf/wd-v1-4-tags +This requires onnxruntime or onnxruntime-gpu to run, so pip install it + +Links to the models are found at the top of the url above +Follow the link to one, e.g. convnextv2, go to Files +Download model.onnx + selected_tags.csv +Rename the model.onnx + csv the same thing, e.g. convnextv2.onnx + convnextv2.csv +Place them in comfy_extras/wd14_models + +Place wd14tagger.py in custom_nodes + +This currently requires this branch of ComfyUI to be of any use: +https://github.com/pythongosssss/ComfyUI/tree/widget2input + +You can right click the CLIPTextEncode node +convert text to input +then feed the results of this node into the text encode \ No newline at end of file diff --git a/wd14-tagger/wd14tagger.py b/wd14-tagger/wd14tagger.py new file mode 100644 index 0000000..da00460 --- /dev/null +++ b/wd14-tagger/wd14tagger.py @@ -0,0 +1,100 @@ +# https://huggingface.co/spaces/SmilingWolf/wd-v1-4-tags + +import numpy as np +from PIL import Image +import csv +import os + +NODE_CLASS_MAPPINGS = {} +valid = False + +try: + import onnxruntime as ort + from onnxruntime import InferenceSession + valid = True +except ImportError: + print("onnxruntime is required for wd14 tagger") + print("to use gpu") + print("pip install onnxruntime-gpu") + print("or to use cpu") + print("pip install onnxruntime") + +models_dir = os.path.realpath(os.path.join(os.path.dirname(os.path.realpath(__file__)), "../comfy_extras/wd14_models")) +if not os.path.exists(models_dir): + print("Place WD14 tagger models + tags in: " + models_dir) + print("You can download them from: https://huggingface.co/spaces/SmilingWolf/wd-v1-4-tags") + print("Name the model.onnx and selected_tags.csv something unique per model, e.g. convnextv2.onnx/.csv") +elif valid: + class WD14Tagger: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE", ), + "model": (sorted(filter(lambda x: x.endswith(".onnx"), os.listdir(models_dir))), ), + "threshold": ("FLOAT", {"default": 0.35, "min": 0.0, "max": 1, "step": 0.05}), + "character_threshold": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1, "step": 0.05}), + }} + + RETURN_TYPES = ("STRING",) + FUNCTION = "tag" + + CATEGORY = "image" + + def tag(self, image, model, threshold, character_threshold): + name = os.path.join(models_dir, model) + model = InferenceSession(name, providers=ort.get_available_providers()) + + input = model.get_inputs()[0] + height = input.shape[1] + + # Read all tags from csv and locate start of each category + tags = [] + general_index = None + character_index = None + with open(os.path.splitext(name)[0] + ".csv") as f: + reader = csv.reader(f) + next(reader) + for row in reader: + if general_index is None and row[2] == "0": + general_index = reader.line_num - 2 + elif character_index is None and row[2] == "4": + character_index = reader.line_num - 2 + tags.append(row[1]) + + tensor = image*255 + tensor = np.array(tensor, dtype=np.uint8) + if np.ndim(tensor) > 3: + assert tensor.shape[0] == 1 + tensor = tensor[0] + + image = Image.fromarray(tensor) + # Reduce to max size and pad with white + ratio = float(height)/max(image.size) + new_size = tuple([int(x*ratio) for x in image.size]) + image = image.resize(new_size, Image.ANTIALIAS) + square = Image.new("RGB", (height, height), (255, 255, 255)) + square.paste(image, ((height-new_size[0])//2, (height-new_size[1])//2)) + + image = np.array(square).astype(np.float32) + image = image[:, :, ::-1] # RGB -> BGR + image = np.expand_dims(image, 0) + + label_name = model.get_outputs()[0].name + probs = model.run([label_name], {input.name: image})[0] + + result = list(zip(tags, probs[0])) + + rating = max(result[:general_index], key=lambda x: x[1]) + general = [item for item in result[general_index:character_index] if item[1] > threshold] + character = [item for item in result[character_index:] if item[1] > character_threshold] + + res = ", ".join((item[0] for item in character + general)).replace(" ", + "_").replace("(", "\\(").replace(")", "\\)") + + print(res) + + return (res,) + + NODE_CLASS_MAPPINGS = { + "WD14Tagger": WD14Tagger, + }