diff --git a/exllama.py b/exllama.py index 052c69e..9ba7200 100644 --- a/exllama.py +++ b/exllama.py @@ -179,7 +179,7 @@ class Tokenizer: return { "required": { "model": ("EXL_MODEL",), - "text": ("STRING", {"forceInput": True, "multiline": True}), + "text": ("STRING", {"forceInput": True}), "add_bos_token": ("BOOLEAN", {"default": True}), "encode_special_tokens": ("BOOLEAN", {"default": True}), }, diff --git a/text.js b/text.js index 40d7af2..75919a5 100644 --- a/text.js +++ b/text.js @@ -1,33 +1,60 @@ -import { app } from "../../../scripts/app.js"; -import { ComfyWidgets } from "../../../scripts/widgets.js"; +import { app } from "../../../scripts/app.js" app.registerExtension({ name: "ZuellniText", async beforeRegisterNodeDef(nodeType, nodeData, app) { - if (nodeData.name === "ZuellniTextPreview") { - const onExecuted = nodeType.prototype.onExecuted; + if (nodeData.category != "Zuellni/Text") + return + + const onNodeCreated = nodeType.prototype.onNodeCreated + const onExecuted = nodeType.prototype.onExecuted + + if (nodeData.name == "ZuellniTextPreview") { + nodeType.prototype.onNodeCreated = function() { + const output = this.widgets.find(w => w.name == "output") + + if (output) { + output.inputEl.placeholder = "" + output.inputEl.readOnly = true + output.inputEl.style.cursor = "default" + output.inputEl.style.opacity = 0.7 + } + + this.setSize(this.computeSize()); + return onNodeCreated?.apply(this, arguments) + } nodeType.prototype.onExecuted = function(message) { - onExecuted?.apply(this, arguments); + const output = this.widgets.find(w => w.name == "output") + output && (output.value = message.text) + return onExecuted?.apply(this, arguments) + } + } else if (nodeData.name == "ZuellniTextReplace") { + nodeType.prototype.onNodeCreated = function() { + const count = this.widgets.find(w => w.name == "count") - if (this.widgets) { - const index = this.widgets.findIndex((w) => w.name === "output"); - - if (index !== -1) { - for (let i = index; i < this.widgets.length; i++) - this.widgets[i].onRemove?.(); - - this.widgets.length = index; - } - - const options = ["STRING", {multiline: true }]; - const widget = ComfyWidgets["STRING"](this, "output", options, app).widget; - - widget.inputEl.readOnly = true; - widget.inputEl.style.opacity = 0.7; - widget.value = message.text; + if (count) { + count.callback = () => this.onChanged(count.value) + this.onChanged(count.value) } - }; + + return onNodeCreated?.apply(this, arguments) + } + + nodeType.prototype.onChanged = function(count) { + !this.inputs && (this.inputs = []) + const current = this.inputs.length + + if (current == count) + return + + if (current < count) + for (let i = current; i < count; i++) + this.addInput(String.fromCharCode(i + 97), "STRING") + else + for (let i = current - 1; i >= count; i--) + this.removeInput(i) + } } - }, -}); + } +}) diff --git a/text.py b/text.py index 0841c08..7592cb3 100644 --- a/text.py +++ b/text.py @@ -1,7 +1,36 @@ +import string + _CATEGORY = "Zuellni/Text" _MAPPING = "ZuellniText" +class Convert: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "text": ("STRING", {"forceInput": True}), + "strip": (("punctuation", "whitespace", "both", False),), + "case": (("lower", "upper", "capitalize", False),), + }, + } + + CATEGORY = _CATEGORY + FUNCTION = "convert" + RETURN_NAMES = ("TEXT",) + RETURN_TYPES = ("STRING",) + + def convert(self, text, strip, case): + if strip == "both": + text = text.strip(string.punctuation + string.whitespace) + elif strip: + text = text.strip(getattr(string, strip)) + if case: + text = getattr(text, case)() + + return (text,) + + class Message: @classmethod def INPUT_TYPES(cls): @@ -25,14 +54,19 @@ class Message: class Preview: @classmethod def INPUT_TYPES(cls): - return {"required": {"text": ("STRING", {"forceInput": True})}} + return { + "required": { + "text": ("STRING", {"forceInput": True}), + "output": ("STRING", {"multiline": True}), + }, + } CATEGORY = _CATEGORY FUNCTION = "preview" OUTPUT_NODE = True RETURN_TYPES = () - def preview(self, text): + def preview(self, text, output): return {"ui": {"text": [text]}} @@ -40,12 +74,9 @@ class Replace: @classmethod def INPUT_TYPES(cls): return { - "required": {"text": ("STRING", {"multiline": True})}, - "optional": { - "a": ("STRING", {"forceInput": True, "multiline": True}), - "b": ("STRING", {"forceInput": True, "multiline": True}), - "c": ("STRING", {"forceInput": True, "multiline": True}), - "d": ("STRING", {"forceInput": True, "multiline": True}), + "required": { + "count": ("INT", {"default": 1, "min": 1, "max": 26}), + "text": ("STRING", {"multiline": True}), }, } @@ -54,24 +85,44 @@ class Replace: RETURN_NAMES = ("TEXT",) RETURN_TYPES = ("STRING",) - def replace(self, text, **inputs): - for key, value in inputs.items(): - if value: - text = text.replace(f"[{key}]", value) + def replace(self, count, text="", **kwargs): + for index in range(count): + key = chr(index + 97) + + if key in kwargs and kwargs[key]: + text = text.replace(f"{{{key}}}", kwargs[key]) return (text,) +class String: + @classmethod + def INPUT_TYPES(cls): + return {"required": {"text": ("STRING", {"multiline": True})}} + + CATEGORY = _CATEGORY + FUNCTION = "get" + RETURN_NAMES = ("TEXT",) + RETURN_TYPES = ("STRING",) + + def get(self, text): + return (text,) + + NODE_CLASS_MAPPINGS = { + f"{_MAPPING}Convert": Convert, f"{_MAPPING}Message": Message, f"{_MAPPING}Preview": Preview, f"{_MAPPING}Replace": Replace, + f"{_MAPPING}String": String, } NODE_DISPLAY_NAME_MAPPINGS = { + f"{_MAPPING}Convert": "Convert", f"{_MAPPING}Message": "Message", f"{_MAPPING}Preview": "Preview", f"{_MAPPING}Replace": "Replace", + f"{_MAPPING}String": "String", } WEB_DIRECTORY = "."