diff --git a/wd14-tagger/wd14tagger.js b/wd14-tagger/wd14tagger.js index cfa47b9..82e2a4d 100644 --- a/wd14-tagger/wd14tagger.js +++ b/wd14-tagger/wd14tagger.js @@ -25,6 +25,36 @@ app.registerExtension({ } }); + return r; + }; + } else if (nodeData.name === "LoadImage") { + const onNodeCreated = nodeType.prototype.onNodeCreated; + + const BUTTON_TEXT = "WD14 Interrogate"; + const BUTTON_TEXT_LOADING = "Loading..."; + nodeType.prototype.onNodeCreated = function () { + const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined; + const btn = this.addWidget("button", BUTTON_TEXT, "interrogate", () => { + if (btn.name === BUTTON_TEXT_LOADING) return; + + btn.name = BUTTON_TEXT_LOADING; + app.canvas.setDirty(true); + (async () => { + try { + const w = this.widgets.find((w) => w.name === "image"); + if (w?.value) { + const tags = await ( + await fetch("/pysssss/wd14tagger?type=input&image=" + encodeURIComponent(w.value)) + ).json(); + alert(tags); + } + } finally { + btn.name = BUTTON_TEXT; + app.canvas.setDirty(true); + } + })(); + }); + return r; }; } diff --git a/wd14-tagger/wd14tagger.py b/wd14-tagger/wd14tagger.py index d6c69e4..b897f8b 100644 --- a/wd14-tagger/wd14tagger.py +++ b/wd14-tagger/wd14tagger.py @@ -5,6 +5,7 @@ from PIL import Image import csv import os from server import PromptServer +from aiohttp import web NODE_CLASS_MAPPINGS = {} valid = False @@ -20,18 +21,95 @@ except ImportError: 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")) +root_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..") +models_dir = os.path.abspath(os.path.join(root_dir, "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: + def get_models(): + return sorted(filter(lambda x: x.endswith(".onnx"), os.listdir(models_dir))) + + def tag(image, model, threshold = 0.35, character_threshold = 0.85, exclude_tags = ""): + name = os.path.join(models_dir, model) + model = InferenceSession(name, providers=ort.get_available_providers()) + + input = model.get_inputs()[0] + height = input.shape[1] + + # 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) + + # 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]) + + 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] + + all = character + general + remove = [s.strip() for s in exclude_tags.lower().split(",")] + all = [tag for tag in all if tag[0] not in remove] + + res = ", ".join((item[0].replace("(", "\\(").replace(")", "\\)") for item in all)) + + print(res) + return res + + @PromptServer.instance.routes.get("/pysssss/wd14tagger") + async def get_tags(request): + if "image" not in request.query: + return web.Response(status=404) + + type = request.rel_url.query.get("type", "output") + if type not in ["output", "input", "temp"]: + return web.Response(status=400) + + target_dir = os.path.abspath(os.path.join(root_dir, type)) + image_path = os.path.abspath(os.path.join(target_dir, request.query["image"])) + c = os.path.commonpath((image_path, target_dir)) + if os.path.commonpath((image_path, target_dir)) != target_dir: + return web.Response(status=403) + + if not os.path.isfile(image_path): + return web.Response(status=404) + + image = Image.open(image_path) + + return web.json_response(tag(image, get_models()[0])) + class WD14Tagger: @classmethod def INPUT_TYPES(s): return {"required": { "image": ("IMAGE", ), - "model": (sorted(filter(lambda x: x.endswith(".onnx"), os.listdir(models_dir))), ), + "model": (get_models(), ), "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}), "exclude_tags": ("STRING", {"default": ""}), @@ -44,26 +122,6 @@ elif valid: CATEGORY = "image" def tag(self, image, model, threshold, character_threshold, exclude_tags = ""): - 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: @@ -71,33 +129,8 @@ elif valid: 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] - - all = character + general - remove = [s.strip() for s in exclude_tags.lower().split(",")] - all = [tag for tag in all if tag[0] not in remove] - - res = ", ".join((item[0].replace("(", "\\(").replace(")", "\\)") for item in all)) - - print(res) + res = tag(image, model, threshold, character_threshold, exclude_tags) if PromptServer.instance.client_id is not None: PromptServer.instance.send_sync("wd14tagger", res, PromptServer.instance.client_id)