This commit is contained in:
pythongosssss
2023-05-14 14:13:14 +01:00
parent 49206233a1
commit 1363997818
4 changed files with 1 additions and 220 deletions
+1 -2
View File
@@ -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
-14
View File
@@ -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
-62
View File
@@ -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;
};
}
},
});
-142
View File
@@ -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,
}