diff --git a/README.md b/README.md index 33d6843..c45238c 100644 --- a/README.md +++ b/README.md @@ -15,7 +15,7 @@ Name | Description :--- | :--- Loader | Used to load EXL2/GPTQ Llama models. You can find a lot of them on [Hugging Face](https://huggingface.co/TheBloke). Clone the model repository or download all the files in it and place them in an empty directory, then specify the path in `model_dir`. The `model.safetensors` file won't work on its own. Generator | Generates a `string` based on the given input for use with other nodes. Default values correspond to the `simple-1` preset from [text-generation-webui](https://github.com/oobabooga/text-generation-webui). -Format | Replaces variables enclosed in brackets, such as `[a]`, with their values. +Replace | Replaces variables enclosed in brackets, such as `[a]`, with their values. Preview | Displays generated outputs in the UI. ## Workflow diff --git a/exllama.py b/exllama.py index 4e2f96e..33ba998 100644 --- a/exllama.py +++ b/exllama.py @@ -3,11 +3,13 @@ import random from time import time import torch -from comfy.model_management import soft_empty_cache -from comfy.utils import ProgressBar from exllamav2 import ExLlamaV2, ExLlamaV2Cache_8bit, ExLlamaV2Config, ExLlamaV2Tokenizer from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator +from comfy.model_management import soft_empty_cache +from comfy.utils import ProgressBar +from nodes import MAX_RESOLUTION as MAX + class Loader: @classmethod @@ -15,7 +17,7 @@ class Loader: return { "required": { "model_dir": ("STRING", {"default": ""}), - "max_seq_len": ("INT", {"default": 2048, "max": 8192}), + "max_seq_len": ("INT", {"default": 1024, "max": MAX}), }, } @@ -50,7 +52,12 @@ class Loader: self.base = ExLlamaV2(self.config) self.base.load() self.cache = ExLlamaV2Cache_8bit(self.base) - self.generator = ExLlamaV2StreamingGenerator(self.base, self.cache, self.tokenizer) + + self.generator = ExLlamaV2StreamingGenerator( + self.base, + self.cache, + self.tokenizer + ) def unload(self): self.base = None @@ -69,7 +76,7 @@ class Generator: "model": ("EXL_MODEL",), "unload": ("BOOLEAN", {"default": False}), "stop_on_newline": ("BOOLEAN", {"default": False}), - "max_new_tokens": ("INT", {"default": 128, "max": 8192}), + "max_new_tokens": ("INT", {"default": 128, "max": MAX}), "temperature": ("FLOAT", {"default": 0.7, "max": 2, "step": 0.01}), "top_k": ("INT", {"default": 20, "max": 200}), "top_p": ("FLOAT", {"default": 0.9, "max": 1, "step": 0.01}), diff --git a/text.js b/text.js index 828c395..3894b88 100644 --- a/text.js +++ b/text.js @@ -14,9 +14,9 @@ app.registerExtension({ const position = this.widgets.findIndex((w) => w.name === "text"); if (position !== -1) { - for (let i = position; i < this.widgets.length; i++) { + for (let i = position; i < this.widgets.length; i++) this.widgets[i].onRemove?.(); - } + this.widgets.length = position; } diff --git a/text.py b/text.py index 8a26f62..6dfc23b 100644 --- a/text.py +++ b/text.py @@ -1,4 +1,22 @@ -class Format: +class Preview: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "text": ("STRING", {"forceInput": True, "multiline": True}), + } + } + + CATEGORY = "Zuellni/Text" + FUNCTION = "preview" + OUTPUT_NODE = True + RETURN_TYPES = () + + def preview(self, text): + return {"ui": {"text": [text]}} + + +class Replace: @classmethod def INPUT_TYPES(cls): return { @@ -14,40 +32,26 @@ class Format: } CATEGORY = "Zuellni/Text" - FUNCTION = "format" + FUNCTION = "replace" RETURN_NAMES = ("TEXT",) RETURN_TYPES = ("STRING",) - def format(self, text, **vars): - for key, value in vars.items(): + def replace(self, text, **inputs): + for key, value in inputs.items(): if value: text = text.replace(f"[{key}]", value) return (text,) -class Preview: - @classmethod - def INPUT_TYPES(cls): - return {"required": {"text": ("STRING", {"forceInput": True})}} - - CATEGORY = "Zuellni/Text" - FUNCTION = "preview" - OUTPUT_NODE = True - RETURN_TYPES = () - - def preview(self, text): - return {"ui": {"text": [text]}} - - NODE_CLASS_MAPPINGS = { - "ZuellniTextFormat": Format, "ZuellniTextPreview": Preview, + "ZuellniTextReplace": Replace, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniTextFormat": "Format", "ZuellniTextPreview": "Preview", + "ZuellniTextReplace": "Replace", } WEB_DIRECTORY = "."