diff --git a/README.md b/README.md index 6d94885..a4577ec 100644 --- a/README.md +++ b/README.md @@ -2,38 +2,28 @@ A simple local text generator for [ComfyUI](https://github.com/comfyanonymous/ComfyUI) using [ExLlamaV2](https://github.com/turboderp/exllamav2). ## Installation -Clone the repository to `custom_nodes`: +Clone the repository to `custom_nodes` and install the requirements: ``` git clone https://github.com/Zuellni/ComfyUI-ExLlama-Nodes custom_nodes/ComfyUI-ExLlamaV2-Nodes -``` - -Install the requirements: -``` pip install -r custom_nodes/ComfyUI-ExLlamaV2-Nodes/requirements.txt ``` -On Windows, install one of the precompiled [wheels](https://github.com/turboderp/exllamav2/releases/latest) instead: +Use wheels for [ExLlamaV2](https://github.com/turboderp/exllamav2/releases/latest) and [Flash Attention](https://github.com/bdashore3/flash-attention/releases/latest) on Windows: ``` -pip install https://github.com/turboderp/exllamav2/releases/download/v0.0.xx/exllamav2-0.0.xx+cuXXX-cpXXX-cpXXX-win_amd64.whl +pip install exllamav2-X.X.X+cuXXX.torch2.X.X-cp3XX-cp3XX-win_amd64.whl +pip install flash_attn-X.X.X+cuXXX.torch2.X.X-cp3XX-cp3XX-win_amd64.whl ``` -Check which one you need with: -``` -python -c "import sys, torch; print(f'cu{torch.version.cuda.replace('.', '')}-cp{sys.version_info[0]}{sys.version_info[1]}')" -``` - -> [!CAUTION] -> If you see errors related to ExLlamaV2 while loading the nodes, try to install it following the [official instructions](https://github.com/turboderp/exllamav2#installation). - ## Usage -Only EXL2, 4-bit GPTQ, and unquantized HF models are supported. You can find them on [Hugging Face](https://huggingface.co). See the model card in each repository for details on instruction formats. +Only EXL2, 4-bit GPTQ and unquantized models are supported. You can find them on [Hugging Face](https://huggingface.co). -To use a model with the nodes, you should clone its repository with git or manually download all the files and place them in `models/llm`. -For example, if you'd like to download the 6-bit [Llama-3-8B-Instruct](https://huggingface.co/turboderp/Llama-3-8B-Instruct-exl2), use the following command: +To use a model with the nodes, you should clone its repository with `git` or manually download all the files and place them in `models/llm`. +For example, if you want to download the 6-bit [Llama-3-8B-Instruct](https://huggingface.co/turboderp/Llama-3-8B-Instruct-exl2), use the following command: ``` git install lfs git clone https://huggingface.co/turboderp/Llama-3-8B-Instruct-exl2 -b 6.0bpw models/llm/Llama-3-8B-Instruct-exl2-6.0bpw ``` + > [!TIP] > You can add your own `llm` path to the [extra_model_paths.yaml](https://github.com/comfyanonymous/ComfyUI/blob/master/extra_model_paths.yaml.example) file and put the models there instead. @@ -46,26 +36,36 @@ git clone https://huggingface.co/turboderp/Llama-3-8B-Instruct-exl2 -b 6.0bpw mo cache_bits - Lower value equals lower VRAM usage but also impacts generation speed and quality. + A lower value reduces VRAM usage, but also affects generation speed and quality. + + + + fast_tensors + Enabling reduces RAM usage and speeds up model loading. + + + + flash_attention + Enabling reduces VRAM usage, not supported on cards with compute capability below 8.0. max_seq_len - Max context, higher value equals higher VRAM usage. 0 will default to config. + Max context, higher value equals higher VRAM usage. 0 will default to model config. Generator - Generates text based on the given prompt. Refer to text-generation-webui for parameters. + Generates text based on the given prompt. Refer to SillyTavern for sampler parameters. unload - Unloads the model after each generation. + Unloads the model after each generation to reduce VRAM usage. - single_line - Stops the generation on newline. + stop_conditions + List of strings to stop generation on, e.g. ["\n"] to stop on newline. Leave empty to only stop on eos token. @@ -78,11 +78,11 @@ git clone https://huggingface.co/turboderp/Llama-3-8B-Instruct-exl2 -b 6.0bpw mo Replacer - Replaces variable names enclosed in brackets, eg [a], with their values. + Replaces variable names in brackets, e.g. [a], with their values. ## Workflow -The example workflow is embedded in the image below and can be opened in ComfyUI. +An example workflow is embedded in the image below and can be opened in ComfyUI. ![workflow](https://github.com/Zuellni/ComfyUI-ExLlama-Nodes/assets/123005779/bf688acb-6f7a-4410-98ff-cf22b6937ae7) diff --git a/exllama.py b/exllama.py index 7661f64..37ef347 100644 --- a/exllama.py +++ b/exllama.py @@ -1,4 +1,5 @@ import gc +import json import random from pathlib import Path from time import time @@ -31,7 +32,9 @@ class Loader: return { "required": { "model": (models, {"default": default}), - "cache_bits": ((4, 8, 16), {"default": 16}), + "cache_bits": ((4, 6, 8, 16), {"default": 16}), + "fast_tensors": ("BOOLEAN", {"default": True}), + "flash_attention": ("BOOLEAN", {"default": True}), "max_seq_len": ("INT", {"default": 2048, "max": 2**20}), }, } @@ -42,10 +45,12 @@ class Loader: RETURN_NAMES = ("MODEL",) RETURN_TYPES = ("EXL_MODEL",) - def setup(self, model, cache_bits, max_seq_len): + def setup(self, model, cache_bits, fast_tensors, flash_attention, max_seq_len): self.unload() self.cache_bits = cache_bits self.config = ExLlamaV2Config(__class__._MODELS[model]) + self.config.fasttensors = fast_tensors + self.config.no_flash_attn = not flash_attention if max_seq_len: self.config.max_seq_len = max_seq_len @@ -75,14 +80,22 @@ class Loader: self.cache = ( ExLlamaV2Cache_Q4(self.model, lazy=True) if self.cache_bits == 4 - else ExLlamaV2Cache_8bit(self.model, lazy=True) + else ExLlamaV2Cache_Q6(self.model, lazy=True) + if self.cache_bits == 6 + else ExLlamaV2Cache_Q8(self.model, lazy=True) if self.cache_bits == 8 else ExLlamaV2Cache(self.model, lazy=True) ) self.tokenizer = ExLlamaV2Tokenizer(self.config) self.model.load_autosplit(self.cache, callback=lambda _, __: progress.update(1)) - self.generator = ExLlamaV2StreamingGenerator(self.model, self.cache, self.tokenizer) + + self.generator = ExLlamaV2DynamicGenerator( + model=self.model, + cache=self.cache, + tokenizer=self.tokenizer, + paged=not self.config.no_flash_attn, + ) def unload(self): if hasattr(self, "model") and self.model: @@ -104,7 +117,7 @@ class Generator: "required": { "model": ("EXL_MODEL",), "unload": ("BOOLEAN", {"default": False}), - "single_line": ("BOOLEAN", {"default": False}), + "stop_conditions": ("STRING", {"default": r'["\n"]'}), "max_tokens": ("INT", {"default": 128, "max": 2**20}), "temperature": ("FLOAT", {"default": 1, "max": 5, "step": 0.01}), "top_k": ("INT", {"max": 200}), @@ -132,7 +145,7 @@ class Generator: self, model, unload, - single_line, + stop_conditions, max_tokens, temperature, top_k, @@ -147,7 +160,7 @@ class Generator: info=None, id=None, ): - if not text: + if not text.strip(): return ("",) if unload: @@ -155,6 +168,7 @@ class Generator: model.unload() model.load() + random.seed(seed) input = model.tokenizer.encode(text, encode_special_tokens=True) input_len = input.shape[-1] max_len = model.config.max_seq_len - input_len @@ -163,11 +177,9 @@ class Generator: if not max_tokens or max_tokens > max_len: max_tokens = max_len - if single_line: - stop.append(model.tokenizer.newline_token_id) - - model.generator.set_stop_conditions(stop) - random.seed(seed) + if stop_conditions.strip(): + stop_conditions = json.loads(stop_conditions) + stop.extend(stop_conditions) settings = ExLlamaV2Sampler.Settings() settings.temperature = temperature @@ -179,21 +191,29 @@ class Generator: settings.token_repetition_penalty = repetition_penalty settings.temperature_last = temperature_last - start = time() - model.generator.begin_stream_ex(input, settings) + job = ExLlamaV2DynamicJob( + input_ids=input, + max_new_tokens=max_tokens, + stop_conditions=stop, + ) + progress = ProgressBar(max_tokens) + model.generator.enqueue(job) + start = time() eos = False - output = "" + chunks = [] tokens = 0 - while not eos and tokens < max_tokens: - response = model.generator.stream_ex() - output += response["chunk"] - eos = response["eos"] - progress.update(1) - tokens += 1 + while not eos: + for response in model.generator.iterate(): + if response["stage"] == "streaming": + chunk = response.get("text", "") + eos = response["eos"] + chunks.append(chunk) + progress.update(1) + tokens += 1 - output = output.strip() + output = "".join(chunks).strip() total = round(time() - start, 2) speed = round(tokens / total, 2) diff --git a/requirements.txt b/requirements.txt index edbf7cb..04c8177 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ -exllamav2>=0.0.17; platform_system == "Linux" +exllamav2>=0.1.5; platform_system == "Linux" +flash-attn>=2.5.7; platform_system == "Linux"