From 136399781810edb4e494d924f56bc29b96a0d170 Mon Sep 17 00:00:00 2001 From: pythongosssss <125205205+pythongosssss@users.noreply.github.com> Date: Sun, 14 May 2023 14:13:14 +0100 Subject: [PATCH] Moved --- README.md | 3 +- wd14-tagger/README.md | 14 ---- wd14-tagger/wd14tagger.js | 62 ----------------- wd14-tagger/wd14tagger.py | 142 -------------------------------------- 4 files changed, 1 insertion(+), 220 deletions(-) delete mode 100644 wd14-tagger/README.md delete mode 100644 wd14-tagger/wd14tagger.js delete mode 100644 wd14-tagger/wd14tagger.py diff --git a/README.md b/README.md index 330645c..a70c5cc 100644 --- a/README.md +++ b/README.md @@ -49,5 +49,4 @@ Takes input from a node that produces a string and displays it, useful for thing Provides basic support for touch screen devices, its not perfect but better than nothing ## WD14 Tagger -![image](https://user-images.githubusercontent.com/125205205/230175199-c1478840-a8c6-4428-8df9-52215e1410bc.png) -Uses WD14 Tagger to interrogate images +Moved to: https://github.com/pythongosssss/ComfyUI-WD14-Tagger \ No newline at end of file diff --git a/wd14-tagger/README.md b/wd14-tagger/README.md deleted file mode 100644 index b501ddd..0000000 --- a/wd14-tagger/README.md +++ /dev/null @@ -1,14 +0,0 @@ -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 - -You can right click the CLIPTextEncode node -convert text to input -then feed the results of this node into the text encode diff --git a/wd14-tagger/wd14tagger.js b/wd14-tagger/wd14tagger.js deleted file mode 100644 index 82e2a4d..0000000 --- a/wd14-tagger/wd14tagger.js +++ /dev/null @@ -1,62 +0,0 @@ -import { app } from "/scripts/app.js"; -import { ComfyWidgets } from "/scripts/widgets.js"; -import { api } from "/scripts/api.js"; - -// Displays the wd14 prompt - -app.registerExtension({ - name: "pysssss.Wd14Tagger", - async beforeRegisterNodeDef(nodeType, nodeData, app) { - if (nodeData.name === "WD14Tagger") { - const onNodeCreated = nodeType.prototype.onNodeCreated; - nodeType.prototype.onNodeCreated = function () { - const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined; - - const w = ComfyWidgets["STRING"](this, "tags", ["STRING", { multiline: true }], app).widget; - w.inputEl.readOnly = true; - w.inputEl.style.opacity = 0.6; - - api.addEventListener("wd14tagger", (e) => { - if (+app.runningNodeId === this.id) { - w.inputEl.value = e.detail; - if (this.size[1] < 180) { - this.size[1] = 180; - } - } - }); - - 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 deleted file mode 100644 index b897f8b..0000000 --- a/wd14-tagger/wd14tagger.py +++ /dev/null @@ -1,142 +0,0 @@ -# https://huggingface.co/spaces/SmilingWolf/wd-v1-4-tags - -import numpy as np -from PIL import Image -import csv -import os -from server import PromptServer -from aiohttp import web - -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") - -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": (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": ""}), - }} - - RETURN_TYPES = ("STRING",) - FUNCTION = "tag" - OUTPUT_NODE = True - - CATEGORY = "image" - - def tag(self, image, model, threshold, character_threshold, exclude_tags = ""): - 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) - - 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) - - return (res,) - - NODE_CLASS_MAPPINGS = { - "WD14Tagger": WD14Tagger, - }