diff --git a/README.md b/README.md index 419f90b..8e9dd49 100644 --- a/README.md +++ b/README.md @@ -13,8 +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. +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.

ExLlama allocates memory based on `max_seq_len`. Lowering it is a good way to save on VRAM. +LoRA | Used to load LoRAs. The directory should contain `adapter_model.bin` or `.safetensors` and `adapter_config.json`. LoRA parameter count has to match the model. +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 93a0e9c..a76ee6b 100644 --- a/exllama.py +++ b/exllama.py @@ -1,14 +1,21 @@ -from gc import collect +import gc +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, ExLlamaV2Config, ExLlamaV2Tokenizer +from exllamav2 import ( + ExLlamaV2, + ExLlamaV2Cache_8bit, + ExLlamaV2Config, + ExLlamaV2Lora, + ExLlamaV2Tokenizer, +) from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator -class Loader: +class Model: @classmethod def INPUT_TYPES(cls): return { @@ -18,30 +25,85 @@ class Loader: }, } + CATEGORY = "Zuellni/ExLlama" + FUNCTION = "prepare" + RETURN_NAMES = ("MODEL",) + RETURN_TYPES = ("EXL_MODEL",) + + def __init__(self): + self.config = None + self.base = None + self.cache = None + self.tokenizer = None + self.generator = None + + def prepare(self, model_dir, max_seq_len): + self.unload() + + self.config = ExLlamaV2Config() + self.config.model_dir = model_dir + self.config.prepare() + + if max_seq_len: + self.config.max_seq_len = max_seq_len + + self.load() + + 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 + + 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_dir, max_seq_len): - collect() - soft_empty_cache() + def load(self, model, lora_dir): + model, loras = model - config = ExLlamaV2Config() - config.model_dir = model_dir - config.prepare() + lora = ExLlamaV2Lora.from_directory(model.load(), lora_dir) + loras = loras.copy() + loras.append(lora) - if max_seq_len: - config.max_seq_len = max_seq_len - - model = ExLlamaV2(config) - model.load() - - cache = ExLlamaV2Cache(model) - tokenizer = ExLlamaV2Tokenizer(config) - generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer) - - return ((tokenizer, generator),) + return ((model, loras),) class Generator: @@ -57,8 +119,10 @@ 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}), - "allowed_strings": ("STRING", {"default": ""}), + "allow_strings": ("BOOLEAN", {"default": False}), + "strings": ("STRING", {"default": ""}), "text": ("STRING", {"multiline": True}), }, "hidden": { @@ -72,6 +136,34 @@ 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, @@ -82,26 +174,28 @@ class Generator: typical_p, penalty, seed, + unload, stop_on_newline, - allowed_strings, + allow_strings, + 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, loras = model + + model.load() + text = 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 - text.shape[-1] if stop_on_newline: - stop_conditions.append(tokenizer.newline_token_id) + stop_conditions.append(model.tokenizer.newline_token_id) settings = ExLlamaV2Sampler.Settings() settings.temperature = temperature @@ -110,47 +204,23 @@ 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 + if strings: + strings = self.format(strings) + tokens = model.tokenizer.encode(strings) + vocab_size = model.config.vocab_size padding = vocab_size + (-vocab_size % 32) - settings.token_bias = torch.full((padding,), float("-inf")) - settings.token_bias[allowed_tokens] = 0 + 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") - torch.manual_seed(seed) - generator.set_stop_conditions(stop_conditions) - generator.begin_stream(text, settings) + random.seed(seed) + model.generator.set_stop_conditions(stop_conditions) + model.generator.begin_stream(text, settings, loras=loras) progress = ProgressBar(max_new_tokens) start = time() eos = False @@ -158,12 +228,12 @@ class Generator: tokens = 0 while not eos and tokens < max_new_tokens: - chunk, eos, _ = generator.stream() + chunk, eos, _ = model.generator.stream() - if allowed_strings: + if strings and allow_strings: c = (output + chunk).strip() - if not any(c in s for s in allowed_strings): + if not any(c in s for s in strings): break progress.update(1) @@ -175,6 +245,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) @@ -186,11 +259,13 @@ class Generator: NODE_CLASS_MAPPINGS = { - "ZuellniExLlamaLoader": Loader, + "ZuellniExLlamaModel": Model, + "ZuellniExLlamaLora": Lora, "ZuellniExLlamaGenerator": Generator, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniExLlamaLoader": "Loader", + "ZuellniExLlamaModel": "Model", + "ZuellniExLlamaLora": "LoRA", "ZuellniExLlamaGenerator": "Generator", } 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"