From a33a1ac3954812d2b591e358974a9d50e7295a22 Mon Sep 17 00:00:00 2001 From: Zuellni <123005779+Zuellni@users.noreply.github.com> Date: Mon, 20 Nov 2023 21:44:52 +0100 Subject: [PATCH] Add gpu split, add 8bit cache toggle, add min_p, encode specal tokens, update requirements versions --- exllama.py | 93 +++++++++++++++++++++++++++++------------------- requirements.txt | 4 +-- text.js | 22 +++++------- 3 files changed, 67 insertions(+), 52 deletions(-) diff --git a/exllama.py b/exllama.py index 33ba998..db07fb5 100644 --- a/exllama.py +++ b/exllama.py @@ -3,12 +3,10 @@ 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 nodes import MAX_RESOLUTION as MAX +from exllamav2 import ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Cache_8bit, ExLlamaV2Config, ExLlamaV2Tokenizer +from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator class Loader: @@ -17,7 +15,9 @@ class Loader: return { "required": { "model_dir": ("STRING", {"default": ""}), - "max_seq_len": ("INT", {"default": 1024, "max": MAX}), + "gpu_split": ("STRING", {"default": ""}), + "cache_8bit": ("BOOLEAN", {"default": False}), + "max_seq_len": ("INT", {"default": 1024, "max": 2**16}), }, } @@ -32,36 +32,45 @@ class Loader: self.cache = None self.tokenizer = None self.generator = None + self.gpu_split = None + self.cache_8bit = False - def process(self, model_dir, max_seq_len): + def process(self, model_dir, gpu_split, cache_8bit, max_seq_len): self.unload() self.config = ExLlamaV2Config() self.config.model_dir = model_dir self.config.prepare() + if gpu_split: + self.gpu_split = [float(a) for a in gpu_split.split(",")] + if max_seq_len: self.config.max_seq_len = max_seq_len - self.tokenizer = ExLlamaV2Tokenizer(self.config) + self.cache_8bit = cache_8bit 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) + if self.base: + return - self.generator = ExLlamaV2StreamingGenerator( - self.base, - self.cache, - self.tokenizer - ) + self.base = ExLlamaV2(self.config) + self.base.load(gpu_split=self.gpu_split) + + if self.cache_8bit: + self.cache = ExLlamaV2Cache_8bit(self.base) + else: + self.cache = ExLlamaV2Cache(self.base) + + self.tokenizer = ExLlamaV2Tokenizer(self.config) + self.generator = ExLlamaV2StreamingGenerator(self.base, self.cache, self.tokenizer) def unload(self): self.base = None self.cache = None + self.tokenizer = None self.generator = None gc.collect() @@ -75,13 +84,14 @@ class Generator: "required": { "model": ("EXL_MODEL",), "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}), + "single_line": ("BOOLEAN", {"default": False}), + "max_tokens": ("INT", {"default": 128, "max": 2**16}), + "temperature": ("FLOAT", {"default": 1, "max": 2, "step": 0.01}), + "min_p": ("FLOAT", {"default": 0.1, "max": 1, "step": 0.01}), + "top_k": ("INT", {"max": 200}), + "top_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), + "typical": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), + "penalty": ("FLOAT", {"default": 1, "min": 1, "max": 2, "step": 0.01}), "seed": ("INT", {"max": 2**64 - 1}), "text": ("STRING", {"multiline": True}), }, @@ -100,12 +110,13 @@ class Generator: self, model, unload, - stop_on_newline, - max_new_tokens, + single_line, + max_tokens, temperature, + min_p, top_k, top_p, - typical_p, + typical, penalty, seed, text, @@ -116,33 +127,37 @@ class Generator: return ("",) model.load() - input = model.tokenizer.encode(text) - stop_conditions = [model.tokenizer.eos_token_id] + input = model.tokenizer.encode(text, encode_special_tokens=True) + input_len = input.shape[-1] + max_len = model.config.max_seq_len - input_len + stop = [model.tokenizer.eos_token_id] - if not max_new_tokens: - max_new_tokens = model.config.max_seq_len - input.shape[-1] + if not max_tokens or max_tokens > max_len: + max_tokens = max_len - if stop_on_newline: - stop_conditions.append(model.tokenizer.newline_token_id) + if single_line: + stop.append(model.tokenizer.newline_token_id) - model.generator.set_stop_conditions(stop_conditions) + model.generator.set_stop_conditions(stop) + torch.manual_seed(seed) random.seed(seed) settings = ExLlamaV2Sampler.Settings() settings.temperature = temperature + settings.min_p = min_p settings.top_k = top_k settings.top_p = top_p - settings.typical = typical_p + settings.typical = typical settings.token_repetition_penalty = penalty model.generator.begin_stream(input, settings, token_healing=True) - progress = ProgressBar(max_new_tokens) + progress = ProgressBar(max_tokens) start = time() eos = False output = "" tokens = 0 - while not eos and tokens < max_new_tokens: + while not eos and tokens < max_tokens: chunk, eos, _ = model.generator.stream() progress.update(1) output += chunk @@ -151,7 +166,11 @@ class Generator: 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)") + + print( + f"Output generated in {total} seconds", + f"({input_len} context, {tokens} tokens, {speed}t/s)", + ) if unload: model.unload() diff --git a/requirements.txt b/requirements.txt index 3815689..f859570 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,2 +1,2 @@ -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" +exllamav2>=0.0.8; platform_system == "Linux" +https://github.com/turboderp/exllamav2/releases/download/v0.0.8/exllamav2-0.0.8+cu121-cp311-cp311-win_amd64.whl; platform_system == "Windows" diff --git a/text.js b/text.js index 3894b88..2c22395 100644 --- a/text.js +++ b/text.js @@ -2,27 +2,23 @@ import { app } from "../../../scripts/app.js"; import { ComfyWidgets } from "../../../scripts/widgets.js"; app.registerExtension({ - name: "ZuellniTextPreview", + name: "ZuellniText", async beforeRegisterNodeDef(nodeType, nodeData, app) { if (nodeData.name === "ZuellniTextPreview") { - const onExecuted = nodeType.prototype.onExecuted; - - nodeType.prototype.onExecuted = function (message) { - onExecuted?.apply(this, arguments); - + nodeType.prototype.onExecuted = function(message) { if (this.widgets) { - const position = this.widgets.findIndex((w) => w.name === "text"); + const index = this.widgets.findIndex((w) => w.name === "output"); - if (position !== -1) { - for (let i = position; i < this.widgets.length; i++) + if (index !== -1) { + for (let i = index; i < this.widgets.length; i++) this.widgets[i].onRemove?.(); - this.widgets.length = position; + this.widgets.length = index; } - const type = ["STRING", { multiline: true }]; - const widget = ComfyWidgets["STRING"](this, "text", type, app).widget; - + this.widgets.length = 1; + const options = ["STRING", {multiline: true }] + const widget = ComfyWidgets["STRING"](this, "output", options, app).widget; widget.inputEl.readOnly = true; widget.inputEl.style.opacity = 0.7; widget.value = message.text;