diff --git a/README.md b/README.md index 0951199..33d6843 100644 --- a/README.md +++ b/README.md @@ -13,10 +13,8 @@ If you see any ExLlama-related errors while loading, install it manually followi ## Nodes Name | Description :--- | :--- -Model | 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. -LoRA | Used to load LoRAs. The directory should contain `adapter_model.bin`/`adapter_model.safetensors` and `adapter_config.json`. LoRA parameter count has to match the model. +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). -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. diff --git a/exllama.py b/exllama.py index a76ee6b..c9b7c05 100644 --- a/exllama.py +++ b/exllama.py @@ -5,17 +5,11 @@ 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, - ExLlamaV2Lora, - ExLlamaV2Tokenizer, -) +from exllamav2 import ExLlamaV2, ExLlamaV2Cache_8bit, ExLlamaV2Config, ExLlamaV2Tokenizer from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator -class Model: +class Loader: @classmethod def INPUT_TYPES(cls): return { @@ -26,7 +20,7 @@ class Model: } CATEGORY = "Zuellni/ExLlama" - FUNCTION = "prepare" + FUNCTION = "process" RETURN_NAMES = ("MODEL",) RETURN_TYPES = ("EXL_MODEL",) @@ -37,9 +31,8 @@ class Model: self.tokenizer = None self.generator = None - def prepare(self, model_dir, max_seq_len): + def process(self, model_dir, max_seq_len): self.unload() - self.config = ExLlamaV2Config() self.config.model_dir = model_dir self.config.prepare() @@ -47,63 +40,25 @@ class Model: if max_seq_len: self.config.max_seq_len = max_seq_len + self.tokenizer = ExLlamaV2Tokenizer(self.config) self.load() - return ((self, []),) + return (self,) def load(self): if not self.base: self.base = ExLlamaV2(self.config) self.base.load() - self.cache = ExLlamaV2Cache_8bit(self.base) - self.tokenizer = ExLlamaV2Tokenizer(self.config) - - self.generator = ExLlamaV2StreamingGenerator( - self.base, - self.cache, - self.tokenizer, - ) - - return self.base + self.generator = ExLlamaV2StreamingGenerator(self.base, self.cache, self.tokenizer) def unload(self): - if self.base: - self.base.unload() - - del self.base, self.cache, self.tokenizer, self.generator - gc.collect() - soft_empty_cache() - self.base = None self.cache = None - self.tokenizer = None self.generator = None - -class Lora: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("EXL_MODEL",), - "lora_dir": ("STRING", {"default": ""}), - }, - } - - CATEGORY = "Zuellni/ExLlama" - FUNCTION = "load" - RETURN_NAMES = ("MODEL",) - RETURN_TYPES = ("EXL_MODEL",) - - def load(self, model, lora_dir): - model, loras = model - - lora = ExLlamaV2Lora.from_directory(model.load(), lora_dir) - loras = loras.copy() - loras.append(lora) - - return ((model, loras),) + gc.collect() + soft_empty_cache() class Generator: @@ -112,6 +67,8 @@ class Generator: return { "required": { "model": ("EXL_MODEL",), + "unload": ("BOOLEAN", {"default": False}), + "stop_on_newline": ("BOOLEAN", {"default": False}), "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}), @@ -119,10 +76,6 @@ class Generator: "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}), - "unload": ("BOOLEAN", {"default": False}), - "stop_on_newline": ("BOOLEAN", {"default": False}), - "allow_strings": ("BOOLEAN", {"default": False}), - "strings": ("STRING", {"default": ""}), "text": ("STRING", {"multiline": True}), }, "hidden": { @@ -136,37 +89,11 @@ class Generator: RETURN_NAMES = ("TEXT",) RETURN_TYPES = ("STRING",) - def format(self, strings): - list = [] - - for string in strings.split(","): - if "-" in string: - start, end = string.split("-") - - if start.isdigit() and end.isdigit(): - start, end = int(start), int(end) - - if start <= end: - list.extend(map(str, range(start, end + 1))) - else: - list.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: - list.extend(map(chr, range(start, end + 1))) - else: - list.extend(map(chr, range(start, end + -1, -1))) - else: - list.append(string) - else: - list.append(string) - - return list - def generate( self, model, + unload, + stop_on_newline, max_new_tokens, temperature, top_k, @@ -174,10 +101,6 @@ class Generator: typical_p, penalty, seed, - unload, - stop_on_newline, - allow_strings, - strings, text, info=None, id=None, @@ -185,18 +108,19 @@ class Generator: if not text: return ("",) - model, loras = model - model.load() - text = model.tokenizer.encode(text) + input = model.tokenizer.encode(text) stop_conditions = [model.tokenizer.eos_token_id] - if not max_new_tokens: - max_new_tokens = model.config.max_seq_len - text.shape[-1] - if stop_on_newline: stop_conditions.append(model.tokenizer.newline_token_id) + if not max_new_tokens: + max_new_tokens = model.config.max_seq_len - input.shape[-1] + + model.generator.set_stop_conditions(stop_conditions) + random.seed(seed) + settings = ExLlamaV2Sampler.Settings() settings.temperature = temperature settings.top_k = top_k @@ -204,23 +128,7 @@ class Generator: settings.typical = typical_p settings.token_repetition_penalty = penalty - if strings: - strings = self.format(strings) - tokens = model.tokenizer.encode(strings) - vocab_size = model.config.vocab_size - padding = vocab_size + (-vocab_size % 32) - - if allow_strings: - settings.token_bias = torch.full((padding,), float("-inf")) - settings.token_bias[tokens] = 0 - max_new_tokens = tokens.shape[-1] - else: - settings.token_bias = torch.zeros((padding,)) - settings.token_bias[tokens] = float("-inf") - - random.seed(seed) - model.generator.set_stop_conditions(stop_conditions) - model.generator.begin_stream(text, settings, loras=loras) + model.generator.begin_stream(input, settings, token_healing=True) progress = ProgressBar(max_new_tokens) start = time() eos = False @@ -229,13 +137,6 @@ class Generator: while not eos and tokens < max_new_tokens: chunk, eos, _ = model.generator.stream() - - if strings and allow_strings: - c = (output + chunk).strip() - - if not any(c in s for s in strings): - break - progress.update(1) output += chunk tokens += 1 @@ -259,13 +160,11 @@ class Generator: NODE_CLASS_MAPPINGS = { - "ZuellniExLlamaModel": Model, - "ZuellniExLlamaLora": Lora, + "ZuellniExLlamaLoader": Loader, "ZuellniExLlamaGenerator": Generator, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniExLlamaModel": "Model", - "ZuellniExLlamaLora": "LoRA", + "ZuellniExLlamaLoader": "Loader", "ZuellniExLlamaGenerator": "Generator", } diff --git a/text.py b/text.py index 43a9497..8a26f62 100644 --- a/text.py +++ b/text.py @@ -1,51 +1,3 @@ -from comfy.model_management import InterruptProcessingException - - -class Condition: - @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}), - }, - } - - CATEGORY = "Zuellni/Text" - FUNCTION = "condition" - OUTPUT_Node = True - RETURN_NAMES = ("TEXT",) - RETURN_TYPES = ("STRING",) - - 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,) - - class Format: @classmethod def INPUT_TYPES(cls): @@ -89,13 +41,11 @@ class Preview: NODE_CLASS_MAPPINGS = { - "ZuellniTextCondition": Condition, "ZuellniTextFormat": Format, "ZuellniTextPreview": Preview, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniTextCondition": "Condition", "ZuellniTextFormat": "Format", "ZuellniTextPreview": "Preview", }