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.

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"