From baed69aba5849a8b0916fb345adec4623b56e83e Mon Sep 17 00:00:00 2001 From: Zuellni <123005779+Zuellni@users.noreply.github.com> Date: Sun, 7 Apr 2024 19:35:32 +0200 Subject: [PATCH] Bump minimum exllamav2 version, add autosplit and q4 cache, some reformatting --- exllama.py | 77 ++++++++++++++++++++------------------ requirements-no-wheels.txt | 2 - requirements-torch-21.txt | 2 - requirements-torch-22.txt | 2 - requirements.txt | 1 + text.js | 4 +- text.py | 30 +++++++-------- 7 files changed, 57 insertions(+), 61 deletions(-) delete mode 100644 requirements-no-wheels.txt delete mode 100644 requirements-torch-21.txt delete mode 100644 requirements-torch-22.txt create mode 100644 requirements.txt diff --git a/exllama.py b/exllama.py index be2af53..3629757 100644 --- a/exllama.py +++ b/exllama.py @@ -10,12 +10,16 @@ from exllamav2 import ( ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Cache_8bit, + ExLlamaV2Cache_Q4, ExLlamaV2Config, ExLlamaV2Tokenizer, ) from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator from folder_paths import add_model_folder_path, get_folder_paths, models_dir +_CATEGORY = "Zuellni/ExLlama" +_MAPPING = "ZuellniExLlama" + class Loader: @classmethod @@ -34,59 +38,57 @@ class Loader: return { "required": { "model": (models, {"default": default}), - "gpu_split": ("STRING", {"default": ""}), - "cache_8bit": ("BOOLEAN", {"default": False}), - "max_seq_len": ("INT", {"default": 1024, "max": 2**16}), + "cache_bits": ((4, 8, 16), {"default": 4}), + "max_seq_len": ("INT", {"default": 2048, "max": 2**20}), }, } _MODELS = {} - CATEGORY = "Zuellni/ExLlama" + CATEGORY = _CATEGORY FUNCTION = "setup" RETURN_NAMES = ("MODEL",) RETURN_TYPES = ("EXL_MODEL",) - def setup(self, model, gpu_split, cache_8bit, max_seq_len): + def setup(self, model, cache_bits, max_seq_len): self.unload() + self.cache_bits = cache_bits + self.config = ExLlamaV2Config() self.config.model_dir = __class__._MODELS[model] self.config.prepare() if max_seq_len: self.config.max_seq_len = max_seq_len - - self.gpu_split = [float(a) for a in gpu_split.split(",") if gpu_split] - self.cache_8bit = cache_8bit + self.config.max_input_len = max_seq_len + self.config.max_attention_len = max_seq_len**2 return (self,) def load(self): if ( - hasattr(self, "model") and - hasattr(self, "cache") and - hasattr(self, "tokenizer") and - hasattr(self, "generator") and - self.model and - self.cache and - self.tokenizer and - self.generator + hasattr(self, "model") + and hasattr(self, "cache") + and hasattr(self, "tokenizer") + and hasattr(self, "generator") + and self.model + and self.cache + and self.tokenizer + and self.generator ): return self.model = ExLlamaV2(self.config) - progress = ProgressBar(len(self.model.modules)) - - self.model.load( - gpu_split=self.gpu_split, - callback=lambda s, _: progress.update_absolute(s), - ) + progress = ProgressBar(len(self.model.modules) + 1) self.cache = ( - ExLlamaV2Cache_8bit(self.model) - if self.cache_8bit - else ExLlamaV2Cache(self.model) + ExLlamaV2Cache_Q4(self.model, lazy=True) + if self.cache_bits == 4 + else ExLlamaV2Cache_8bit(self.model, lazy=True) + if self.cache_bits == 8 + else ExLlamaV2Cache(self.model, lazy=True) ) + self.model.load_autosplit(self.cache, callback=lambda _, __: progress.update(1)) self.tokenizer = ExLlamaV2Tokenizer(self.config) self.generator = ExLlamaV2StreamingGenerator( @@ -116,14 +118,14 @@ class Generator: "model": ("EXL_MODEL",), "unload": ("BOOLEAN", {"default": False}), "single_line": ("BOOLEAN", {"default": False}), - "max_tokens": ("INT", {"default": 128, "max": 2**16}), + "max_tokens": ("INT", {"default": 128, "max": 2**20}), "temperature": ("FLOAT", {"default": 1, "max": 5, "step": 0.01}), "top_k": ("INT", {"max": 200}), "top_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), "typical_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), "min_p": ("FLOAT", {"max": 1, "step": 0.01}), "top_a": ("FLOAT", {"max": 1, "step": 0.01}), - "penalty": ("FLOAT", {"default": 1, "min": 1, "max": 3, "step": 0.01}), + "repetition_penalty": ("FLOAT", {"default": 1, "min": 1, "max": 3, "step": 0.01}), "temperature_last": ("BOOLEAN", {"default": True}), "seed": ("INT", {"max": 2**64 - 1}), "text": ("STRING", {"multiline": True}), @@ -134,7 +136,7 @@ class Generator: }, } - CATEGORY = "Zuellni/ExLlama" + CATEGORY = _CATEGORY FUNCTION = "generate" RETURN_NAMES = ("TEXT",) RETURN_TYPES = ("STRING",) @@ -151,7 +153,7 @@ class Generator: typical_p, min_p, top_a, - penalty, + repetition_penalty, temperature_last, seed, text, @@ -187,20 +189,21 @@ class Generator: settings.typical = typical_p settings.min_p = min_p settings.top_a = top_a - settings.token_repetition_penalty = penalty + settings.token_repetition_penalty = repetition_penalty settings.temperature_last = temperature_last start = time() - model.generator.begin_stream(input, settings) + model.generator.begin_stream_ex(input, settings) progress = ProgressBar(max_tokens) eos = False output = "" tokens = 0 while not eos and tokens < max_tokens: - chunk, eos, _ = model.generator.stream() + response = model.generator.stream_ex() + output += response["chunk"] + eos = response["eos"] progress.update(1) - output += chunk tokens += 1 output = output.strip() @@ -226,11 +229,11 @@ class Generator: NODE_CLASS_MAPPINGS = { - "ZuellniExLlamaLoader": Loader, - "ZuellniExLlamaGenerator": Generator, + f"{_MAPPING}Loader": Loader, + f"{_MAPPING}Generator": Generator, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniExLlamaLoader": "Loader", - "ZuellniExLlamaGenerator": "Generator", + f"{_MAPPING}Loader": "Loader", + f"{_MAPPING}Generator": "Generator", } diff --git a/requirements-no-wheels.txt b/requirements-no-wheels.txt deleted file mode 100644 index 744601d..0000000 --- a/requirements-no-wheels.txt +++ /dev/null @@ -1,2 +0,0 @@ -exllamav2 -flash-attn diff --git a/requirements-torch-21.txt b/requirements-torch-21.txt deleted file mode 100644 index 44430f0..0000000 --- a/requirements-torch-21.txt +++ /dev/null @@ -1,2 +0,0 @@ -https://github.com/turboderp/exllamav2/releases/download/v0.0.12/exllamav2-0.0.12+cu121-cp311-cp311-win_amd64.whl -https://github.com/bdashore3/flash-attention/releases/download/v2.4.2/flash_attn-2.4.2+cu122torch2.1.2cxx11abiFALSE-cp311-cp311-win_amd64.whl diff --git a/requirements-torch-22.txt b/requirements-torch-22.txt deleted file mode 100644 index a32d28e..0000000 --- a/requirements-torch-22.txt +++ /dev/null @@ -1,2 +0,0 @@ -https://github.com/turboderp/exllamav2/releases/download/v0.0.14/exllamav2-0.0.14+cu121-cp311-cp311-win_amd64.whl -https://github.com/bdashore3/flash-attention/releases/download/v2.5.2/flash_attn-2.5.2+cu122torch2.2.0cxx11abiFALSE-cp311-cp311-win_amd64.whl diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..0baa9d6 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +exllamav2>=0.0.16 diff --git a/text.js b/text.js index f182000..68097f3 100644 --- a/text.js +++ b/text.js @@ -4,7 +4,7 @@ import { ComfyWidgets } from "../../../scripts/widgets.js"; app.registerExtension({ name: "ZuellniText", async beforeRegisterNodeDef(nodeType, nodeData, app) { - if (nodeData.name === "ZuellniTextPreview") { + if (nodeData.name === "ZuellniTextPreviewer") { const onExecuted = nodeType.prototype.onExecuted; nodeType.prototype.onExecuted = function(message) { @@ -20,7 +20,7 @@ app.registerExtension({ this.widgets.length = index; } - const options = ["STRING", {multiline: true }] + const options = ["STRING", {multiline: true }]; const widget = ComfyWidgets["STRING"](this, "output", options, app).widget; widget.inputEl.readOnly = true; diff --git a/text.py b/text.py index 36439d7..346211b 100644 --- a/text.py +++ b/text.py @@ -1,13 +1,13 @@ -class Preview: +_CATEGORY = "Zuellni/Text" +_MAPPING = "ZuellniText" + + +class Previewer: @classmethod def INPUT_TYPES(cls): - return { - "required": { - "text": ("STRING", {"forceInput": True}), - } - } + return {"required": {"text": ("STRING", {"forceInput": True})}} - CATEGORY = "Zuellni/Text" + CATEGORY = _CATEGORY FUNCTION = "preview" OUTPUT_NODE = True RETURN_TYPES = () @@ -16,13 +16,11 @@ class Preview: return {"ui": {"text": [text]}} -class Replace: +class Replacer: @classmethod def INPUT_TYPES(cls): return { - "required": { - "text": ("STRING", {"multiline": True}), - }, + "required": {"text": ("STRING", {"multiline": True})}, "optional": { "a": ("STRING", {"forceInput": True, "multiline": True}), "b": ("STRING", {"forceInput": True, "multiline": True}), @@ -31,7 +29,7 @@ class Replace: }, } - CATEGORY = "Zuellni/Text" + CATEGORY = _CATEGORY FUNCTION = "replace" RETURN_NAMES = ("TEXT",) RETURN_TYPES = ("STRING",) @@ -45,13 +43,13 @@ class Replace: NODE_CLASS_MAPPINGS = { - "ZuellniTextPreview": Preview, - "ZuellniTextReplace": Replace, + f"{_MAPPING}Previewer": Previewer, + f"{_MAPPING}Replacer": Replacer, } NODE_DISPLAY_NAME_MAPPINGS = { - "ZuellniTextPreview": "Preview", - "ZuellniTextReplace": "Replace", + f"{_MAPPING}Previewer": "Preview", + f"{_MAPPING}Replacer": "Replace", } WEB_DIRECTORY = "."