diff --git a/README.md b/README.md index d9b5d31..fb2a462 100644 --- a/README.md +++ b/README.md @@ -15,10 +15,11 @@ 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.

ExLlama allocates memory based on `max_seq_len`. Lowering it is a good way to save on VRAM. It's currently not possible to offload the model to RAM. 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).

ExLlama isn't deterministic, so the outputs may differ even with the same seed. -Previewer | Displays generated outputs in the UI and appends them to workflow metadata. -Replacer | Replaces variables enclosed in brackets, such as `[a]`, with their values. +Condition | Checks if the input meets some condition, interrupts processing otherwise. +Format | Replaces variables enclosed in brackets, such as `[a]`, with their values. +Preview | Displays generated outputs in the UI. ## Workflow -The image below can be opened in ComfyUI. The [model](https://huggingface.co/turboderp/Mistral-7B-instruct-exl2/tree/2.5bpw) uses around 3-4GB of VRAM depending on sequence length. +The image below can be opened in ComfyUI. ![workflow](https://github.com/Zuellni/ComfyUI-ExLlama-Nodes/assets/123005779/b68549c1-233a-4199-bb1a-7004e0638299) diff --git a/exllama.py b/exllama.py index 961407a..93a0e9c 100644 --- a/exllama.py +++ b/exllama.py @@ -14,7 +14,7 @@ class Loader: return { "required": { "model_dir": ("STRING", {"default": ""}), - "max_seq_len": ("INT", {"default": 2048, "min": 1, "max": 8192}), + "max_seq_len": ("INT", {"default": 2048, "max": 8192}), }, } @@ -23,28 +23,25 @@ class Loader: RETURN_NAMES = ("MODEL",) RETURN_TYPES = ("EXL_MODEL",) - def __init__(self): - self.model = None - def load(self, model_dir, max_seq_len): - del self.model collect() soft_empty_cache() config = ExLlamaV2Config() config.model_dir = model_dir config.prepare() - config.max_seq_len = max_seq_len - self.model = ExLlamaV2(config) - self.model.load() + if max_seq_len: + config.max_seq_len = max_seq_len - cache = ExLlamaV2Cache(self.model) + model = ExLlamaV2(config) + model.load() + + cache = ExLlamaV2Cache(model) tokenizer = ExLlamaV2Tokenizer(config) - generator = ExLlamaV2StreamingGenerator(self.model, cache, tokenizer) - settings = ExLlamaV2Sampler.Settings() + generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer) - return ((tokenizer, generator, settings),) + return ((tokenizer, generator),) class Generator: @@ -53,16 +50,21 @@ class Generator: return { "required": { "model": ("EXL_MODEL",), - "stop_on_newline": ("BOOLEAN", {"default": False}), - "max_tokens": ("INT", {"default": 128, "min": 1, "max": 8192}), + "max_new_tokens": ("INT", {"default": 128, "max": 8192}), "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}), - "typical": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), + "typical_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), "penalty": ("FLOAT", {"default": 1.15, "min": 1, "max": 2, "step": 0.01}), "seed": ("INT", {"max": 2**64 - 1}), + "stop_on_newline": ("BOOLEAN", {"default": False}), + "allowed_strings": ("STRING", {"default": ""}), "text": ("STRING", {"multiline": True}), }, + "hidden": { + "info": "EXTRA_PNGINFO", + "id": "UNIQUE_ID", + }, } CATEGORY = "Zuellni/ExLlama" @@ -73,51 +75,114 @@ class Generator: def generate( self, model, - stop_on_newline, - max_tokens, + max_new_tokens, temperature, top_k, top_p, - typical, + typical_p, penalty, seed, + stop_on_newline, + allowed_strings, text, + info=None, + id=None, ): + text = text.strip() + if not text: return ("",) - tokenizer, generator, settings = model - progress = ProgressBar(max_tokens) - prompt = tokenizer.encode(text) - + tokenizer, generator = model + text = tokenizer.encode(text) stop_conditions = [tokenizer.eos_token_id] - stop_on_newline and stop_conditions.append(tokenizer.newline_token_id) - generator.set_stop_conditions(stop_conditions) + if not max_new_tokens: + max_new_tokens = tokenizer.config.max_seq_len - text.shape[-1] + + if stop_on_newline: + stop_conditions.append(tokenizer.newline_token_id) + + settings = ExLlamaV2Sampler.Settings() settings.temperature = temperature settings.top_k = top_k settings.top_p = top_p - settings.typical = typical + settings.typical = typical_p settings.token_repetition_penalty = penalty + if allowed_strings: + strings = [] + + for string in allowed_strings.split(","): + string = string.strip() + + if "-" in string: + start, end = string.split("-") + + if start.isdigit() and end.isdigit(): + start, end = int(start), int(end) + + if start <= end: + strings.extend(map(str, range(start, end + 1))) + else: + strings.extend(map(str, range(start, end - 1, -1))) + elif len(start) == 1 and len(end) == 1: + start, end = ord(start), ord(end) + + if start <= end: + strings.extend(map(chr, range(start, end + 1))) + else: + strings.extend(map(chr, range(start, end + -1, -1))) + else: + strings.append(string) + else: + strings.append(string) + + allowed_strings = strings + allowed_tokens = tokenizer.encode(allowed_strings) + max_new_tokens = allowed_tokens.shape[-1] + + vocab_size = tokenizer.config.vocab_size + padding = vocab_size + (-vocab_size % 32) + + settings.token_bias = torch.full((padding,), float("-inf")) + settings.token_bias[allowed_tokens] = 0 + torch.manual_seed(seed) - generator.begin_stream(prompt, settings) + generator.set_stop_conditions(stop_conditions) + generator.begin_stream(text, settings) + progress = ProgressBar(max_new_tokens) start = time() eos = False output = "" tokens = 0 - while not eos and tokens < max_tokens: + while not eos and tokens < max_new_tokens: chunk, eos, _ = generator.stream() + + if allowed_strings: + c = (output + chunk).strip() + + if not any(c in s for s in allowed_strings): + break + progress.update(1) output += chunk tokens += 1 + output = output.strip() total = round(time() - start, 2) speed = round(tokens / total, 2) print(f"Output generated in {total} seconds ({tokens} tokens, {speed} tokens/s)") - return (output.strip(),) + if id and info and "workflow" in info: + nodes = info["workflow"]["nodes"] + node = next((n for n in nodes if str(n["id"]) == id), None) + + if node: + node["widgets_values"] = [output] + + return (output,) NODE_CLASS_MAPPINGS = { diff --git a/text.js b/text.js index 79eef99..828c395 100644 --- a/text.js +++ b/text.js @@ -2,9 +2,9 @@ import { app } from "../../../scripts/app.js"; import { ComfyWidgets } from "../../../scripts/widgets.js"; app.registerExtension({ - name: "ZuellniTextPreviewer", + name: "ZuellniTextPreview", async beforeRegisterNodeDef(nodeType, nodeData, app) { - if (nodeData.name === "ZuellniTextPreviewer") { + if (nodeData.name === "ZuellniTextPreview") { const onExecuted = nodeType.prototype.onExecuted; nodeType.prototype.onExecuted = function (message) { diff --git a/text.py b/text.py index 6488c82..43a9497 100644 --- a/text.py +++ b/text.py @@ -1,33 +1,52 @@ -class Previewer: +from comfy.model_management import InterruptProcessingException + + +class Condition: @classmethod def INPUT_TYPES(cls): return { "required": { - "text": ("STRING", {"forceInput": True}), + "a": ("STRING", {"forceInput": True}), + "condition": (["==", "!=", ">", ">=", "<", "<=", "in", "sw", "ew"],), + "b": ("STRING", {"default": ""}), }, - "hidden": { - "info": "EXTRA_PNGINFO", - "id": "UNIQUE_ID", + "optional": { + "text": ("STRING", {"forceInput": True, "multiline": True}), }, } CATEGORY = "Zuellni/Text" - FUNCTION = "preview" - OUTPUT_NODE = True - RETURN_TYPES = () + FUNCTION = "condition" + OUTPUT_Node = True + RETURN_NAMES = ("TEXT",) + RETURN_TYPES = ("STRING",) - def preview(self, text, info=None, id=None): - if id and info and "workflow" in info: - nodes = info["workflow"]["nodes"] - node = next((n for n in nodes if str(n["id"]) == id), None) + def condition(self, a, condition, b, text=None): + try: + a = float(a) + b = float(b) + except: + pass - if node: - node["widgets_values"] = [text] + conditions = { + "==": lambda: a == b, + "!=": lambda: a != b, + ">": lambda: a > b, + ">=": lambda: a >= b, + "<": lambda: a < b, + "<=": lambda: a <= b, + "in": lambda: str(a) in str(b), + "sw": lambda: str(a).startswith(str(b)), + "ew": lambda: str(a).endswith(str(b)), + } - return {"ui": {"text": [text]}} + if not conditions[condition](): + raise InterruptProcessingException() + + return (text,) -class Replacer: +class Format: @classmethod def INPUT_TYPES(cls): return { @@ -43,25 +62,42 @@ class Replacer: } CATEGORY = "Zuellni/Text" - FUNCTION = "replace" + FUNCTION = "format" RETURN_NAMES = ("TEXT",) RETURN_TYPES = ("STRING",) - def replace(self, text, **vars): + def format(self, text, **vars): for key, value in vars.items(): - text = text.replace(f"[{key}]", value) + 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 = { - "ZuellniTextPreviewer": Previewer, - "ZuellniTextReplacer": Replacer, + "ZuellniTextCondition": Condition, + "ZuellniTextFormat": Format, + "ZuellniTextPreview": Preview, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniTextPreviewer": "Preview Text", - "ZuellniTextReplacer": "Replace Text", + "ZuellniTextCondition": "Condition", + "ZuellniTextFormat": "Format", + "ZuellniTextPreview": "Preview", } WEB_DIRECTORY = "."