diff --git a/README.md b/README.md index 419f90b..c45238c 100644 --- a/README.md +++ b/README.md @@ -13,10 +13,9 @@ If you see any ExLlama-related errors while loading, install it manually followi ## Nodes 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. -Condition | Checks if the input meets some condition, interrupts processing otherwise. -Format | Replaces variables enclosed in brackets, such as `[a]`, with their values. +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). +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 93a0e9c..33ba998 100644 --- a/exllama.py +++ b/exllama.py @@ -1,11 +1,14 @@ -from gc import collect +import gc +import random from time import time import torch +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 exllamav2 import ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Config, ExLlamaV2Tokenizer -from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator +from nodes import MAX_RESOLUTION as MAX class Loader: @@ -14,34 +17,55 @@ class Loader: return { "required": { "model_dir": ("STRING", {"default": ""}), - "max_seq_len": ("INT", {"default": 2048, "max": 8192}), + "max_seq_len": ("INT", {"default": 1024, "max": MAX}), }, } CATEGORY = "Zuellni/ExLlama" - FUNCTION = "load" + FUNCTION = "process" RETURN_NAMES = ("MODEL",) RETURN_TYPES = ("EXL_MODEL",) - def load(self, model_dir, max_seq_len): - collect() - soft_empty_cache() + def __init__(self): + self.config = None + self.base = None + self.cache = None + self.tokenizer = None + self.generator = None - config = ExLlamaV2Config() - config.model_dir = model_dir - config.prepare() + def process(self, model_dir, max_seq_len): + self.unload() + self.config = ExLlamaV2Config() + self.config.model_dir = model_dir + self.config.prepare() if max_seq_len: - config.max_seq_len = max_seq_len + self.config.max_seq_len = max_seq_len - model = ExLlamaV2(config) - model.load() + self.tokenizer = ExLlamaV2Tokenizer(self.config) + self.load() - cache = ExLlamaV2Cache(model) - tokenizer = ExLlamaV2Tokenizer(config) - generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer) + return (self,) - return ((tokenizer, generator),) + def load(self): + if not self.base: + self.base = ExLlamaV2(self.config) + self.base.load() + self.cache = ExLlamaV2Cache_8bit(self.base) + + self.generator = ExLlamaV2StreamingGenerator( + self.base, + self.cache, + self.tokenizer + ) + + def unload(self): + self.base = None + self.cache = None + self.generator = None + + gc.collect() + soft_empty_cache() class Generator: @@ -50,15 +74,15 @@ class Generator: return { "required": { "model": ("EXL_MODEL",), - "max_new_tokens": ("INT", {"default": 128, "max": 8192}), + "unload": ("BOOLEAN", {"default": False}), + "stop_on_newline": ("BOOLEAN", {"default": False}), + "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}), "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": { @@ -75,6 +99,8 @@ class Generator: def generate( self, model, + unload, + stop_on_newline, max_new_tokens, temperature, top_k, @@ -82,26 +108,25 @@ class Generator: typical_p, penalty, seed, - stop_on_newline, - allowed_strings, text, info=None, id=None, ): - text = text.strip() - if not text: return ("",) - tokenizer, generator = model - text = tokenizer.encode(text) - stop_conditions = [tokenizer.eos_token_id] + model.load() + input = model.tokenizer.encode(text) + stop_conditions = [model.tokenizer.eos_token_id] if not max_new_tokens: - max_new_tokens = tokenizer.config.max_seq_len - text.shape[-1] + max_new_tokens = model.config.max_seq_len - input.shape[-1] if stop_on_newline: - stop_conditions.append(tokenizer.newline_token_id) + stop_conditions.append(model.tokenizer.newline_token_id) + + model.generator.set_stop_conditions(stop_conditions) + random.seed(seed) settings = ExLlamaV2Sampler.Settings() settings.temperature = temperature @@ -110,47 +135,7 @@ class Generator: 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.set_stop_conditions(stop_conditions) - generator.begin_stream(text, settings) + model.generator.begin_stream(input, settings, token_healing=True) progress = ProgressBar(max_new_tokens) start = time() eos = False @@ -158,14 +143,7 @@ class Generator: tokens = 0 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 - + chunk, eos, _ = model.generator.stream() progress.update(1) output += chunk tokens += 1 @@ -175,6 +153,9 @@ class Generator: speed = round(tokens / total, 2) print(f"Output generated in {total} seconds ({tokens} tokens, {speed} tokens/s)") + if unload: + model.unload() + 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) diff --git a/requirements.txt b/requirements.txt index 7aea264..3815689 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ -exllamav2 +exllamav2>=0.0.7; platform_system == "Linux" +https://github.com/turboderp/exllamav2/releases/download/v0.0.7/exllamav2-0.0.7+cu121-cp311-cp311-win_amd64.whl; platform_system == "Windows" 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 43a9497..36439d7 100644 --- a/text.py +++ b/text.py @@ -1,52 +1,22 @@ -from comfy.model_management import InterruptProcessingException - - -class Condition: +class Preview: @classmethod def INPUT_TYPES(cls): return { "required": { - "a": ("STRING", {"forceInput": True}), - "condition": (["==", "!=", ">", ">=", "<", "<=", "in", "sw", "ew"],), - "b": ("STRING", {"default": ""}), - }, - "optional": { - "text": ("STRING", {"forceInput": True, "multiline": True}), - }, + "text": ("STRING", {"forceInput": True}), + } } CATEGORY = "Zuellni/Text" - FUNCTION = "condition" - OUTPUT_Node = True - RETURN_NAMES = ("TEXT",) - RETURN_TYPES = ("STRING",) + FUNCTION = "preview" + OUTPUT_NODE = True + RETURN_TYPES = () - def condition(self, a, condition, b, text=None): - try: - a = float(a) - b = float(b) - except: - pass - - 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)), - } - - if not conditions[condition](): - raise InterruptProcessingException() - - return (text,) + def preview(self, text): + return {"ui": {"text": [text]}} -class Format: +class Replace: @classmethod def INPUT_TYPES(cls): return { @@ -62,42 +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 = { - "ZuellniTextCondition": Condition, - "ZuellniTextFormat": Format, "ZuellniTextPreview": Preview, + "ZuellniTextReplace": Replace, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniTextCondition": "Condition", - "ZuellniTextFormat": "Format", "ZuellniTextPreview": "Preview", + "ZuellniTextReplace": "Replace", } WEB_DIRECTORY = "."